From d33d36ce862459e144e814dcdab049a8bde3b06f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 07:56:46 -0700 Subject: [PATCH] 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 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> --- .github/workflows/test-redis-compat.yml | 6 + litellm/caching/evicted_client_closer.py | 2 + litellm/caching/redis_cache.py | 62 ++++++- .../auth_cache_invalidation_pubsub.py | 12 -- .../proxy/common_utils/config_sync_pubsub.py | 29 +-- .../caching/test_evicted_client_closer.py | 29 +++ .../caching/test_redis_cluster_cache.py | 173 +++++++++++++++++- .../test_mcp_client.py | 6 +- .../mcp_server/test_byok_credential_cache.py | 2 +- .../proxy/auth/test_auth_checks.py | 4 +- .../test_auth_cache_invalidation_pubsub.py | 2 +- .../common_utils/test_config_sync_pubsub.py | 49 +++-- tests/test_litellm/proxy/test_proxy_server.py | 2 +- 13 files changed, 319 insertions(+), 59 deletions(-) diff --git a/.github/workflows/test-redis-compat.yml b/.github/workflows/test-redis-compat.yml index 0b58cf9d486..25fb8f8bce3 100644 --- a/.github/workflows/test-redis-compat.yml +++ b/.github/workflows/test-redis-compat.yml @@ -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 \ diff --git a/litellm/caching/evicted_client_closer.py b/litellm/caching/evicted_client_closer.py index eee7e2ea289..6e4635dd83a 100644 --- a/litellm/caching/evicted_client_closer.py +++ b/litellm/caching/evicted_client_closer.py @@ -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: diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 7cec84e0ebb..0b56c28f9b1 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -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: diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index 11cb66d1a7f..3e09ad7157f 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -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)) diff --git a/litellm/proxy/common_utils/config_sync_pubsub.py b/litellm/proxy/common_utils/config_sync_pubsub.py index b20c0d9c9a5..b4ebb5fa876 100644 --- a/litellm/proxy/common_utils/config_sync_pubsub.py +++ b/litellm/proxy/common_utils/config_sync_pubsub.py @@ -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)) diff --git a/tests/test_litellm/caching/test_evicted_client_closer.py b/tests/test_litellm/caching/test_evicted_client_closer.py index 939be5f3d6b..7679e276621 100644 --- a/tests/test_litellm/caching/test_evicted_client_closer.py +++ b/tests/test_litellm/caching/test_evicted_client_closer.py @@ -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) diff --git a/tests/test_litellm/caching/test_redis_cluster_cache.py b/tests/test_litellm/caching/test_redis_cluster_cache.py index 0763b5110d5..ba1deabd0e1 100644 --- a/tests/test_litellm/caching/test_redis_cluster_cache.py +++ b/tests/test_litellm/caching/test_redis_cluster_cache.py @@ -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()) diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index ae30b086c6e..368e34c455d 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -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 = ( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py index 8ec5b8642bc..0ec4b431276 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py @@ -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() diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index b811d4453ca..e42a47a1091 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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() diff --git a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py index 96770ee01c4..34d9741bcc0 100644 --- a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py @@ -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 diff --git a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py index 0f64ef2b4ca..83ed3afa293 100644 --- a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py @@ -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: diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 250556c9281..884a9c81500 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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()