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)