refactor(services): migrate datasource_provider_service to SQLAlchemy 2.0 select() API (#34974)

This commit is contained in:
wdeveloper16
2026-04-12 01:23:24 +00:00
committed by GitHub
parent 7515eee0a8
commit 7bd5e80323
2 changed files with 205 additions and 170 deletions
+166 -120
View File
@@ -4,7 +4,7 @@ from collections.abc import Mapping
from typing import Any
from graphon.model_runtime.entities.provider_entities import FormType
from sqlalchemy import func, select
from sqlalchemy import delete, func, select, update
from sqlalchemy.orm import Session, sessionmaker
from configs import dify_config
@@ -54,11 +54,13 @@ class DatasourceProviderService:
remove oauth custom client params
"""
with sessionmaker(bind=db.engine).begin() as session:
session.query(DatasourceOauthTenantParamConfig).filter_by(
tenant_id=tenant_id,
provider=datasource_provider_id.provider_name,
plugin_id=datasource_provider_id.plugin_id,
).delete()
session.execute(
delete(DatasourceOauthTenantParamConfig).where(
DatasourceOauthTenantParamConfig.tenant_id == tenant_id,
DatasourceOauthTenantParamConfig.provider == datasource_provider_id.provider_name,
DatasourceOauthTenantParamConfig.plugin_id == datasource_provider_id.plugin_id,
)
)
def decrypt_datasource_provider_credentials(
self,
@@ -110,15 +112,21 @@ class DatasourceProviderService:
"""
with sessionmaker(bind=db.engine).begin() as session:
if credential_id:
datasource_provider = (
session.query(DatasourceProvider).filter_by(tenant_id=tenant_id, id=credential_id).first()
datasource_provider = session.scalar(
select(DatasourceProvider)
.where(DatasourceProvider.tenant_id == tenant_id, DatasourceProvider.id == credential_id)
.limit(1)
)
else:
datasource_provider = (
session.query(DatasourceProvider)
.filter_by(tenant_id=tenant_id, provider=provider, plugin_id=plugin_id)
datasource_provider = session.scalar(
select(DatasourceProvider)
.where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.provider == provider,
DatasourceProvider.plugin_id == plugin_id,
)
.order_by(DatasourceProvider.is_default.desc(), DatasourceProvider.created_at.asc())
.first()
.limit(1)
)
if not datasource_provider:
return {}
@@ -173,12 +181,15 @@ class DatasourceProviderService:
get all datasource credentials by provider
"""
with sessionmaker(bind=db.engine).begin() as session:
datasource_providers = (
session.query(DatasourceProvider)
.filter_by(tenant_id=tenant_id, provider=provider, plugin_id=plugin_id)
datasource_providers = session.scalars(
select(DatasourceProvider)
.where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.provider == provider,
DatasourceProvider.plugin_id == plugin_id,
)
.order_by(DatasourceProvider.is_default.desc(), DatasourceProvider.created_at.asc())
.all()
)
).all()
if not datasource_providers:
return []
current_user = get_current_user()
@@ -232,15 +243,15 @@ class DatasourceProviderService:
update datasource provider name
"""
with sessionmaker(bind=db.engine).begin() as session:
target_provider = (
session.query(DatasourceProvider)
.filter_by(
tenant_id=tenant_id,
id=credential_id,
provider=datasource_provider_id.provider_name,
plugin_id=datasource_provider_id.plugin_id,
target_provider = session.scalar(
select(DatasourceProvider)
.where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.id == credential_id,
DatasourceProvider.provider == datasource_provider_id.provider_name,
DatasourceProvider.plugin_id == datasource_provider_id.plugin_id,
)
.first()
.limit(1)
)
if target_provider is None:
raise ValueError("provider not found")
@@ -250,16 +261,16 @@ class DatasourceProviderService:
# check name is exist
if (
session.query(DatasourceProvider)
.filter_by(
tenant_id=tenant_id,
name=name,
provider=datasource_provider_id.provider_name,
plugin_id=datasource_provider_id.plugin_id,
session.scalar(
select(func.count(DatasourceProvider.id)).where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.name == name,
DatasourceProvider.provider == datasource_provider_id.provider_name,
DatasourceProvider.plugin_id == datasource_provider_id.plugin_id,
)
)
.count()
> 0
):
or 0
) > 0:
raise ValueError("Authorization name is already exists")
target_provider.name = name
@@ -273,26 +284,31 @@ class DatasourceProviderService:
"""
with sessionmaker(bind=db.engine).begin() as session:
# get provider
target_provider = (
session.query(DatasourceProvider)
.filter_by(
tenant_id=tenant_id,
id=credential_id,
provider=datasource_provider_id.provider_name,
plugin_id=datasource_provider_id.plugin_id,
target_provider = session.scalar(
select(DatasourceProvider)
.where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.id == credential_id,
DatasourceProvider.provider == datasource_provider_id.provider_name,
DatasourceProvider.plugin_id == datasource_provider_id.plugin_id,
)
.first()
.limit(1)
)
if target_provider is None:
raise ValueError("provider not found")
# clear default provider
session.query(DatasourceProvider).filter_by(
tenant_id=tenant_id,
provider=target_provider.provider,
plugin_id=target_provider.plugin_id,
is_default=True,
).update({"is_default": False})
session.execute(
update(DatasourceProvider)
.where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.provider == target_provider.provider,
DatasourceProvider.plugin_id == target_provider.plugin_id,
DatasourceProvider.is_default.is_(True),
)
.values(is_default=False)
.execution_options(synchronize_session=False)
)
# set new default provider
target_provider.is_default = True
@@ -311,14 +327,14 @@ class DatasourceProviderService:
if client_params is None and enabled is None:
return
with sessionmaker(bind=db.engine).begin() as session:
tenant_oauth_client_params = (
session.query(DatasourceOauthTenantParamConfig)
.filter_by(
tenant_id=tenant_id,
provider=datasource_provider_id.provider_name,
plugin_id=datasource_provider_id.plugin_id,
tenant_oauth_client_params = session.scalar(
select(DatasourceOauthTenantParamConfig)
.where(
DatasourceOauthTenantParamConfig.tenant_id == tenant_id,
DatasourceOauthTenantParamConfig.provider == datasource_provider_id.provider_name,
DatasourceOauthTenantParamConfig.plugin_id == datasource_provider_id.plugin_id,
)
.first()
.limit(1)
)
if not tenant_oauth_client_params:
@@ -351,9 +367,14 @@ class DatasourceProviderService:
"""
with Session(db.engine).no_autoflush as session:
return (
session.query(DatasourceOauthParamConfig)
.filter_by(provider=datasource_provider_id.provider_name, plugin_id=datasource_provider_id.plugin_id)
.first()
session.scalar(
select(DatasourceOauthParamConfig)
.where(
DatasourceOauthParamConfig.provider == datasource_provider_id.provider_name,
DatasourceOauthParamConfig.plugin_id == datasource_provider_id.plugin_id,
)
.limit(1)
)
is not None
)
@@ -423,15 +444,15 @@ class DatasourceProviderService:
plugin_id = datasource_provider_id.plugin_id
with Session(db.engine).no_autoflush as session:
# get tenant oauth client params
tenant_oauth_client_params = (
session.query(DatasourceOauthTenantParamConfig)
.filter_by(
tenant_id=tenant_id,
provider=provider,
plugin_id=plugin_id,
enabled=True,
tenant_oauth_client_params = session.scalar(
select(DatasourceOauthTenantParamConfig)
.where(
DatasourceOauthTenantParamConfig.tenant_id == tenant_id,
DatasourceOauthTenantParamConfig.provider == provider,
DatasourceOauthTenantParamConfig.plugin_id == plugin_id,
DatasourceOauthTenantParamConfig.enabled.is_(True),
)
.first()
.limit(1)
)
if tenant_oauth_client_params:
encrypter, _ = self.get_oauth_encrypter(tenant_id, datasource_provider_id)
@@ -443,8 +464,13 @@ class DatasourceProviderService:
is_verified = PluginService.is_plugin_verified(tenant_id, provider_controller.plugin_unique_identifier)
if is_verified:
# fallback to system oauth client params
oauth_client_params = (
session.query(DatasourceOauthParamConfig).filter_by(provider=provider, plugin_id=plugin_id).first()
oauth_client_params = session.scalar(
select(DatasourceOauthParamConfig)
.where(
DatasourceOauthParamConfig.provider == provider,
DatasourceOauthParamConfig.plugin_id == plugin_id,
)
.limit(1)
)
if oauth_client_params:
return oauth_client_params.system_credentials
@@ -455,15 +481,13 @@ class DatasourceProviderService:
def generate_next_datasource_provider_name(
session: Session, tenant_id: str, provider_id: DatasourceProviderID, credential_type: CredentialType
) -> str:
db_providers = (
session.query(DatasourceProvider)
.filter_by(
tenant_id=tenant_id,
provider=provider_id.provider_name,
plugin_id=provider_id.plugin_id,
db_providers = session.scalars(
select(DatasourceProvider).where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.provider == provider_id.provider_name,
DatasourceProvider.plugin_id == provider_id.plugin_id,
)
.all()
)
).all()
return generate_incremental_name(
[provider.name for provider in db_providers],
f"{credential_type.get_name()}",
@@ -485,8 +509,10 @@ class DatasourceProviderService:
with sessionmaker(bind=db.engine).begin() as session:
lock = f"datasource_provider_create_lock:{tenant_id}_{provider_id}_{CredentialType.OAUTH2.value}"
with redis_client.lock(lock, timeout=20):
target_provider = (
session.query(DatasourceProvider).filter_by(id=credential_id, tenant_id=tenant_id).first()
target_provider = session.scalar(
select(DatasourceProvider)
.where(DatasourceProvider.id == credential_id, DatasourceProvider.tenant_id == tenant_id)
.limit(1)
)
if target_provider is None:
raise ValueError("provider not found")
@@ -496,25 +522,28 @@ class DatasourceProviderService:
db_provider_name = target_provider.name
else:
name_conflict = (
session.query(DatasourceProvider)
.filter_by(
tenant_id=tenant_id,
name=db_provider_name,
provider=provider_id.provider_name,
plugin_id=provider_id.plugin_id,
auth_type=CredentialType.OAUTH2.value,
session.scalar(
select(func.count(DatasourceProvider.id)).where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.name == db_provider_name,
DatasourceProvider.provider == provider_id.provider_name,
DatasourceProvider.plugin_id == provider_id.plugin_id,
DatasourceProvider.auth_type == CredentialType.OAUTH2.value,
)
)
.count()
or 0
)
if name_conflict > 0:
db_provider_name = generate_incremental_name(
[
provider.name
for provider in session.query(DatasourceProvider).filter_by(
tenant_id=tenant_id,
provider=provider_id.provider_name,
plugin_id=provider_id.plugin_id,
)
for provider in session.scalars(
select(DatasourceProvider).where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.provider == provider_id.provider_name,
DatasourceProvider.plugin_id == provider_id.plugin_id,
)
).all()
],
db_provider_name,
)
@@ -556,25 +585,27 @@ class DatasourceProviderService:
)
else:
if (
session.query(DatasourceProvider)
.filter_by(
tenant_id=tenant_id,
name=db_provider_name,
provider=provider_id.provider_name,
plugin_id=provider_id.plugin_id,
auth_type=credential_type.value,
session.scalar(
select(func.count(DatasourceProvider.id)).where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.name == db_provider_name,
DatasourceProvider.provider == provider_id.provider_name,
DatasourceProvider.plugin_id == provider_id.plugin_id,
DatasourceProvider.auth_type == credential_type.value,
)
)
.count()
> 0
):
or 0
) > 0:
db_provider_name = generate_incremental_name(
[
provider.name
for provider in session.query(DatasourceProvider).filter_by(
tenant_id=tenant_id,
provider=provider_id.provider_name,
plugin_id=provider_id.plugin_id,
)
for provider in session.scalars(
select(DatasourceProvider).where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.provider == provider_id.provider_name,
DatasourceProvider.plugin_id == provider_id.plugin_id,
)
).all()
],
db_provider_name,
)
@@ -627,11 +658,16 @@ class DatasourceProviderService:
# check name is exist
if (
session.query(DatasourceProvider)
.filter_by(tenant_id=tenant_id, plugin_id=plugin_id, provider=provider_name, name=db_provider_name)
.count()
> 0
):
session.scalar(
select(func.count(DatasourceProvider.id)).where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.plugin_id == plugin_id,
DatasourceProvider.provider == provider_name,
DatasourceProvider.name == db_provider_name,
)
)
or 0
) > 0:
raise ValueError("Authorization name is already exists")
try:
@@ -918,21 +954,31 @@ class DatasourceProviderService:
"""
with sessionmaker(bind=db.engine).begin() as session:
datasource_provider = (
session.query(DatasourceProvider)
.filter_by(tenant_id=tenant_id, id=auth_id, provider=provider, plugin_id=plugin_id)
.first()
datasource_provider = session.scalar(
select(DatasourceProvider)
.where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.id == auth_id,
DatasourceProvider.provider == provider,
DatasourceProvider.plugin_id == plugin_id,
)
.limit(1)
)
if not datasource_provider:
raise ValueError("Datasource provider not found")
# update name
if name and name != datasource_provider.name:
if (
session.query(DatasourceProvider)
.filter_by(tenant_id=tenant_id, name=name, provider=provider, plugin_id=plugin_id)
.count()
> 0
):
session.scalar(
select(func.count(DatasourceProvider.id)).where(
DatasourceProvider.tenant_id == tenant_id,
DatasourceProvider.name == name,
DatasourceProvider.provider == provider,
DatasourceProvider.plugin_id == plugin_id,
)
)
or 0
) > 0:
raise ValueError("Authorization name is already exists")
datasource_provider.name = name