mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(cache): apply Redis namespace to all key operations (#31288)
The namespace configured under cache_params was only applied to get/set/
increment paths. Operations that take keys through other code paths (the Lua
scripts registered via async_register_script, delete, scan_iter, rpush, lpop,
get_ttl, and the sync increment_cache) hit raw keys. With a namespace set, the
rate limiter ({key}:tokens/requests/window), pod-lock release, and budget
limiters wrote keys outside the configured prefix, breaking multi-tenant key
isolation and leaving those operations reading keys the namespaced writes never
created.
check_and_fix_namespace is now applied uniformly across every key-taking
RedisCache operation. It is a no-op when no namespace is configured, so
deployments without a namespace are unaffected. The prefix is prepended ahead of
any {hash-tag}, so Redis Cluster slotting is preserved.
Resolves LIT-3374
This commit is contained in:
parent
62f93a3343
commit
01035499da
2 changed files with 203 additions and 4 deletions
|
|
@ -435,6 +435,7 @@ class RedisCache(BaseCache):
|
|||
_redis_client = self.redis_client
|
||||
start_time = time.time()
|
||||
set_ttl = self.get_ttl(ttl=ttl)
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
try:
|
||||
start_time = time.time()
|
||||
result: int = _redis_client.incr(name=key, amount=value) # type: ignore
|
||||
|
|
@ -498,6 +499,7 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
return []
|
||||
|
||||
pattern = self.check_and_fix_namespace(key=pattern)
|
||||
async for key in _redis_client.scan_iter(match=pattern + "*", count=count): # type: ignore
|
||||
keys.append(key)
|
||||
if len(keys) >= count:
|
||||
|
|
@ -538,6 +540,11 @@ class RedisCache(BaseCache):
|
|||
Register a Lua script with Redis asynchronously.
|
||||
Works with both standalone Redis and Redis Cluster.
|
||||
|
||||
The returned callable namespaces every key it is invoked with, so Lua
|
||||
scripts hit the same prefixed keys as get/set/increment. Without this,
|
||||
scripts would operate on raw keys while the rest of the cache uses the
|
||||
namespace, leaving rate-limit and lock keys outside the configured prefix.
|
||||
|
||||
Args:
|
||||
script (str): The Lua script to register
|
||||
|
||||
|
|
@ -548,7 +555,15 @@ class RedisCache(BaseCache):
|
|||
_redis_client = self.init_async_client()
|
||||
# For standalone Redis
|
||||
if hasattr(_redis_client, "register_script"):
|
||||
return _redis_client.register_script(script) # type: ignore
|
||||
registered_script = _redis_client.register_script(script) # type: ignore
|
||||
|
||||
async def namespaced_script(
|
||||
keys: list[str], args: list[Any], client: Any = None
|
||||
) -> Any:
|
||||
keys = [self.check_and_fix_namespace(key=key) for key in keys]
|
||||
return await registered_script(keys=keys, args=args, client=client)
|
||||
|
||||
return namespaced_script
|
||||
# For Redis Cluster
|
||||
elif hasattr(_redis_client, "script_load"):
|
||||
# Load the script and get its SHA
|
||||
|
|
@ -556,6 +571,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
# Return a callable that uses evalsha
|
||||
async def script_callable(keys: List[str], args: List[Any]) -> Any:
|
||||
keys = [self.check_and_fix_namespace(key=key) for key in keys]
|
||||
return _redis_client.evalsha(script_sha, len(keys), *keys, *args) # type: ignore
|
||||
|
||||
return script_callable
|
||||
|
|
@ -1257,6 +1273,7 @@ class RedisCache(BaseCache):
|
|||
async def delete_cache_keys(self, keys):
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
|
||||
_redis_client: Any = self.init_async_client()
|
||||
keys = [self.check_and_fix_namespace(key=key) for key in keys]
|
||||
# keys is a list, unpack it so it gets passed as individual elements to delete
|
||||
await _redis_client.delete(*keys)
|
||||
|
||||
|
|
@ -1322,10 +1339,12 @@ class RedisCache(BaseCache):
|
|||
async def async_delete_cache(self, key: str):
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
|
||||
_redis_client: Any = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
# keys is str
|
||||
return await _redis_client.delete(key)
|
||||
|
||||
def delete_cache(self, key):
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
self.redis_client.delete(key)
|
||||
|
||||
async def _pipeline_increment_helper(
|
||||
|
|
@ -1432,6 +1451,7 @@ class RedisCache(BaseCache):
|
|||
try:
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ttl`
|
||||
_redis_client: Any = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
ttl = await _redis_client.ttl(key)
|
||||
if ttl <= -1: # -1 means the key does not exist, -2 key does not exist
|
||||
return None
|
||||
|
|
@ -1460,6 +1480,7 @@ class RedisCache(BaseCache):
|
|||
int: The length of the list after the push operation
|
||||
"""
|
||||
_redis_client: Any = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
start_time = time.time()
|
||||
try:
|
||||
response = await _redis_client.rpush(key, *values)
|
||||
|
|
@ -1499,7 +1520,8 @@ class RedisCache(BaseCache):
|
|||
) -> List[int]:
|
||||
"""Helper function for pipeline rpush operations"""
|
||||
for rpush_op in rpush_list:
|
||||
pipe.rpush(rpush_op["key"], *rpush_op["values"])
|
||||
key = self.check_and_fix_namespace(key=rpush_op["key"])
|
||||
pipe.rpush(key, *rpush_op["values"])
|
||||
results = await pipe.execute()
|
||||
# Preserve positional correspondence — raise on per-command errors
|
||||
for r in results:
|
||||
|
|
@ -1586,6 +1608,7 @@ class RedisCache(BaseCache):
|
|||
**kwargs,
|
||||
) -> Union[Any, List[Any]]:
|
||||
_redis_client: Any = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
start_time = time.time()
|
||||
print_verbose(f"LPOP from Redis list: key: {key}, count: {count}")
|
||||
try:
|
||||
|
|
@ -1658,17 +1681,19 @@ class RedisCache(BaseCache):
|
|||
|
||||
if major_version >= 7:
|
||||
for lpop_op in lpop_list:
|
||||
pipe.lpop(lpop_op["key"], lpop_op["count"])
|
||||
key = self.check_and_fix_namespace(key=lpop_op["key"])
|
||||
pipe.lpop(key, lpop_op["count"])
|
||||
raw_results = await pipe.execute()
|
||||
else:
|
||||
# For Redis < 7, LPOP doesn't support count param.
|
||||
# Issue `count` individual LPOP commands per key, all in one pipeline.
|
||||
counts: List[int] = []
|
||||
for lpop_op in lpop_list:
|
||||
key = self.check_and_fix_namespace(key=lpop_op["key"])
|
||||
count = lpop_op["count"] or 1
|
||||
counts.append(count)
|
||||
for _ in range(count):
|
||||
pipe.lpop(lpop_op["key"])
|
||||
pipe.lpop(key)
|
||||
flat_results = await pipe.execute()
|
||||
|
||||
# Re-group the flat results back into per-key lists
|
||||
|
|
|
|||
|
|
@ -517,3 +517,177 @@ async def test_async_lpop_with_float_redis_version(
|
|||
|
||||
# Verify the method completed without error
|
||||
assert result is not None
|
||||
|
||||
|
||||
# LIT-3374: the namespace must be applied uniformly across every key-taking
|
||||
# Redis operation, not just get/set/increment. Before the fix these paths wrote
|
||||
# or read raw keys, so with a namespace configured the prefixed keys other
|
||||
# operations created were silently missed.
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"namespace, raw_keys, expected_keys",
|
||||
[
|
||||
(None, ["{k:v}:tokens", "{k:v}:requests"], ["{k:v}:tokens", "{k:v}:requests"]),
|
||||
(
|
||||
"litellm_sandbox",
|
||||
["{k:v}:tokens", "{k:v}:requests"],
|
||||
["litellm_sandbox:{k:v}:tokens", "litellm_sandbox:{k:v}:requests"],
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_namespaces_keys(
|
||||
namespace, raw_keys, expected_keys, monkeypatch, redis_no_ping
|
||||
):
|
||||
"""The callable returned by async_register_script (used by the rate limiter
|
||||
Lua scripts, pod-lock release, and budget limiters) must namespace every key
|
||||
it is invoked with. The hash tag is preserved so cluster slotting is intact."""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
|
||||
registered_script = AsyncMock(return_value="ok")
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.register_script = MagicMock(return_value=registered_script)
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
script = redis_cache.async_register_script("return 1")
|
||||
result = await script(keys=raw_keys, args=[60])
|
||||
|
||||
assert result == "ok"
|
||||
registered_script.assert_awaited_once_with(
|
||||
keys=expected_keys, args=[60], client=None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_cache_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
await redis_cache.async_delete_cache("k")
|
||||
mock_redis_instance.delete.assert_awaited_once_with(expected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_cache_keys_namespaces_keys(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
await redis_cache.delete_cache_keys(["k"])
|
||||
mock_redis_instance.delete.assert_awaited_once_with(expected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_ttl_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.ttl = AsyncMock(return_value=42)
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
ttl = await redis_cache.async_get_ttl("k")
|
||||
assert ttl == 42
|
||||
mock_redis_instance.ttl.assert_awaited_once_with(expected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_lpop_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.lpop = AsyncMock(return_value=b"value")
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
await redis_cache.async_lpop(key="k")
|
||||
mock_redis_instance.lpop.assert_awaited_once_with(expected, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.rpush = AsyncMock(return_value=1)
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
await redis_cache.async_rpush("k", ["v"])
|
||||
mock_redis_instance.rpush.assert_awaited_once_with(expected, "v")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected_match", [(None, "k*"), ("ns", "ns:k*")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_scan_iter_namespaces_pattern(
|
||||
namespace, expected_match, monkeypatch, redis_no_ping
|
||||
):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
|
||||
captured = {}
|
||||
|
||||
def scan_iter(match, count):
|
||||
captured["match"] = match
|
||||
|
||||
async def gen():
|
||||
for _ in ():
|
||||
yield _
|
||||
|
||||
return gen()
|
||||
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.scan_iter = scan_iter
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
await redis_cache.async_scan_iter(pattern="k")
|
||||
assert captured["match"] == expected_match
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
def test_increment_cache_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_client = MagicMock()
|
||||
mock_client.incr.return_value = 5
|
||||
mock_client.ttl.return_value = 100
|
||||
redis_cache.redis_client = mock_client
|
||||
redis_cache.increment_cache(key="k", value=1)
|
||||
mock_client.incr.assert_called_once_with(name=expected, amount=1)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
def test_delete_cache_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_client = MagicMock()
|
||||
redis_cache.redis_client = mock_client
|
||||
redis_cache.delete_cache(key="k")
|
||||
mock_client.delete.assert_called_once_with(expected)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue