fix(api): delete custom models stored with legacy model_type values (#39708)

Co-authored-by: Crazywoola <100913391+crazywoola@users.noreply.github.com>
This commit is contained in:
Souravrajvi0 2026-07-30 06:42:26 +05:30 committed by GitHub
parent 89ceba027e
commit fdaa073d53
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 53 additions and 5 deletions

View File

@ -57,6 +57,19 @@ logger = logging.getLogger(__name__)
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.
"""
values = [model_type.value]
origin = model_type.to_origin_model_type()
if origin not in values:
values.append(origin)
return tuple(values)
class ProviderConfiguration(BaseModel):
"""
Provider configuration entity for managing model provider settings.
@ -846,7 +859,7 @@ class ProviderConfiguration(BaseModel):
ProviderModel.tenant_id == self.tenant_id,
ProviderModel.provider_name.in_(provider_names),
ProviderModel.model_name == model,
ProviderModel.model_type == model_type,
ProviderModel.model_type.in_(_model_type_db_values(model_type)),
)
return session.execute(stmt).scalar_one_or_none()
@ -871,7 +884,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 == model_type,
ProviderModelCredential.model_type.in_(_model_type_db_values(model_type)),
)
credential_record = session.execute(stmt).scalar_one_or_none()
@ -1184,7 +1197,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 == model_type,
ProviderModelCredential.model_type.in_(_model_type_db_values(model_type)),
)
credential_record = session.execute(stmt).scalar_one_or_none()
if not credential_record:
@ -1228,7 +1241,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 == model_type,
ProviderModelCredential.model_type.in_(_model_type_db_values(model_type)),
)
available_credentials_count = session.execute(count_stmt).scalar() or 0
session.delete(credential_record)
@ -1394,7 +1407,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 == model_type,
ProviderModelSetting.model_type.in_(_model_type_db_values(model_type)),
ProviderModelSetting.model_name == model,
)
return session.execute(stmt).scalars().first()

View File

@ -1156,6 +1156,41 @@ def test_get_custom_model_record_supports_plugin_id_alias() -> None:
assert result is custom_model_record
def test_model_type_db_values_includes_pre_1_15_aliases() -> None:
from core.entities.provider_configuration import _model_type_db_values
assert _model_type_db_values(ModelType.LLM) == ("llm", "text-generation")
assert _model_type_db_values(ModelType.TEXT_EMBEDDING) == ("text-embedding", "embeddings")
assert _model_type_db_values(ModelType.RERANK) == ("rerank", "reranking")
assert _model_type_db_values(ModelType.TTS) == ("tts",)
def test_get_custom_model_record_uses_legacy_aware_model_type_filter(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regression for #39559: lookups must include pre-1.15 model_type aliases."""
import core.entities.provider_configuration as provider_configuration_module
captured: dict[str, tuple[str, ...]] = {}
original = provider_configuration_module._model_type_db_values
def _capture(model_type: ModelType) -> tuple[str, ...]:
values = original(model_type)
captured["values"] = values
return values
monkeypatch.setattr(provider_configuration_module, "_model_type_db_values", _capture)
configuration = _build_provider_configuration(provider_name="langgenius/ollama/ollama")
session = Mock()
session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace(id="legacy-model")
result = configuration._get_custom_model_record(ModelType.LLM, "llama3", session)
assert result.id == "legacy-model"
assert captured["values"] == ("llm", "text-generation")
def test_get_specific_custom_model_credential_success_and_not_found() -> None:
configuration = _build_provider_configuration()
configuration.provider.model_credential_schema = _build_secret_model_schema()