From 321be018772754d21b1f9a033488452a798ba081 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 30 Sep 2026 20:39:48 -0700 Subject: [PATCH] fix(caching): write the response-cache SET to Redis at once instead of on the post-call batch (#43973) * fix(caching): write the response-cache SET to Redis at once instead of on the post-call batch The post-call Redis batch only goes out after every success callback finishes or the 1s deadline, so an identical request sent right after the first response missed the cache and went to the provider again. Response-cache writes go straight to Redis again, the counters, rate limits, TPM and slot releases keep riding the post-call batch Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): re-export DualCache explicitly instead of through a noqa Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): keep the DualCache re-export as a reasoned noqa for the strict ruff gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/caching.py | 48 +------------ litellm/caching/dual_cache.py | 6 -- .../test_request_redis_batch_post_call.py | 69 +++++-------------- 3 files changed, 20 insertions(+), 103 deletions(-) diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 85fd56c01e3..9e04ca79822 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -8,7 +8,6 @@ # Thank you users! We ❤️ you! - Krrish & Ishaan import ast -import asyncio import hashlib import json import logging @@ -31,11 +30,10 @@ from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg from .azure_blob_cache import AzureBlobCache from .base_cache import BaseCache from .disk_cache import DiskCache -from .dual_cache import DualCache +from .dual_cache import DualCache # noqa: F401 # re-exported, callers import DualCache from litellm.caching.caching from .gcs_cache import GCSCache from .in_memory_cache import InMemoryCache from .qdrant_semantic_cache import QdrantSemanticCache -from .redis_batch import active_post_call_redis_batch from .redis_cache import RedisCache, log_redis_failure from .redis_cluster_cache import RedisClusterCache from .redis_semantic_cache import RedisSemanticCache @@ -70,15 +68,6 @@ def print_verbose(print_statement): pass -def _ttl_seconds(raw: object) -> int | None: - if not isinstance(raw, (int, float, str)): - return None - try: - return int(raw) - except ValueError: - return None - - class CacheMode(str, Enum): default_on = "default_on" default_off = "default_off" @@ -770,8 +759,6 @@ class Cache: await self.batch_cache_write(result, **kwargs) else: cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs) - if await self._defer_set_to_post_call_batch(cache_key, cached_data, kwargs, dynamic_cache_object): - return if dynamic_cache_object is not None: await dynamic_cache_object.async_set_cache(cache_key, cached_data, **kwargs) else: @@ -779,39 +766,6 @@ class Cache: except Exception as e: self._log_add_cache_failure(e) - async def _defer_set_to_post_call_batch( - self, - cache_key: str, - cached_data: object, - kwargs: Mapping[str, object], - dynamic_cache_object: BaseCache | None, - ) -> bool: - """A plain SET on the Redis response cache rides the request's post-call pipeline with the counters, - instead of its own round trip. Anything with SET options keeps the direct path.""" - if kwargs.get("nx"): - return False - ttl: Final = _ttl_seconds(kwargs.get("ttl")) - if isinstance(dynamic_cache_object, DualCache): - deferred: Final = await dynamic_cache_object.async_set_cache_post_call(cache_key, cached_data, ttl) - if deferred is None: - return False - deferred.on_settled(self._log_deferred_add_cache_failure) - return True - if dynamic_cache_object is not None or not isinstance(self.cache, RedisCache): - return False - batch: Final = active_post_call_redis_batch(self.cache) - if batch is None: - return False - batch.set(cache_key, cached_data, ttl).on_settled(self._log_deferred_add_cache_failure) - return True - - def _log_deferred_add_cache_failure(self, future: asyncio.Future[None]) -> None: - if future.cancelled(): - return - failure: Final = future.exception() - if isinstance(failure, Exception): - self._log_add_cache_failure(failure) - def _convert_to_cached_embedding( self, embedding_response: Any, diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 47ce1d35895..042d27eb553 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -525,12 +525,6 @@ class DualCache(BaseCache): batch: Final = None if self.redis_cache is None else active_request_redis_batch(self.redis_cache) return None if batch is None else await self._set_on_batch(batch, key, value, ttl) - async def async_set_cache_post_call(self, key: str, value: object, ttl: float | None) -> BatchResult[None] | None: - """Memory now, the Redis SET on the request's post-call pipeline; None when no pipeline is open, so the - caller takes its direct path.""" - batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache) - return None if batch is None else await self._set_on_batch(batch, key, value, ttl) - async def async_delete_cache_pre_call(self, key: str) -> BatchResult[None] | None: """Memory now, the Redis DEL on the request's pipeline; None when no pipeline is open, so the caller takes its direct path.""" diff --git a/tests/unit/caching/test_request_redis_batch_post_call.py b/tests/unit/caching/test_request_redis_batch_post_call.py index 2b5d3b3dbbb..fdd328a9a57 100644 --- a/tests/unit/caching/test_request_redis_batch_post_call.py +++ b/tests/unit/caching/test_request_redis_batch_post_call.py @@ -1,6 +1,7 @@ """One Redis pipeline per backend for the post-call writes of a request: spend counters, rate-limit token -scripts and slot releases, deployment TPM and the response-cache SET all ride the post-call batch, which -goes out once the success/failure callbacks have run (or at the deadline when no callback phase closes it).""" +scripts and slot releases and deployment TPM all ride the post-call batch, which goes out once the success/failure +callbacks have run (or at the deadline when no callback phase closes it). The response-cache SET stays direct so the +next identical request can hit it while the callbacks are still running.""" from __future__ import annotations @@ -103,7 +104,9 @@ def _limiter(redis_cache: FakeRedisCache) -> _PROXY_MaxParallelRequestsHandler_v def _slot_stash(slot_id: str, *counter_keys: str) -> RequestRateLimiterStash: - return RequestRateLimiterStash(parallel_slot=ParallelSlotAcquisition(slot_id=slot_id, counter_keys=list(counter_keys))) + return RequestRateLimiterStash( + parallel_slot=ParallelSlotAcquisition(slot_id=slot_id, counter_keys=list(counter_keys)) + ) def _token_ops(*keys: str) -> list[RedisPipelineIncrementOperation]: @@ -140,13 +143,9 @@ async def test_every_post_call_owner_rides_one_pipeline_that_goes_out_when_the_c client = FakeClient(_ok_replies) redis_cache = PostCallFakeRedisCache(client) limiter = _limiter(redis_cache) - response_cache = _response_cache(redis_cache) tpm, router_cache = _tpm_router(redis_cache) with request_redis_batch_scope(): - await response_cache.async_add_cache( - {"id": "resp"}, messages=[{"role": "user", "content": "hi"}], model="gpt", ttl=120 - ) await tpm.async_log_success_event(_tpm_kwargs(), None, None, None) await limiter.async_increment_tokens_with_ttl_preservation(_token_ops("{api_key:k1}:tokens")) await limiter._release_stashed_parallel_slot( @@ -156,7 +155,7 @@ async def test_every_post_call_owner_rides_one_pipeline_that_goes_out_when_the_c await flush_post_call_redis_batches() assert len(client.pipelines) == 1 - assert _names(client) == ["SET", "INCRBYFLOAT", "EXPIRE", "EVALSHA", "EVALSHA"] + assert _names(client) == ["INCRBYFLOAT", "EXPIRE", "EVALSHA", "EVALSHA"] evalshas = [c for c in client.pipelines[0].commands if c[0] == "EVALSHA"] assert [c[1] for c in evalshas] == [sha_of(TOKEN_INCREMENT_SCRIPT), sha_of(PARALLEL_RELEASE_SCRIPT)] assert redis_cache.alone == [] @@ -169,40 +168,42 @@ async def test_every_post_call_owner_rides_one_pipeline_that_goes_out_when_the_c @pytest.mark.asyncio -async def test_the_response_cache_write_is_the_same_set_the_direct_path_issues(): +async def test_the_response_cache_set_reaches_redis_before_the_post_call_pipeline_goes_out(): client = FakeClient(_ok_replies) redis_cache = PostCallFakeRedisCache(client) response_cache = _response_cache(redis_cache) kwargs = {"messages": [{"role": "user", "content": "hi"}], "model": "gpt", "ttl": 120} + cache_key = response_cache.get_cache_key(**kwargs) with request_redis_batch_scope(): await response_cache.async_add_cache({"id": "resp"}, **kwargs) + assert redis_cache.store[cache_key]["response"] == {"id": "resp"} await flush_post_call_redis_batches() - cache_key = response_cache.get_cache_key(**kwargs) - (command,) = client.pipelines[0].commands - assert (command[0], command[1], command[3]) == ("SET", cache_key, 120) - assert json.loads(command[2])["response"] == {"id": "resp"} + assert client.pipelines == [] + (direct_set,) = redis_cache.alone + assert (direct_set[0], direct_set[1], direct_set[2]["ttl"]) == ("SET", cache_key, 120) @pytest.mark.asyncio -async def test_a_chat_response_written_through_the_handler_dual_cache_lands_in_memory_and_rides_the_pipeline(): +async def test_a_chat_response_written_through_the_handler_dual_cache_is_in_memory_and_redis_at_once(): client = FakeClient(_ok_replies) redis_cache = PostCallFakeRedisCache(client) response_cache = _response_cache(redis_cache) handler_cache = DualCache(redis_cache=redis_cache, in_memory_cache=InMemoryCache()) kwargs = {"messages": [{"role": "user", "content": "hi"}], "model": "gpt", "ttl": 120} + cache_key = response_cache.get_cache_key(**kwargs) with request_redis_batch_scope(): await response_cache.async_add_cache('{"id": "resp"}', dynamic_cache_object=handler_cache, **kwargs) - cache_key = response_cache.get_cache_key(**kwargs) in_memory = await handler_cache.in_memory_cache.async_get_cache(cache_key) assert in_memory["response"] == '{"id": "resp"}' - assert redis_cache.alone == [] + assert redis_cache.store[cache_key]["response"] == '{"id": "resp"}' await flush_post_call_redis_batches() - (command,) = client.pipelines[0].commands - assert (command[0], command[1], command[3]) == ("SET", cache_key, 120) + assert client.pipelines == [] + (direct_set,) = redis_cache.alone + assert (direct_set[0], direct_set[1], direct_set[2]["ttl"]) == ("SET", cache_key, 120) @pytest.mark.asyncio @@ -215,10 +216,8 @@ async def test_a_failed_operation_fails_only_its_owner_and_the_owner_applies_its client = FakeClient(replies) redis_cache = PostCallFakeRedisCache(client) limiter = _limiter(redis_cache) - response_cache = _response_cache(redis_cache) with request_redis_batch_scope(): - await response_cache.async_add_cache({"id": "resp"}, messages=[{"role": "user", "content": "hi"}], model="gpt") await limiter.async_increment_tokens_with_ttl_preservation(_token_ops("{api_key:k1}:tokens")) await limiter.async_increment_tokens_with_ttl_preservation(_token_ops("{team:t1}:tokens")) await flush_post_call_redis_batches() @@ -279,22 +278,6 @@ async def test_a_slot_released_before_the_response_reaches_redis_at_once_not_on_ assert client.pipelines == [] -@pytest.mark.asyncio -async def test_a_deferred_response_cache_set_without_a_ttl_expires_in_redis_like_the_direct_path(): - client = FakeClient(_ok_replies) - redis_cache = PostCallFakeRedisCache(client) - dual_cache = DualCache(redis_cache=redis_cache, in_memory_cache=InMemoryCache(), default_in_memory_ttl=300) - - 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] == 300 - - @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: @@ -384,20 +367,6 @@ async def test_two_backends_get_one_post_call_pipeline_each(): assert [c[1] for c in a_client.pipelines[0].commands if c[0] == "INCRBYFLOAT"] == ["x", "z"] -@pytest.mark.asyncio -async def test_a_numeric_string_ttl_reaches_redis_as_the_direct_path_would_send_it(): - client = FakeClient(_ok_replies) - response_cache = _response_cache(PostCallFakeRedisCache(client)) - kwargs = {"messages": [{"role": "user", "content": "hi"}], "model": "gpt", "ttl": "3600"} - - with request_redis_batch_scope(): - await response_cache.async_add_cache({"id": "resp"}, **kwargs) - await flush_post_call_redis_batches() - - (command,) = client.pipelines[0].commands - assert (command[0], command[3]) == ("SET", 3600) - - @pytest.mark.asyncio async def test_post_call_writes_still_waiting_on_their_callbacks_are_drained_at_shutdown(): client = FakeClient(_ok_replies)