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

This commit is contained in:
WH-2099
2026-07-02 04:46:10 +00:00
committed by GitHub
parent 5d0576de0a
commit 7fb797c121
4 changed files with 407 additions and 262 deletions
+134 -120
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]: