From cdf35ba6f62b173437aa1826f397f58ac6f3dcf0 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Thu, 24 Sep 2026 10:58:10 -0700 Subject: [PATCH] fix(caching): address review on RedisCache pub/sub and script registration get_message(timeout=None) kept returning None after one empty poll slice, so an unbounded wait ended after a second. It now keeps polling until a message arrives async_subscribe closes the pubsub when SUBSCRIBE fails, so a checked-out pool connection is handed back instead of stranded The script_load/evalsha branch in script registration was unreachable: the async RedisCluster client has register_script in every supported redis-py release, and the branch would have passed an unawaited coroutine as the SHA. It is removed along with the mock test that pinned it, and the lock tests now also run against a real one-node Redis Cluster The Redis compat workflow now also triggers on changes to RedisCache and the lock Co-Authored-By: Claude Opus 5.5 --- .github/workflows/test-redis-compat.yml | 3 + litellm/caching/redis_cache.py | 44 ++++----- .../test_litellm/caching/test_redis_cache.py | 95 +++++++++++++++---- .../test_redis_distributed_lock.py | 79 +++++++++++---- 4 files changed, 155 insertions(+), 66 deletions(-) diff --git a/.github/workflows/test-redis-compat.yml b/.github/workflows/test-redis-compat.yml index 633a4c99ced..a01e55c4258 100644 --- a/.github/workflows/test-redis-compat.yml +++ b/.github/workflows/test-redis-compat.yml @@ -11,6 +11,9 @@ on: - "tests/test_litellm/test_redis.py" - "tests/local_testing/test_caching.py" - "tests/test_litellm/caching/test_redis_connection_pool.py" + - "litellm/caching/redis_cache.py" + - "litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py" + - "tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py" - ".github/workflows/test-redis-compat.yml" - "pyproject.toml" - "uv.lock" diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 40315dcac36..a3d7b9b18b7 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -631,6 +631,8 @@ class RedisSubscription: 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 and deadline is None: + continue if frame is None: return None message = _redis_message(frame) @@ -1065,33 +1067,19 @@ class RedisCache(BaseCache): executor; see that method for why the binding must be per loop. """ _redis_client: Final[Any] = self.init_async_client() - if hasattr(_redis_client, "register_script"): - registered_script: Final = _redis_client.register_script(script) + if not hasattr(_redis_client, "register_script"): + raise ValueError("Redis client does not support Lua script registration") + registered_script: Final = _redis_client.register_script(script) - async def standalone_executor( - keys: Sequence[str], - args: Sequence[str | bytes | int | float], - client: object = None, - ) -> object: - namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys) - return await registered_script(keys=namespaced_keys, args=args, client=client) + async def executor( + keys: Sequence[str], + args: Sequence[str | bytes | int | float], + client: object = None, + ) -> object: + namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys) + return await registered_script(keys=namespaced_keys, args=args, client=client) - return standalone_executor - - if hasattr(_redis_client, "script_load"): - script_sha: Final = _redis_client.script_load(script) - - async def cluster_executor( - keys: Sequence[str], - args: Sequence[str | bytes | int | float], - client: object = None, - ) -> object: - namespaced_keys: Final = tuple(self.check_and_fix_namespace(key=key) for key in keys) - return await _redis_client.evalsha(script_sha, len(namespaced_keys), *namespaced_keys, *args) - - return cluster_executor - - raise ValueError("Redis client does not support Lua script registration") + return executor @_redis_circuit_breaker_guard async def async_set_cache(self, key, value, **kwargs): @@ -1918,7 +1906,11 @@ class RedisCache(BaseCache): @_redis_circuit_breaker_guard async def async_subscribe(self, *channels: str) -> RedisSubscription: pubsub: Final = self._standalone_async_client().pubsub() - await pubsub.subscribe(*channels) + try: + await pubsub.subscribe(*channels) + except BaseException: + await pubsub.aclose() # pyright: ignore[reportAttributeAccessIssue] # types-redis 4.6 stubs predate PubSub.aclose + raise return RedisSubscription(pubsub) def connection_pool_status(self) -> "RedisPoolStatus": diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index d8b4f3f50be..0460a9cda49 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -1,6 +1,6 @@ import asyncio import time -from collections.abc import Iterator +from collections.abc import Iterable, Iterator from datetime import timedelta from unittest.mock import AsyncMock, MagicMock, patch @@ -283,26 +283,6 @@ async def test_async_register_script_not_shared_across_namespaces(monkeypatch, r reg_b.assert_awaited_once_with(keys=("ns_b:k",), args=[], client=None) -@pytest.mark.asyncio -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") - redis_cache = RedisCache(namespace="ns") - - cluster_client = MagicMock(spec=["script_load", "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): - 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) - - @pytest.mark.asyncio 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 @@ -1611,3 +1591,76 @@ def test_connection_pool_status_reports_the_sync_pool(socket_redis_cache: RedisC 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__ + + +class _FrameSequencePubSub: + """A redis-py pubsub whose polls yield canned frames; a None frame is an empty poll slice.""" + + def __init__(self, frames: Iterable[object]) -> None: + self.frames = iter(frames) + + async def get_message(self, *, timeout: float | None) -> object: + return next(self.frames) + + async def aclose(self) -> None: + pass + + +@pytest.mark.asyncio +async def test_get_message_without_timeout_waits_through_empty_polls() -> None: + from litellm.caching.redis_cache import RedisMessage, RedisSubscription + + late_message = {"type": "message", "channel": b"events", "data": b"late"} + subscription = RedisSubscription(_FrameSequencePubSub((None, None, late_message))) # pyright: ignore[reportArgumentType] # duck-typed fake + + assert await subscription.get_message(timeout=None) == RedisMessage(channel="events", payload=b"late") + + +class _RefusingPubSub: + """A redis-py pubsub whose SUBSCRIBE fails after it has checked a connection out of the pool.""" + + def __init__(self) -> None: + self.closed = False + + async def subscribe(self, *channels: str) -> None: + raise ConnectionError("socket dropped after checkout") + + async def aclose(self) -> None: + self.closed = True + + +class _RefusingClient: + def __init__(self, pubsub: _RefusingPubSub) -> None: + self._pubsub = pubsub + + def pubsub(self) -> _RefusingPubSub: + return self._pubsub + + async def ping(self) -> bool: + return True + + +@pytest.mark.asyncio +async def test_async_subscribe_closes_the_pubsub_when_subscribe_fails( + fake_redis_port: int, monkeypatch: pytest.MonkeyPatch +) -> None: + for name in ( + "REDIS_URL", + "REDIS_HOST", + "REDIS_PORT", + "REDIS_PASSWORD", + "REDIS_CLUSTER_NODES", + "REDIS_SENTINEL_NODES", + ): + monkeypatch.delenv(name, raising=False) + pubsub = _RefusingPubSub() + + class _RefusingClientCache(RedisCache): + def init_async_client(self): + return _RefusingClient(pubsub) + + cache = _RefusingClientCache(host="127.0.0.1", port=fake_redis_port) + + with pytest.raises(ConnectionError): + await cache.async_subscribe("events") + assert pubsub.closed, "a subscription that never completed must hand its connection back to the pool" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py index cb3f25714cc..e94858105e3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py @@ -1,6 +1,7 @@ """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 +The happy paths run against a real ``redis-server``, once standalone and once as a one-node cluster, +because every operation is a Lua script and the cluster client registers scripts on its own path; the degrade paths use a fake ``RedisCache`` whose scripts fail. """ @@ -9,6 +10,7 @@ import shutil import subprocess import time from collections.abc import Callable, Iterator, Sequence +from dataclasses import dataclass from pathlib import Path from typing import Final @@ -23,29 +25,62 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_c LockAcquisition, ) +_ALL_HASH_SLOTS: Final = (0, 16383) -@pytest.fixture -def redis_port(tmp_path: Path, unused_tcp_port_factory: Callable[[], int]) -> Iterator[int]: + +@dataclass(frozen=True, slots=True) +class _Server: + port: int + cluster: bool + + +def _wait_until(ready: Callable[[], bool], failure: str, log_path: Path) -> None: + for _ in range(100): + if ready(): + return + time.sleep(0.1) + pytest.fail(f"{failure}: {log_path.read_text()}") + + +def _pings(admin: redis.Redis) -> bool: + try: + return bool(admin.ping()) + except redis.ConnectionError: + return False + + +def _cluster_is_ok(admin: redis.Redis) -> bool: + return b"cluster_state:ok" in admin.execute_command("CLUSTER", "INFO") + + +@pytest.fixture(params=("standalone", "cluster")) +def redis_server( + request: pytest.FixtureRequest, tmp_path: Path, unused_tcp_port_factory: Callable[[], int] +) -> Iterator[_Server]: server: Final = shutil.which("redis-server") if server is None: pytest.skip("redis-server is required for the lock's Lua scripts") + cluster: Final = request.param == "cluster" 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') + config.write_text( + f'bind 127.0.0.1\nport {port}\ndir "{tmp_path}"\nsave ""\nappendonly no\n' + + ( + f"cluster-enabled yes\ncluster-config-file nodes.conf\ncluster-port {unused_tcp_port_factory()}\n" + if cluster + else "" + ) + ) 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 + _wait_until(lambda: _pings(admin), "Redis did not start", log_path) + if cluster: + admin.execute_command("CLUSTER", "ADDSLOTSRANGE", *_ALL_HASH_SLOTS) + _wait_until(lambda: _cluster_is_ok(admin), "Cluster never became ok", log_path) + yield _Server(port=port, cluster=cluster) finally: process.terminate() try: @@ -56,7 +91,7 @@ def redis_port(tmp_path: Path, unused_tcp_port_factory: Callable[[], int]) -> It @pytest.fixture -def redis_cache(redis_port: int, monkeypatch: pytest.MonkeyPatch) -> RedisCache: +def redis_cache(redis_server: _Server, monkeypatch: pytest.MonkeyPatch) -> RedisCache: for name in ( "REDIS_URL", "REDIS_HOST", @@ -66,12 +101,18 @@ def redis_cache(redis_port: int, monkeypatch: pytest.MonkeyPatch) -> RedisCache: "REDIS_SENTINEL_NODES", ): monkeypatch.delenv(name, raising=False) - return RedisCache(host="127.0.0.1", port=redis_port, namespace="tenant") + if redis_server.cluster: + return RedisCache(startup_nodes=[{"host": "127.0.0.1", "port": redis_server.port}], namespace="tenant") + return RedisCache(host="127.0.0.1", port=redis_server.port, namespace="tenant") @pytest.fixture -def raw(redis_port: int) -> Iterator[redis.Redis]: - with redis.Redis(host="127.0.0.1", port=redis_port) as client: +def raw(redis_server: _Server) -> Iterator[redis.Redis | redis.RedisCluster]: + if redis_server.cluster: + with redis.RedisCluster(host="127.0.0.1", port=redis_server.port) as cluster_client: + yield cluster_client + return + with redis.Redis(host="127.0.0.1", port=redis_server.port) as client: yield client @@ -85,7 +126,7 @@ class _FailingScriptsCache: return run -async def test_acquire_wins_once_and_reports_held_to_the_next_caller(redis_cache: RedisCache, raw: redis.Redis) -> None: +async def test_acquire_wins_once_and_reports_held_to_the_next_caller(redis_cache: RedisCache, raw: redis.Redis | redis.RedisCluster) -> None: lock = RedisDistributedLock(redis_cache) assert await lock.acquire("k", "tok-1", 10.0) is LockAcquisition.ACQUIRED @@ -111,7 +152,7 @@ async def test_release_deletes_only_when_the_token_matches(redis_cache: RedisCac assert await lock.is_held("k") is False -async def test_extend_refreshes_ttl_only_when_the_token_matches(redis_cache: RedisCache, raw: redis.Redis) -> None: +async def test_extend_refreshes_ttl_only_when_the_token_matches(redis_cache: RedisCache, raw: redis.Redis | redis.RedisCluster) -> None: lock = RedisDistributedLock(redis_cache) await lock.acquire("k", "owner-B", 1.0)