fix(api): keep provider refresh waiters single-flight (#38226)

This commit is contained in:
WH-2099 2026-07-02 12:46:10 +08:00 committed by GitHub
parent 5d0576de0a
commit 7fb797c121
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 407 additions and 262 deletions

View File

@ -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]:

View File

@ -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)

View File

@ -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"]

View File

@ -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: