mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
perf(api): reduce workflow startup latency for chatflow (#36773)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
@@ -433,10 +433,9 @@ def test_get_model_type_instance_and_schema_delegate_to_factory() -> None:
|
||||
mock_model_type_instance = Mock()
|
||||
mock_schema = _build_ai_model("gpt-4o")
|
||||
mock_factory = Mock()
|
||||
mock_factory.get_provider_schema.return_value = configuration.provider
|
||||
mock_factory.get_model_schema.return_value = mock_schema
|
||||
mock_assembly = Mock()
|
||||
mock_assembly.model_runtime = Mock()
|
||||
mock_assembly.model_runtime.get_model_schema.return_value = mock_schema
|
||||
mock_assembly.model_provider_factory = mock_factory
|
||||
|
||||
with (
|
||||
@@ -455,13 +454,12 @@ def test_get_model_type_instance_and_schema_delegate_to_factory() -> None:
|
||||
assert model_type_instance is mock_model_type_instance
|
||||
assert model_schema is mock_schema
|
||||
assert mock_assembly_builder.call_count == 2
|
||||
mock_factory.get_provider_schema.assert_called_once_with(provider="openai")
|
||||
mock_model_builder.assert_called_once_with(
|
||||
runtime=mock_assembly.model_runtime,
|
||||
provider_schema=configuration.provider,
|
||||
model_type=ModelType.LLM,
|
||||
)
|
||||
mock_factory.get_model_schema.assert_called_once_with(
|
||||
mock_assembly.model_runtime.get_model_schema.assert_called_once_with(
|
||||
provider="openai",
|
||||
model_type=ModelType.LLM,
|
||||
model="gpt-4o",
|
||||
@@ -472,18 +470,13 @@ def test_get_model_type_instance_and_schema_delegate_to_factory() -> None:
|
||||
def test_get_model_type_instance_and_schema_reuse_bound_runtime_factory() -> None:
|
||||
configuration = _build_provider_configuration()
|
||||
bound_runtime = Mock()
|
||||
bound_runtime.get_model_schema.return_value = _build_ai_model("gpt-4o")
|
||||
configuration.bind_model_runtime(bound_runtime)
|
||||
|
||||
mock_model_type_instance = Mock()
|
||||
mock_schema = _build_ai_model("gpt-4o")
|
||||
mock_factory = Mock()
|
||||
mock_factory.get_provider_schema.return_value = configuration.provider
|
||||
mock_factory.get_model_schema.return_value = mock_schema
|
||||
|
||||
with (
|
||||
patch(
|
||||
"core.entities.provider_configuration.ModelProviderFactory", return_value=mock_factory
|
||||
) as mock_factory_cls,
|
||||
patch("core.entities.provider_configuration.ModelProviderFactory") as mock_factory_cls,
|
||||
patch("core.entities.provider_configuration.create_plugin_model_assembly") as mock_assembly_builder,
|
||||
patch(
|
||||
"core.entities.provider_configuration.create_model_type_instance",
|
||||
@@ -494,16 +487,20 @@ def test_get_model_type_instance_and_schema_reuse_bound_runtime_factory() -> Non
|
||||
model_schema = configuration.get_model_schema(ModelType.LLM, "gpt-4o", {"api_key": "x"})
|
||||
|
||||
assert model_type_instance is mock_model_type_instance
|
||||
assert model_schema is mock_schema
|
||||
assert mock_factory_cls.call_count == 2
|
||||
mock_factory_cls.assert_called_with(runtime=bound_runtime)
|
||||
assert model_schema == bound_runtime.get_model_schema.return_value
|
||||
mock_factory_cls.assert_not_called()
|
||||
mock_assembly_builder.assert_not_called()
|
||||
mock_factory.get_provider_schema.assert_called_once_with(provider="openai")
|
||||
mock_model_builder.assert_called_once_with(
|
||||
runtime=bound_runtime,
|
||||
provider_schema=configuration.provider,
|
||||
model_type=ModelType.LLM,
|
||||
)
|
||||
bound_runtime.get_model_schema.assert_called_once_with(
|
||||
provider="openai",
|
||||
model_type=ModelType.LLM,
|
||||
model="gpt-4o",
|
||||
credentials={"api_key": "x"},
|
||||
)
|
||||
|
||||
|
||||
def test_get_provider_model_returns_none_when_model_not_found() -> None:
|
||||
@@ -544,6 +541,99 @@ def test_get_provider_models_system_deduplicates_sorts_and_filters_active() -> N
|
||||
assert [model.model for model in active_models] == ["b-model"]
|
||||
|
||||
|
||||
def test_get_provider_models_system_filters_requested_model() -> None:
|
||||
configuration = _build_provider_configuration()
|
||||
provider_schema = ProviderEntity(
|
||||
provider="openai",
|
||||
label=I18nObject(en_US="OpenAI"),
|
||||
supported_model_types=[ModelType.LLM],
|
||||
configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL],
|
||||
models=[_build_ai_model("a-model"), _build_ai_model("target-model"), _build_ai_model("b-model")],
|
||||
)
|
||||
mock_factory = Mock()
|
||||
mock_factory.get_provider_schema.return_value = provider_schema
|
||||
|
||||
with patch(
|
||||
"core.entities.provider_configuration.create_plugin_model_assembly",
|
||||
return_value=SimpleNamespace(model_runtime=Mock(), model_provider_factory=mock_factory),
|
||||
):
|
||||
models = configuration.get_provider_models(
|
||||
model_type=ModelType.LLM,
|
||||
only_active=False,
|
||||
model="target-model",
|
||||
)
|
||||
|
||||
assert [model.model for model in models] == ["target-model"]
|
||||
|
||||
|
||||
def test_get_provider_models_system_customizable_filters_requested_restricted_model() -> None:
|
||||
provider = ProviderEntity(
|
||||
provider="openai",
|
||||
label=I18nObject(en_US="OpenAI"),
|
||||
supported_model_types=[ModelType.LLM],
|
||||
configurate_methods=[ConfigurateMethod.CUSTOMIZABLE_MODEL],
|
||||
)
|
||||
system_configuration = SystemConfiguration(
|
||||
enabled=True,
|
||||
credentials={"api_key": "test-key"},
|
||||
current_quota_type=ProviderQuotaType.TRIAL,
|
||||
quota_configurations=[
|
||||
QuotaConfiguration(
|
||||
quota_type=ProviderQuotaType.TRIAL,
|
||||
quota_unit=QuotaUnit.TOKENS,
|
||||
quota_limit=1_000,
|
||||
quota_used=0,
|
||||
is_valid=True,
|
||||
restrict_models=[
|
||||
RestrictModel(model="target-model", base_model_name="base-model", model_type=ModelType.LLM),
|
||||
RestrictModel(model="other-model", base_model_name="base-model", model_type=ModelType.LLM),
|
||||
],
|
||||
)
|
||||
],
|
||||
)
|
||||
provider_schema = ProviderEntity(
|
||||
provider="openai",
|
||||
label=I18nObject(en_US="OpenAI"),
|
||||
supported_model_types=[ModelType.LLM],
|
||||
configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL],
|
||||
models=[],
|
||||
)
|
||||
mock_factory = Mock()
|
||||
mock_factory.get_provider_schema.return_value = provider_schema
|
||||
|
||||
with patch("core.entities.provider_configuration.original_provider_configurate_methods", {}):
|
||||
configuration = ProviderConfiguration(
|
||||
tenant_id="tenant-1",
|
||||
provider=provider,
|
||||
preferred_provider_type=ProviderType.SYSTEM,
|
||||
using_provider_type=ProviderType.SYSTEM,
|
||||
system_configuration=system_configuration,
|
||||
custom_configuration=CustomConfiguration(provider=None, models=[]),
|
||||
model_settings=[],
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"core.entities.provider_configuration.create_plugin_model_assembly",
|
||||
return_value=SimpleNamespace(model_runtime=Mock(), model_provider_factory=mock_factory),
|
||||
),
|
||||
patch.object(
|
||||
ProviderConfiguration,
|
||||
"get_model_schema",
|
||||
side_effect=lambda *args, **kwargs: _build_ai_model(kwargs["model"]),
|
||||
) as mock_get_model_schema,
|
||||
):
|
||||
models = configuration.get_provider_models(
|
||||
model_type=ModelType.LLM,
|
||||
only_active=False,
|
||||
model="target-model",
|
||||
)
|
||||
|
||||
assert [model.model for model in models] == ["target-model"]
|
||||
mock_get_model_schema.assert_called_once()
|
||||
assert mock_get_model_schema.call_args.kwargs["model"] == "target-model"
|
||||
|
||||
|
||||
def test_get_custom_provider_models_sets_status_for_removed_credentials_and_invalid_lb_configs() -> None:
|
||||
configuration = _build_provider_configuration()
|
||||
configuration.using_provider_type = ProviderType.CUSTOM
|
||||
@@ -611,6 +701,48 @@ def test_get_custom_provider_models_sets_status_for_removed_credentials_and_inva
|
||||
assert invalid_lb_map["custom-model"] is True
|
||||
|
||||
|
||||
def test_get_custom_provider_models_filters_requested_base_model() -> None:
|
||||
configuration = _build_provider_configuration()
|
||||
configuration.using_provider_type = ProviderType.CUSTOM
|
||||
configuration.custom_configuration.provider = CustomProviderConfiguration(credentials={"api_key": "provider-key"})
|
||||
provider_schema = ProviderEntity(
|
||||
provider="openai",
|
||||
label=I18nObject(en_US="OpenAI"),
|
||||
supported_model_types=[ModelType.LLM],
|
||||
configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL],
|
||||
models=[_build_ai_model("base-model"), _build_ai_model("target-model")],
|
||||
)
|
||||
|
||||
models = configuration._get_custom_provider_models(
|
||||
model_types=[ModelType.LLM],
|
||||
provider_schema=provider_schema,
|
||||
model_setting_map={},
|
||||
model="target-model",
|
||||
)
|
||||
|
||||
assert [model.model for model in models] == ["target-model"]
|
||||
|
||||
|
||||
def test_get_provider_models_reuses_cached_provider_schema() -> None:
|
||||
configuration = _build_provider_configuration()
|
||||
provider_schema = ProviderEntity(
|
||||
provider="openai",
|
||||
label=I18nObject(en_US="OpenAI"),
|
||||
supported_model_types=[ModelType.LLM],
|
||||
configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL],
|
||||
models=[_build_ai_model("a-model"), _build_ai_model("b-model")],
|
||||
)
|
||||
configuration.provider = provider_schema
|
||||
|
||||
with patch(
|
||||
"core.entities.provider_configuration.create_plugin_model_assembly",
|
||||
) as mock_assembly_builder:
|
||||
configuration.get_provider_models(model_type=ModelType.LLM, model="a-model")
|
||||
configuration.get_provider_models(model_type=ModelType.LLM, model="b-model")
|
||||
|
||||
mock_assembly_builder.assert_not_called()
|
||||
|
||||
|
||||
def test_validator_adds_predefined_model_for_customizable_provider_with_restrictions() -> None:
|
||||
provider = ProviderEntity(
|
||||
provider="openai",
|
||||
@@ -1402,25 +1534,22 @@ def test_system_and_custom_provider_model_helpers_cover_remaining_skip_paths() -
|
||||
return _build_ai_model("embed-model", model_type=ModelType.TEXT_EMBEDDING)
|
||||
return _build_ai_model("target")
|
||||
|
||||
with patch(
|
||||
"core.entities.provider_configuration.original_provider_configurate_methods",
|
||||
{"openai": [ConfigurateMethod.CUSTOMIZABLE_MODEL]},
|
||||
):
|
||||
with patch.object(ProviderConfiguration, "get_model_schema", side_effect=_system_schema):
|
||||
system_models = configuration._get_system_provider_models(
|
||||
model_types=[ModelType.LLM],
|
||||
provider_schema=provider_schema,
|
||||
model_setting_map={
|
||||
ModelType.LLM: {
|
||||
"target": ModelSettings(
|
||||
model="target",
|
||||
model_type=ModelType.LLM,
|
||||
enabled=False,
|
||||
load_balancing_configs=[],
|
||||
)
|
||||
}
|
||||
},
|
||||
)
|
||||
configuration._original_provider_configurate_methods = (ConfigurateMethod.CUSTOMIZABLE_MODEL,)
|
||||
with patch.object(ProviderConfiguration, "get_model_schema", side_effect=_system_schema):
|
||||
system_models = configuration._get_system_provider_models(
|
||||
model_types=[ModelType.LLM],
|
||||
provider_schema=provider_schema,
|
||||
model_setting_map={
|
||||
ModelType.LLM: {
|
||||
"target": ModelSettings(
|
||||
model="target",
|
||||
model_type=ModelType.LLM,
|
||||
enabled=False,
|
||||
load_balancing_configs=[],
|
||||
)
|
||||
}
|
||||
},
|
||||
)
|
||||
assert any(model.model == "target" and model.status == ModelStatus.DISABLED for model in system_models)
|
||||
|
||||
configuration.using_provider_type = ProviderType.CUSTOM
|
||||
|
||||
@@ -28,6 +28,9 @@ class _FakeRedis:
|
||||
def get(self, key: str) -> str | None:
|
||||
return self._values.get(key)
|
||||
|
||||
def mget(self, keys: list[str]) -> list[str | None]:
|
||||
return [self.get(key) for key in keys]
|
||||
|
||||
def setex(self, key: str, ttl: int, value: str) -> None:
|
||||
self._values[key] = value
|
||||
self.setex_calls.append((key, ttl, value))
|
||||
@@ -36,6 +39,13 @@ class _FakeRedis:
|
||||
self._values.pop(key, None)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_plugin_model_provider_memory_cache() -> None:
|
||||
PluginService._plugin_model_providers_memory_cache.clear()
|
||||
yield
|
||||
PluginService._plugin_model_providers_memory_cache.clear()
|
||||
|
||||
|
||||
def _build_model_schema() -> AIModelEntity:
|
||||
return AIModelEntity(
|
||||
model="gpt-4o-mini",
|
||||
@@ -329,6 +339,7 @@ class TestPluginModelRuntime:
|
||||
"redis_client",
|
||||
SimpleNamespace(
|
||||
get=Mock(return_value=None),
|
||||
mget=Mock(return_value=[None, None]),
|
||||
delete=Mock(),
|
||||
setex=Mock(),
|
||||
),
|
||||
|
||||
@@ -345,6 +345,62 @@ def test_fetch_model_config_hydrates_model_instance_runtime_settings(model_confi
|
||||
provider_model.raise_for_status.assert_called_once()
|
||||
|
||||
|
||||
def test_fetch_model_config_reuses_validated_provider_model_from_dify_credentials_provider(
|
||||
model_config: ModelConfigWithCredentialsEntity,
|
||||
):
|
||||
mock_provider_manager = mock.MagicMock()
|
||||
mock_configurations = mock.MagicMock()
|
||||
mock_provider_configuration = mock.MagicMock()
|
||||
mock_provider_model = mock.MagicMock()
|
||||
mock_model_factory = mock.MagicMock(spec=DifyModelFactory)
|
||||
|
||||
mock_configurations.get.return_value = mock_provider_configuration
|
||||
mock_provider_configuration.get_provider_model.return_value = mock_provider_model
|
||||
mock_provider_configuration.get_current_credentials.return_value = {"api_key": "test"}
|
||||
mock_provider_manager.get_configurations.return_value = mock_configurations
|
||||
|
||||
run_context = DifyRunContext(
|
||||
tenant_id="tenant",
|
||||
app_id="app",
|
||||
user_id="user",
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
)
|
||||
credentials_provider = DifyCredentialsProvider(
|
||||
run_context=run_context,
|
||||
provider_manager=mock_provider_manager,
|
||||
)
|
||||
|
||||
model_instance = mock.MagicMock(
|
||||
model_type_instance=model_config.provider_model_bundle.model_type_instance,
|
||||
provider_model_bundle=model_config.provider_model_bundle,
|
||||
)
|
||||
mock_model_factory.init_model_instance.return_value = model_instance
|
||||
|
||||
with mock.patch.object(
|
||||
model_instance.model_type_instance.__class__,
|
||||
"get_model_schema",
|
||||
return_value=model_config.model_schema,
|
||||
autospec=True,
|
||||
):
|
||||
fetch_model_config(
|
||||
node_data_model=ModelConfig(
|
||||
provider="openai",
|
||||
name="gpt-3.5-turbo",
|
||||
mode="chat",
|
||||
completion_params={},
|
||||
),
|
||||
credentials_provider=credentials_provider,
|
||||
model_factory=mock_model_factory,
|
||||
)
|
||||
|
||||
mock_provider_configuration.get_provider_model.assert_called_once_with(
|
||||
model_type=ModelType.LLM,
|
||||
model="gpt-3.5-turbo",
|
||||
)
|
||||
mock_provider_model.raise_for_status.assert_called_once()
|
||||
|
||||
|
||||
def test_dify_model_access_adapters_call_managers():
|
||||
mock_provider_manager = mock.MagicMock()
|
||||
mock_model_manager = mock.MagicMock()
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
import datetime
|
||||
import uuid
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
from unittest.mock import MagicMock, Mock, call, patch
|
||||
|
||||
import pytest
|
||||
from pydantic import TypeAdapter
|
||||
from redis import RedisError
|
||||
|
||||
@@ -13,6 +14,15 @@ from graphon.model_runtime.entities.provider_entities import ConfigurateMethod,
|
||||
MODULE = "core.plugin.plugin_service"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_plugin_model_provider_memory_cache() -> None:
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
PluginService._plugin_model_providers_memory_cache.clear()
|
||||
yield
|
||||
PluginService._plugin_model_providers_memory_cache.clear()
|
||||
|
||||
|
||||
class _FakeSession:
|
||||
def __init__(self) -> None:
|
||||
self.execute = Mock()
|
||||
@@ -68,6 +78,17 @@ def _build_install_task(*, task_id: str = "task-1", status: PluginInstallTaskSta
|
||||
)
|
||||
|
||||
|
||||
def _provider_cache_key(tenant_id: str, generation: int | None = None) -> str:
|
||||
if generation is None:
|
||||
return f"plugin_model_providers:tenant_id:{tenant_id}"
|
||||
|
||||
return f"plugin_model_providers:tenant_id:{tenant_id}:generation:{generation}"
|
||||
|
||||
|
||||
def _provider_generation_key(tenant_id: str) -> str:
|
||||
return f"plugin_model_providers_generation:tenant_id:{tenant_id}"
|
||||
|
||||
|
||||
class TestFetchLatestPluginVersion:
|
||||
def test_skips_marketplace_fetch_when_disabled(self) -> None:
|
||||
"""Cache misses stay None; marketplace is never called when disabled."""
|
||||
@@ -120,9 +141,13 @@ class TestPluginModelProviderCache:
|
||||
"""A valid tenant cache entry is reused across runtime calls without plugin daemon access."""
|
||||
cached_provider = _build_provider_entity()
|
||||
cached_payload = TypeAdapter(list[ProviderEntity]).dump_json([cached_provider]).decode("utf-8")
|
||||
generation_key = _provider_generation_key("tenant-1")
|
||||
cache_key = _provider_cache_key("tenant-1", 0)
|
||||
legacy_cache_key = _provider_cache_key("tenant-1")
|
||||
|
||||
with patch(f"{MODULE}.redis_client") as redis_client:
|
||||
redis_client.get.return_value = cached_payload
|
||||
redis_client.get.return_value = None
|
||||
redis_client.mget.return_value = [cached_payload, None]
|
||||
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
@@ -132,14 +157,20 @@ class TestPluginModelProviderCache:
|
||||
assert [provider.provider for provider in result] == ["langgenius/openai/openai"]
|
||||
client.fetch_model_providers.assert_not_called()
|
||||
redis_client.setex.assert_not_called()
|
||||
redis_client.get.assert_called_once_with(generation_key)
|
||||
redis_client.mget.assert_called_once_with([cache_key, legacy_cache_key])
|
||||
|
||||
def test_fetch_plugin_model_providers_deletes_invalid_cache_and_refetches(self) -> None:
|
||||
"""Invalid cache payloads are tenant-scoped invalidated before falling back to the daemon."""
|
||||
"""Invalid generation-scoped cache payloads are removed before falling back to the daemon."""
|
||||
generation_key = _provider_generation_key("tenant-1")
|
||||
cache_key = _provider_cache_key("tenant-1", 0)
|
||||
legacy_cache_key = _provider_cache_key("tenant-1")
|
||||
with (
|
||||
patch(f"{MODULE}.redis_client") as redis_client,
|
||||
patch(f"{MODULE}.dify_config") as mock_config,
|
||||
):
|
||||
redis_client.get.return_value = "not-json"
|
||||
redis_client.get.side_effect = [None, None]
|
||||
redis_client.mget.return_value = ["not-json", None]
|
||||
mock_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL = 86400
|
||||
client = Mock()
|
||||
client.fetch_model_providers.return_value = [_build_plugin_model_provider()]
|
||||
@@ -148,12 +179,13 @@ class TestPluginModelProviderCache:
|
||||
|
||||
result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client)
|
||||
|
||||
cache_key = "plugin_model_providers:tenant_id:tenant-1"
|
||||
redis_client.delete.assert_called_once_with(cache_key)
|
||||
redis_client.setex.assert_called_once()
|
||||
assert redis_client.setex.call_args.args[0] == cache_key
|
||||
assert redis_client.setex.call_args.args[1] == 86400
|
||||
assert [provider.provider for provider in result] == ["langgenius/openai/openai"]
|
||||
redis_client.get.assert_has_calls([call(generation_key), call(generation_key)])
|
||||
redis_client.mget.assert_called_once_with([cache_key, legacy_cache_key])
|
||||
|
||||
def test_fetch_plugin_model_providers_refetches_when_cache_read_fails(self) -> None:
|
||||
"""Redis read failures do not block provider discovery for the tenant."""
|
||||
@@ -169,10 +201,29 @@ class TestPluginModelProviderCache:
|
||||
client.fetch_model_providers.assert_called_once_with("tenant-1")
|
||||
assert [provider.provider for provider in result] == ["langgenius/openai/openai"]
|
||||
|
||||
def test_fetch_plugin_model_providers_refetches_when_cached_payload_batch_read_fails(self) -> None:
|
||||
"""Redis mget failures do not block provider discovery for the tenant."""
|
||||
cache_key = _provider_cache_key("tenant-1", 0)
|
||||
legacy_cache_key = _provider_cache_key("tenant-1")
|
||||
with patch(f"{MODULE}.redis_client") as redis_client:
|
||||
redis_client.get.return_value = None
|
||||
redis_client.mget.side_effect = RedisError("redis unavailable")
|
||||
client = Mock()
|
||||
client.fetch_model_providers.return_value = [_build_plugin_model_provider()]
|
||||
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client)
|
||||
|
||||
client.fetch_model_providers.assert_called_once_with("tenant-1")
|
||||
redis_client.mget.assert_called_once_with([cache_key, legacy_cache_key])
|
||||
assert [provider.provider for provider in result] == ["langgenius/openai/openai"]
|
||||
|
||||
def test_fetch_plugin_model_providers_returns_fresh_result_when_cache_write_fails(self) -> None:
|
||||
"""Redis write failures are non-fatal after fresh provider data has been fetched."""
|
||||
with patch(f"{MODULE}.redis_client") as redis_client:
|
||||
redis_client.get.return_value = None
|
||||
redis_client.mget.return_value = [None, None]
|
||||
redis_client.setex.side_effect = RedisError("redis unavailable")
|
||||
client = Mock()
|
||||
client.fetch_model_providers.return_value = [_build_plugin_model_provider()]
|
||||
@@ -191,6 +242,7 @@ class TestPluginModelProviderCache:
|
||||
patch(f"{MODULE}.PluginModelClient") as client_cls,
|
||||
):
|
||||
redis_client.get.return_value = None
|
||||
redis_client.mget.return_value = [None, None]
|
||||
client = client_cls.return_value
|
||||
client.fetch_model_providers.return_value = [_build_plugin_model_provider()]
|
||||
|
||||
@@ -202,23 +254,98 @@ class TestPluginModelProviderCache:
|
||||
client.fetch_model_providers.assert_called_once_with("tenant-1")
|
||||
assert [provider.provider for provider in result] == ["langgenius/openai/openai"]
|
||||
|
||||
def test_invalidate_plugin_model_providers_cache_uses_tenant_cache_key(self) -> None:
|
||||
with patch(f"{MODULE}.redis_client") as redis_client:
|
||||
def test_fetch_plugin_model_providers_reuses_process_local_cache(self) -> None:
|
||||
generation_key = _provider_generation_key("tenant-1")
|
||||
with (
|
||||
patch(f"{MODULE}.redis_client") as redis_client,
|
||||
patch(f"{MODULE}.PluginModelClient") as client_cls,
|
||||
):
|
||||
redis_client.get.side_effect = [None, None, None]
|
||||
redis_client.mget.return_value = [None, None]
|
||||
client = client_cls.return_value
|
||||
client.fetch_model_providers.return_value = [_build_plugin_model_provider()]
|
||||
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
PluginService.invalidate_plugin_model_providers_cache("tenant-1")
|
||||
first_result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1")
|
||||
redis_client.get.reset_mock()
|
||||
redis_client.mget.reset_mock()
|
||||
redis_client.setex.reset_mock()
|
||||
client.fetch_model_providers.reset_mock()
|
||||
|
||||
redis_client.delete.assert_called_once_with("plugin_model_providers:tenant_id:tenant-1")
|
||||
second_result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1")
|
||||
|
||||
def test_invalidate_plugin_model_providers_cache_ignores_redis_delete_failure(self) -> None:
|
||||
redis_client.get.assert_called_once_with(generation_key)
|
||||
redis_client.mget.assert_not_called()
|
||||
redis_client.setex.assert_not_called()
|
||||
client.fetch_model_providers.assert_not_called()
|
||||
assert [provider.provider for provider in second_result] == ["langgenius/openai/openai"]
|
||||
assert second_result[0] == first_result[0]
|
||||
assert second_result[0] is not first_result[0]
|
||||
|
||||
def test_invalidate_plugin_model_providers_cache_uses_redis_pipeline(self) -> None:
|
||||
with patch(f"{MODULE}.redis_client") as redis_client:
|
||||
redis_client.delete.side_effect = RedisError("redis unavailable")
|
||||
pipe = redis_client.pipeline.return_value
|
||||
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
PluginService.invalidate_plugin_model_providers_cache("tenant-1")
|
||||
|
||||
redis_client.delete.assert_called_once_with("plugin_model_providers:tenant_id:tenant-1")
|
||||
redis_client.pipeline.assert_called_once_with(transaction=False)
|
||||
pipe.delete.assert_called_once_with(_provider_cache_key("tenant-1"))
|
||||
pipe.incr.assert_called_once_with(_provider_generation_key("tenant-1"))
|
||||
pipe.execute.assert_called_once_with()
|
||||
|
||||
def test_invalidate_plugin_model_providers_cache_ignores_redis_pipeline_failure(self) -> None:
|
||||
with patch(f"{MODULE}.redis_client") as redis_client:
|
||||
pipe = redis_client.pipeline.return_value
|
||||
pipe.execute.side_effect = RedisError("redis unavailable")
|
||||
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
PluginService.invalidate_plugin_model_providers_cache("tenant-1")
|
||||
|
||||
redis_client.pipeline.assert_called_once_with(transaction=False)
|
||||
pipe.delete.assert_called_once_with(_provider_cache_key("tenant-1"))
|
||||
pipe.incr.assert_called_once_with(_provider_generation_key("tenant-1"))
|
||||
pipe.execute.assert_called_once_with()
|
||||
|
||||
def test_invalidate_plugin_model_providers_cache_clears_process_local_cache(self) -> None:
|
||||
with patch(f"{MODULE}.redis_client") as redis_client:
|
||||
pipe = redis_client.pipeline.return_value
|
||||
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
PluginService._store_in_memory_plugin_model_providers("tenant-1", 0, [_build_provider_entity()])
|
||||
PluginService.invalidate_plugin_model_providers_cache("tenant-1")
|
||||
|
||||
assert PluginService._plugin_model_providers_memory_cache == {}
|
||||
redis_client.pipeline.assert_called_once_with(transaction=False)
|
||||
pipe.delete.assert_called_once_with(_provider_cache_key("tenant-1"))
|
||||
pipe.incr.assert_called_once_with(_provider_generation_key("tenant-1"))
|
||||
pipe.execute.assert_called_once_with()
|
||||
|
||||
def test_fetch_plugin_model_providers_ignores_stale_process_local_cache_after_generation_bump(self) -> None:
|
||||
generation_key = _provider_generation_key("tenant-1")
|
||||
new_cache_key = _provider_cache_key("tenant-1", 1)
|
||||
with patch(f"{MODULE}.redis_client") as redis_client:
|
||||
redis_client.get.side_effect = [b"1", b"1"]
|
||||
redis_client.mget.return_value = [None]
|
||||
client = Mock()
|
||||
client.fetch_model_providers.return_value = [_build_plugin_model_provider(provider="anthropic")]
|
||||
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
PluginService._store_in_memory_plugin_model_providers("tenant-1", 0, [_build_provider_entity()])
|
||||
result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client)
|
||||
|
||||
client.fetch_model_providers.assert_called_once_with("tenant-1")
|
||||
redis_client.get.assert_has_calls([call(generation_key), call(generation_key)])
|
||||
redis_client.mget.assert_called_once_with([new_cache_key])
|
||||
redis_client.setex.assert_called_once()
|
||||
assert redis_client.setex.call_args.args[0] == new_cache_key
|
||||
assert PluginService._plugin_model_providers_memory_cache["tenant-1"][0] == 1
|
||||
assert [provider.provider for provider in result] == ["langgenius/anthropic/anthropic"]
|
||||
|
||||
|
||||
class TestPluginListEndpointCounts:
|
||||
|
||||
Reference in New Issue
Block a user