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 <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-30 20:39:48 -07:00 • committed by GitHub
parent f67caac8d4
commit 321be01877
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 20 additions and 103 deletions

View file

@ -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,

View file

@ -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."""

View file

@ -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)