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)