mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
feat(caching): give RedisCache a pub/sub and lease surface and drop raw client reaches
Add async_publish and async_subscribe to RedisCache, returning a RedisSubscription that yields RedisMessage values and skips redis-py's subscribe acks, plus connection_pool_status for the debug endpoint. Cluster clients raise NotImplementedError, matching the existing skip behaviour. Move the config sync and auth cache invalidation modules onto that surface and rewrite the MCP refresh lock onto async_register_script, so no proxy code reaches into init_async_client or redis_client any more. This defines the surface a native RedisCache has to match.
This commit is contained in:
parent
1ceb8fb08e
commit
12a28c4708
12 changed files with 678 additions and 531 deletions
1
.github/workflows/test-redis-compat.yml
vendored
1
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -82,6 +82,7 @@ jobs:
|
|||
uv run --no-sync pytest \
|
||||
tests/test_litellm/test_redis.py \
|
||||
tests/test_litellm/caching/test_redis_connection_pool.py \
|
||||
tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py \
|
||||
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
|
||||
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
|
||||
--tb=short -vv \
|
||||
|
|
|
|||
|
|
@ -23,7 +23,8 @@ from dataclasses import dataclass
|
|||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -52,6 +53,7 @@ if TYPE_CHECKING:
|
|||
from prometheus_client import Gauge as _PromGauge
|
||||
from redis.asyncio import Redis, RedisCluster
|
||||
from redis.asyncio.client import Pipeline
|
||||
from redis.asyncio.client import PubSub as AsyncPubSub
|
||||
from redis.asyncio.cluster import ClusterPipeline
|
||||
|
||||
pipeline = Pipeline
|
||||
|
|
@ -577,6 +579,68 @@ def _redis_circuit_breaker_guard_sync(method: Callable[..., _RedisCallResult]) -
|
|||
)
|
||||
|
||||
|
||||
class _PubSubFrame(TypedDict):
|
||||
type: ReadOnly[str]
|
||||
channel: ReadOnly[str | bytes]
|
||||
data: ReadOnly[str | bytes | int]
|
||||
|
||||
|
||||
_PUBSUB_FRAME: Final = TypeAdapter(_PubSubFrame)
|
||||
|
||||
|
||||
class RedisPoolStatus(TypedDict):
|
||||
max_connections: ReadOnly[int | None]
|
||||
connection_class: ReadOnly[str | None]
|
||||
|
||||
|
||||
_PUBSUB_WAIT_SLICE_SECONDS: Final = 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RedisMessage:
|
||||
channel: str
|
||||
payload: bytes
|
||||
|
||||
|
||||
def _redis_message(frame: object) -> RedisMessage | None:
|
||||
"""The published message a pub/sub frame carries; ``None`` for subscribe acks and anything malformed."""
|
||||
try:
|
||||
parsed: Final = _PUBSUB_FRAME.validate_python(frame)
|
||||
except ValidationError:
|
||||
return None
|
||||
channel: Final = parsed["channel"]
|
||||
payload: Final = parsed["data"]
|
||||
if parsed["type"] not in ("message", "pmessage") or isinstance(payload, int):
|
||||
return None
|
||||
return RedisMessage(
|
||||
channel=channel.decode("utf-8", errors="replace") if isinstance(channel, bytes) else channel,
|
||||
payload=payload.encode("utf-8") if isinstance(payload, str) else payload,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RedisSubscription:
|
||||
"""One SUBSCRIBE session. ``get_message`` waits up to ``timeout`` seconds (``None`` waits
|
||||
indefinitely) for a published message and skips the acks redis-py interleaves with them."""
|
||||
|
||||
pubsub: "AsyncPubSub"
|
||||
|
||||
async def get_message(self, *, timeout: float | None) -> RedisMessage | None:
|
||||
clock: Final = asyncio.get_running_loop().time
|
||||
deadline: Final = None if timeout is None else clock() + timeout
|
||||
while True:
|
||||
remaining = _PUBSUB_WAIT_SLICE_SECONDS if deadline is None else max(deadline - clock(), 0.0)
|
||||
frame: object = await self.pubsub.get_message(timeout=remaining)
|
||||
if frame is None:
|
||||
return None
|
||||
message = _redis_message(frame)
|
||||
if message is not None:
|
||||
return message
|
||||
|
||||
async def aclose(self) -> None:
|
||||
await self.pubsub.aclose() # pyright: ignore[reportAttributeAccessIssue] # types-redis 4.6 stubs predate PubSub.aclose
|
||||
|
||||
|
||||
class RedisCache(BaseCache):
|
||||
# if users don't provider one, use the default litellm cache
|
||||
|
||||
|
|
@ -1839,6 +1903,34 @@ class RedisCache(BaseCache):
|
|||
key = self.check_and_fix_namespace(key=key)
|
||||
self.redis_client.delete(key)
|
||||
|
||||
def _standalone_async_client(self) -> "Redis[bytes] | Redis[str]":
|
||||
from redis.asyncio import RedisCluster
|
||||
|
||||
client: Final[object] = self.init_async_client()
|
||||
if isinstance(client, RedisCluster):
|
||||
raise NotImplementedError("Redis Cluster clients have no pub/sub support")
|
||||
return cast("Redis[bytes] | Redis[str]", client) # cast-ok: decode_responses picks the reply type at runtime
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_publish(self, channel: str, message: str | bytes) -> int:
|
||||
return await self._standalone_async_client().publish(channel, message)
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_subscribe(self, *channels: str) -> RedisSubscription:
|
||||
pubsub: Final = self._standalone_async_client().pubsub()
|
||||
await pubsub.subscribe(*channels)
|
||||
return RedisSubscription(pubsub)
|
||||
|
||||
def connection_pool_status(self) -> "RedisPoolStatus":
|
||||
pool: Final = getattr(self.redis_client, "connection_pool", None)
|
||||
max_connections: Final = getattr(pool, "max_connections", None)
|
||||
connection_class: Final = getattr(getattr(pool, "connection_class", None), "__name__", None)
|
||||
status: Final[RedisPoolStatus] = {
|
||||
"max_connections": max_connections if isinstance(max_connections, int) else None,
|
||||
"connection_class": connection_class if isinstance(connection_class, str) else None,
|
||||
}
|
||||
return status
|
||||
|
||||
async def _pipeline_increment_helper(
|
||||
self,
|
||||
pipe: pipeline,
|
||||
|
|
|
|||
|
|
@ -1,94 +1,84 @@
|
|||
"""Concrete ``DistributedLock`` over a Redis client: ``SET NX PX`` / owner-only renew / delete.
|
||||
"""Concrete ``DistributedLock`` over ``RedisCache`` scripts: ``SET NX PX`` / owner-only renew / delete.
|
||||
|
||||
The cross-replica lock the ``RedisRefreshCoordinator`` elects refreshers with. ``acquire`` is an
|
||||
atomic ``SET key token NX PX ttl`` (only the first caller wins; the entry self-expires so a crashed
|
||||
holder can't wedge refresh). ``extend`` renews the lease only when the token still matches, and
|
||||
``release`` deletes the key only when it still holds this caller's token, so a holder whose lock already
|
||||
PX-expired and was re-acquired by another worker cannot delete the new holder's lock. ``is_held`` is
|
||||
``EXISTS``. Every key is run through the injected ``namespace_key`` before it reaches Redis, so lock
|
||||
keys carry the same namespace as cache keys and cannot collide with another deployment sharing Redis.
|
||||
``EXISTS``. Every operation runs through ``RedisCache.async_register_script``, which prefixes the key
|
||||
with the cache namespace, so lock keys carry the same namespace as cache keys and cannot collide with
|
||||
another deployment sharing Redis.
|
||||
|
||||
The Redis client is injected (in production the async client from LiteLLM's ``RedisCache``), so the
|
||||
lock is unit-testable with a fake. A transport error on ``acquire`` returns ``LockAcquisition.ERROR`` -
|
||||
distinct from ``HELD`` - so the coordinator refreshes anyway instead of mistaking a dead backend for a
|
||||
busy holder; a Redis blip degrades to an extra refresh, never a stale bearer.
|
||||
A transport error on ``acquire`` returns ``LockAcquisition.ERROR`` - distinct from ``HELD`` - so the
|
||||
coordinator refreshes anyway instead of mistaking a dead backend for a busy holder; a Redis blip
|
||||
degrades to an extra refresh, never a stale bearer.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import KW_ONLY, dataclass
|
||||
from typing import Final, Protocol
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
|
||||
LockAcquisition,
|
||||
)
|
||||
|
||||
# Delete the key only if it still holds this caller's token, so a holder whose lock already expired
|
||||
# (PX) and was re-acquired by another worker cannot delete the new holder's lock.
|
||||
_RELEASE_IF_OWNER = "if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) else return 0 end"
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
_ACQUIRE_IF_ABSENT: Final = "return redis.call('set', KEYS[1], ARGV[1], 'NX', 'PX', ARGV[2])"
|
||||
_EXTEND_IF_OWNER: Final = (
|
||||
"if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('pexpire', KEYS[1], ARGV[2]) else return 0 end"
|
||||
)
|
||||
_RELEASE_IF_OWNER: Final = (
|
||||
"if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) else return 0 end"
|
||||
)
|
||||
_IS_HELD: Final = "return redis.call('exists', KEYS[1])"
|
||||
_SCRIPT_REPLY: Final[TypeAdapter[int | bytes | str | None]] = TypeAdapter(int | bytes | str | None)
|
||||
|
||||
|
||||
class RedisCommands(Protocol):
|
||||
"""The slice of the async Redis client this lock needs."""
|
||||
|
||||
async def set(self, name: str, value: str, *, nx: bool = False, px: int | None = None) -> object | None: ...
|
||||
|
||||
async def eval(self, script: str, numkeys: int, *keys_and_args: str) -> object: ...
|
||||
|
||||
async def exists(self, *names: str) -> int: ...
|
||||
def _milliseconds(seconds: float) -> str:
|
||||
return str(int(seconds * 1000))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RedisDistributedLock:
|
||||
client: RedisCommands
|
||||
_: KW_ONLY
|
||||
namespace_key: Callable[[str], str] = lambda key: key
|
||||
redis_cache: RedisCache
|
||||
|
||||
async def _run(self, script: str, key: str, *args: str) -> int | bytes | str | None:
|
||||
return _SCRIPT_REPLY.validate_python(
|
||||
await self.redis_cache.async_register_script(script)(keys=(key,), args=args)
|
||||
)
|
||||
|
||||
async def acquire(self, key: str, token: str, ttl_seconds: float) -> LockAcquisition:
|
||||
try:
|
||||
result: Final = await self.client.set(self.namespace_key(key), token, nx=True, px=int(ttl_seconds * 1000))
|
||||
# Degrade on any Redis client error: redis.exceptions narrows only via an import that
|
||||
# is Unknown under basedpyright, and the lock must never crash the resolve path.
|
||||
except Exception as exc: # noqa: BLE001
|
||||
result: Final = await self._run(_ACQUIRE_IF_ABSENT, key, token, _milliseconds(ttl_seconds))
|
||||
except Exception as exc: # noqa: BLE001 # the lock must never crash the resolve path
|
||||
verbose_logger.warning("RedisDistributedLock.acquire failed: %s", exc)
|
||||
return LockAcquisition.ERROR
|
||||
return LockAcquisition.ACQUIRED if result is not None else LockAcquisition.HELD
|
||||
|
||||
async def extend(self, key: str, token: str, ttl_seconds: float) -> bool:
|
||||
try:
|
||||
result: Final = await self.client.eval(
|
||||
_EXTEND_IF_OWNER,
|
||||
1,
|
||||
self.namespace_key(key),
|
||||
token,
|
||||
str(int(ttl_seconds * 1000)),
|
||||
)
|
||||
# Degrade on any Redis client error: redis.exceptions narrows only via an import that
|
||||
# is Unknown under basedpyright, and the lock must never crash the resolve path.
|
||||
except Exception as exc: # noqa: BLE001
|
||||
result: Final = await self._run(_EXTEND_IF_OWNER, key, token, _milliseconds(ttl_seconds))
|
||||
except Exception as exc: # noqa: BLE001 # the lock must never crash the resolve path
|
||||
verbose_logger.warning("RedisDistributedLock.extend failed: %s", exc)
|
||||
return False
|
||||
return result == 1
|
||||
|
||||
async def release(self, key: str, token: str) -> None:
|
||||
try:
|
||||
await self.client.eval(_RELEASE_IF_OWNER, 1, self.namespace_key(key), token)
|
||||
# Degrade on any Redis client error: redis.exceptions narrows only via an import that
|
||||
# is Unknown under basedpyright, and the lock must never crash the resolve path.
|
||||
except Exception as exc: # noqa: BLE001
|
||||
await self._run(_RELEASE_IF_OWNER, key, token)
|
||||
except Exception as exc: # noqa: BLE001 # the lock must never crash the resolve path
|
||||
verbose_logger.warning("RedisDistributedLock.release failed: %s", exc)
|
||||
|
||||
async def is_held(self, key: str) -> bool:
|
||||
try:
|
||||
return await self.client.exists(self.namespace_key(key)) > 0
|
||||
# Degrade on any Redis client error: redis.exceptions narrows only via an import that
|
||||
# is Unknown under basedpyright, and the lock must never crash the resolve path.
|
||||
except Exception as exc: # noqa: BLE001
|
||||
# On error, report "not held" so a waiter stops waiting and re-reads rather than blocking.
|
||||
result: Final = await self._run(_IS_HELD, key)
|
||||
except Exception as exc: # noqa: BLE001 # a waiter stops waiting and re-reads rather than blocking
|
||||
verbose_logger.warning("RedisDistributedLock.is_held failed: %s", exc)
|
||||
return False
|
||||
return isinstance(result, int) and result > 0
|
||||
|
|
|
|||
|
|
@ -31,11 +31,4 @@ def runtime_refresh_coordinator() -> RefreshCoordinator | None:
|
|||
redis_cache: Final = user_api_key_cache.redis_cache
|
||||
if redis_cache is None:
|
||||
return None
|
||||
# The Redis client from init_async_client() is only partially typed; the lock validates every
|
||||
# reply it depends on, so the untyped boundary is contained here.
|
||||
redis_client: Final = redis_cache.init_async_client() # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType] # litellm redis wrapper is untyped
|
||||
lock: Final = RedisDistributedLock(
|
||||
redis_client, # pyright: ignore[reportArgumentType,reportUnknownArgumentType] # litellm redis wrapper is untyped
|
||||
namespace_key=redis_cache.check_and_fix_namespace,
|
||||
)
|
||||
return RedisRefreshCoordinator(lock)
|
||||
return RedisRefreshCoordinator(RedisDistributedLock(redis_cache))
|
||||
|
|
|
|||
|
|
@ -5,15 +5,11 @@ from dataclasses import asdict, dataclass
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.config_sync_pubsub import (
|
||||
_ConfigSyncPubSub,
|
||||
_pubsub_capable_client,
|
||||
coordination_redis_cache,
|
||||
)
|
||||
from litellm.proxy.common_utils.config_sync_pubsub import coordination_redis_cache
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.caching.redis_cache import RedisCache, RedisMessage, RedisSubscription
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
AUTH_CACHE_INVALIDATION_CHANNEL: Final = "litellm_proxy.auth_cache_invalidation"
|
||||
|
|
@ -73,15 +69,13 @@ def _message_from_data(data: object) -> _CacheInvalidationMessage | None:
|
|||
|
||||
async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None:
|
||||
try:
|
||||
client: Final = _pubsub_capable_client(redis_cache)
|
||||
if client is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support",
|
||||
cache_key,
|
||||
)
|
||||
return
|
||||
async with _in_flight_publishes:
|
||||
await client.publish(auth_cache_invalidation_channel(redis_cache), message)
|
||||
await redis_cache.async_publish(auth_cache_invalidation_channel(redis_cache), message)
|
||||
except NotImplementedError:
|
||||
verbose_proxy_logger.debug(
|
||||
"auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support",
|
||||
cache_key,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors
|
||||
verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e)
|
||||
|
||||
|
|
@ -184,20 +178,14 @@ class AuthCacheInvalidationSubscriber:
|
|||
backoff_seconds = _BACKOFF_INITIAL_SECONDS # rebind-ok: exponential backoff accumulator across reconnects
|
||||
while True:
|
||||
try:
|
||||
client = _pubsub_capable_client(self._redis_cache)
|
||||
if client is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"auth cache invalidation subscriber disabled: cluster redis client has no pub/sub support; "
|
||||
"cross-worker eviction falls back to the local cache TTL"
|
||||
)
|
||||
subscription = await self._open_subscription()
|
||||
if subscription is None:
|
||||
return
|
||||
pubsub = client.pubsub()
|
||||
try:
|
||||
await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache))
|
||||
backoff_seconds = _BACKOFF_INITIAL_SECONDS
|
||||
await self._consume(pubsub)
|
||||
await self._consume(subscription)
|
||||
finally:
|
||||
await self._close_pubsub(pubsub)
|
||||
await self._close_subscription(subscription)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as e: # noqa: BLE001 # any redis failure falls through to backoff and reconnect
|
||||
|
|
@ -209,16 +197,25 @@ class AuthCacheInvalidationSubscriber:
|
|||
await asyncio.sleep(backoff_seconds)
|
||||
backoff_seconds = min(backoff_seconds * 2, _BACKOFF_MAX_SECONDS)
|
||||
|
||||
async def _consume(self, pubsub: _ConfigSyncPubSub) -> None:
|
||||
async def _open_subscription(self) -> "RedisSubscription | None":
|
||||
try:
|
||||
return await self._redis_cache.async_subscribe(auth_cache_invalidation_channel(self._redis_cache))
|
||||
except NotImplementedError:
|
||||
verbose_proxy_logger.warning(
|
||||
"auth cache invalidation subscriber disabled: cluster redis client has no pub/sub support; "
|
||||
"cross-worker eviction falls back to the local cache TTL"
|
||||
)
|
||||
return None
|
||||
|
||||
async def _consume(self, subscription: "RedisSubscription") -> None:
|
||||
while True:
|
||||
message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS)
|
||||
message = await subscription.get_message(timeout=_POLL_TIMEOUT_SECONDS)
|
||||
if message is None:
|
||||
continue
|
||||
self._apply_message(message)
|
||||
|
||||
def _apply_message(self, message: object) -> None:
|
||||
data: Final = message.get("data") if isinstance(message, dict) else None
|
||||
parsed: Final = _message_from_data(data)
|
||||
def _apply_message(self, message: "RedisMessage") -> None:
|
||||
parsed: Final = _message_from_data(message.payload)
|
||||
if parsed is None:
|
||||
return
|
||||
if parsed.new_value is not None:
|
||||
|
|
@ -230,8 +227,8 @@ class AuthCacheInvalidationSubscriber:
|
|||
additional_cache.delete_cache(parsed.cache_key)
|
||||
|
||||
@staticmethod
|
||||
async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None:
|
||||
async def _close_subscription(subscription: "RedisSubscription") -> None:
|
||||
try:
|
||||
await pubsub.aclose()
|
||||
await subscription.aclose()
|
||||
except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection
|
||||
verbose_proxy_logger.debug("auth cache invalidation pubsub close failed: %s", e)
|
||||
|
|
|
|||
|
|
@ -4,27 +4,13 @@ import random
|
|||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import TYPE_CHECKING, Final, Protocol, cast # noqa: TID251 # untyped prisma/redis boundary needs cast
|
||||
from typing import TYPE_CHECKING, Final, cast # noqa: TID251 # untyped prisma boundary needs cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.repositories.prisma_protocols import RowT_co, TableActions
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
|
||||
class _ConfigSyncPubSub(Protocol):
|
||||
def subscribe(self, *channels: str) -> Awaitable[object]: ...
|
||||
|
||||
def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Awaitable[object]: ...
|
||||
|
||||
def aclose(self) -> Awaitable[object]: ...
|
||||
|
||||
|
||||
class _ConfigSyncPubSubClient(Protocol):
|
||||
def publish(self, channel: str, message: str) -> Awaitable[int]: ...
|
||||
|
||||
def pubsub(self) -> _ConfigSyncPubSub: ...
|
||||
from litellm.caching.redis_cache import RedisCache, RedisSubscription
|
||||
|
||||
|
||||
CONFIG_SYNC_CHANNEL: Final = "litellm_proxy.config_change"
|
||||
|
|
@ -82,22 +68,6 @@ def config_sync_channel(redis_cache: "RedisCache") -> str:
|
|||
return f"{redis_cache.namespace}:{CONFIG_SYNC_CHANNEL}"
|
||||
|
||||
|
||||
def _raw_async_client(redis_cache: "RedisCache") -> object:
|
||||
return cast( # cast-ok: redis-py generics leave the client type partially unknown
|
||||
object,
|
||||
redis_cache.init_async_client(), # pyright: ignore[reportUnknownMemberType] # redis generics
|
||||
)
|
||||
|
||||
|
||||
def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient | None:
|
||||
from redis.asyncio import Redis
|
||||
|
||||
client: Final = _raw_async_client(redis_cache)
|
||||
if isinstance(client, Redis):
|
||||
return cast(_ConfigSyncPubSubClient, client) # cast-ok: protocol view of the standalone redis client
|
||||
return None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ConfigChangeMessage:
|
||||
object_type: str
|
||||
|
|
@ -111,14 +81,12 @@ async def publish_config_change(redis_cache: "RedisCache | None", object_type: s
|
|||
if redis_cache is None:
|
||||
return
|
||||
try:
|
||||
client: Final = _pubsub_capable_client(redis_cache)
|
||||
if client is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"config sync publish for %s skipped: cluster redis client has no pub/sub support",
|
||||
object_type,
|
||||
)
|
||||
return
|
||||
await client.publish(config_sync_channel(redis_cache), _config_change_message_json(object_type))
|
||||
await redis_cache.async_publish(config_sync_channel(redis_cache), _config_change_message_json(object_type))
|
||||
except NotImplementedError:
|
||||
verbose_proxy_logger.debug(
|
||||
"config sync publish for %s skipped: cluster redis client has no pub/sub support",
|
||||
object_type,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # best-effort publish; writes must never fail on redis errors
|
||||
verbose_proxy_logger.warning("config sync publish for %s failed: %s", object_type, e)
|
||||
|
||||
|
|
@ -237,20 +205,14 @@ class ConfigSyncSubscriber:
|
|||
backoff_seconds = self._backoff_initial_seconds
|
||||
while True:
|
||||
try:
|
||||
client = _pubsub_capable_client(self._redis_cache)
|
||||
if client is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"config sync subscriber disabled: cluster redis client has no pub/sub support; "
|
||||
"interval polling remains the only sync mechanism"
|
||||
)
|
||||
subscription = await self._open_subscription()
|
||||
if subscription is None:
|
||||
return
|
||||
pubsub = client.pubsub()
|
||||
try:
|
||||
await pubsub.subscribe(config_sync_channel(self._redis_cache))
|
||||
backoff_seconds = self._backoff_initial_seconds
|
||||
await self._consume(pubsub)
|
||||
await self._consume(subscription)
|
||||
finally:
|
||||
await self._close_pubsub(pubsub)
|
||||
await self._close_subscription(subscription)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as e: # noqa: BLE001 # any redis failure falls through to backoff and reconnect
|
||||
|
|
@ -262,14 +224,24 @@ class ConfigSyncSubscriber:
|
|||
await self._sleep(backoff_seconds)
|
||||
backoff_seconds = min(backoff_seconds * 2, self._backoff_max_seconds)
|
||||
|
||||
async def _consume(self, pubsub: _ConfigSyncPubSub) -> None:
|
||||
async def _open_subscription(self) -> "RedisSubscription | None":
|
||||
try:
|
||||
return await self._redis_cache.async_subscribe(config_sync_channel(self._redis_cache))
|
||||
except NotImplementedError:
|
||||
verbose_proxy_logger.warning(
|
||||
"config sync subscriber disabled: cluster redis client has no pub/sub support; "
|
||||
"interval polling remains the only sync mechanism"
|
||||
)
|
||||
return None
|
||||
|
||||
async def _consume(self, subscription: "RedisSubscription") -> None:
|
||||
while True:
|
||||
message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS)
|
||||
message = await subscription.get_message(timeout=_POLL_TIMEOUT_SECONDS)
|
||||
if message is None:
|
||||
continue
|
||||
await self._sleep(self._debounce_seconds + self._rng.uniform(0.0, self._jitter_max_seconds))
|
||||
await self._wait_for_min_resync_interval()
|
||||
await self._drain_pending(pubsub)
|
||||
await self._drain_pending(subscription)
|
||||
await self._run_resync_callbacks()
|
||||
self._last_resync_at = self._monotonic()
|
||||
|
||||
|
|
@ -286,8 +258,8 @@ class ConfigSyncSubscriber:
|
|||
await self._sleep(seconds_until_next_resync)
|
||||
|
||||
@staticmethod
|
||||
async def _drain_pending(pubsub: _ConfigSyncPubSub) -> None:
|
||||
while await pubsub.get_message(ignore_subscribe_messages=True, timeout=0) is not None:
|
||||
async def _drain_pending(subscription: "RedisSubscription") -> None:
|
||||
while await subscription.get_message(timeout=0) is not None:
|
||||
pass
|
||||
|
||||
async def _run_resync_callbacks(self) -> None:
|
||||
|
|
@ -298,8 +270,8 @@ class ConfigSyncSubscriber:
|
|||
verbose_proxy_logger.warning("config sync resync callback failed: %s", e)
|
||||
|
||||
@staticmethod
|
||||
async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None:
|
||||
async def _close_subscription(subscription: "RedisSubscription") -> None:
|
||||
try:
|
||||
await pubsub.aclose()
|
||||
await subscription.aclose()
|
||||
except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection
|
||||
verbose_proxy_logger.debug("config sync pubsub close failed: %s", e)
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ import sys
|
|||
import tracemalloc
|
||||
from collections import Counter
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Annotated, Any, Final, NamedTuple, Protocol, TypedDict
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Final, NamedTuple, Protocol, TypedDict
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from typing_extensions import ReadOnly
|
||||
|
|
@ -22,6 +22,9 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
from litellm.proxy.bug_report_config import build_proxy_environment_report
|
||||
from litellm.proxy.common_utils.resource_ownership import is_proxy_admin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
|
|
@ -482,7 +485,7 @@ def _get_uncollectable_objects_info() -> Mapping[str, object]:
|
|||
|
||||
|
||||
def _get_cache_memory_stats(
|
||||
user_api_key_cache, llm_router, proxy_logging_obj, redis_usage_cache
|
||||
user_api_key_cache, llm_router, proxy_logging_obj, redis_usage_cache: "RedisCache | None"
|
||||
) -> Mapping[str, object]:
|
||||
"""Calculate memory usage for all caches."""
|
||||
cache_stats: Final[dict[str, object]] = {}
|
||||
|
|
@ -529,22 +532,8 @@ def _get_cache_memory_stats(
|
|||
cache_stats["redis_usage_cache"] = {
|
||||
"enabled": True,
|
||||
"cache_type": type(redis_usage_cache).__name__,
|
||||
"connection_pool": redis_usage_cache.connection_pool_status(),
|
||||
}
|
||||
# Try to get Redis connection pool info if available
|
||||
try:
|
||||
if hasattr(redis_usage_cache, "redis_client") and redis_usage_cache.redis_client:
|
||||
if hasattr(redis_usage_cache.redis_client, "connection_pool"):
|
||||
pool_info: Final = redis_usage_cache.redis_client.connection_pool
|
||||
cache_stats["redis_usage_cache"]["connection_pool"] = {
|
||||
"max_connections": (
|
||||
pool_info.max_connections if hasattr(pool_info, "max_connections") else None
|
||||
),
|
||||
"connection_class": (
|
||||
pool_info.connection_class.__name__ if hasattr(pool_info, "connection_class") else None
|
||||
),
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Error getting Redis pool info: %s", e)
|
||||
else:
|
||||
cache_stats["redis_usage_cache"] = {"enabled": False}
|
||||
|
||||
|
|
|
|||
|
|
@ -58,9 +58,7 @@ def test_check_and_fix_namespace_prefixes_keys_sharing_the_namespace_prefix(
|
|||
|
||||
@pytest.mark.parametrize("namespace", [None, "litellm"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_cache_applies_namespace(
|
||||
namespace, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_delete_cache_applies_namespace(namespace, monkeypatch, redis_no_ping):
|
||||
"""async_delete_cache must prefix keys with the namespace, matching every
|
||||
other cache operation. Without this, Redis NOPERM errors occur when an
|
||||
ACL restricts DEL to the litellm:* pattern."""
|
||||
|
|
@ -68,9 +66,7 @@ async def test_async_delete_cache_applies_namespace(
|
|||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_delete_cache(key="3997c4abcdef")
|
||||
|
||||
expected_key = "litellm:3997c4abcdef" if namespace else "3997c4abcdef"
|
||||
|
|
@ -133,9 +129,7 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
|
|||
]
|
||||
|
||||
# Test the helper method
|
||||
result = await redis_cache.handle_lpop_count_for_older_redis_versions(
|
||||
pipe=mock_pipeline, key="test_key", count=2
|
||||
)
|
||||
result = await redis_cache.handle_lpop_count_for_older_redis_versions(pipe=mock_pipeline, key="test_key", count=2)
|
||||
|
||||
# Verify results
|
||||
assert result == [b"value1", b"value2"]
|
||||
|
|
@ -144,18 +138,14 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_pipeline_empty_list_returns_empty(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_rpush_pipeline_empty_list_returns_empty(monkeypatch, redis_no_ping):
|
||||
"""Empty rpush_list should return empty list without touching Redis"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
result = await redis_cache.async_rpush_pipeline(rpush_list=[])
|
||||
|
||||
assert result == []
|
||||
|
|
@ -170,9 +160,7 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
|
|||
|
||||
mock_redis_instance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
result = await redis_cache.async_lpop_pipeline(lpop_list=[])
|
||||
|
||||
assert result == []
|
||||
|
|
@ -197,9 +185,7 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
|
|||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_namespaces_keys(
|
||||
namespace, raw_keys, expected_keys, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_register_script_namespaces_keys(namespace, raw_keys, expected_keys, monkeypatch, redis_no_ping):
|
||||
"""The callable returned by async_register_script (used by the rate limiter
|
||||
Lua scripts, pod-lock release, and budget limiters) must namespace every key
|
||||
it is invoked with. The hash tag is preserved so cluster slotting is intact."""
|
||||
|
|
@ -210,16 +196,12 @@ async def test_async_register_script_namespaces_keys(
|
|||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.register_script = MagicMock(return_value=registered_script)
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
script = redis_cache.async_register_script("return 1")
|
||||
result = await script(keys=raw_keys, args=[60])
|
||||
|
||||
assert result == "ok"
|
||||
registered_script.assert_awaited_once_with(
|
||||
keys=tuple(expected_keys), args=[60], client=None
|
||||
)
|
||||
registered_script.assert_awaited_once_with(keys=tuple(expected_keys), args=[60], client=None)
|
||||
|
||||
|
||||
# LIT-3298: rate limits tripped at ~40M instead of 80M. async_register_script
|
||||
|
|
@ -257,12 +239,8 @@ def test_async_register_script_binds_per_event_loop(namespace, monkeypatch):
|
|||
loop_a = asyncio.new_event_loop()
|
||||
loop_b = asyncio.new_event_loop()
|
||||
try:
|
||||
result_a = loop_a.run_until_complete(
|
||||
script(keys=["{k:v}:tokens"], args=[60])
|
||||
)
|
||||
result_b = loop_b.run_until_complete(
|
||||
script(keys=["{k:v}:tokens"], args=[60])
|
||||
)
|
||||
result_a = loop_a.run_until_complete(script(keys=["{k:v}:tokens"], args=[60]))
|
||||
result_b = loop_b.run_until_complete(script(keys=["{k:v}:tokens"], args=[60]))
|
||||
finally:
|
||||
loop_a.close()
|
||||
loop_b.close()
|
||||
|
|
@ -275,9 +253,7 @@ def test_async_register_script_binds_per_event_loop(namespace, monkeypatch):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_not_shared_across_namespaces(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_register_script_not_shared_across_namespaces(monkeypatch, redis_no_ping):
|
||||
"""Two caches with different namespaces registering the SAME script must
|
||||
each run against their own client and key prefix. A content-only executor
|
||||
cache would let the second cache reuse the first's executor and namespace."""
|
||||
|
|
@ -293,9 +269,10 @@ async def test_async_register_script_not_shared_across_namespaces(
|
|||
client_b.register_script = MagicMock(return_value=reg_b)
|
||||
|
||||
same_script = "return redis.call('GET', KEYS[1])"
|
||||
with patch.object(
|
||||
cache_a, "init_async_client", return_value=client_a
|
||||
), patch.object(cache_b, "init_async_client", return_value=client_b):
|
||||
with (
|
||||
patch.object(cache_a, "init_async_client", return_value=client_a),
|
||||
patch.object(cache_b, "init_async_client", return_value=client_b),
|
||||
):
|
||||
script_a = cache_a.async_register_script(same_script)
|
||||
script_b = cache_b.async_register_script(same_script)
|
||||
result_a = await script_a(keys=["k"], args=[])
|
||||
|
|
@ -307,9 +284,7 @@ async def test_async_register_script_not_shared_across_namespaces(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_cluster_path_uses_evalsha(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_register_script_cluster_path_uses_evalsha(monkeypatch, redis_no_ping):
|
||||
"""Redis Cluster exposes script_load/evalsha rather than register_script.
|
||||
The script is loaded once and invoked via evalsha with namespaced keys."""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
|
|
@ -319,23 +294,17 @@ async def test_async_register_script_cluster_path_uses_evalsha(
|
|||
cluster_client.script_load = MagicMock(return_value="sha123")
|
||||
cluster_client.evalsha = AsyncMock(return_value="cluster-ok")
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=cluster_client
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=cluster_client):
|
||||
script = redis_cache.async_register_script("return 'cluster'")
|
||||
result = await script(keys=["{k:v}:tokens"], args=[5, 60])
|
||||
|
||||
assert result == "cluster-ok"
|
||||
cluster_client.script_load.assert_called_once_with("return 'cluster'")
|
||||
cluster_client.evalsha.assert_awaited_once_with(
|
||||
"sha123", 1, "ns:{k:v}:tokens", 5, 60
|
||||
)
|
||||
cluster_client.evalsha.assert_awaited_once_with("sha123", 1, "ns:{k:v}:tokens", 5, 60)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_raises_for_unsupported_client(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_register_script_raises_for_unsupported_client(monkeypatch, redis_no_ping):
|
||||
"""A client exposing neither register_script nor script_load fails loudly
|
||||
rather than silently returning a no-op callable."""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
|
|
@ -350,46 +319,34 @@ async def test_async_register_script_raises_for_unsupported_client(
|
|||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_cache_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_delete_cache_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_delete_cache("k")
|
||||
mock_redis_instance.delete.assert_awaited_once_with(expected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_cache_keys_namespaces_keys(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_delete_cache_keys_namespaces_keys(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.delete_cache_keys(["k"])
|
||||
mock_redis_instance.delete.assert_awaited_once_with(expected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_ttl_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_get_ttl_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.ttl = AsyncMock(return_value=42)
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
ttl = await redis_cache.async_get_ttl("k")
|
||||
assert ttl == 42
|
||||
mock_redis_instance.ttl.assert_awaited_once_with(expected)
|
||||
|
|
@ -397,41 +354,31 @@ async def test_async_get_ttl_namespaces_key(
|
|||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_lpop_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_lpop_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.lpop = AsyncMock(return_value=b"value")
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_lpop(key="k")
|
||||
mock_redis_instance.lpop.assert_awaited_once_with(expected, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_rpush_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.rpush = AsyncMock(return_value=1)
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_rpush("k", ["v"])
|
||||
mock_redis_instance.rpush.assert_awaited_once_with(expected, "v")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected_match", [(None, "k*"), ("ns", "ns:k*")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_scan_iter_namespaces_pattern(
|
||||
namespace, expected_match, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_scan_iter_namespaces_pattern(namespace, expected_match, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
|
||||
|
|
@ -448,17 +395,13 @@ async def test_async_scan_iter_namespaces_pattern(
|
|||
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.scan_iter = scan_iter
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_scan_iter(pattern="k")
|
||||
assert captured["match"] == expected_match
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
def test_increment_cache_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
def test_increment_cache_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_client = MagicMock()
|
||||
|
|
@ -1550,3 +1493,121 @@ async def test_async_rpush_and_trim_runs_push_and_trim_in_one_transaction(monkey
|
|||
assert pushed_len == 4
|
||||
assert rows == ["b", "c", "d"]
|
||||
assert pipe.queued == [("rpush", "ns:buf", "c", "d"), ("ltrim", "ns:buf", "-3", "-1")]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_redis_port() -> Iterator[int]:
|
||||
import threading
|
||||
|
||||
import fakeredis
|
||||
|
||||
server = fakeredis.TcpFakeServer(("127.0.0.1", 0), server_type="redis")
|
||||
worker = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
worker.start()
|
||||
try:
|
||||
yield server.server_address[1]
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
worker.join(timeout=5)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def socket_redis_cache(fake_redis_port: int, monkeypatch: pytest.MonkeyPatch) -> RedisCache:
|
||||
for name in (
|
||||
"REDIS_URL",
|
||||
"REDIS_HOST",
|
||||
"REDIS_PORT",
|
||||
"REDIS_PASSWORD",
|
||||
"REDIS_CLUSTER_NODES",
|
||||
"REDIS_SENTINEL_NODES",
|
||||
):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
return RedisCache(host="127.0.0.1", port=fake_redis_port)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_subscribe_receives_only_messages_published_after_it(socket_redis_cache: RedisCache) -> None:
|
||||
from litellm.caching.redis_cache import RedisMessage
|
||||
|
||||
assert await socket_redis_cache.async_publish("events", "before anyone listens") == 0
|
||||
|
||||
subscription = await socket_redis_cache.async_subscribe("events")
|
||||
assert await subscription.get_message(timeout=0) is None, "the subscribe ack is not a message"
|
||||
assert await socket_redis_cache.async_publish("events", "hello") == 1
|
||||
assert await socket_redis_cache.async_publish("other", b"ignored") == 0
|
||||
|
||||
assert await subscription.get_message(timeout=2) == RedisMessage(channel="events", payload=b"hello")
|
||||
assert await subscription.get_message(timeout=0) is None
|
||||
|
||||
await subscription.aclose()
|
||||
for _ in range(50):
|
||||
if await socket_redis_cache.async_publish("events", "nobody") == 0:
|
||||
break
|
||||
await asyncio.sleep(0.02)
|
||||
else:
|
||||
pytest.fail("a closed subscription still counts as a receiver")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_subscribe_covers_every_named_channel(socket_redis_cache: RedisCache) -> None:
|
||||
subscription = await socket_redis_cache.async_subscribe("first", "second")
|
||||
try:
|
||||
await socket_redis_cache.async_publish("second", "two")
|
||||
await socket_redis_cache.async_publish("first", "one")
|
||||
|
||||
received = [await subscription.get_message(timeout=2) for _ in range(2)]
|
||||
assert [(message.channel, message.payload) for message in received if message] == [
|
||||
("second", b"two"),
|
||||
("first", b"one"),
|
||||
]
|
||||
finally:
|
||||
await subscription.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pubsub_refuses_cluster_clients(fake_redis_port: int, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from redis.asyncio import RedisCluster
|
||||
|
||||
for name in (
|
||||
"REDIS_URL",
|
||||
"REDIS_HOST",
|
||||
"REDIS_PORT",
|
||||
"REDIS_PASSWORD",
|
||||
"REDIS_CLUSTER_NODES",
|
||||
"REDIS_SENTINEL_NODES",
|
||||
):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
|
||||
class _ClusterClientCache(RedisCache):
|
||||
def init_async_client(self):
|
||||
return RedisCluster.__new__(RedisCluster)
|
||||
|
||||
cache = _ClusterClientCache(host="127.0.0.1", port=fake_redis_port)
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await cache.async_publish("events", "hello")
|
||||
with pytest.raises(NotImplementedError):
|
||||
await cache.async_subscribe("events")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pubsub_calls_feed_the_circuit_breaker(socket_redis_cache: RedisCache) -> None:
|
||||
socket_redis_cache._circuit_breaker.record_failure()
|
||||
socket_redis_cache._circuit_breaker.record_failure()
|
||||
socket_redis_cache._circuit_breaker.record_failure()
|
||||
socket_redis_cache._circuit_breaker.record_failure()
|
||||
socket_redis_cache._circuit_breaker.record_failure()
|
||||
assert socket_redis_cache._circuit_breaker.is_open()
|
||||
|
||||
with pytest.raises(RedisCircuitBreakerOpenError):
|
||||
await socket_redis_cache.async_publish("events", "hello")
|
||||
with pytest.raises(RedisCircuitBreakerOpenError):
|
||||
await socket_redis_cache.async_subscribe("events")
|
||||
|
||||
|
||||
def test_connection_pool_status_reports_the_sync_pool(socket_redis_cache: RedisCache) -> None:
|
||||
status = socket_redis_cache.connection_pool_status()
|
||||
|
||||
assert status["max_connections"] == socket_redis_cache.redis_client.connection_pool.max_connections
|
||||
assert status["connection_class"] == socket_redis_cache.redis_client.connection_pool.connection_class.__name__
|
||||
|
|
|
|||
|
|
@ -1,7 +1,21 @@
|
|||
"""Tests for the Redis lock: acquire NX/PX with a token, compare-and-delete release, namespacing."""
|
||||
"""Tests for the Redis lock: acquire NX/PX with a token, owner-only extend and release, namespacing.
|
||||
|
||||
The happy paths run against a real ``redis-server`` because every operation is a Lua script; the
|
||||
degrade paths use a fake ``RedisCache`` whose scripts fail.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import shutil
|
||||
import subprocess
|
||||
import time
|
||||
from collections.abc import Callable, Iterator, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import redis
|
||||
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_distributed_lock import (
|
||||
RedisDistributedLock,
|
||||
)
|
||||
|
|
@ -10,117 +24,122 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_c
|
|||
)
|
||||
|
||||
|
||||
class _FakeRedis:
|
||||
"""Models just enough Redis to exercise SET NX (token store) and the compare-and-delete EVAL."""
|
||||
|
||||
def __init__(self, set_returns=True, exists_returns=1, raise_on=()):
|
||||
self._set_returns = set_returns
|
||||
self._exists_returns = exists_returns
|
||||
self._raise_on = set(raise_on)
|
||||
self.set_calls = []
|
||||
self.eval_calls = []
|
||||
self.expire_calls = []
|
||||
self.deleted = []
|
||||
self.store: dict = {}
|
||||
|
||||
async def set(self, name, value, *, nx=False, px=None):
|
||||
if "set" in self._raise_on:
|
||||
raise RuntimeError("redis down")
|
||||
self.set_calls.append((name, value, nx, px))
|
||||
if self._set_returns:
|
||||
self.store[name] = value
|
||||
return self._set_returns
|
||||
|
||||
async def eval(self, script, numkeys, *keys_and_args):
|
||||
if "eval" in self._raise_on:
|
||||
raise RuntimeError("redis down")
|
||||
self.eval_calls.append((numkeys, keys_and_args))
|
||||
key, token = keys_and_args[0], keys_and_args[1]
|
||||
if len(keys_and_args) == 3:
|
||||
ttl_ms = keys_and_args[2]
|
||||
if self.store.get(key) == token:
|
||||
self.expire_calls.append((key, ttl_ms))
|
||||
return 1
|
||||
return 0
|
||||
if self.store.get(key) == token: # compare-and-delete: only the owner deletes
|
||||
del self.store[key]
|
||||
self.deleted.append(key)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
async def exists(self, *names):
|
||||
if "exists" in self._raise_on:
|
||||
raise RuntimeError("redis down")
|
||||
return self._exists_returns
|
||||
@pytest.fixture
|
||||
def redis_port(tmp_path: Path, unused_tcp_port_factory: Callable[[], int]) -> Iterator[int]:
|
||||
server: Final = shutil.which("redis-server")
|
||||
if server is None:
|
||||
pytest.skip("redis-server is required for the lock's Lua scripts")
|
||||
port: Final = unused_tcp_port_factory()
|
||||
log_path: Final = tmp_path / "redis.log"
|
||||
config: Final = tmp_path / "redis.conf"
|
||||
config.write_text(f'bind 127.0.0.1\nport {port}\ndir "{tmp_path}"\nsave ""\nappendonly no\n')
|
||||
with log_path.open("w") as log:
|
||||
process: Final = subprocess.Popen((server, str(config)), stdout=log, stderr=subprocess.STDOUT)
|
||||
try:
|
||||
with redis.Redis(host="127.0.0.1", port=port, socket_timeout=1, socket_connect_timeout=1) as admin:
|
||||
for _ in range(100):
|
||||
try:
|
||||
admin.ping()
|
||||
break
|
||||
except redis.ConnectionError:
|
||||
time.sleep(0.1)
|
||||
else:
|
||||
pytest.fail(f"Redis did not start: {log_path.read_text()}")
|
||||
yield port
|
||||
finally:
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acquire_sets_token_with_nx_px_and_reports_acquired():
|
||||
redis = _FakeRedis(set_returns=True)
|
||||
assert await RedisDistributedLock(redis).acquire("k", "tok-1", 10.0) is (LockAcquisition.ACQUIRED)
|
||||
assert redis.set_calls == [("k", "tok-1", True, 10000)] # token value, NX, px in ms
|
||||
@pytest.fixture
|
||||
def redis_cache(redis_port: int, monkeypatch: pytest.MonkeyPatch) -> RedisCache:
|
||||
for name in (
|
||||
"REDIS_URL",
|
||||
"REDIS_HOST",
|
||||
"REDIS_PORT",
|
||||
"REDIS_PASSWORD",
|
||||
"REDIS_CLUSTER_NODES",
|
||||
"REDIS_SENTINEL_NODES",
|
||||
):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
return RedisCache(host="127.0.0.1", port=redis_port, namespace="tenant")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acquire_reports_held_when_key_already_held():
|
||||
# redis SET NX returns None when the key exists -> another worker holds it.
|
||||
assert await RedisDistributedLock(_FakeRedis(set_returns=None)).acquire("k", "tok", 10.0) is LockAcquisition.HELD
|
||||
@pytest.fixture
|
||||
def raw(redis_port: int) -> Iterator[redis.Redis]:
|
||||
with redis.Redis(host="127.0.0.1", port=redis_port) as client:
|
||||
yield client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acquire_reports_error_on_redis_error_distinct_from_held():
|
||||
# A dead backend must be distinguishable from a busy holder so the coordinator refreshes anyway.
|
||||
assert await RedisDistributedLock(_FakeRedis(raise_on=["set"])).acquire("k", "tok", 10.0) is LockAcquisition.ERROR
|
||||
class _FailingScriptsCache:
|
||||
"""A RedisCache whose every registered script fails at run time, as a dead Redis does."""
|
||||
|
||||
def async_register_script(self, script: str) -> Callable[..., object]:
|
||||
async def run(keys: Sequence[str], args: Sequence[str]) -> object:
|
||||
raise ConnectionError("redis down")
|
||||
|
||||
return run
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_deletes_only_when_the_token_matches():
|
||||
# Regression: a holder whose lock PX-expired and was re-acquired by another worker must not be
|
||||
# able to delete the new holder's lock. release with a stale token is a no-op.
|
||||
redis = _FakeRedis()
|
||||
lock = RedisDistributedLock(redis)
|
||||
await lock.acquire("k", "owner-B", 10.0) # B currently holds the lock
|
||||
await lock.release("k", "owner-A") # A's stale token
|
||||
assert redis.deleted == [] and redis.store.get("k") == "owner-B" # B's lock survives
|
||||
await lock.release("k", "owner-B") # the real owner releases
|
||||
assert redis.deleted == ["k"] and "k" not in redis.store
|
||||
async def test_acquire_wins_once_and_reports_held_to_the_next_caller(redis_cache: RedisCache, raw: redis.Redis) -> None:
|
||||
lock = RedisDistributedLock(redis_cache)
|
||||
|
||||
assert await lock.acquire("k", "tok-1", 10.0) is LockAcquisition.ACQUIRED
|
||||
assert await lock.acquire("k", "tok-2", 10.0) is LockAcquisition.HELD
|
||||
assert raw.get("tenant:k") == b"tok-1", "the key carries the cache namespace and the winner's token"
|
||||
assert 0 < raw.pttl("tenant:k") <= 10_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_keys_are_namespaced_before_reaching_redis():
|
||||
redis = _FakeRedis()
|
||||
lock = RedisDistributedLock(redis, namespace_key=lambda key: f"ns:{key}")
|
||||
await lock.acquire("k", "tok", 10.0)
|
||||
await lock.extend("k", "tok", 10.0)
|
||||
await lock.release("k", "tok")
|
||||
await lock.is_held("k")
|
||||
assert redis.set_calls[0][0] == "ns:k" # acquire namespaced
|
||||
assert redis.eval_calls[0][1][0] == "ns:k" # extend (EVAL KEYS[1]) namespaced
|
||||
assert redis.eval_calls[1][1][0] == "ns:k" # release (EVAL KEYS[1]) namespaced
|
||||
assert redis.deleted == ["ns:k"]
|
||||
async def test_acquire_reports_error_on_redis_error_distinct_from_held() -> None:
|
||||
lock = RedisDistributedLock(_FailingScriptsCache()) # pyright: ignore[reportArgumentType] # duck-typed fake
|
||||
|
||||
assert await lock.acquire("k", "tok", 10.0) is LockAcquisition.ERROR
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extend_refreshes_ttl_only_when_the_token_matches():
|
||||
redis = _FakeRedis()
|
||||
lock = RedisDistributedLock(redis)
|
||||
async def test_release_deletes_only_when_the_token_matches(redis_cache: RedisCache) -> None:
|
||||
lock = RedisDistributedLock(redis_cache)
|
||||
await lock.acquire("k", "owner-B", 10.0)
|
||||
assert await lock.extend("k", "owner-A", 10.0) is False
|
||||
assert await lock.extend("k", "owner-B", 10.0) is True
|
||||
assert redis.expire_calls == [("k", "10000")]
|
||||
|
||||
await lock.release("k", "owner-A")
|
||||
assert await lock.is_held("k") is True, "a stale token must not delete another worker's lock"
|
||||
|
||||
await lock.release("k", "owner-B")
|
||||
assert await lock.is_held("k") is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extend_degrades_to_false_on_redis_error():
|
||||
assert await RedisDistributedLock(_FakeRedis(raise_on=["eval"])).extend("k", "tok", 10.0) is False
|
||||
async def test_extend_refreshes_ttl_only_when_the_token_matches(redis_cache: RedisCache, raw: redis.Redis) -> None:
|
||||
lock = RedisDistributedLock(redis_cache)
|
||||
await lock.acquire("k", "owner-B", 1.0)
|
||||
|
||||
assert await lock.extend("k", "owner-A", 30.0) is False
|
||||
assert raw.pttl("tenant:k") <= 1_000
|
||||
assert await lock.extend("k", "owner-B", 30.0) is True
|
||||
assert raw.pttl("tenant:k") > 1_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_is_held_reflects_exists():
|
||||
assert await RedisDistributedLock(_FakeRedis(exists_returns=1)).is_held("k") is True
|
||||
assert await RedisDistributedLock(_FakeRedis(exists_returns=0)).is_held("k") is False
|
||||
async def test_extend_and_release_degrade_on_redis_error() -> None:
|
||||
lock = RedisDistributedLock(_FailingScriptsCache()) # pyright: ignore[reportArgumentType] # duck-typed fake
|
||||
|
||||
assert await lock.extend("k", "tok", 10.0) is False
|
||||
await lock.release("k", "tok")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_is_held_degrades_to_false_on_redis_error():
|
||||
assert await RedisDistributedLock(_FakeRedis(raise_on=["exists"])).is_held("k") is False
|
||||
async def test_lock_expires_on_its_own_after_the_ttl(redis_cache: RedisCache) -> None:
|
||||
lock = RedisDistributedLock(redis_cache)
|
||||
await lock.acquire("k", "tok", 0.05)
|
||||
assert await lock.is_held("k") is True
|
||||
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
assert await lock.is_held("k") is False
|
||||
assert await lock.acquire("k", "tok-2", 10.0) is LockAcquisition.ACQUIRED
|
||||
|
||||
|
||||
async def test_is_held_degrades_to_false_on_redis_error() -> None:
|
||||
lock = RedisDistributedLock(_FailingScriptsCache()) # pyright: ignore[reportArgumentType] # duck-typed fake
|
||||
|
||||
assert await lock.is_held("k") is False
|
||||
|
|
|
|||
|
|
@ -2,14 +2,14 @@ import asyncio
|
|||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Iterable
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from redis.asyncio import Redis
|
||||
|
||||
import litellm.proxy.common_utils.auth_cache_invalidation_pubsub as pubsub_module
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cache import RedisMessage
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
|
||||
AUTH_CACHE_INVALIDATION_CHANNEL,
|
||||
|
|
@ -20,16 +20,7 @@ from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
|
|||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
|
||||
class _RecordingRedisClient(Redis):
|
||||
def __init__(self) -> None:
|
||||
self.published: list[tuple[str, str]] = []
|
||||
|
||||
async def publish(self, channel: str, message: str) -> int:
|
||||
self.published.append((channel, message))
|
||||
return 1
|
||||
|
||||
|
||||
class _WedgedPublishRedisClient(Redis):
|
||||
class _WedgedPublisher:
|
||||
def __init__(self) -> None:
|
||||
self.attempted: list[str] = []
|
||||
self.in_flight = 0
|
||||
|
|
@ -45,26 +36,17 @@ class _WedgedPublishRedisClient(Redis):
|
|||
return 1
|
||||
|
||||
|
||||
class _FailingPublishRedisClient(Redis):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
async def publish(self, channel: str, message: str) -> int:
|
||||
raise ConnectionError("redis down")
|
||||
|
||||
|
||||
class _QueuePubSub:
|
||||
def __init__(self, initial_messages: Iterable[object] = ()) -> None:
|
||||
self.queue: asyncio.Queue[object] = asyncio.Queue()
|
||||
"""A subscription fed from a queue of messages, standing in for RedisSubscription."""
|
||||
|
||||
def __init__(self, initial_messages: Iterable[RedisMessage] = ()) -> None:
|
||||
self.queue: asyncio.Queue[RedisMessage] = asyncio.Queue()
|
||||
for message in initial_messages:
|
||||
self.queue.put_nowait(message)
|
||||
self.subscribed_channels: list[str] = []
|
||||
self.closed = False
|
||||
|
||||
async def subscribe(self, *channels: str) -> None:
|
||||
self.subscribed_channels.extend(channels)
|
||||
|
||||
async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> object | None:
|
||||
async def get_message(self, *, timeout: float | None) -> RedisMessage | None:
|
||||
try:
|
||||
return await asyncio.wait_for(self.queue.get(), timeout)
|
||||
except asyncio.TimeoutError:
|
||||
|
|
@ -74,49 +56,70 @@ class _QueuePubSub:
|
|||
self.closed = True
|
||||
|
||||
|
||||
class _ScriptedPubSubRedisClient(Redis):
|
||||
def __init__(self, pubsubs: Iterable[_QueuePubSub]) -> None:
|
||||
self._scripted_pubsubs = iter(pubsubs)
|
||||
|
||||
def pubsub(self) -> _QueuePubSub:
|
||||
return next(self._scripted_pubsubs)
|
||||
|
||||
|
||||
class _FakeRedisCache:
|
||||
def __init__(self, client: object, namespace: str | None = None) -> None:
|
||||
self._client = client
|
||||
"""The slice of RedisCache the module uses: namespace, async_publish and async_subscribe."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
subscriptions: Iterable[_QueuePubSub] = (),
|
||||
namespace: str | None = None,
|
||||
publish: Callable[[str, str], Awaitable[int]] | None = None,
|
||||
publish_error: Exception | None = None,
|
||||
) -> None:
|
||||
self._subscriptions = iter(subscriptions)
|
||||
self._publish = publish
|
||||
self._publish_error = publish_error
|
||||
self.namespace = namespace
|
||||
self.published: list[tuple[str, str]] = []
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
return self._client
|
||||
async def async_publish(self, channel: str, message: str) -> int:
|
||||
if self._publish_error is not None:
|
||||
raise self._publish_error
|
||||
if self._publish is not None:
|
||||
return await self._publish(channel, message)
|
||||
self.published.append((channel, message))
|
||||
return 1
|
||||
|
||||
async def async_subscribe(self, *channels: str) -> _QueuePubSub:
|
||||
subscription = next(self._subscriptions)
|
||||
subscription.subscribed_channels.extend(channels)
|
||||
return subscription
|
||||
|
||||
|
||||
def _invalidation_message(cache_key: str) -> dict:
|
||||
return {"type": "message", "data": json.dumps({"cache_key": cache_key}).encode()}
|
||||
class _ClusterRedisCache(_FakeRedisCache):
|
||||
async def async_publish(self, channel: str, message: str) -> int:
|
||||
raise NotImplementedError("Redis Cluster clients have no pub/sub support")
|
||||
|
||||
async def async_subscribe(self, *channels: str) -> _QueuePubSub:
|
||||
raise NotImplementedError("Redis Cluster clients have no pub/sub support")
|
||||
|
||||
|
||||
def _invalidation_message(cache_key: str) -> RedisMessage:
|
||||
return RedisMessage(channel=AUTH_CACHE_INVALIDATION_CHANNEL, payload=json.dumps({"cache_key": cache_key}).encode())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_publish_sends_cache_key_json_on_channel() -> None:
|
||||
client = _RecordingRedisClient()
|
||||
cache = _FakeRedisCache()
|
||||
with patch(
|
||||
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
|
||||
return_value=_FakeRedisCache(client=client),
|
||||
return_value=cache,
|
||||
):
|
||||
await publish_auth_cache_invalidation(cache_key="project_id:p-1")
|
||||
|
||||
assert client.published == [(AUTH_CACHE_INVALIDATION_CHANNEL, json.dumps({"cache_key": "project_id:p-1"}))]
|
||||
assert cache.published == [(AUTH_CACHE_INVALIDATION_CHANNEL, json.dumps({"cache_key": "project_id:p-1"}))]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_publish_uses_namespaced_channel() -> None:
|
||||
client = _RecordingRedisClient()
|
||||
cache = _FakeRedisCache(namespace="ns1")
|
||||
with patch(
|
||||
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
|
||||
return_value=_FakeRedisCache(client=client, namespace="ns1"),
|
||||
return_value=cache,
|
||||
):
|
||||
await publish_auth_cache_invalidation(cache_key="project_id:p-1")
|
||||
|
||||
assert client.published[0][0] == f"ns1:{AUTH_CACHE_INVALIDATION_CHANNEL}"
|
||||
assert cache.published[0][0] == f"ns1:{AUTH_CACHE_INVALIDATION_CHANNEL}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -132,11 +135,33 @@ async def test_publish_noops_without_coordination_redis() -> None:
|
|||
async def test_publish_swallows_redis_errors() -> None:
|
||||
with patch(
|
||||
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
|
||||
return_value=_FakeRedisCache(client=_FailingPublishRedisClient()),
|
||||
return_value=_FakeRedisCache(publish_error=ConnectionError("redis down")),
|
||||
):
|
||||
await publish_auth_cache_invalidation(cache_key="project_id:p-1")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_publish_skips_clients_without_pubsub_support() -> None:
|
||||
with patch(
|
||||
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
|
||||
return_value=_ClusterRedisCache(),
|
||||
):
|
||||
await publish_auth_cache_invalidation(cache_key="project_id:p-1")
|
||||
await asyncio.gather(*pubsub_module._pending_publishes) # pyright: ignore[reportPrivateUsage] # drain module-level tasks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subscriber_disables_itself_without_pubsub_support() -> None:
|
||||
subscriber = AuthCacheInvalidationSubscriber(redis_cache=_ClusterRedisCache(), user_api_key_cache=UserApiKeyCache())
|
||||
|
||||
subscriber.start()
|
||||
task = subscriber._task
|
||||
assert task is not None
|
||||
await asyncio.wait_for(task, timeout=5)
|
||||
|
||||
assert task.done() is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subscriber_deletes_local_cache_entry_on_message() -> None:
|
||||
"""
|
||||
|
|
@ -150,7 +175,7 @@ async def test_subscriber_deletes_local_cache_entry_on_message() -> None:
|
|||
|
||||
pubsub = _QueuePubSub(initial_messages=[_invalidation_message("project_id:p-1")])
|
||||
subscriber = AuthCacheInvalidationSubscriber(
|
||||
redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])),
|
||||
redis_cache=_FakeRedisCache(subscriptions=[pubsub]),
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
subscriber.start()
|
||||
|
|
@ -180,7 +205,7 @@ async def test_subscriber_deletes_key_object_partition_entry_on_message() -> Non
|
|||
|
||||
pubsub = _QueuePubSub(initial_messages=[_invalidation_message(hashed_token)])
|
||||
subscriber = AuthCacheInvalidationSubscriber(
|
||||
redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])),
|
||||
redis_cache=_FakeRedisCache(subscriptions=[pubsub]),
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
subscriber.start()
|
||||
|
|
@ -210,7 +235,7 @@ async def test_subscriber_deletes_additional_in_memory_cache_entry_on_message()
|
|||
|
||||
pubsub = _QueuePubSub(initial_messages=[_invalidation_message("spend:team_member:u-1:t-1")])
|
||||
subscriber = AuthCacheInvalidationSubscriber(
|
||||
redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])),
|
||||
redis_cache=_FakeRedisCache(subscriptions=[pubsub]),
|
||||
user_api_key_cache=cache,
|
||||
additional_in_memory_caches=(spend_counter_in_memory_cache,),
|
||||
)
|
||||
|
|
@ -232,13 +257,13 @@ async def test_subscriber_ignores_malformed_messages() -> None:
|
|||
cache.in_memory_cache.set_cache("project_id:p-1", {"models": []})
|
||||
|
||||
subscriber = AuthCacheInvalidationSubscriber(
|
||||
redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[_QueuePubSub()])),
|
||||
redis_cache=_FakeRedisCache(subscriptions=[_QueuePubSub()]),
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
subscriber._apply_message({"type": "message", "data": b"not json"})
|
||||
subscriber._apply_message({"type": "message", "data": json.dumps({"other": "x"}).encode()})
|
||||
subscriber._apply_message("raw string")
|
||||
subscriber._apply_message(None)
|
||||
subscriber._apply_message(RedisMessage(channel=AUTH_CACHE_INVALIDATION_CHANNEL, payload=b"not json"))
|
||||
subscriber._apply_message(
|
||||
RedisMessage(channel=AUTH_CACHE_INVALIDATION_CHANNEL, payload=json.dumps({"other": "x"}).encode())
|
||||
)
|
||||
|
||||
assert cache.in_memory_cache.get_cache("project_id:p-1") is not None
|
||||
|
||||
|
|
@ -247,11 +272,11 @@ async def test_subscriber_ignores_malformed_messages() -> None:
|
|||
async def test_evict_and_broadcast_evicts_locally_and_returns_while_redis_publish_never_answers() -> None:
|
||||
cache = UserApiKeyCache()
|
||||
cache.set_cache("user-wedged", UserAPIKeyAuth(user_id="user-wedged"), model_type=UserAPIKeyAuth)
|
||||
client = _WedgedPublishRedisClient()
|
||||
client = _WedgedPublisher()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
|
||||
return_value=_FakeRedisCache(client=client),
|
||||
return_value=_FakeRedisCache(publish=client.publish),
|
||||
):
|
||||
started = time.monotonic()
|
||||
await evict_and_broadcast(cache_keys=("user-wedged",), user_api_key_cache=cache)
|
||||
|
|
@ -270,11 +295,11 @@ async def test_publish_holds_at_most_sixteen_redis_connections_while_redis_is_we
|
|||
) -> None:
|
||||
monkeypatch.setattr(pubsub_module, "_in_flight_publishes", asyncio.Semaphore(16))
|
||||
monkeypatch.setattr(pubsub_module, "_pending_publishes", set())
|
||||
client = _WedgedPublishRedisClient()
|
||||
client = _WedgedPublisher()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
|
||||
return_value=_FakeRedisCache(client=client),
|
||||
return_value=_FakeRedisCache(publish=client.publish),
|
||||
):
|
||||
for i in range(64):
|
||||
await publish_auth_cache_invalidation(cache_key=f"user-{i}")
|
||||
|
|
|
|||
|
|
@ -1,22 +1,22 @@
|
|||
import asyncio
|
||||
import json
|
||||
import random
|
||||
from typing import Callable, Coroutine, Iterable, List, Optional, Tuple
|
||||
from collections.abc import Callable, Coroutine, Iterable
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from redis.asyncio import Redis
|
||||
|
||||
import litellm
|
||||
from litellm.caching.redis_cache import RedisMessage
|
||||
from litellm.proxy.common_utils.config_sync_pubsub import (
|
||||
_CONFIG_SYNCED_TABLE_NAMES,
|
||||
_RESYNC_APPLIED_CONFIG_PARAM_NAMES,
|
||||
_WRITE_ACTION_NAMES,
|
||||
CONFIG_SYNC_CHANNEL,
|
||||
CONFIG_SYNC_JITTER_MAX_SECONDS,
|
||||
CONFIG_SYNC_MIN_RESYNC_INTERVAL_SECONDS,
|
||||
ConfigSyncSubscriber,
|
||||
_CONFIG_SYNCED_TABLE_NAMES,
|
||||
_PublishOnWriteActions,
|
||||
_RESYNC_APPLIED_CONFIG_PARAM_NAMES,
|
||||
_WRITE_ACTION_NAMES,
|
||||
publish_config_change,
|
||||
wrap_table_actions_for_config_sync,
|
||||
)
|
||||
|
|
@ -64,60 +64,36 @@ _EXPECTED_RESYNC_APPLIED_CONFIG_PARAM_NAMES = frozenset(
|
|||
_STARTUP_ONLY_CONFIG_PARAM_NAMES = ("environment_variables",)
|
||||
|
||||
|
||||
class _RecordingRedisClient(Redis):
|
||||
def __init__(self) -> None:
|
||||
self.published: List[Tuple[str, str]] = []
|
||||
|
||||
async def publish(self, channel: str, message: str) -> int:
|
||||
self.published.append((channel, message))
|
||||
return 1
|
||||
|
||||
|
||||
class _FailingPublishRedisClient(Redis):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
async def publish(self, channel: str, message: str) -> int:
|
||||
raise ConnectionError("redis down")
|
||||
|
||||
|
||||
class _NotRedisClient:
|
||||
def __init__(self) -> None:
|
||||
self.published: List[Tuple[str, str]] = []
|
||||
|
||||
async def publish(self, channel: str, message: str) -> int:
|
||||
self.published.append((channel, message))
|
||||
return 1
|
||||
|
||||
|
||||
class _QueuePubSub:
|
||||
"""A subscription fed from a queue of payloads, standing in for RedisSubscription."""
|
||||
|
||||
def __init__(self, initial_messages: Iterable[str] = ()) -> None:
|
||||
self.queue: "asyncio.Queue[str]" = asyncio.Queue()
|
||||
self.queue: asyncio.Queue[str] = asyncio.Queue()
|
||||
for message in initial_messages:
|
||||
self.queue.put_nowait(message)
|
||||
self.subscribed_channels: List[str] = []
|
||||
self.subscribed_channels: list[str] = []
|
||||
self.closed = False
|
||||
|
||||
async def subscribe(self, *channels: str) -> None:
|
||||
self.subscribed_channels.extend(channels)
|
||||
|
||||
async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Optional[str]:
|
||||
async def get_message(self, *, timeout: float | None) -> RedisMessage | None:
|
||||
if timeout == 0:
|
||||
try:
|
||||
return self.queue.get_nowait()
|
||||
return self._message(self.queue.get_nowait())
|
||||
except asyncio.QueueEmpty:
|
||||
return None
|
||||
try:
|
||||
return await asyncio.wait_for(self.queue.get(), timeout)
|
||||
return self._message(await asyncio.wait_for(self.queue.get(), timeout))
|
||||
except asyncio.TimeoutError:
|
||||
return None
|
||||
|
||||
def _message(self, payload: str) -> RedisMessage:
|
||||
return RedisMessage(channel=self.subscribed_channels[0], payload=payload.encode())
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
class _BrokenPubSub(_QueuePubSub):
|
||||
async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Optional[str]:
|
||||
async def get_message(self, *, timeout: float | None) -> RedisMessage | None:
|
||||
raise ConnectionError("connection lost")
|
||||
|
||||
|
||||
|
|
@ -131,11 +107,11 @@ class _EmptyPollsThenMessagePubSub(_QueuePubSub):
|
|||
super().__init__(initial_messages=initial_messages)
|
||||
self.remaining_empty_polls = empty_polls
|
||||
|
||||
async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Optional[str]:
|
||||
async def get_message(self, *, timeout: float | None) -> RedisMessage | None:
|
||||
if timeout != 0 and self.remaining_empty_polls > 0:
|
||||
self.remaining_empty_polls -= 1
|
||||
return None
|
||||
return await super().get_message(ignore_subscribe_messages=ignore_subscribe_messages, timeout=timeout)
|
||||
return await super().get_message(timeout=timeout)
|
||||
|
||||
|
||||
class _FakeClock:
|
||||
|
|
@ -146,32 +122,44 @@ class _FakeClock:
|
|||
return self.now
|
||||
|
||||
|
||||
class _ScriptedPubSubRedisClient(Redis):
|
||||
def __init__(self, pubsubs: Iterable[_QueuePubSub]) -> None:
|
||||
self._scripted_pubsubs = iter(pubsubs)
|
||||
|
||||
def pubsub(self) -> _QueuePubSub:
|
||||
return next(self._scripted_pubsubs)
|
||||
|
||||
|
||||
class _FakeRedisCache:
|
||||
def __init__(self, client: object, namespace: Optional[str] = None) -> None:
|
||||
self._client = client
|
||||
"""The slice of RedisCache the module uses: namespace, async_publish and async_subscribe."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
subscriptions: Iterable[_QueuePubSub] = (),
|
||||
namespace: str | None = None,
|
||||
publish_error: Exception | None = None,
|
||||
) -> None:
|
||||
self._subscriptions = iter(subscriptions)
|
||||
self._publish_error = publish_error
|
||||
self.namespace = namespace
|
||||
self.published: list[tuple[str, str]] = []
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
return self._client
|
||||
async def async_publish(self, channel: str, message: str) -> int:
|
||||
if self._publish_error is not None:
|
||||
raise self._publish_error
|
||||
self.published.append((channel, message))
|
||||
return 1
|
||||
|
||||
async def async_subscribe(self, *channels: str) -> _QueuePubSub:
|
||||
subscription = next(self._subscriptions)
|
||||
subscription.subscribed_channels.extend(channels)
|
||||
return subscription
|
||||
|
||||
|
||||
class _ExplodingRedisCache:
|
||||
namespace: Optional[str] = None
|
||||
class _ClusterRedisCache(_FakeRedisCache):
|
||||
"""RedisCache over a cluster client raises NotImplementedError for both pub/sub methods."""
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
raise ConnectionError("cannot connect")
|
||||
async def async_publish(self, channel: str, message: str) -> int:
|
||||
raise NotImplementedError("Redis Cluster clients have no pub/sub support")
|
||||
|
||||
async def async_subscribe(self, *channels: str) -> _QueuePubSub:
|
||||
raise NotImplementedError("Redis Cluster clients have no pub/sub support")
|
||||
|
||||
|
||||
def _recording_callback(
|
||||
events: List[str], name: str, fired: asyncio.Event
|
||||
events: list[str], name: str, fired: asyncio.Event
|
||||
) -> Callable[[], Coroutine[None, None, None]]:
|
||||
async def callback() -> None:
|
||||
events.append(name)
|
||||
|
|
@ -185,49 +173,44 @@ async def test_publish_noops_when_redis_cache_is_none() -> None:
|
|||
|
||||
|
||||
async def test_publish_sends_object_type_json_on_channel() -> None:
|
||||
client = _RecordingRedisClient()
|
||||
cache = _FakeRedisCache(client)
|
||||
cache = _FakeRedisCache()
|
||||
|
||||
await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable")
|
||||
|
||||
assert len(client.published) == 1
|
||||
channel, message = client.published[0]
|
||||
assert len(cache.published) == 1
|
||||
channel, message = cache.published[0]
|
||||
assert channel == "litellm_proxy.config_change"
|
||||
assert json.loads(message) == {"object_type": "litellm_proxymodeltable"}
|
||||
|
||||
|
||||
async def test_publish_uses_namespaced_channel() -> None:
|
||||
client = _RecordingRedisClient()
|
||||
cache = _FakeRedisCache(client, namespace="prod-eu")
|
||||
cache = _FakeRedisCache(namespace="prod-eu")
|
||||
|
||||
await publish_config_change(redis_cache=cache, object_type="litellm_credentialstable")
|
||||
|
||||
assert client.published[0][0] == "prod-eu:litellm_proxy.config_change"
|
||||
assert cache.published[0][0] == "prod-eu:litellm_proxy.config_change"
|
||||
|
||||
|
||||
async def test_publish_swallows_redis_publish_errors() -> None:
|
||||
cache = _FakeRedisCache(_FailingPublishRedisClient())
|
||||
cache = _FakeRedisCache(publish_error=ConnectionError("redis down"))
|
||||
|
||||
await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable")
|
||||
|
||||
|
||||
async def test_publish_swallows_client_init_errors() -> None:
|
||||
await publish_config_change(redis_cache=_ExplodingRedisCache(), object_type="litellm_proxymodeltable")
|
||||
assert cache.published == []
|
||||
|
||||
|
||||
async def test_publish_skips_clients_without_pubsub_support() -> None:
|
||||
client = _NotRedisClient()
|
||||
cache = _FakeRedisCache(client)
|
||||
cache = _ClusterRedisCache()
|
||||
|
||||
await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable")
|
||||
|
||||
assert client.published == []
|
||||
assert cache.published == []
|
||||
|
||||
|
||||
async def test_subscriber_runs_injected_callbacks_in_order_on_message() -> None:
|
||||
pubsub = _QueuePubSub()
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
events: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
events: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
subscriber = ConfigSyncSubscriber(
|
||||
redis_cache=cache,
|
||||
|
|
@ -252,8 +235,8 @@ async def test_subscriber_runs_injected_callbacks_in_order_on_message() -> None:
|
|||
async def test_burst_within_debounce_window_coalesces_into_one_resync() -> None:
|
||||
burst = [json.dumps({"object_type": "litellm_proxymodeltable"}) for _ in range(5)]
|
||||
pubsub = _QueuePubSub(initial_messages=burst)
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
resyncs: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
resyncs: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
subscriber = ConfigSyncSubscriber(
|
||||
redis_cache=cache,
|
||||
|
|
@ -273,8 +256,8 @@ async def test_burst_within_debounce_window_coalesces_into_one_resync() -> None:
|
|||
|
||||
async def test_subscriber_subscribes_on_namespaced_channel_and_resyncs() -> None:
|
||||
pubsub = _QueuePubSub()
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]), namespace="prod-eu")
|
||||
resyncs: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub], namespace="prod-eu")
|
||||
resyncs: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
subscriber = ConfigSyncSubscriber(
|
||||
redis_cache=cache,
|
||||
|
|
@ -299,8 +282,8 @@ class _MaxJitterRandom(random.Random):
|
|||
|
||||
async def test_debounce_sleep_adds_jitter_from_injected_rng() -> None:
|
||||
pubsub = _QueuePubSub(initial_messages=[json.dumps({"object_type": "litellm_proxymodeltable"})])
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
sleeps: List[float] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
sleeps: list[float] = []
|
||||
fired = asyncio.Event()
|
||||
|
||||
async def recording_sleep(seconds: float) -> None:
|
||||
|
|
@ -332,7 +315,7 @@ def test_default_min_resync_interval_caps_reload_rate() -> None:
|
|||
|
||||
def _throttled_subscriber(
|
||||
cache: object,
|
||||
events: List[str],
|
||||
events: list[str],
|
||||
fired: asyncio.Event,
|
||||
clock: _FakeClock,
|
||||
min_resync_interval_seconds: float = 10.0,
|
||||
|
|
@ -358,8 +341,8 @@ def _throttled_subscriber(
|
|||
|
||||
async def test_resync_arriving_inside_min_interval_waits_out_the_remainder() -> None:
|
||||
pubsub = _QueuePubSub()
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
events: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
events: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
clock = _FakeClock()
|
||||
subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock)
|
||||
|
|
@ -378,8 +361,8 @@ async def test_resync_arriving_inside_min_interval_waits_out_the_remainder() ->
|
|||
|
||||
async def test_resync_after_min_interval_elapsed_is_not_throttled() -> None:
|
||||
pubsub = _QueuePubSub()
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
events: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
events: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
clock = _FakeClock()
|
||||
subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock)
|
||||
|
|
@ -398,8 +381,8 @@ async def test_resync_after_min_interval_elapsed_is_not_throttled() -> None:
|
|||
|
||||
async def test_writes_during_the_throttle_wait_collapse_into_the_next_resync() -> None:
|
||||
pubsub = _QueuePubSub()
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
events: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
events: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
clock = _FakeClock()
|
||||
subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock)
|
||||
|
|
@ -420,8 +403,8 @@ async def test_writes_during_the_throttle_wait_collapse_into_the_next_resync() -
|
|||
|
||||
async def test_polls_without_messages_do_not_trigger_resyncs() -> None:
|
||||
pubsub = _EmptyPollsThenMessagePubSub(empty_polls=3, initial_messages=["change"])
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
resyncs: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
resyncs: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
subscriber = ConfigSyncSubscriber(
|
||||
redis_cache=cache,
|
||||
|
|
@ -442,8 +425,8 @@ async def test_polls_without_messages_do_not_trigger_resyncs() -> None:
|
|||
async def test_failing_pubsub_close_still_reconnects() -> None:
|
||||
broken = _CloseFailingBrokenPubSub()
|
||||
healthy = _QueuePubSub(initial_messages=["change"])
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([broken, healthy]))
|
||||
resyncs: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[broken, healthy])
|
||||
resyncs: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
subscriber = ConfigSyncSubscriber(
|
||||
redis_cache=cache,
|
||||
|
|
@ -464,7 +447,7 @@ async def test_failing_pubsub_close_still_reconnects() -> None:
|
|||
|
||||
async def test_second_start_does_not_open_a_second_subscription() -> None:
|
||||
pubsub = _QueuePubSub()
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
subscriber = ConfigSyncSubscriber(redis_cache=cache, resync_callbacks=(), debounce_seconds=0.01)
|
||||
|
||||
subscriber.start()
|
||||
|
|
@ -481,8 +464,8 @@ async def test_second_start_does_not_open_a_second_subscription() -> None:
|
|||
async def test_redis_error_leads_to_backoff_and_resubscribe() -> None:
|
||||
broken = _BrokenPubSub()
|
||||
healthy = _QueuePubSub(initial_messages=[json.dumps({"object_type": "litellm_credentialstable"})])
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([broken, healthy]))
|
||||
resyncs: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[broken, healthy])
|
||||
resyncs: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
subscriber = ConfigSyncSubscriber(
|
||||
redis_cache=cache,
|
||||
|
|
@ -508,8 +491,8 @@ async def test_redis_error_leads_to_backoff_and_resubscribe() -> None:
|
|||
|
||||
async def test_failing_resync_callback_does_not_kill_subscriber() -> None:
|
||||
pubsub = _QueuePubSub()
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
resyncs: List[str] = []
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
resyncs: list[str] = []
|
||||
fired = asyncio.Event()
|
||||
|
||||
async def failing_callback() -> None:
|
||||
|
|
@ -536,7 +519,7 @@ async def test_failing_resync_callback_does_not_kill_subscriber() -> None:
|
|||
|
||||
async def test_stop_cancels_subscriber_cleanly() -> None:
|
||||
pubsub = _QueuePubSub()
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]))
|
||||
cache = _FakeRedisCache(subscriptions=[pubsub])
|
||||
subscriber = ConfigSyncSubscriber(redis_cache=cache, resync_callbacks=(), debounce_seconds=0.01)
|
||||
|
||||
subscriber.start()
|
||||
|
|
@ -552,15 +535,15 @@ async def test_stop_cancels_subscriber_cleanly() -> None:
|
|||
|
||||
|
||||
async def test_stop_before_start_is_a_noop() -> None:
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([]))
|
||||
cache = _FakeRedisCache(subscriptions=[])
|
||||
subscriber = ConfigSyncSubscriber(redis_cache=cache, resync_callbacks=())
|
||||
|
||||
await subscriber.stop()
|
||||
|
||||
|
||||
async def test_subscriber_exits_without_callbacks_when_client_lacks_pubsub() -> None:
|
||||
cache = _FakeRedisCache(_NotRedisClient())
|
||||
resyncs: List[str] = []
|
||||
cache = _ClusterRedisCache()
|
||||
resyncs: list[str] = []
|
||||
subscriber = ConfigSyncSubscriber(
|
||||
redis_cache=cache,
|
||||
resync_callbacks=(_recording_callback(resyncs, "resync", asyncio.Event()),),
|
||||
|
|
@ -575,7 +558,7 @@ async def test_subscriber_exits_without_callbacks_when_client_lacks_pubsub() ->
|
|||
|
||||
|
||||
class _FakeTableActions:
|
||||
def __init__(self, calls: List[Tuple[str, str]]) -> None:
|
||||
def __init__(self, calls: list[tuple[str, str]]) -> None:
|
||||
self._calls = calls
|
||||
|
||||
async def create(self, **kwargs: object) -> object:
|
||||
|
|
@ -588,7 +571,7 @@ class _FakeTableActions:
|
|||
|
||||
|
||||
class _AllWritesTableActions:
|
||||
def __init__(self, calls: List[str]) -> None:
|
||||
def __init__(self, calls: list[str]) -> None:
|
||||
self._calls = calls
|
||||
|
||||
def __getattr__(self, name: str) -> Callable[..., Coroutine[None, None, str]]:
|
||||
|
|
@ -599,7 +582,7 @@ class _AllWritesTableActions:
|
|||
return action
|
||||
|
||||
|
||||
def _recording_publish(calls: List[Tuple[str, str]]) -> Callable[[str], Coroutine[None, None, None]]:
|
||||
def _recording_publish(calls: list[tuple[str, str]]) -> Callable[[str], Coroutine[None, None, None]]:
|
||||
async def publish(object_type: str) -> None:
|
||||
calls.append(("publish", object_type))
|
||||
|
||||
|
|
@ -615,7 +598,7 @@ def test_wrapper_passes_through_unsynced_tables() -> None:
|
|||
|
||||
|
||||
async def test_wrapper_publishes_table_name_after_write() -> None:
|
||||
calls: List[Tuple[str, str]] = []
|
||||
calls: list[tuple[str, str]] = []
|
||||
wrapped = wrap_table_actions_for_config_sync(
|
||||
actions=_FakeTableActions(calls),
|
||||
table_name="litellm_proxymodeltable",
|
||||
|
|
@ -629,7 +612,7 @@ async def test_wrapper_publishes_table_name_after_write() -> None:
|
|||
|
||||
|
||||
async def test_wrapper_does_not_publish_on_reads() -> None:
|
||||
calls: List[Tuple[str, str]] = []
|
||||
calls: list[tuple[str, str]] = []
|
||||
wrapped = wrap_table_actions_for_config_sync(
|
||||
actions=_FakeTableActions(calls),
|
||||
table_name="litellm_proxymodeltable",
|
||||
|
|
@ -660,8 +643,8 @@ def test_tool_telemetry_table_writes_pass_through_unwrapped() -> None:
|
|||
|
||||
@pytest.mark.parametrize("action_name", _EXPECTED_WRITE_ACTION_NAMES)
|
||||
async def test_wrapper_publishes_for_every_write_action(action_name: str) -> None:
|
||||
write_calls: List[str] = []
|
||||
publish_calls: List[Tuple[str, str]] = []
|
||||
write_calls: list[str] = []
|
||||
publish_calls: list[tuple[str, str]] = []
|
||||
wrapped = wrap_table_actions_for_config_sync(
|
||||
actions=_AllWritesTableActions(write_calls),
|
||||
table_name="litellm_guardrailstable",
|
||||
|
|
@ -680,7 +663,7 @@ async def test_model_repository_write_publishes_via_live_coordination_cache() ->
|
|||
from litellm.proxy.proxy_server import _set_redis_usage_cache
|
||||
from litellm.repositories.model_repository import ModelRepository
|
||||
|
||||
client = _RecordingRedisClient()
|
||||
client = _FakeRedisCache()
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_proxymodeltable.update = AsyncMock(return_value={"model_id": "m-1"})
|
||||
repository = ModelRepository(prisma_client)
|
||||
|
|
@ -688,7 +671,7 @@ async def test_model_repository_write_publishes_via_live_coordination_cache() ->
|
|||
assert isinstance(table, _PublishOnWriteActions)
|
||||
|
||||
previous_cache = proxy_server.redis_usage_cache
|
||||
_set_redis_usage_cache(_FakeRedisCache(client))
|
||||
_set_redis_usage_cache(client)
|
||||
try:
|
||||
await table.update(where={"model_id": "m-1"}, data={"model_name": "gpt-5.2"})
|
||||
finally:
|
||||
|
|
@ -708,14 +691,14 @@ async def test_ui_settings_write_publishes_via_live_coordination_cache() -> None
|
|||
from litellm.proxy.proxy_server import _set_redis_usage_cache
|
||||
from litellm.repositories.table_repositories import UISettingsRepository
|
||||
|
||||
client = _RecordingRedisClient()
|
||||
client = _FakeRedisCache()
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_uisettings.upsert = AsyncMock(return_value={"id": "ui_settings"})
|
||||
table = UISettingsRepository(prisma_client).table
|
||||
assert isinstance(table, _PublishOnWriteActions)
|
||||
|
||||
previous_cache = proxy_server.redis_usage_cache
|
||||
_set_redis_usage_cache(_FakeRedisCache(client))
|
||||
_set_redis_usage_cache(client)
|
||||
try:
|
||||
await table.upsert(
|
||||
where={"id": "ui_settings"},
|
||||
|
|
@ -730,14 +713,14 @@ async def test_ui_settings_write_publishes_via_live_coordination_cache() -> None
|
|||
assert json.loads(message) == {"object_type": "litellm_uisettings"}
|
||||
|
||||
|
||||
async def _publish_calls_for_invalidated_param(param_name: str) -> List[Tuple[str, str]]:
|
||||
async def _publish_calls_for_invalidated_param(param_name: str) -> list[tuple[str, str]]:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.proxy_server import _set_redis_usage_cache
|
||||
from litellm.proxy.utils import invalidate_config_param
|
||||
|
||||
client = _RecordingRedisClient()
|
||||
client = _FakeRedisCache()
|
||||
previous_cache = proxy_server.redis_usage_cache
|
||||
_set_redis_usage_cache(_FakeRedisCache(client))
|
||||
_set_redis_usage_cache(client)
|
||||
try:
|
||||
await invalidate_config_param(param_name)
|
||||
finally:
|
||||
|
|
@ -771,9 +754,9 @@ async def test_evict_config_param_does_not_publish() -> None:
|
|||
from litellm.proxy.proxy_server import _set_redis_usage_cache
|
||||
from litellm.proxy.utils import evict_config_param
|
||||
|
||||
client = _RecordingRedisClient()
|
||||
client = _FakeRedisCache()
|
||||
previous_cache = proxy_server.redis_usage_cache
|
||||
_set_redis_usage_cache(_FakeRedisCache(client))
|
||||
_set_redis_usage_cache(client)
|
||||
try:
|
||||
await evict_config_param("model_cost_map_reload_config")
|
||||
finally:
|
||||
|
|
@ -803,10 +786,10 @@ async def test_model_cost_map_reload_does_not_publish_config_change() -> None:
|
|||
|
||||
litellm_config_cache.flush_cache()
|
||||
prisma_client = _reload_config_prisma_client()
|
||||
client = _RecordingRedisClient()
|
||||
client = _FakeRedisCache()
|
||||
previous_cache = proxy_server.redis_usage_cache
|
||||
original_model_cost = litellm.model_cost.copy()
|
||||
_set_redis_usage_cache(_FakeRedisCache(client))
|
||||
_set_redis_usage_cache(client)
|
||||
try:
|
||||
from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
||||
|
||||
|
|
@ -833,9 +816,9 @@ async def test_anthropic_beta_headers_reload_does_not_publish_config_change() ->
|
|||
|
||||
litellm_config_cache.flush_cache()
|
||||
prisma_client = _reload_config_prisma_client()
|
||||
client = _RecordingRedisClient()
|
||||
client = _FakeRedisCache()
|
||||
previous_cache = proxy_server.redis_usage_cache
|
||||
_set_redis_usage_cache(_FakeRedisCache(client))
|
||||
_set_redis_usage_cache(client)
|
||||
try:
|
||||
with patch("litellm.anthropic_beta_headers_manager.reload_beta_headers_config") as mock_reload:
|
||||
mock_reload.return_value = {}
|
||||
|
|
@ -855,11 +838,11 @@ class _StopFailingSubscriber(ConfigSyncSubscriber):
|
|||
async def test_proxy_config_subscriber_resyncs_deployments_only() -> None:
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([_QueuePubSub()]))
|
||||
cache = _FakeRedisCache(subscriptions=[_QueuePubSub()])
|
||||
config = ProxyConfig()
|
||||
prisma_client = MagicMock()
|
||||
proxy_logging_obj = MagicMock()
|
||||
calls: List[Tuple[str, object, object]] = []
|
||||
calls: list[tuple[str, object, object]] = []
|
||||
|
||||
async def fake_add_deployment(prisma_client: object, proxy_logging_obj: object) -> None:
|
||||
calls.append(("add_deployment", prisma_client, proxy_logging_obj))
|
||||
|
|
@ -902,7 +885,7 @@ async def test_proxy_config_does_not_start_subscriber_without_coordination_redis
|
|||
async def test_proxy_config_keeps_the_first_subscriber_on_repeat_start() -> None:
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
cache = _FakeRedisCache(_ScriptedPubSubRedisClient([_QueuePubSub()]))
|
||||
cache = _FakeRedisCache(subscriptions=[_QueuePubSub()])
|
||||
config = ProxyConfig()
|
||||
|
||||
config.start_config_sync_subscriber(prisma_client=MagicMock(), proxy_logging_obj=MagicMock(), redis_cache=cache)
|
||||
|
|
@ -920,7 +903,7 @@ async def test_proxy_config_shutdown_survives_a_failing_subscriber_stop() -> Non
|
|||
|
||||
config = ProxyConfig()
|
||||
config.config_sync_subscriber = _StopFailingSubscriber(
|
||||
redis_cache=_FakeRedisCache(_ScriptedPubSubRedisClient([])),
|
||||
redis_cache=_FakeRedisCache(subscriptions=[]),
|
||||
resync_callbacks=(),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -146,3 +146,28 @@ def test_debug_report_returns_what_the_bug_report_link_carries_and_nothing_from_
|
|||
"model_list[*].provider = [azure]",
|
||||
]
|
||||
assert not any(hostile in response.text for hostile in HOSTILE_STRINGS), response.text
|
||||
|
||||
|
||||
def test_cache_stats_report_the_redis_pool_through_the_cache_object() -> None:
|
||||
from types import SimpleNamespace
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.common_utils.debug_utils import _get_cache_memory_stats
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
class _PooledRedisCache:
|
||||
def connection_pool_status(self) -> Mapping[str, object]:
|
||||
return {"max_connections": 7, "connection_class": "Connection"}
|
||||
|
||||
stats = _get_cache_memory_stats(
|
||||
user_api_key_cache=UserApiKeyCache(),
|
||||
llm_router=None,
|
||||
proxy_logging_obj=SimpleNamespace(internal_usage_cache=SimpleNamespace(dual_cache=DualCache())),
|
||||
redis_usage_cache=_PooledRedisCache(),
|
||||
)
|
||||
|
||||
assert stats["redis_usage_cache"] == {
|
||||
"enabled": True,
|
||||
"cache_type": "_PooledRedisCache",
|
||||
"connection_pool": {"max_connections": 7, "connection_class": "Connection"},
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue