From 7fb797c12113d01dc22ff7d60b2f45a449f23c5b Mon Sep 17 00:00:00 2001 From: WH-2099 Date: Thu, 2 Jul 2026 12:46:10 +0800 Subject: [PATCH] fix(api): keep provider refresh waiters single-flight (#38226) --- api/core/plugin/plugin_service.py | 254 ++++++------ .../core/plugin/test_model_runtime_adapter.py | 23 +- .../services/plugin/test_plugin_service.py | 384 ++++++++++++------ .../services/test_agent_config_service.py | 8 +- 4 files changed, 407 insertions(+), 262 deletions(-) diff --git a/api/core/plugin/plugin_service.py b/api/core/plugin/plugin_service.py index b39032894e4..694599a5c2d 100644 --- a/api/core/plugin/plugin_service.py +++ b/api/core/plugin/plugin_service.py @@ -14,9 +14,10 @@ metadata. import logging import time -from collections.abc import Mapping, Sequence +from collections.abc import Iterator, Mapping, Sequence +from contextlib import contextmanager from mimetypes import guess_type -from typing import Any, ClassVar, Literal +from typing import Literal, Protocol from pydantic import BaseModel, TypeAdapter, ValidationError from redis import RedisError @@ -68,9 +69,13 @@ logger = logging.getLogger(__name__) _provider_entities_adapter: TypeAdapter[list[ProviderEntity]] = TypeAdapter(list[ProviderEntity]) -class PluginService: - _plugin_model_providers_memory_cache: ClassVar[dict[str, tuple[int, float, tuple[ProviderEntity, ...]]]] = {} +class _RedisLock(Protocol): + def acquire(self, *, blocking: bool = True, blocking_timeout: float | None = None) -> bool: ... + def release(self) -> None: ... + + +class PluginService: class LatestPluginCache(BaseModel): plugin_id: str version: str @@ -137,10 +142,6 @@ class PluginService: declaration.provider_name = cls._get_provider_short_name_alias(provider) return declaration - @classmethod - def _copy_provider_entities(cls, providers: Sequence[ProviderEntity]) -> tuple[ProviderEntity, ...]: - return tuple(provider.model_copy(deep=True) for provider in providers) - @classmethod def _load_plugin_model_providers_generation(cls, tenant_id: str) -> int | None: cache_key = cls._get_plugin_model_providers_generation_cache_key(tenant_id) @@ -171,63 +172,14 @@ class PluginService: ) return None - @classmethod - def _load_in_memory_plugin_model_providers( - cls, memory_cache_key: str, generation: int - ) -> tuple[ProviderEntity, ...] | None: - cached_entry = cls._plugin_model_providers_memory_cache.get(memory_cache_key) - if cached_entry is None: - return None - - cached_generation, expires_at, providers = cached_entry - if cached_generation != generation or time.monotonic() >= expires_at: - cls._plugin_model_providers_memory_cache.pop(memory_cache_key, None) - return None - - return cls._copy_provider_entities(providers) - - @classmethod - def _store_in_memory_plugin_model_providers( - cls, memory_cache_key: str, generation: int, providers: Sequence[ProviderEntity] - ) -> None: - ttl = dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL - if ttl <= 0: - cls._plugin_model_providers_memory_cache.pop(memory_cache_key, None) - return - - cls._plugin_model_providers_memory_cache[memory_cache_key] = ( - generation, - time.monotonic() + ttl, - cls._copy_provider_entities(providers), - ) - - @classmethod - def _load_cached_plugin_model_providers( - cls, tenant_id: str, *, client: PluginModelClient | None = None - ) -> tuple[ProviderEntity, ...] | None: - generation = cls._load_plugin_model_providers_generation(tenant_id) - cached_providers, _ = cls._load_cached_plugin_model_providers_for_generation(tenant_id, generation) - return cached_providers - @classmethod def _load_cached_plugin_model_providers_for_generation( cls, tenant_id: str, generation: int | None ) -> tuple[tuple[ProviderEntity, ...] | None, bool]: - if generation is not None: - in_memory_cached_providers = cls._load_in_memory_plugin_model_providers(tenant_id, generation) - if in_memory_cached_providers is not None: - return in_memory_cached_providers, True - if generation is None: return None, False - cache_keys = [] - cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id, generation)) - if generation == 0: - cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id)) - - if not cache_keys: - return None, True + cache_keys = [cls._get_plugin_model_providers_cache_key(tenant_id, generation)] try: cached_provider_entries = redis_client.mget(cache_keys) @@ -248,8 +200,6 @@ class PluginService: try: providers = tuple(_provider_entities_adapter.validate_json(cached_providers)) - if generation is not None: - cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers) return providers, True except (TypeError, ValueError, ValidationError): logger.warning( @@ -275,58 +225,92 @@ class PluginService: ) -> None: cache_key = cls._get_plugin_model_providers_cache_key(tenant_id, generation) try: - payload = _provider_entities_adapter.dump_json(list(providers)).decode("utf-8") + payload = _provider_entities_adapter.dump_json(list(providers)) redis_client.setex(cache_key, dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL, payload) except (RedisError, RuntimeError): logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True) @classmethod - def _try_acquire_plugin_model_providers_lock(cls, tenant_id: str, generation: int) -> tuple[Any | None, bool]: + @contextmanager + def _plugin_model_providers_refresh_lock( + cls, tenant_id: str, generation: int, *, wait_timeout: float + ) -> Iterator[bool]: lock_key = cls._get_plugin_model_providers_lock_key(tenant_id, generation) try: - lock = redis_client.lock(lock_key, timeout=cls.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, blocking=False) - acquired = lock.acquire(blocking=False) + refresh_lock: _RedisLock = redis_client.lock( + lock_key, + timeout=cls.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, + sleep=cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL, + ) except (RedisError, RuntimeError): + logger.warning( + "Failed to create plugin model providers refresh lock for tenant %s.", + tenant_id, + exc_info=True, + ) + yield False + return + + try: + lock_acquired = refresh_lock.acquire(blocking=True, blocking_timeout=wait_timeout) + except LockError: + logger.warning( + "Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s", + tenant_id, + generation, + exc_info=True, + ) + yield False + return + except (RedisError, RuntimeError): + # Redis failures should not block provider discovery; callers fetch directly from the daemon. logger.warning( "Failed to acquire plugin model providers refresh lock for tenant %s.", tenant_id, exc_info=True, ) - return None, False + yield False + return - if not acquired: - return None, True - - return lock, True - - @classmethod - def _release_plugin_model_providers_lock(cls, tenant_id: str, lock: Any) -> None: - try: - lock.release() - except (LockError, RedisError, RuntimeError): + if not lock_acquired: logger.warning( - "Failed to release plugin model providers refresh lock for tenant %s.", + "Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s", tenant_id, - exc_info=True, + generation, ) + yield False + return + + try: + yield True + finally: + try: + refresh_lock.release() + except (LockError, RedisError, RuntimeError): + # Release failures must not hide the daemon result or the original exception. + logger.warning( + "Failed to release plugin model providers refresh lock for tenant %s generation %s.", + tenant_id, + generation, + exc_info=True, + ) @classmethod - def _wait_for_plugin_model_providers_refresh( - cls, tenant_id: str, *, client: PluginModelClient | None = None - ) -> tuple[ProviderEntity, ...] | None: - deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT - while time.monotonic() < deadline: - time.sleep(cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL) - cached_providers = cls._load_cached_plugin_model_providers(tenant_id, client=client) - if cached_providers is not None: - return cached_providers - - return None + def _fetch_and_cache_plugin_model_providers( + cls, tenant_id: str, client: PluginModelClient | None, *, refresh_generation: int | None + ) -> tuple[ProviderEntity, ...]: + model_client = client or PluginModelClient() + providers = tuple( + cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id) + ) + generation = cls._load_plugin_model_providers_generation(tenant_id) + if generation is not None and generation == refresh_generation: + cls._store_cached_plugin_model_providers(tenant_id, generation, providers) + return providers @classmethod def invalidate_plugin_model_providers_cache(cls, tenant_id: str) -> None: - """Invalidate tenant-scoped provider metadata across Redis and worker-local mirrors.""" - cls._plugin_model_providers_memory_cache.pop(tenant_id, None) + """Invalidate tenant-scoped provider metadata stored in Redis.""" cache_key = cls._get_plugin_model_providers_cache_key(tenant_id) generation_key = cls._get_plugin_model_providers_generation_cache_key(tenant_id) try: @@ -348,38 +332,68 @@ class PluginService: are intentionally owned by this service so tenant isolation and cache expiry are handled in one place. """ - generation = cls._load_plugin_model_providers_generation(tenant_id) - cached_providers, cache_available = cls._load_cached_plugin_model_providers_for_generation( - tenant_id, generation - ) - if cached_providers is not None: - return cached_providers + deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT - refresh_lock: Any | None = None - refresh_generation = generation - if generation is not None and cache_available: - lock_wait_deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_TTL - while time.monotonic() < lock_wait_deadline: - refresh_lock, lock_available = cls._try_acquire_plugin_model_providers_lock(tenant_id, generation) - if refresh_lock is not None or not lock_available: - break - refreshed_providers = cls._wait_for_plugin_model_providers_refresh(tenant_id, client=client) - if refreshed_providers is not None: - return refreshed_providers - - model_client = client or PluginModelClient() - try: - providers = tuple( - cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id) - ) + while True: generation = cls._load_plugin_model_providers_generation(tenant_id) - if generation is not None and generation == refresh_generation: - cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers) - cls._store_cached_plugin_model_providers(tenant_id, generation, providers) - return providers - finally: - if refresh_lock is not None: - cls._release_plugin_model_providers_lock(tenant_id, refresh_lock) + cached_providers, cache_available = cls._load_cached_plugin_model_providers_for_generation( + tenant_id, generation + ) + if cached_providers is not None: + return cached_providers + + if generation is None or not cache_available: + return cls._fetch_and_cache_plugin_model_providers( + tenant_id, + client, + refresh_generation=generation, + ) + + wait_timeout = deadline - time.monotonic() + if wait_timeout < 0: + logger.warning( + "Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s", + tenant_id, + generation, + ) + return cls._fetch_and_cache_plugin_model_providers( + tenant_id, + client, + refresh_generation=generation, + ) + + with cls._plugin_model_providers_refresh_lock( + tenant_id, + generation, + wait_timeout=wait_timeout, + ) as lock_acquired: + if not lock_acquired: + return cls._fetch_and_cache_plugin_model_providers( + tenant_id, + client, + refresh_generation=generation, + ) + + latest_generation = cls._load_plugin_model_providers_generation(tenant_id) + cached_providers, cache_available = cls._load_cached_plugin_model_providers_for_generation( + tenant_id, latest_generation + ) + if cached_providers is not None: + return cached_providers + if latest_generation is None or not cache_available: + return cls._fetch_and_cache_plugin_model_providers( + tenant_id, + client, + refresh_generation=latest_generation, + ) + if latest_generation != generation: + continue + + return cls._fetch_and_cache_plugin_model_providers( + tenant_id, + client, + refresh_generation=generation, + ) @staticmethod def fetch_latest_plugin_version(plugin_ids: Sequence[str]) -> Mapping[str, LatestPluginCache | None]: diff --git a/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py b/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py index c3ee4227d25..7b723acc812 100644 --- a/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py +++ b/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py @@ -3,7 +3,7 @@ import datetime import uuid from types import SimpleNamespace -from unittest.mock import Mock, patch, sentinel +from unittest.mock import MagicMock, Mock, patch, sentinel import pytest @@ -44,7 +44,13 @@ class _FakeRedis: def delete(self, key: str) -> None: self._values.pop(key, None) - def lock(self, key: str, *, timeout: int, blocking: bool) -> "_FakeRedisLock": + def lock( + self, + key: str, + *, + timeout: int, + sleep: float, + ) -> "_FakeRedisLock": return _FakeRedisLock(self, key) @@ -54,7 +60,7 @@ class _FakeRedisLock: self._key = key self._acquired = False - def acquire(self, *, blocking: bool) -> bool: + def acquire(self, *, blocking: bool = True, blocking_timeout: float | None = None) -> bool: if self._key in self._redis._values: return False @@ -68,13 +74,6 @@ class _FakeRedisLock: self._acquired = False -@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", @@ -436,10 +435,10 @@ class TestPluginModelRuntime: "redis_client", SimpleNamespace( get=Mock(return_value=None), - mget=Mock(return_value=[None, None]), + mget=Mock(return_value=[None]), delete=Mock(), setex=Mock(), - lock=Mock(return_value=SimpleNamespace(acquire=Mock(return_value=True), release=Mock())), + lock=Mock(return_value=MagicMock()), ), ) monkeypatch.setattr(plugin_service_module.dify_config, "PLUGIN_MODEL_PROVIDERS_CACHE_TTL", 0) diff --git a/api/tests/unit_tests/services/plugin/test_plugin_service.py b/api/tests/unit_tests/services/plugin/test_plugin_service.py index c7bd4dff08e..f847b868e89 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_service.py @@ -14,15 +14,6 @@ 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() @@ -140,14 +131,13 @@ class TestPluginModelProviderCache: def test_fetch_plugin_model_providers_returns_cached_provider_without_calling_daemon(self) -> None: """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") + cached_payload = TypeAdapter(list[ProviderEntity]).dump_json([cached_provider]) 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 = None - redis_client.mget.return_value = [cached_payload, None] + redis_client.mget.return_value = [cached_payload] from core.plugin.plugin_service import PluginService @@ -158,19 +148,18 @@ class TestPluginModelProviderCache: 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]) + redis_client.mget.assert_called_once_with([cache_key]) def test_fetch_plugin_model_providers_deletes_invalid_cache_and_refetches(self) -> None: """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.side_effect = [None, None] - redis_client.mget.return_value = ["not-json", None] + redis_client.get.side_effect = [None, None, None] + redis_client.mget.side_effect = [["not-json"], [None]] mock_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL = 86400 client = Mock() client.fetch_model_providers.return_value = [_build_plugin_model_provider()] @@ -184,8 +173,11 @@ class TestPluginModelProviderCache: 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]) + redis_client.get.assert_has_calls([call(generation_key), call(generation_key), call(generation_key)]) + assert redis_client.mget.call_args_list == [ + call([cache_key]), + call([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.""" @@ -204,7 +196,6 @@ class TestPluginModelProviderCache: 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") @@ -216,14 +207,14 @@ class TestPluginModelProviderCache: 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]) + redis_client.mget.assert_called_once_with([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.mget.return_value = [None] redis_client.setex.side_effect = RedisError("redis unavailable") client = Mock() client.fetch_model_providers.return_value = [_build_plugin_model_provider()] @@ -238,17 +229,15 @@ class TestPluginModelProviderCache: def test_fetch_plugin_model_providers_waits_for_concurrent_refresh_cache_fill(self) -> None: """A cache miss waits for the active tenant refresh instead of stampeding the daemon.""" cached_provider = _build_provider_entity() - cached_payload = TypeAdapter(list[ProviderEntity]).dump_json([cached_provider]).decode("utf-8") + cached_payload = TypeAdapter(list[ProviderEntity]).dump_json([cached_provider]) 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}.time.sleep") as sleep, + patch(f"{MODULE}.time.monotonic", return_value=100.0), ): redis_client.get.return_value = None - redis_client.mget.side_effect = [[None, None], [cached_payload, None]] - redis_client.lock.return_value.acquire.return_value = False + redis_client.mget.side_effect = [[None], [cached_payload]] client = Mock() client.fetch_model_providers.return_value = [_build_plugin_model_provider(provider="anthropic")] @@ -259,51 +248,31 @@ class TestPluginModelProviderCache: redis_client.lock.assert_called_once_with( PluginService._get_plugin_model_providers_lock_key("tenant-1", 0), timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, - blocking=False, + sleep=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL, + ) + redis_client.lock.return_value.acquire.assert_called_once_with( + blocking=True, + blocking_timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT, ) - redis_client.lock.return_value.acquire.assert_called_once_with(blocking=False) assert redis_client.mget.call_args_list == [ - call([cache_key, legacy_cache_key]), - call([cache_key, legacy_cache_key]), + call([cache_key]), + call([cache_key]), ] - sleep.assert_called() + redis_client.lock.return_value.release.assert_called_once() client.fetch_model_providers.assert_not_called() assert [provider.provider for provider in result] == ["langgenius/openai/openai"] - def test_fetch_plugin_model_providers_retries_lock_after_wait_timeout(self) -> None: - """Only a lock owner should refresh the daemon when the first refresh takes too long.""" + def test_fetch_plugin_model_providers_falls_back_when_refresh_lock_wait_times_out(self) -> None: + """A request should stop waiting and fetch directly instead of surfacing lock contention.""" + cache_key = _provider_cache_key("tenant-1", 0) with ( patch(f"{MODULE}.redis_client") as redis_client, - patch(f"{MODULE}.time.sleep"), patch(f"{MODULE}.PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT", 0), + patch(f"{MODULE}.time.monotonic", return_value=100.0), ): redis_client.get.return_value = None - redis_client.mget.return_value = [None, None] - redis_client.lock.return_value.acquire.side_effect = [False, True] - 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) - - assert redis_client.lock.return_value.acquire.call_args_list == [ - call(blocking=False), - call(blocking=False), - ] - client.fetch_model_providers.assert_called_once_with("tenant-1") - redis_client.lock.return_value.release.assert_called_once_with() - assert [provider.provider for provider in result] == ["langgenius/openai/openai"] - - def test_fetch_plugin_model_providers_releases_owned_refresh_lock_after_store(self) -> None: - """The refresh owner releases only its token after storing provider metadata.""" - 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.return_value = [None, None] - redis_client.lock.return_value.acquire.return_value = True + redis_client.mget.return_value = [None] + redis_client.lock.return_value.acquire.return_value = False client = Mock() client.fetch_model_providers.return_value = [_build_plugin_model_provider()] @@ -314,15 +283,220 @@ class TestPluginModelProviderCache: redis_client.lock.assert_called_once_with( PluginService._get_plugin_model_providers_lock_key("tenant-1", 0), timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, - blocking=False, + sleep=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL, ) - redis_client.lock.return_value.acquire.assert_called_once_with(blocking=False) - redis_client.mget.assert_called_once_with([cache_key, legacy_cache_key]) - redis_client.lock.return_value.release.assert_called_once_with() + redis_client.lock.return_value.acquire.assert_called_once_with(blocking=True, blocking_timeout=0) + redis_client.lock.return_value.release.assert_not_called() + client.fetch_model_providers.assert_called_once_with("tenant-1") + redis_client.setex.assert_called_once() + assert redis_client.setex.call_args.args[0] == cache_key + assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + + def test_fetch_plugin_model_providers_restarts_lock_path_after_generation_changes(self) -> None: + """Waiters re-read provider generation before trying to become the next refresh owner.""" + generation_key = _provider_generation_key("tenant-1") + stale_cache_key = _provider_cache_key("tenant-1", 0) + new_cache_key = _provider_cache_key("tenant-1", 1) + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.monotonic", side_effect=[100.0, 100.0, 100.5, 101.0]), + ): + redis_client.get.side_effect = [None, b"1", b"1", b"1", b"1"] + redis_client.mget.side_effect = [[None], [None], [None], [None]] + client = Mock() + client.fetch_model_providers.return_value = [_build_plugin_model_provider(provider="anthropic")] + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + assert redis_client.get.call_args_list == [ + call(generation_key), + call(generation_key), + call(generation_key), + call(generation_key), + call(generation_key), + ] + assert redis_client.mget.call_args_list == [ + call([stale_cache_key]), + call([new_cache_key]), + call([new_cache_key]), + call([new_cache_key]), + ] + assert redis_client.lock.call_args_list == [ + call( + PluginService._get_plugin_model_providers_lock_key("tenant-1", 0), + timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, + sleep=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL, + ), + call( + PluginService._get_plugin_model_providers_lock_key("tenant-1", 1), + timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, + sleep=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL, + ), + ] + assert redis_client.lock.return_value.acquire.call_args_list == [ + call(blocking=True, blocking_timeout=2.0), + call(blocking=True, blocking_timeout=1.5), + ] + assert redis_client.lock.return_value.release.call_count == 2 + client.fetch_model_providers.assert_called_once_with("tenant-1") + redis_client.setex.assert_called_once() + assert redis_client.setex.call_args.args[0] == new_cache_key + assert [provider.provider for provider in result] == ["langgenius/anthropic/anthropic"] + + def test_fetch_plugin_model_providers_falls_back_when_generation_retries_exhaust_wait_budget(self) -> None: + """Generation retry loops share one request-local lock wait deadline before direct fetch fallback.""" + generation_key = _provider_generation_key("tenant-1") + stale_cache_key = _provider_cache_key("tenant-1", 0) + new_cache_key = _provider_cache_key("tenant-1", 1) + + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.monotonic", side_effect=[100.0, 100.0, 102.1]), + ): + redis_client.get.side_effect = [None, b"1", b"1", b"1"] + redis_client.mget.side_effect = [[None], [None], [None]] + client = Mock() + client.fetch_model_providers.return_value = [_build_plugin_model_provider(provider="anthropic")] + + from core.plugin.plugin_service import PluginService + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + assert redis_client.get.call_args_list == [ + call(generation_key), + call(generation_key), + call(generation_key), + call(generation_key), + ] + assert redis_client.mget.call_args_list == [ + call([stale_cache_key]), + call([new_cache_key]), + call([new_cache_key]), + ] + redis_client.lock.assert_called_once_with( + PluginService._get_plugin_model_providers_lock_key("tenant-1", 0), + timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, + sleep=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL, + ) + redis_client.lock.return_value.acquire.assert_called_once_with(blocking=True, blocking_timeout=2.0) + redis_client.lock.return_value.release.assert_called_once() + client.fetch_model_providers.assert_called_once_with("tenant-1") + redis_client.setex.assert_called_once() + assert redis_client.setex.call_args.args[0] == new_cache_key + assert [provider.provider for provider in result] == ["langgenius/anthropic/anthropic"] + + def test_fetch_plugin_model_providers_releases_owned_refresh_lock_after_store(self) -> None: + """The refresh owner releases only its token after storing provider metadata.""" + cache_key = _provider_cache_key("tenant-1", 0) + + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.monotonic", return_value=100.0), + ): + redis_client.get.return_value = None + redis_client.mget.return_value = [None] + 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) + + redis_client.lock.assert_called_once_with( + PluginService._get_plugin_model_providers_lock_key("tenant-1", 0), + timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_TTL, + sleep=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL, + ) + redis_client.lock.return_value.acquire.assert_called_once_with( + blocking=True, + blocking_timeout=PluginService.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT, + ) + assert redis_client.mget.call_args_list == [ + call([cache_key]), + call([cache_key]), + ] + redis_client.lock.return_value.release.assert_called_once() redis_client.eval.assert_not_called() 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_returns_fresh_result_when_refresh_lock_release_fails(self) -> None: + """Release failures are logged, not allowed to hide a successful daemon refresh.""" + cache_key = _provider_cache_key("tenant-1", 0) + + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.monotonic", return_value=100.0), + ): + redis_client.get.return_value = None + redis_client.mget.return_value = [None] + redis_client.lock.return_value.release.side_effect = RedisError("release failed") + 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) + + assert redis_client.mget.call_args_list == [ + call([cache_key]), + call([cache_key]), + ] + client.fetch_model_providers.assert_called_once_with("tenant-1") + redis_client.setex.assert_called_once() + redis_client.lock.return_value.release.assert_called_once() + assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + + def test_fetch_plugin_model_providers_releases_owned_refresh_lock_when_fetch_fails(self) -> None: + """Release failures must not hide the daemon failure that happened while owning the lock.""" + cache_key = _provider_cache_key("tenant-1", 0) + + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.monotonic", return_value=100.0), + ): + redis_client.get.return_value = None + redis_client.mget.return_value = [None] + redis_client.lock.return_value.release.side_effect = RedisError("release failed") + client = Mock() + client.fetch_model_providers.side_effect = RuntimeError("daemon failed") + + from core.plugin.plugin_service import PluginService + + with pytest.raises(RuntimeError, match="daemon failed"): + PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + assert redis_client.mget.call_args_list == [ + call([cache_key]), + call([cache_key]), + ] + client.fetch_model_providers.assert_called_once_with("tenant-1") + redis_client.setex.assert_not_called() + redis_client.lock.return_value.release.assert_called_once() + + def test_fetch_plugin_model_providers_falls_back_when_refresh_lock_acquire_fails(self) -> None: + """Redis acquire failures degrade to a direct daemon fetch instead of hiding provider data.""" + with ( + patch(f"{MODULE}.redis_client") as redis_client, + patch(f"{MODULE}.time.monotonic", return_value=100.0), + ): + redis_client.get.return_value = None + redis_client.mget.return_value = [None] + redis_client.lock.return_value.acquire.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) + + redis_client.lock.return_value.acquire.assert_called_once() + redis_client.lock.return_value.release.assert_not_called() + 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_skips_wait_when_refresh_lock_fails(self) -> None: """Lock API failures should fall back directly instead of adding timeout latency.""" with ( @@ -330,7 +504,7 @@ class TestPluginModelProviderCache: patch(f"{MODULE}.time.sleep") as sleep, ): redis_client.get.return_value = None - redis_client.mget.return_value = [None, None] + redis_client.mget.return_value = [None] redis_client.lock.side_effect = RedisError("redis unavailable") redis_client.set.side_effect = AssertionError("raw redis set must not be used for refresh locks") client = Mock() @@ -350,8 +524,7 @@ class TestPluginModelProviderCache: cache_key = _provider_cache_key("tenant-1", 0) with patch(f"{MODULE}.redis_client") as redis_client: redis_client.get.return_value = None - redis_client.mget.return_value = [None, None] - redis_client.lock.return_value.acquire.return_value = True + redis_client.mget.return_value = [None] client = Mock() client.fetch_model_providers.return_value = [] @@ -362,14 +535,13 @@ class TestPluginModelProviderCache: assert result == () redis_client.setex.assert_called_once() assert redis_client.setex.call_args.args[0] == cache_key - redis_client.lock.return_value.release.assert_called_once_with() + redis_client.lock.return_value.release.assert_called_once() def test_fetch_plugin_model_providers_skips_cache_write_when_generation_changes_during_refresh(self) -> None: """A refresh that started before invalidation must not populate the newer generation cache.""" with patch(f"{MODULE}.redis_client") as redis_client: - redis_client.get.side_effect = [None, "1"] - redis_client.mget.return_value = [None, None] - redis_client.lock.return_value.acquire.return_value = True + redis_client.get.side_effect = [None, None, "1"] + redis_client.mget.return_value = [None] client = Mock() client.fetch_model_providers.return_value = [] @@ -380,17 +552,16 @@ class TestPluginModelProviderCache: assert result == () client.fetch_model_providers.assert_called_once_with("tenant-1") redis_client.setex.assert_not_called() - redis_client.lock.return_value.release.assert_called_once_with() + redis_client.lock.return_value.release.assert_called_once() def test_fetch_plugin_model_providers_reuses_cached_empty_provider_list(self) -> None: """A cached empty list should prevent repeated daemon fetches for tenants without plugin models.""" - empty_payload = TypeAdapter(list[ProviderEntity]).dump_json([]).decode("utf-8") + empty_payload = TypeAdapter(list[ProviderEntity]).dump_json([]) 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.return_value = [empty_payload, None] + redis_client.mget.return_value = [empty_payload] client = Mock() from core.plugin.plugin_service import PluginService @@ -398,7 +569,7 @@ class TestPluginModelProviderCache: result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) assert result == () - redis_client.mget.assert_called_once_with([cache_key, legacy_cache_key]) + redis_client.mget.assert_called_once_with([cache_key]) client.fetch_model_providers.assert_not_called() def test_fetch_plugin_model_providers_creates_default_client_on_cache_miss(self) -> None: @@ -408,7 +579,7 @@ class TestPluginModelProviderCache: patch(f"{MODULE}.PluginModelClient") as client_cls, ): redis_client.get.return_value = None - redis_client.mget.return_value = [None, None] + redis_client.mget.return_value = [None] client = client_cls.return_value client.fetch_model_providers.return_value = [_build_plugin_model_provider()] @@ -420,35 +591,6 @@ 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_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 - - 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() - - second_result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1") - - 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: pipe = redis_client.pipeline.return_value @@ -476,41 +618,29 @@ class TestPluginModelProviderCache: 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: + def test_fetch_plugin_model_providers_uses_new_generation_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.get.side_effect = [b"1", 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.get.assert_has_calls([call(generation_key), call(generation_key), call(generation_key)]) + assert redis_client.mget.call_args_list == [ + call([new_cache_key]), + call([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 + redis_client.lock.return_value.acquire.assert_called_once() + redis_client.lock.return_value.release.assert_called_once() assert [provider.provider for provider in result] == ["langgenius/anthropic/anthropic"] diff --git a/api/tests/unit_tests/services/test_agent_config_service.py b/api/tests/unit_tests/services/test_agent_config_service.py index e122bcb066b..57b55f0a49f 100644 --- a/api/tests/unit_tests/services/test_agent_config_service.py +++ b/api/tests/unit_tests/services/test_agent_config_service.py @@ -75,7 +75,9 @@ def _zip_bytes(members: dict[str, bytes]) -> bytes: buffer = io.BytesIO() with zipfile.ZipFile(buffer, "w") as archive: for name, payload in members.items(): - archive.writestr(name, payload) + zip_info = zipfile.ZipInfo(filename=name) + zip_info.date_time = (1980, 1, 1, 0, 0, 0) + archive.writestr(zip_info, payload) return buffer.getvalue() @@ -418,8 +420,8 @@ def test_apply_env_text_supports_delete_comments_export_and_keeps_unmentioned_va @pytest.mark.parametrize( "archive_bytes", [ - b"not-a-zip-archive", - _zip_bytes({"README.md": b"missing skill md"}), + pytest.param(b"not-a-zip-archive", id="not-a-zip-archive"), + pytest.param(_zip_bytes({"README.md": b"missing skill md"}), id="missing-skill-md"), ], ) def test_inspect_skill_maps_invalid_archives_to_service_errors(archive_bytes: bytes) -> None: