From fdaa073d5353d9d5aa605cfd17a023c9a2be6c2d Mon Sep 17 00:00:00 2001 From: Souravrajvi0 <144546710+Souravrajvi0@users.noreply.github.com> Date: Thu, 30 Jul 2026 06:42:26 +0530 Subject: [PATCH] fix(api): delete custom models stored with legacy model_type values (#39708) Co-authored-by: Crazywoola <100913391+crazywoola@users.noreply.github.com> --- api/core/entities/provider_configuration.py | 23 +++++++++--- .../test_entities_provider_configuration.py | 35 +++++++++++++++++++ 2 files changed, 53 insertions(+), 5 deletions(-) diff --git a/api/core/entities/provider_configuration.py b/api/core/entities/provider_configuration.py index 95b8b686e05..7154b0db9cd 100644 --- a/api/core/entities/provider_configuration.py +++ b/api/core/entities/provider_configuration.py @@ -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() diff --git a/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py b/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py index b5a48918079..bff91d2aea3 100644 --- a/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py +++ b/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py @@ -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()