mirror of
https://github.com/langgenius/dify.git
synced 2026-07-30 16:59:35 +08:00
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:
parent
89ceba027e
commit
fdaa073d53
@ -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()
|
||||
|
||||
@ -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()
|
||||
|
||||
Loading…
Reference in New Issue
Block a user