mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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 <noreply@anthropic.com>
This commit is contained in:
parent
12a28c4708
commit
cdf35ba6f6
4 changed files with 155 additions and 66 deletions
3
.github/workflows/test-redis-compat.yml
vendored
3
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue