mirror of
https://github.com/langgenius/dify.git
synced 2026-09-19 10:11:30 +08:00
test: use SQLite sessions in services core (#39058)
This commit is contained in:
@@ -9,8 +9,9 @@ from json import JSONDecodeError
|
||||
from typing import Any, override
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, model_validator
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy import String, func, literal, select
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.sql.elements import BindParameter
|
||||
|
||||
from constants import HIDDEN_VALUE
|
||||
from core.entities import PluginCredentialType
|
||||
@@ -60,8 +61,8 @@ original_provider_configurate_methods: dict[str, list[ConfigurateMethod]] = {}
|
||||
def _model_type_db_values(model_type: ModelType) -> tuple[str, ...]:
|
||||
"""Return DB values that may represent ``model_type`` after pre-1.15 upgrades.
|
||||
|
||||
Reads normalize legacy values (``text-generation`` → ``llm``) via ``EnumText``,
|
||||
but SQL equality against the canonical value misses unmigrated rows. Match both.
|
||||
SQL equality against the canonical value misses unmigrated rows, so legacy
|
||||
lookups need to match both the current and provider-native spellings.
|
||||
"""
|
||||
values = [model_type.value]
|
||||
origin = model_type.to_origin_model_type()
|
||||
@@ -70,6 +71,16 @@ def _model_type_db_values(model_type: ModelType) -> tuple[str, ...]:
|
||||
return tuple(values)
|
||||
|
||||
|
||||
def _model_type_db_literals(model_type: ModelType) -> tuple[BindParameter[str], ...]:
|
||||
"""Return string-typed literals for legacy-aware model type filters.
|
||||
|
||||
``EnumText`` rejects legacy spellings during normal binding so they cannot
|
||||
be written back. Explicit string literals bypass that bind processor only
|
||||
for compatibility lookups of rows that predate the canonical enum values.
|
||||
"""
|
||||
return tuple(literal(value, type_=String) for value in _model_type_db_values(model_type))
|
||||
|
||||
|
||||
class ProviderConfiguration(BaseModel):
|
||||
"""
|
||||
Provider configuration entity for managing model provider settings.
|
||||
@@ -859,7 +870,7 @@ class ProviderConfiguration(BaseModel):
|
||||
ProviderModel.tenant_id == self.tenant_id,
|
||||
ProviderModel.provider_name.in_(provider_names),
|
||||
ProviderModel.model_name == model,
|
||||
ProviderModel.model_type.in_(_model_type_db_values(model_type)),
|
||||
ProviderModel.model_type.in_(_model_type_db_literals(model_type)),
|
||||
)
|
||||
|
||||
return session.execute(stmt).scalar_one_or_none()
|
||||
@@ -884,7 +895,7 @@ class ProviderConfiguration(BaseModel):
|
||||
ProviderModelCredential.tenant_id == self.tenant_id,
|
||||
ProviderModelCredential.provider_name.in_(self._get_provider_names()),
|
||||
ProviderModelCredential.model_name == model,
|
||||
ProviderModelCredential.model_type.in_(_model_type_db_values(model_type)),
|
||||
ProviderModelCredential.model_type.in_(_model_type_db_literals(model_type)),
|
||||
)
|
||||
|
||||
credential_record = session.execute(stmt).scalar_one_or_none()
|
||||
@@ -1197,7 +1208,7 @@ class ProviderConfiguration(BaseModel):
|
||||
ProviderModelCredential.tenant_id == self.tenant_id,
|
||||
ProviderModelCredential.provider_name.in_(self._get_provider_names()),
|
||||
ProviderModelCredential.model_name == model,
|
||||
ProviderModelCredential.model_type.in_(_model_type_db_values(model_type)),
|
||||
ProviderModelCredential.model_type.in_(_model_type_db_literals(model_type)),
|
||||
)
|
||||
credential_record = session.execute(stmt).scalar_one_or_none()
|
||||
if not credential_record:
|
||||
@@ -1241,7 +1252,7 @@ class ProviderConfiguration(BaseModel):
|
||||
ProviderModelCredential.tenant_id == self.tenant_id,
|
||||
ProviderModelCredential.provider_name.in_(self._get_provider_names()),
|
||||
ProviderModelCredential.model_name == model,
|
||||
ProviderModelCredential.model_type.in_(_model_type_db_values(model_type)),
|
||||
ProviderModelCredential.model_type.in_(_model_type_db_literals(model_type)),
|
||||
)
|
||||
available_credentials_count = session.execute(count_stmt).scalar() or 0
|
||||
session.delete(credential_record)
|
||||
@@ -1407,7 +1418,7 @@ class ProviderConfiguration(BaseModel):
|
||||
stmt = select(ProviderModelSetting).where(
|
||||
ProviderModelSetting.tenant_id == self.tenant_id,
|
||||
ProviderModelSetting.provider_name.in_(self._get_provider_names()),
|
||||
ProviderModelSetting.model_type.in_(_model_type_db_values(model_type)),
|
||||
ProviderModelSetting.model_type.in_(_model_type_db_literals(model_type)),
|
||||
ProviderModelSetting.model_name == model,
|
||||
)
|
||||
return session.execute(stmt).scalars().first()
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user