mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(caching): propagate auth cache invalidation over Redis Cluster via a node-level pub/sub client (#43110)
* fix(caching): give Redis Cluster clients a node-level pub/sub client so auth invalidation propagates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): shorten pub/sub client docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(caching): noqa BLE001 on best-effort pubsub client close Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): satisfy LIT002/LIT006 in pubsub client derivation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(auth): rename fake pubsub hook to init_pubsub_client Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): rename fake pubsub hook to init_pubsub_client Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(caching): drop class-level health ping patch from pub/sub client tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): close cluster pubsub pools and cover failure paths * fix(caching): clean up expired cluster pubsub clients safely * test(mcp): arm cancellation deadline after requests start --------- Co-authored-by: joshua <joshua@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
61953318bf
commit
d33d36ce86
13 changed files with 319 additions and 59 deletions
6
.github/workflows/test-redis-compat.yml
vendored
6
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -8,9 +8,13 @@ on:
|
|||
paths:
|
||||
- "litellm/_redis.py"
|
||||
- "litellm/_redis_credential_provider.py"
|
||||
- "litellm/caching/redis_cache.py"
|
||||
- "litellm/caching/evicted_client_closer.py"
|
||||
- "tests/test_litellm/test_redis.py"
|
||||
- "tests/local_testing/test_caching.py"
|
||||
- "tests/test_litellm/caching/test_redis_connection_pool.py"
|
||||
- "tests/test_litellm/caching/test_redis_cluster_cache.py"
|
||||
- "tests/test_litellm/caching/test_evicted_client_closer.py"
|
||||
- ".github/workflows/test-redis-compat.yml"
|
||||
- "pyproject.toml"
|
||||
- "uv.lock"
|
||||
|
|
@ -82,6 +86,8 @@ jobs:
|
|||
uv run --no-sync pytest \
|
||||
tests/test_litellm/test_redis.py \
|
||||
tests/test_litellm/caching/test_redis_connection_pool.py \
|
||||
tests/test_litellm/caching/test_redis_cluster_cache.py \
|
||||
tests/test_litellm/caching/test_evicted_client_closer.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 \
|
||||
|
|
|
|||
|
|
@ -136,6 +136,8 @@ def _has_connection_in_flight(client: object) -> bool:
|
|||
window as the only guard, exactly as it was before this check existed.
|
||||
"""
|
||||
try:
|
||||
if getattr(getattr(client, "connection_pool", None), "_in_use_connections", None):
|
||||
return True
|
||||
transport: Final = _transport_of(client)
|
||||
pooled_busy: Final = _pool_has_busy_connection(transport)
|
||||
if pooled_busy is not None:
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from collections.abc import Awaitable, Callable, Iterator, Sequence
|
|||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
|
@ -290,6 +291,29 @@ def _opaque_kwarg_key(value: object) -> str:
|
|||
return f"{type(value).__name__}-{id(value)}"
|
||||
|
||||
|
||||
_CLUSTER_ONLY_CONNECTION_KWARGS: Final[frozenset[str]] = frozenset({"response_callbacks"})
|
||||
|
||||
|
||||
def _cluster_node_pubsub_client( # pyright: ignore[reportUnknownParameterType] # redis generics
|
||||
cluster: async_redis_cluster_client, # pyright: ignore[reportUnknownParameterType] # redis generics
|
||||
) -> async_redis_client:
|
||||
"""Plain async client on one cluster node; classic PUBLISH/SUBSCRIBE is broadcast cluster-wide."""
|
||||
from redis.asyncio import ConnectionPool, Redis
|
||||
|
||||
node: Final = cluster.get_default_node() or next(iter(cluster.nodes_manager.startup_nodes.values()), None)
|
||||
if node is None: # pyright: ignore[reportUnnecessaryComparison] # get_default_node is None before cluster init
|
||||
raise ValueError("cannot derive a pub/sub client: redis cluster has no default node and no startup nodes")
|
||||
node_kwargs: Final = MappingProxyType(
|
||||
{
|
||||
key: value # pyright: ignore[reportAny] # connection_kwargs values are Any in redis stubs
|
||||
for key, value in cluster.connection_kwargs.items() # pyright: ignore[reportAny] # connection_kwargs values are Any in redis stubs
|
||||
if key not in _CLUSTER_ONLY_CONNECTION_KWARGS
|
||||
}
|
||||
)
|
||||
pool: Final = ConnectionPool(host=node.host, port=node.port, **node_kwargs) # pyright: ignore[reportCallIssue, reportArgumentType] # cluster kwargs validated by redis-py at runtime
|
||||
return Redis.from_pool(pool) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # redis generics
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _redis_health_error_types() -> tuple[type, ...]:
|
||||
"""Exception types that mean the Redis backend itself is unhealthy.
|
||||
|
|
@ -738,6 +762,28 @@ class RedisCache(BaseCache):
|
|||
self.redis_async_client = redis_async_client
|
||||
return redis_async_client
|
||||
|
||||
def init_pubsub_client(self) -> async_redis_client: # pyright: ignore[reportUnknownParameterType] # redis generics
|
||||
from redis.asyncio import RedisCluster
|
||||
|
||||
from litellm import in_memory_llm_clients_cache
|
||||
|
||||
client: Final = self.init_async_client() # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # redis generics
|
||||
if not isinstance(client, RedisCluster):
|
||||
return client # pyright: ignore[reportUnknownVariableType] # redis generics
|
||||
cache_key: Final = f"{self._get_async_client_cache_key()}-pubsub"
|
||||
cached_client: Final = in_memory_llm_clients_cache.get_cache( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped in-memory client cache
|
||||
key=cache_key
|
||||
)
|
||||
if cached_client is not None:
|
||||
return cast( # cast-ok: per-loop pub/sub client stored by this method # pyright: ignore[reportUnknownVariableType] # redis generics
|
||||
async_redis_client, cached_client
|
||||
)
|
||||
pubsub_client: Final = _cluster_node_pubsub_client( # pyright: ignore[reportUnknownVariableType] # redis generics
|
||||
cluster=client
|
||||
)
|
||||
in_memory_llm_clients_cache.set_cache(key=cache_key, value=pubsub_client, litellm_owned_client=True) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped in-memory client cache
|
||||
return pubsub_client # pyright: ignore[reportUnknownVariableType] # redis generics
|
||||
|
||||
def _async_commands(self) -> _AsyncRedisCommands:
|
||||
return self.init_async_client()
|
||||
|
||||
|
|
@ -1785,7 +1831,21 @@ class RedisCache(BaseCache):
|
|||
self.redis_client.flushall()
|
||||
|
||||
async def disconnect(self):
|
||||
await self.async_redis_conn_pool.disconnect(inuse_connections=True)
|
||||
from litellm import in_memory_llm_clients_cache
|
||||
|
||||
if self.async_redis_conn_pool is not None:
|
||||
await self.async_redis_conn_pool.disconnect(inuse_connections=True)
|
||||
cached_pubsub_client: Final = cast( # cast-ok: only this module stores clients under this key # pyright: ignore[reportUnknownVariableType] # redis generics
|
||||
async_redis_client | None,
|
||||
in_memory_llm_clients_cache.get_cache( # pyright: ignore[reportUnknownMemberType] # untyped in-memory client cache
|
||||
key=f"{self._get_async_client_cache_key()}-pubsub"
|
||||
),
|
||||
)
|
||||
if cached_pubsub_client is not None:
|
||||
try:
|
||||
await cached_pubsub_client.aclose() # pyright: ignore[reportUnknownMemberType, reportAttributeAccessIssue] # redis stubs leave aclose unknown
|
||||
except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection
|
||||
verbose_logger.debug("Error closing cached pub/sub Redis client: %s", e)
|
||||
try:
|
||||
self.redis_client.close()
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -74,12 +74,6 @@ 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)
|
||||
except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors
|
||||
|
|
@ -185,12 +179,6 @@ class AuthCacheInvalidationSubscriber:
|
|||
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"
|
||||
)
|
||||
return
|
||||
pubsub = client.pubsub()
|
||||
try:
|
||||
await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache))
|
||||
|
|
|
|||
|
|
@ -82,22 +82,13 @@ 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:
|
||||
return cast( # cast-ok: protocol view of the pub/sub-capable async redis client
|
||||
_ConfigSyncPubSubClient,
|
||||
redis_cache.init_pubsub_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
|
||||
|
|
@ -112,12 +103,6 @@ async def publish_config_change(redis_cache: "RedisCache | None", object_type: s
|
|||
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))
|
||||
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)
|
||||
|
|
@ -238,12 +223,6 @@ class ConfigSyncSubscriber:
|
|||
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"
|
||||
)
|
||||
return
|
||||
pubsub = client.pubsub()
|
||||
try:
|
||||
await pubsub.subscribe(config_sync_channel(self._redis_cache))
|
||||
|
|
|
|||
|
|
@ -10,9 +10,11 @@ never closed, because litellm does not own its lifecycle.
|
|||
import asyncio
|
||||
import gc
|
||||
import weakref
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from redis.asyncio import ConnectionPool, Redis
|
||||
|
||||
from litellm.caching.evicted_client_closer import EvictedClientCloser
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
|
@ -72,6 +74,33 @@ def make_closer(clock: FakeClock, grace_seconds: float = 60.0) -> EvictedClientC
|
|||
return EvictedClientCloser(grace_seconds=grace_seconds, clock=clock)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_client_is_closed_only_after_its_subscription_releases_the_connection():
|
||||
closer = EvictedClientCloser(grace_seconds=0)
|
||||
pool = ConnectionPool()
|
||||
client = Redis.from_pool(pool)
|
||||
closed = asyncio.Event()
|
||||
connection = AsyncMock()
|
||||
connection.disconnect.side_effect = closed.set
|
||||
pool._available_connections.append(connection)
|
||||
borrowed = pool.get_available_connection()
|
||||
closer.mark_owned(client)
|
||||
closer.schedule(client)
|
||||
|
||||
closer.reap()
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
connection.disconnect.assert_not_awaited()
|
||||
assert closer.pending_count == 1
|
||||
|
||||
await pool.release(borrowed)
|
||||
closer.reap()
|
||||
await asyncio.wait_for(closed.wait(), timeout=1)
|
||||
|
||||
connection.disconnect.assert_awaited_once()
|
||||
assert closer.pending_count == 0
|
||||
|
||||
|
||||
async def _trickling_upstream(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
|
||||
"""Serves a chunked body slowly, so a request stays on the wire long enough to observe."""
|
||||
await reader.read(4096)
|
||||
|
|
|
|||
|
|
@ -1,13 +1,20 @@
|
|||
import asyncio
|
||||
from importlib import import_module
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
import ssl
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from redis.asyncio import Redis, RedisCluster
|
||||
from redis.asyncio.cluster import ClusterNode
|
||||
from redis.asyncio.connection import SSLConnection
|
||||
|
||||
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||||
from litellm.caching.evicted_client_closer import EvictedClientCloser
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
||||
|
||||
@patch("litellm._redis.init_redis_cluster")
|
||||
|
|
@ -175,3 +182,167 @@ def test_router_create_redis_cache_cluster_detection(
|
|||
with patch.object(RedisCache, "__init__", _mock_redis_cache_init):
|
||||
redis_cache = Router._create_redis_cache(cache_config)
|
||||
assert isinstance(redis_cache, expected_cache_type)
|
||||
|
||||
|
||||
def _isolated_redis_cache(host: str) -> RedisCache:
|
||||
"""RedisCache whose sync client and pool are stubbed out."""
|
||||
with (
|
||||
patch("litellm._redis.get_redis_client", return_value=MagicMock()),
|
||||
patch("litellm._redis.get_redis_connection_pool", return_value=MagicMock()),
|
||||
):
|
||||
return RedisCache(host=host, port=6379)
|
||||
|
||||
|
||||
def _cluster_for_pubsub(startup_node_host: str = "10.9.9.9") -> RedisCluster:
|
||||
"""Uninitialized RedisCluster carrying the connection kwargs a real one would."""
|
||||
return RedisCluster(
|
||||
startup_nodes=[ClusterNode(host=startup_node_host, port=7000)],
|
||||
password="cluster-secret",
|
||||
socket_timeout=7.0,
|
||||
)
|
||||
|
||||
|
||||
def test_init_pubsub_client_derives_a_node_client_for_cluster_backend() -> None:
|
||||
"""LIT-8543: a cluster-backed cache must return a pub/sub-capable client.
|
||||
|
||||
The derived client pins a plain Redis connection pool to the cluster's
|
||||
default node, inheriting the connection kwargs minus cluster-only keys.
|
||||
"""
|
||||
cache = _isolated_redis_cache("cluster-pubsub-default-node")
|
||||
cluster = _cluster_for_pubsub()
|
||||
node = ClusterNode(host="10.1.2.3", port=7001)
|
||||
cluster.nodes_manager.default_node = node
|
||||
cache.init_async_client = MagicMock(return_value=cluster)
|
||||
|
||||
client = cache.init_pubsub_client()
|
||||
|
||||
assert isinstance(client, Redis) and not isinstance(client, RedisCluster)
|
||||
kwargs = client.connection_pool.connection_kwargs
|
||||
assert kwargs["host"] == "10.1.2.3"
|
||||
assert kwargs["port"] == 7001
|
||||
assert kwargs["password"] == "cluster-secret"
|
||||
assert kwargs["socket_timeout"] == 7.0
|
||||
assert "response_callbacks" not in kwargs
|
||||
|
||||
|
||||
def test_init_pubsub_client_falls_back_to_first_startup_node() -> None:
|
||||
"""Before cluster initialization there is no default node; the first
|
||||
startup node is a valid pub/sub target."""
|
||||
cache = _isolated_redis_cache("cluster-pubsub-startup-fallback")
|
||||
cluster = _cluster_for_pubsub(startup_node_host="10.8.8.8")
|
||||
cache.init_async_client = MagicMock(return_value=cluster)
|
||||
|
||||
client = cache.init_pubsub_client()
|
||||
|
||||
assert isinstance(client, Redis)
|
||||
assert client.connection_pool.connection_kwargs["host"] == "10.8.8.8"
|
||||
|
||||
|
||||
def test_init_pubsub_client_returns_the_same_cached_client_on_repeat_calls() -> None:
|
||||
cache = _isolated_redis_cache("cluster-pubsub-caching")
|
||||
cluster = _cluster_for_pubsub()
|
||||
cluster.nodes_manager.default_node = ClusterNode(host="10.1.2.3", port=7001)
|
||||
cache.init_async_client = MagicMock(return_value=cluster)
|
||||
|
||||
first = cache.init_pubsub_client()
|
||||
second = cache.init_pubsub_client()
|
||||
|
||||
assert first is second
|
||||
|
||||
|
||||
def test_init_pubsub_client_returns_the_shared_async_client_for_standalone() -> None:
|
||||
cache = _isolated_redis_cache("standalone-pubsub")
|
||||
standalone = Redis()
|
||||
cache.init_async_client = MagicMock(return_value=standalone)
|
||||
|
||||
assert cache.init_pubsub_client() is standalone
|
||||
|
||||
|
||||
def test_init_pubsub_client_preserves_tls_and_authentication() -> None:
|
||||
cache = _isolated_redis_cache("cluster-pubsub-tls")
|
||||
cluster = RedisCluster(
|
||||
startup_nodes=[ClusterNode(host="redis.example.test", port=7000)],
|
||||
ssl=True,
|
||||
ssl_cert_reqs="required",
|
||||
ssl_check_hostname=True,
|
||||
username="pubsub-user",
|
||||
password="test-password",
|
||||
socket_connect_timeout=3.0,
|
||||
socket_keepalive=True,
|
||||
)
|
||||
cache.init_async_client = MagicMock(return_value=cluster)
|
||||
|
||||
client = cache.init_pubsub_client()
|
||||
connection = client.connection_pool.make_connection()
|
||||
|
||||
assert isinstance(connection, SSLConnection)
|
||||
assert connection.ssl_context.cert_reqs == ssl.CERT_REQUIRED
|
||||
assert connection.ssl_context.check_hostname is True
|
||||
assert connection.username == "pubsub-user"
|
||||
assert connection.password == "test-password"
|
||||
assert connection.socket_connect_timeout == 3.0
|
||||
assert connection.socket_keepalive is True
|
||||
|
||||
|
||||
def test_init_pubsub_client_rejects_missing_nodes_and_can_retry() -> None:
|
||||
cache = _isolated_redis_cache("cluster-pubsub-no-nodes")
|
||||
cluster = _cluster_for_pubsub()
|
||||
cluster.nodes_manager.startup_nodes = {}
|
||||
cache.init_async_client = MagicMock(return_value=cluster)
|
||||
|
||||
with pytest.raises(ValueError, match="no default node and no startup nodes"):
|
||||
cache.init_pubsub_client()
|
||||
|
||||
cluster.nodes_manager.default_node = ClusterNode(host="recovered.example.test", port=7001)
|
||||
client = cache.init_pubsub_client()
|
||||
|
||||
assert client.connection_pool.connection_kwargs["host"] == "recovered.example.test"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("close_fails", [False, True])
|
||||
@pytest.mark.parametrize("has_shared_pool", [False, True])
|
||||
def test_disconnect_closes_derived_pubsub_connections_even_when_pool_close_fails(
|
||||
close_fails: bool, has_shared_pool: bool
|
||||
) -> None:
|
||||
cache = _isolated_redis_cache(f"cluster-pubsub-close-{close_fails}-{has_shared_pool}")
|
||||
cache.async_redis_conn_pool = AsyncMock() if has_shared_pool else None
|
||||
cache.init_async_client = MagicMock(return_value=_cluster_for_pubsub())
|
||||
|
||||
async def exercise() -> None:
|
||||
client = cache.init_pubsub_client()
|
||||
connection = AsyncMock()
|
||||
connection.disconnect.side_effect = ConnectionError("connection close failed") if close_fails else None
|
||||
client.connection_pool._available_connections.append(connection)
|
||||
|
||||
await cache.disconnect()
|
||||
|
||||
connection.disconnect.assert_awaited_once()
|
||||
if has_shared_pool:
|
||||
cache.async_redis_conn_pool.disconnect.assert_awaited_once_with(inuse_connections=True)
|
||||
cache.redis_client.close.assert_called_once()
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
||||
|
||||
def test_expired_pubsub_client_closes_connections_after_eviction() -> None:
|
||||
cache = _isolated_redis_cache("cluster-pubsub-expired")
|
||||
cache.init_async_client = MagicMock(return_value=_cluster_for_pubsub())
|
||||
clients = LLMClientCache(evicted_client_closer=EvictedClientCloser(grace_seconds=0))
|
||||
|
||||
async def exercise() -> None:
|
||||
client = cache.init_pubsub_client()
|
||||
closed = asyncio.Event()
|
||||
connection = AsyncMock()
|
||||
connection.disconnect.side_effect = closed.set
|
||||
client.connection_pool._available_connections.append(connection)
|
||||
cache_key = clients.update_cache_key_with_event_loop(f"{cache._get_async_client_cache_key()}-pubsub")
|
||||
clients.ttl_dict[cache_key] = 0
|
||||
|
||||
replacement = cache.init_pubsub_client()
|
||||
await asyncio.wait_for(closed.wait(), timeout=1)
|
||||
|
||||
assert replacement is not client
|
||||
connection.disconnect.assert_awaited_once()
|
||||
|
||||
with patch("litellm.in_memory_llm_clients_cache", clients):
|
||||
asyncio.run(exercise())
|
||||
|
|
|
|||
|
|
@ -2762,6 +2762,7 @@ async def test_cancellation_delivers_termination_over_tcp(
|
|||
cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool
|
||||
) -> None:
|
||||
started: Final = asyncio.Event()
|
||||
scope_ready: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future()
|
||||
terminations: Final[list[bytes]] = []
|
||||
starts: Final[list[bytes]] = []
|
||||
stop: Final = asyncio.Event()
|
||||
|
|
@ -2858,13 +2859,16 @@ async def test_cancellation_delivers_termination_over_tcp(
|
|||
|
||||
async def invoke():
|
||||
if cancel_mode == "scope":
|
||||
with anyio.fail_after(0.2):
|
||||
with anyio.fail_after(None) as scope:
|
||||
scope_ready.set_result(scope)
|
||||
return await calls()
|
||||
return await calls()
|
||||
|
||||
try:
|
||||
task: Final = asyncio.create_task(invoke())
|
||||
await asyncio.wait_for(started.wait(), 3)
|
||||
if cancel_mode == "scope":
|
||||
(await scope_ready).deadline = anyio.current_time() + 0.2
|
||||
if cancel_mode == "task":
|
||||
task.cancel()
|
||||
expected_error: Final = (
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|||
class _FakeRedisCache:
|
||||
namespace = None
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
def init_pubsub_client(self) -> object:
|
||||
return object()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8998,7 +8998,7 @@ async def test_invalidate_team_member_spend_state_broadcasts_the_spend_counter_t
|
|||
def __init__(self) -> None:
|
||||
self.namespace = None
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
def init_pubsub_client(self) -> object:
|
||||
return _RecordingRedisClient()
|
||||
|
||||
local_spend_counter_cache = DualCache()
|
||||
|
|
@ -9072,7 +9072,7 @@ async def test_invalidate_team_member_spend_state_self_delivered_broadcast_does_
|
|||
def __init__(self) -> None:
|
||||
self.namespace = None
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
def init_pubsub_client(self) -> object:
|
||||
return _RecordingRedisClient()
|
||||
|
||||
local_spend_counter_cache = DualCache()
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ class _FakeRedisCache:
|
|||
self._client = client
|
||||
self.namespace = namespace
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
def init_pubsub_client(self) -> object:
|
||||
return self._client
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -81,14 +81,25 @@ class _FailingPublishRedisClient(Redis):
|
|||
raise ConnectionError("redis down")
|
||||
|
||||
|
||||
class _NotRedisClient:
|
||||
def __init__(self) -> None:
|
||||
class _ScriptedPubSubClient:
|
||||
"""Pub/sub-capable client that is not a redis.asyncio.Redis.
|
||||
|
||||
Mirrors what RedisCache.init_pubsub_client returns for a cluster backend:
|
||||
a node-level client exposing publish/pubsub without being an instance of
|
||||
the standalone Redis class.
|
||||
"""
|
||||
|
||||
def __init__(self, pubsubs: Iterable["_QueuePubSub"]) -> None:
|
||||
self._scripted_pubsubs = iter(pubsubs)
|
||||
self.published: List[Tuple[str, str]] = []
|
||||
|
||||
async def publish(self, channel: str, message: str) -> int:
|
||||
self.published.append((channel, message))
|
||||
return 1
|
||||
|
||||
def pubsub(self) -> "_QueuePubSub":
|
||||
return next(self._scripted_pubsubs)
|
||||
|
||||
|
||||
class _QueuePubSub:
|
||||
def __init__(self, initial_messages: Iterable[str] = ()) -> None:
|
||||
|
|
@ -159,14 +170,14 @@ class _FakeRedisCache:
|
|||
self._client = client
|
||||
self.namespace = namespace
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
def init_pubsub_client(self) -> object:
|
||||
return self._client
|
||||
|
||||
|
||||
class _ExplodingRedisCache:
|
||||
namespace: Optional[str] = None
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
def init_pubsub_client(self) -> object:
|
||||
raise ConnectionError("cannot connect")
|
||||
|
||||
|
||||
|
|
@ -215,13 +226,17 @@ async def test_publish_swallows_client_init_errors() -> None:
|
|||
await publish_config_change(redis_cache=_ExplodingRedisCache(), object_type="litellm_proxymodeltable")
|
||||
|
||||
|
||||
async def test_publish_skips_clients_without_pubsub_support() -> None:
|
||||
client = _NotRedisClient()
|
||||
async def test_publish_reaches_cluster_derived_pubsub_clients() -> None:
|
||||
"""LIT-8543: a cluster-backed cache returns a node-level client from
|
||||
init_pubsub_client; publishes must go out on it instead of being skipped."""
|
||||
client = _ScriptedPubSubClient(pubsubs=[])
|
||||
cache = _FakeRedisCache(client)
|
||||
|
||||
await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable")
|
||||
|
||||
assert client.published == []
|
||||
assert client.published == [
|
||||
(CONFIG_SYNC_CHANNEL, json.dumps({"object_type": "litellm_proxymodeltable"}))
|
||||
]
|
||||
|
||||
|
||||
async def test_subscriber_runs_injected_callbacks_in_order_on_message() -> None:
|
||||
|
|
@ -558,20 +573,26 @@ async def test_stop_before_start_is_a_noop() -> None:
|
|||
await subscriber.stop()
|
||||
|
||||
|
||||
async def test_subscriber_exits_without_callbacks_when_client_lacks_pubsub() -> None:
|
||||
cache = _FakeRedisCache(_NotRedisClient())
|
||||
async def test_subscriber_subscribes_on_cluster_derived_pubsub_client() -> None:
|
||||
"""LIT-8543: the subscriber used to disable itself on cluster caches; now it
|
||||
subscribes on the node-level client init_pubsub_client returns."""
|
||||
pubsub = _QueuePubSub(initial_messages=[json.dumps({"object_type": "litellm_proxymodeltable"})])
|
||||
cache = _FakeRedisCache(_ScriptedPubSubClient(pubsubs=[pubsub]))
|
||||
resyncs: List[str] = []
|
||||
fired = asyncio.Event()
|
||||
subscriber = ConfigSyncSubscriber(
|
||||
redis_cache=cache,
|
||||
resync_callbacks=(_recording_callback(resyncs, "resync", asyncio.Event()),),
|
||||
resync_callbacks=(_recording_callback(resyncs, "resync", fired),),
|
||||
debounce_seconds=0.01,
|
||||
jitter_max_seconds=0.0,
|
||||
)
|
||||
|
||||
subscriber.start()
|
||||
task = subscriber._task
|
||||
assert task is not None
|
||||
await asyncio.wait_for(task, timeout=5)
|
||||
await asyncio.wait_for(fired.wait(), timeout=5)
|
||||
await subscriber.stop()
|
||||
|
||||
assert resyncs == []
|
||||
assert pubsub.subscribed_channels == [CONFIG_SYNC_CHANNEL]
|
||||
assert resyncs == ["resync"]
|
||||
|
||||
|
||||
class _FakeTableActions:
|
||||
|
|
|
|||
|
|
@ -15127,7 +15127,7 @@ async def test_auth_cache_invalidation_subscriber_evicts_byok_credentials_cached
|
|||
def __init__(self, client: object) -> None:
|
||||
self._client = client
|
||||
|
||||
def init_async_client(self) -> object:
|
||||
def init_pubsub_client(self) -> object:
|
||||
return self._client
|
||||
|
||||
byok_credential_cache.flush_cache()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue