test: use SQLite sessions in services core (#39058)

This commit is contained in:
Asuka Minato
2026-08-06 04:07:01 +00:00
committed by GitHub
parent 2443df2e42
commit f08cc41546
2 changed files with 422 additions and 768 deletions
+19 -8
View File
@@ -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