fix(api): prevent plugin provider cache stampedes (#37388)

Co-authored-by: VeraPyuyi <204892921+VeraPyuyi@users.noreply.github.com>
This commit is contained in:
Pyuyi
2026-06-30 10:09:05 +00:00
committed by GitHub
co-authored by VeraPyuyi
parent 44e85c0023
commit 200f8b800f
3 changed files with 291 additions and 23 deletions
+100 -22
View File
@@ -16,10 +16,11 @@ import logging
import time
from collections.abc import Mapping, Sequence
from mimetypes import guess_type
from typing import ClassVar
from typing import Any, ClassVar
from pydantic import BaseModel, TypeAdapter, ValidationError
from redis import RedisError
from redis.exceptions import LockError
from sqlalchemy import delete, select, update
from sqlalchemy.orm import Session
from yarl import URL
@@ -82,6 +83,10 @@ class PluginService:
REDIS_TTL = 60 * 5 # 5 minutes
PLUGIN_MODEL_PROVIDERS_REDIS_KEY_PREFIX = "plugin_model_providers:tenant_id:"
PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX = "plugin_model_providers_generation:tenant_id:"
PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX = "plugin_model_providers_refresh_lock:tenant_id:"
PLUGIN_MODEL_PROVIDERS_LOCK_TTL = 30
PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT = 2.0
PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL = 0.05
PLUGIN_INSTALL_TASK_TERMINAL_STATUSES = (PluginInstallTaskStatus.Success, PluginInstallTaskStatus.Failed)
# Mirror the detail-panel endpoint query size so list reconciliation and
# the visible endpoint drawer exercise the same daemon pagination path.
@@ -98,6 +103,10 @@ class PluginService:
def _get_plugin_model_providers_generation_cache_key(cls, tenant_id: str) -> str:
return f"{cls.PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX}{tenant_id}"
@classmethod
def _get_plugin_model_providers_lock_key(cls, tenant_id: str, generation: int) -> str:
return f"{cls.PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX}{tenant_id}:generation:{generation}"
@staticmethod
def _get_provider_short_name_alias(provider: PluginModelProviderEntity) -> str:
"""
@@ -197,32 +206,41 @@ class PluginService:
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
return in_memory_cached_providers, True
if generation is None:
return None, False
cache_keys = []
if generation is not None:
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))
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
return None, True
try:
cached_provider_entries = redis_client.mget(cache_keys)
except (RedisError, RuntimeError):
except (LockError, RedisError, RuntimeError):
logger.warning("Failed to read cached plugin model providers for tenant %s.", tenant_id, exc_info=True)
return None
return None, False
if len(cached_provider_entries) != len(cache_keys):
logger.warning(
"Unexpected cached plugin model providers response size for tenant %s.",
tenant_id,
)
return None
return None, False
for cache_key, cached_providers in zip(cache_keys, cached_provider_entries):
if not cached_providers:
@@ -232,7 +250,7 @@ class PluginService:
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
return providers, True
except (TypeError, ValueError, ValidationError):
logger.warning(
"Invalid cached plugin model providers for tenant %s; deleting cache key %s.",
@@ -249,7 +267,7 @@ class PluginService:
exc_info=True,
)
return None
return None, True
@classmethod
def _store_cached_plugin_model_providers(
@@ -262,6 +280,49 @@ class PluginService:
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]:
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)
except (RedisError, RuntimeError):
logger.warning(
"Failed to acquire plugin model providers refresh lock for tenant %s.",
tenant_id,
exc_info=True,
)
return None, False
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):
logger.warning(
"Failed to release plugin model providers refresh lock for tenant %s.",
tenant_id,
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
@classmethod
def invalidate_plugin_model_providers_cache(cls, tenant_id: str) -> None:
"""Invalidate tenant-scoped provider metadata across Redis and worker-local mirrors."""
@@ -287,21 +348,38 @@ class PluginService:
are intentionally owned by this service so tenant isolation and cache
expiry are handled in one place.
"""
cached_providers = cls._load_cached_plugin_model_providers(tenant_id, client=client)
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
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()
providers = tuple(
cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id)
)
if not providers:
try:
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_in_memory_plugin_model_providers(tenant_id, generation, providers)
cls._store_cached_plugin_model_providers(tenant_id, generation, providers)
return providers
generation = cls._load_plugin_model_providers_generation(tenant_id)
if generation is not None:
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)
@staticmethod
def fetch_latest_plugin_version(plugin_ids: Sequence[str]) -> Mapping[str, LatestPluginCache | None]: