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:
呆萌闷油瓶
2026-06-11 07:05:35 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent 632df88228
commit c4a8d79be9
7 changed files with 620 additions and 99 deletions
@@ -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: