mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
f67caac8d4
commit
321be01877
3 changed files with 20 additions and 103 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue