This commit is contained in:
dingdangmao 2026-09-30 15:08:59 +00:00 • committed by GitHub
commit 309c099cf7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 148 additions and 20 deletions

View file

@ -151,13 +151,21 @@ class DualCache(BaseCache):
# Update both Redis and in-memory cache
try:
if self.in_memory_cache is not None:
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
kwargs["ttl"] = self.default_in_memory_ttl
mem_kwargs = dict(kwargs) # mutable-ok: per-tier copy; tiers must not share a ttl
if "ttl" not in mem_kwargs and self.default_in_memory_ttl is not None:
mem_kwargs["ttl"] = self.default_in_memory_ttl
self.in_memory_cache.set_cache(key, value, **kwargs)
self.in_memory_cache.set_cache(key, value, **mem_kwargs)
if self.redis_cache is not None and local_only is False:
self.redis_cache.set_cache(key, value, **kwargs)
redis_kwargs = dict(kwargs) # mutable-ok: per-tier copy; tiers must not share a ttl
if "ttl" not in redis_kwargs:
redis_ttl = self.default_redis_ttl
if redis_ttl is None:
redis_ttl = self.default_in_memory_ttl
if redis_ttl is not None:
redis_kwargs["ttl"] = redis_ttl
self.redis_cache.set_cache(key, value, **redis_kwargs)
except Exception as e:
print_verbose(e)
@ -508,12 +516,20 @@ class DualCache(BaseCache):
print_verbose(f"async set cache: cache key: {key}; local_only: {local_only}; value: {value}")
try:
if self.in_memory_cache is not None:
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
kwargs["ttl"] = self.default_in_memory_ttl
await self.in_memory_cache.async_set_cache(key, value, **kwargs)
mem_kwargs = dict(kwargs) # mutable-ok: per-tier copy; tiers must not share a ttl
if "ttl" not in mem_kwargs and self.default_in_memory_ttl is not None:
mem_kwargs["ttl"] = self.default_in_memory_ttl
await self.in_memory_cache.async_set_cache(key, value, **mem_kwargs)
if self.redis_cache is not None and local_only is False:
await self.redis_cache.async_set_cache(key, value, **kwargs)
redis_kwargs = dict(kwargs) # mutable-ok: per-tier copy; tiers must not share a ttl
if "ttl" not in redis_kwargs:
redis_ttl = self.default_redis_ttl
if redis_ttl is None:
redis_ttl = self.default_in_memory_ttl
if redis_ttl is not None:
redis_kwargs["ttl"] = redis_ttl
await self.redis_cache.async_set_cache(key, value, **redis_kwargs)
except Exception as e:
log_redis_failure(
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in async add_cache", e, with_traceback=True
@ -545,7 +561,10 @@ class DualCache(BaseCache):
effective_ttl: Final = self.default_in_memory_ttl if ttl is None else ttl
if self.in_memory_cache is not None:
await self.in_memory_cache.async_set_cache(key, value, ttl=effective_ttl)
return batch.set(key, value, effective_ttl)
redis_ttl = self.default_redis_ttl if ttl is None else ttl
if redis_ttl is None:
redis_ttl = effective_ttl
return batch.set(key, value, redis_ttl)
# async_batch_set_cache
async def async_set_cache_pipeline(
@ -557,13 +576,21 @@ class DualCache(BaseCache):
print_verbose(f"async batch set cache: cache keys: {cache_list}; local_only: {local_only}")
try:
if self.in_memory_cache is not None:
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
kwargs["ttl"] = self.default_in_memory_ttl
await self.in_memory_cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
mem_kwargs = dict(kwargs) # mutable-ok: per-tier copy; tiers must not share a ttl
if "ttl" not in mem_kwargs and self.default_in_memory_ttl is not None:
mem_kwargs["ttl"] = self.default_in_memory_ttl
await self.in_memory_cache.async_set_cache_pipeline(cache_list=cache_list, **mem_kwargs)
if self.redis_cache is not None and local_only is False:
redis_kwargs = dict(kwargs) # mutable-ok: per-tier copy; tiers must not share a ttl
if "ttl" not in redis_kwargs:
redis_ttl = self.default_redis_ttl
if redis_ttl is None:
redis_ttl = self.default_in_memory_ttl
if redis_ttl is not None:
redis_kwargs["ttl"] = redis_ttl
await self.redis_cache.async_set_cache_pipeline(
cache_list=cache_list, ttl=kwargs.pop("ttl", None), **kwargs
cache_list=cache_list, ttl=redis_kwargs.pop("ttl", None), **redis_kwargs
)
except Exception as e:
log_redis_failure(

View file

@ -69,17 +69,21 @@ class UserApiKeyCache(DualCache):
default_redis_ttl: float | None = None,
key_object_in_memory_cache: InMemoryCache | None = None,
) -> None:
# The auth cache contract is that Redis mirrors the in-memory TTL when no
# explicit Redis TTL is configured (see #43187). Compute the effective
# Redis TTL here so DualCache stays generic and does not special-case us.
effective_redis_ttl: Final = default_redis_ttl if default_redis_ttl is not None else default_in_memory_ttl
super().__init__(
in_memory_cache=in_memory_cache,
redis_cache=redis_cache,
default_in_memory_ttl=default_in_memory_ttl,
default_redis_ttl=default_redis_ttl,
default_redis_ttl=effective_redis_ttl,
)
self.key_object_cache: Final = DualCache(
in_memory_cache=key_object_in_memory_cache or InMemoryCache(),
redis_cache=redis_cache,
default_in_memory_ttl=default_in_memory_ttl,
default_redis_ttl=default_redis_ttl,
default_redis_ttl=effective_redis_ttl,
)
def in_memory_cache_for(self, key: str) -> InMemoryCache:

View file

@ -4899,10 +4899,13 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache:
default_redis_ttl=CLI_SSO_SESSION_TTL_SECONDS,
)
if enable_redis_auth_cache is True:
user_api_key_cache.attach_redis_cache(
redis_cache,
default_redis_ttl=litellm.default_redis_ttl,
)
# The auth cache mirrors its own in-memory TTL on the Redis tier (see
# UserApiKeyCache.__init__ and #43187). Do not pass the global
# default_redis_ttl here: that would override the auth cache's shorter
# TTL whenever default_redis_ttl is set before Redis is attached.
# Operators who want a longer auth-cache TTL should set
# general_settings.user_api_key_cache_ttl.
user_api_key_cache.attach_redis_cache(redis_cache)
verbose_proxy_logger.info(
"enable_redis_auth_cache=True: attached Redis to "
"user_api_key_cache — virtual-key lookups are now "
@ -5893,9 +5896,18 @@ class ProxyConfig:
litellm.default_in_memory_ttl = cache_params["default_in_memory_ttl"]
if "default_redis_ttl" in cache_params:
# default_redis_ttl is a DualCache/global setting, not a redis-py
# Redis() constructor kwarg. Promote it to the global and filter it
# out when constructing Cache (do NOT pop the caller's dict: the
# caller snapshots cache_params to detect DB reloads and a mutation
# here would force a cache rebuild on every poll).
litellm.default_redis_ttl = cache_params["default_redis_ttl"]
litellm.cache = Cache(**cache_params)
# Copy first: the caller snapshots cache_params to detect DB reloads, so
# mutating it here would force a cache rebuild on every poll.
cache_kwargs = dict(cache_params) # mutable-ok: local filter; caller's dict is left untouched
cache_kwargs.pop("default_redis_ttl", None)
litellm.cache = Cache(**cache_kwargs)
resolved_usage_cache = redis_usage_cache
cache_backend: Final = litellm.cache.cache if litellm.cache is not None else None

View file

@ -189,3 +189,27 @@ class TestRedisAuthCacheFlag:
ps._attach_redis_usage_cache(fake_redis, enable_redis_auth_cache=False)
assert limiter_cache.redis_cache is fake_redis
assert ps.user_api_key_cache.redis_cache is None
def test_auth_cache_redis_tier_keeps_in_memory_ttl_when_default_redis_ttl_is_set(monkeypatch):
"""
#43187: user_api_key_cache mirrors its in-memory TTL on Redis. A global
default_redis_ttl, from litellm_settings or cache_params, must not override
that when Redis is attached.
"""
from types import SimpleNamespace
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
monkeypatch.setattr(litellm, "default_redis_ttl", 3600)
auth_cache = UserApiKeyCache(default_in_memory_ttl=60)
with (
patch.object(ps, "user_api_key_cache", auth_cache),
patch.object(ps, "spend_counter_cache", DualCache()),
patch.object(ps, "cli_sso_session_cache", DualCache()),
patch.object(ps, "litellm_config_cache", SimpleNamespace(redis_cache=None)),
):
ps._attach_redis_usage_cache(_FakeRedisCache(), enable_redis_auth_cache=True)
assert auth_cache.redis_cache is not None
assert auth_cache.default_redis_ttl == 60
assert auth_cache.key_object_cache.default_redis_ttl == 60

View file

@ -925,3 +925,43 @@ async def test_shared_batch_read_keeps_a_caches_own_tier_failure_to_itself_like_
assert shared == separate == [None, None, [3]]
assert redis.async_batch_get_cache.await_args_list[0].args[0] == ["b1", "c1"]
async def _write_through(dual_cache: DualCache, write_path: str, **kwargs) -> None:
if write_path == "set_cache":
dual_cache.set_cache("ttl_key", "v", **kwargs)
elif write_path == "async_set_cache":
await dual_cache.async_set_cache("ttl_key", "v", **kwargs)
else:
await dual_cache.async_set_cache_pipeline([("ttl_key", "v")], **kwargs)
@pytest.mark.asyncio
@pytest.mark.parametrize("write_path", ["set_cache", "async_set_cache", "async_set_cache_pipeline"])
@pytest.mark.parametrize(
("kwargs", "memory_ttl", "redis_ttl"),
[({}, 60, 3600), ({"ttl": 99}, 99, 99)],
ids=["tier_defaults", "explicit_ttl"],
)
async def test_dual_cache_writes_each_tier_with_its_own_default_ttl(write_path, kwargs, memory_ttl, redis_ttl):
"""
Regression for #43187: the Redis tier was written with default_in_memory_ttl,
so a configured default_redis_ttl never took effect. An explicit ttl still
reaches both tiers unchanged.
"""
in_memory_cache = InMemoryCache(default_ttl=600)
mock_redis = MagicMock()
mock_redis.async_set_cache = AsyncMock()
mock_redis.async_set_cache_pipeline = AsyncMock()
dual_cache = DualCache(
in_memory_cache=in_memory_cache,
redis_cache=mock_redis,
default_in_memory_ttl=60,
default_redis_ttl=3600,
)
before = time.time()
await _write_through(dual_cache, write_path, **kwargs)
after = time.time()
assert getattr(mock_redis, write_path).call_args.kwargs["ttl"] == redis_ttl
expiry = in_memory_cache.ttl_dict["ttl_key"]
assert before + memory_ttl <= expiry <= after + memory_ttl

View file

@ -295,6 +295,27 @@ async def test_a_deferred_response_cache_set_without_a_ttl_expires_in_redis_like
assert command[3] == 300
@pytest.mark.asyncio
async def test_a_deferred_response_cache_set_without_a_ttl_uses_default_redis_ttl_when_configured():
client = FakeClient(_ok_replies)
redis_cache = PostCallFakeRedisCache(client)
dual_cache = DualCache(
redis_cache=redis_cache,
in_memory_cache=InMemoryCache(),
default_in_memory_ttl=300,
default_redis_ttl=3600,
)
await dual_cache.async_set_cache("direct", {"id": "resp"})
with request_redis_batch_scope():
await dual_cache.async_set_cache_post_call("deferred", {"id": "resp"}, None)
await flush_post_call_redis_batches()
(command,) = client.pipelines[0].commands
assert (command[0], command[1], command[3]) == ("SET", "deferred", redis_cache.alone[0][2]["ttl"])
assert command[3] == 3600
@pytest.mark.asyncio
async def test_a_released_slot_is_free_locally_at_once_and_the_older_redis_count_does_not_overwrite_the_gauge():
def replies(command: tuple[object, ...]) -> object: