From 9525452d3706e3b0a39e2e7fa5f4c12b50439c78 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 17:26:05 -0700 Subject: [PATCH] perf(proxy): one post-call Redis pipeline per backend for spend, rate-limit, routing and response-cache writes (#43779) Post-call owners declare into one request-scoped RedisBatch per Redis backend: spend counter increments and reservation reconciliation, rate-limit token Lua updates and refunds, parallel-slot release (freed locally at once), deployment TPM, and compatible async response-cache SETs. The batch is sent once the success and failure callbacks have run, or on a deadline, and pending batches are drained at shutdown before Redis disconnects. nx writes, non-Redis caches and calls outside a request stay direct; numeric string TTLs keep the direct-path coercion. Resolves LIT-8883 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yassin --- litellm/caching/caching.py | 48 +- litellm/caching/dual_cache.py | 90 ++- litellm/caching/redis_batch.py | 107 ++- litellm/litellm_core_utils/litellm_logging.py | 3 + .../hooks/parallel_request_limiter_v3.py | 118 +++- .../proxy/hooks/proxy_track_cost_callback.py | 14 + litellm/proxy/proxy_server.py | 84 ++- litellm/router_strategy/lowest_tpm_rpm_v2.py | 2 +- .../test_request_redis_batch_post_call.py | 661 ++++++++++++++++++ 9 files changed, 1091 insertions(+), 36 deletions(-) create mode 100644 tests/unit/caching/test_request_redis_batch_post_call.py diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index e157730779b..85fd56c01e3 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -8,6 +8,7 @@ # Thank you users! We ❤️ you! - Krrish & Ishaan import ast +import asyncio import hashlib import json import logging @@ -30,10 +31,11 @@ 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 # noqa: F401 +from .dual_cache import DualCache 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 @@ -68,6 +70,15 @@ 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" @@ -759,6 +770,8 @@ 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: @@ -766,6 +779,39 @@ 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 1d7afcbee8f..996273d558a 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -8,23 +8,23 @@ Has 4 primary methods: - async_get_cache """ +import asyncio import itertools import logging import time -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from dataclasses import dataclass from threading import Lock from typing import TYPE_CHECKING, Any, Final -if TYPE_CHECKING: - from litellm.types.caching import RedisPipelineIncrementOperation - import litellm from litellm._logging import print_verbose, verbose_logger from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE +from litellm.types.caching import RedisPipelineIncrementOperation from .base_cache import BaseCache from .in_memory_cache import DEFAULT_MAX_SIZE_IN_MEMORY, InMemoryCache +from .redis_batch import BatchResult, RedisBatch, active_post_call_redis_batch from .redis_cache import RedisCache, RedisCircuitBreakerOpenError, log_redis_failure if TYPE_CHECKING: @@ -59,6 +59,24 @@ class PendingBatchRead: previous_access_times: dict[str, float | None] +@dataclass(frozen=True, slots=True) +class DeclaredBatchRead: + """A ``async_batch_get_cache`` split in two: the memory half done, the Redis half declared on a ``RedisBatch`` + so it rides that batch's next round trip, resolved later with ``async_resolve_batch_get``.""" + + keys: tuple[str, ...] + pending: PendingBatchRead + result: BatchResult[Mapping[str, object]] | None + + +def _log_deferred_increment_failure(future: asyncio.Future[float]) -> None: + if future.cancelled(): + return + failure: Final = future.exception() + if failure is not None: + log_redis_failure(verbose_logger, logging.WARNING, "post-call Redis increment failed", failure) + + class DualCache(BaseCache): """ DualCache is a cache implementation that updates both Redis and an in-memory cache simultaneously. @@ -335,7 +353,7 @@ class DualCache(BaseCache): ) async def _apply_batch_get( - self, pending: PendingBatchRead, redis_result: dict[str, object] | None, **kwargs: object + self, pending: PendingBatchRead, redis_result: Mapping[str, object] | None, **kwargs: object ) -> list[object | None]: if redis_result is None or all(v is None for v in redis_result.values()): return pending.result @@ -349,6 +367,22 @@ class DualCache(BaseCache): await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs)) return merged + async def declare_batch_get(self, keys: Sequence[str], batch: RedisBatch) -> DeclaredBatchRead: + pending: Final = await self._prepare_batch_get( + list(keys), # mutable-ok: the shared batch read takes a list + local_only=False, + throttle_redis=False, + ) + return DeclaredBatchRead( + keys=tuple(keys), + pending=pending, + result=batch.mget(pending.redis_keys) if pending.redis_keys else None, + ) + + async def async_resolve_batch_get(self, declared: DeclaredBatchRead) -> list[object | None]: + redis_result: Final = None if declared.result is None else await declared.result + return await self._apply_batch_get(declared.pending, redis_result) + async def async_batch_get_cache( self, keys: list, @@ -468,6 +502,17 @@ class DualCache(BaseCache): verbose_logger, logging.ERROR, "LiteLLM Cache: exception in async add_cache", e, with_traceback=True ) + 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) + if batch is None: + return None + 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) + # async_batch_set_cache async def async_set_cache_pipeline( self, cache_list: Sequence[tuple[str, object]], local_only: bool = False, **kwargs @@ -535,6 +580,41 @@ class DualCache(BaseCache): ) return result + async def async_increment_cache_post_call( + self, + key: str, + value: float, + ttl: int | None, + parent_otel_span: Span | None = None, + ) -> None: + """Memory is incremented now; the Redis increment rides the request's post-call pipeline when one is + open, and runs on its own as ``async_increment_cache`` otherwise.""" + await self.async_increment_cache_pipeline_post_call( + (RedisPipelineIncrementOperation(key=key, increment_value=value, ttl=ttl),), parent_otel_span + ) + + async def async_increment_cache_pipeline_post_call( + self, + increment_list: Sequence["RedisPipelineIncrementOperation"], + parent_otel_span: Span | None = None, + ) -> None: + batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache) + operations: Final = list(increment_list) # mutable-ok: both increment pipelines take a list + if batch is None: + await self.async_increment_cache_pipeline(operations, parent_otel_span=parent_otel_span) + return + try: + if self.in_memory_cache is not None: + await self.in_memory_cache.async_increment_pipeline( + increment_list=operations, parent_otel_span=parent_otel_span + ) + except Exception as e: # noqa: BLE001 # same tolerance as async_increment_cache_pipeline + log_redis_failure(verbose_logger, logging.WARNING, "in-memory increment failed", e) + for operation in increment_list: + batch.increment(operation["key"], operation["increment_value"], operation["ttl"]).on_settled( + _log_deferred_increment_failure + ) + async def async_increment_cache_pipeline( self, increment_list: list["RedisPipelineIncrementOperation"], diff --git a/litellm/caching/redis_batch.py b/litellm/caching/redis_batch.py index 8bd9554ec55..f6052192685 100644 --- a/litellm/caching/redis_batch.py +++ b/litellm/caching/redis_batch.py @@ -14,6 +14,7 @@ import hashlib import json import logging import time +import weakref from collections.abc import Awaitable, Callable, Generator, Mapping, Sequence from contextvars import ContextVar, Token from dataclasses import dataclass, field @@ -32,6 +33,8 @@ from litellm.types.services import ServiceTypes _T = TypeVar("_T") _ScriptArg = str | bytes | int | float +SettledHook = Callable[[asyncio.Future[_T]], Awaitable[None] | None] # mutable-ok: Callable params +POST_CALL_FLUSH_DEADLINE_SECONDS: Final = 1.0 class RegisteredScript(Protocol): @@ -52,11 +55,24 @@ class _Op(Generic[_T]): how to run on its own when the batch cannot pipeline (cluster client, or a reply the pipeline cannot settle, like NOSCRIPT).""" - __slots__ = ("future",) + __slots__ = ("future", "settled_hooks") def __init__(self) -> None: self.future: Final[asyncio.Future[_T]] = asyncio.get_running_loop().create_future() self.future.add_done_callback(_mark_retrieved) + self.settled_hooks: Final[list[SettledHook[_T]]] = [] # mutable-ok: append-only registry + + async def run_settled_hooks(self) -> None: + for hook in self.settled_hooks: + await self._run_settled_hook(hook) + + async def _run_settled_hook(self, hook: SettledHook[_T]) -> None: + try: + follow_up: Final = hook(self.future) + if follow_up is not None: + await follow_up + except Exception as e: # noqa: BLE001 # one owner's follow-up must not stop the others + verbose_logger.warning("redis batch settled hook failed: %s", e) def enqueue(self, pipe: _RedisPipeline) -> int: raise NotImplementedError @@ -238,6 +254,11 @@ class BatchResult(Generic[_T]): def done(self) -> bool: return self._op.future.done() + def on_settled(self, hook: SettledHook[_T]) -> None: + """For an owner that does not await: runs inside the flush once this operation has its result or + failure (or was cancelled with the pipeline), so the flush completes with the follow-up done.""" + self._op.settled_hooks.append(hook) + @dataclass(slots=True) class RedisBatch: @@ -294,6 +315,7 @@ class RedisBatch: for op in ops: if not op.future.done(): op.future.cancel() + await asyncio.gather(*(op.run_settled_hooks() for op in ops)) async def _flush_pipeline(self, ops: Sequence[_Op[object]]) -> None: start_time: Final = time.time() @@ -354,14 +376,34 @@ def _backend_key(redis_cache: RedisCache) -> object: return (type(redis_cache), redis_cache.namespace, settings) +_open_post_call: Final[weakref.WeakSet[RequestRedisBatches]] = weakref.WeakSet() +"""Requests whose post-call batch still holds declared ops, so a shutdown can send them before Redis goes away.""" + + class RequestRedisBatches: """One ``RedisBatch`` per Redis backend for the current request, so readers of different caches that - share a server (the proxy's and the router's) share the pipeline.""" + share a server (the proxy's and the router's) share the pipeline. - __slots__ = ("_batches", "prefetched") + The post-call batches hold the writes nothing waits on (counters, token scripts, the response cache). + They flush once, when the success or failure callbacks have all run, or at ``post_call_deadline`` + seconds after the first declaration when no callback phase closes them.""" - def __init__(self) -> None: + __slots__ = ( + "__weakref__", + "_batches", + "_deadline", + "_deadline_flush", + "_post_call", + "post_call_deadline", + "prefetched", + ) + + def __init__(self, post_call_deadline: float = POST_CALL_FLUSH_DEADLINE_SECONDS) -> None: self._batches: Final[dict[object, RedisBatch]] = {} # mutable-ok: lazily filled per backend + self._post_call: Final[dict[object, RedisBatch]] = {} # mutable-ok: lazily filled per backend + self.post_call_deadline: Final = post_call_deadline + self._deadline: asyncio.TimerHandle | None = None + self._deadline_flush: asyncio.Task[None] | None = None # Reads declared early for a consumer that runs later in the request, keyed by consumer name. self.prefetched: Final[dict[str, object]] = {} # mutable-ok: armed pre-admission, taken at use @@ -373,14 +415,44 @@ class RequestRedisBatches: self._batches[key] = batch return batch + def post_call(self, redis_cache: RedisCache) -> RedisBatch: + key: Final = _backend_key(redis_cache) + existing: Final = self._post_call.get(key) + batch: Final = ( + existing + if existing is not None + else self._post_call.setdefault(key, RedisBatch(redis_cache, name="post_call_redis_batch")) + ) + if self._deadline is None: + self._deadline = asyncio.get_running_loop().call_later(self.post_call_deadline, self._flush_on_deadline) + _open_post_call.add(self) + return batch + + def _flush_on_deadline(self) -> None: + self._deadline = None + self._deadline_flush = asyncio.ensure_future(self.flush_post_call()) + async def flush_all(self) -> None: """Send whatever is still declared (write-backs nobody awaits) before the request scope closes.""" await asyncio.gather(*(batch.flush() for batch in self._batches.values() if batch.pending)) + async def flush_post_call(self) -> None: + """One pipeline per backend for the post-call writes; the deadline is disarmed since this is that flush.""" + if self._deadline is not None: + self._deadline.cancel() + self._deadline = None + await asyncio.gather(*(batch.flush() for batch in self._post_call.values() if batch.pending)) + if not any(batch.pending for batch in self._post_call.values()): + _open_post_call.discard(self) + @property def batches(self) -> tuple[RedisBatch, ...]: return tuple(self._batches.values()) + @property + def post_call_batches(self) -> tuple[RedisBatch, ...]: + return tuple(self._post_call.values()) + _active_request_batches: Final[ContextVar[RequestRedisBatches | None]] = ContextVar( "request_redis_batches", default=None @@ -399,19 +471,40 @@ def active_request_redis_batches() -> RequestRedisBatches | None: return _active_request_batches.get() +def active_post_call_redis_batch(redis_cache: RedisCache) -> RedisBatch | None: + """The request's post-call batch for this backend, or None outside a ``request_redis_batch_scope``.""" + batches: Final = _active_request_batches.get() + if batches is None: + return None + return batches.post_call(redis_cache) + + +async def flush_post_call_redis_batches() -> None: + """Called where the success and failure callbacks of a request have all run.""" + batches: Final = _active_request_batches.get() + if batches is not None: + await batches.flush_post_call() + + +async def drain_post_call_redis_batches() -> None: + """Sends every post-call batch still waiting on its callbacks or deadline; for the shutdown path.""" + await asyncio.gather(*(batches.flush_post_call() for batches in tuple(_open_post_call))) + + class request_redis_batch_scope: """Redis reads declared inside share one pipeline per backend; nested scopes join the outer one.""" - __slots__ = ("_token",) + __slots__ = ("_post_call_deadline", "_token") - def __init__(self) -> None: + def __init__(self, post_call_deadline: float = POST_CALL_FLUSH_DEADLINE_SECONDS) -> None: self._token: Token[RequestRedisBatches | None] | None = None + self._post_call_deadline: Final = post_call_deadline def __enter__(self) -> RequestRedisBatches: outer: Final = _active_request_batches.get() if outer is not None: return outer - batches: Final = RequestRedisBatches() + batches: Final = RequestRedisBatches(post_call_deadline=self._post_call_deadline) self._token = _active_request_batches.set(batches) return batches diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index c393fa3caef..e292ab7b2ec 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -35,6 +35,7 @@ from litellm._uuid import uuid from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final from litellm.caching.caching import DualCache from litellm.caching.caching_handler import LLMCachingHandler +from litellm.caching.redis_batch import flush_post_call_redis_batches from litellm.constants import ( DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT, @@ -3552,6 +3553,7 @@ class Logging(LiteLLMLoggingBaseClass): traceback.format_exc(), ) self._handle_callback_failure(callback=callback) + await flush_post_call_redis_batches() def _handle_callback_failure(self, callback: object): """ @@ -3937,6 +3939,7 @@ class Logging(LiteLLMLoggingBaseClass): ) # Track callback logging failures in Prometheus self._handle_callback_failure(callback=callback) + await flush_post_call_redis_batches() def _get_trace_id(self, service_name: Literal["langfuse"]) -> str | None: """ diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 589aa7da7b5..e509f03d458 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -33,7 +33,12 @@ from typing_extensions import NotRequired, ReadOnly from litellm import DualCache from litellm._logging import verbose_proxy_logger -from litellm.caching.redis_batch import BatchResult, RegisteredScript, active_request_redis_batch +from litellm.caching.redis_batch import ( + BatchResult, + RegisteredScript, + active_post_call_redis_batch, + active_request_redis_batch, +) from litellm.caching.redis_cache import log_redis_failure from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger @@ -1893,6 +1898,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self, stash: RequestRateLimiterStash | None, parent_otel_span: Span | None, + *, + in_logging_callback: bool = False, ) -> None: if stash is None: return @@ -1900,7 +1907,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): acquisition: Final = stash.parallel_slot if acquisition is None: return - await self._release_parallel_request_slots(acquisition, parent_otel_span) + deferred: Final = in_logging_callback and await self._defer_parallel_slot_release( + acquisition, parent_otel_span + ) + if not deferred: + await self._release_parallel_request_slots(acquisition, parent_otel_span) stash.parallel_slot = None # rebind-ok: marks this request's slot as released async def _release_parallel_request_slots( @@ -1926,14 +1937,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): keys=counter_keys, args=[slot_id for _ in counter_keys], ) - for counter_key, remaining in zip(counter_keys, raw): - await self.internal_usage_cache.async_set_cache( - key=counter_key, - value=max(0, int(remaining)), - ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, - litellm_parent_otel_span=parent_otel_span, - local_only=True, - ) + await self._mirror_released_parallel_slots(counter_keys, raw, parent_otel_span) return except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the in-memory release, never a 500 log_redis_failure( @@ -1942,7 +1946,55 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): "parallel_release_script failed, falling back to in-memory release", e, ) + await self._release_parallel_request_slots_in_memory(counter_keys, slot_id, parent_otel_span) + async def _defer_parallel_slot_release( + self, acquisition: ParallelSlotAcquisition, parent_otel_span: Span | None + ) -> bool: + """Only for a release from the logging callbacks: the response has left and the callbacks' end flushes + the pipeline. A release before the response goes to Redis at once, so another worker's next acquire + never counts a finished request. The local gauge frees the slot at once, so admission on this worker + sees the capacity before the pipeline goes out. The count Redis returns from the pipeline is not + mirrored: by then a newer acquire on this worker may have written a fresher count, and the next + acquire refreshes the gauge anyway.""" + counter_keys: Final = acquisition["counter_keys"] + slot_id: Final = acquisition["slot_id"] + redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache + script: Final = self.parallel_release_script + batch: Final = None if redis_cache is None else active_post_call_redis_batch(redis_cache) + if batch is None or script is None or not counter_keys or not slot_id: + return False + await self._release_parallel_request_slots_in_memory(counter_keys, slot_id, parent_otel_span) + + async def settle(future: asyncio.Future[object]) -> None: + if future.cancelled() or future.exception() is not None: + log_redis_failure( + verbose_proxy_logger, + logging.WARNING, + "parallel_release_script failed, the slot stays released in memory only", + future.exception() if not future.cancelled() else asyncio.CancelledError(), + ) + + batch.script(PARALLEL_RELEASE_SCRIPT, script, counter_keys, (slot_id,) * len(counter_keys)).on_settled(settle) + return True + + async def _mirror_released_parallel_slots( + self, counter_keys: list[str], remaining_by_key: Sequence[object], parent_otel_span: Span | None + ) -> None: + for counter_key, remaining in zip(counter_keys, remaining_by_key): + if not isinstance(remaining, (int, float, str, bytes)): + continue + await self.internal_usage_cache.async_set_cache( + key=counter_key, + value=max(0, int(remaining)), + ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + + async def _release_parallel_request_slots_in_memory( + self, counter_keys: list[str], slot_id: str, parent_otel_span: Span | None + ) -> None: async with self._check_and_increment_lock: for counter_key in counter_keys: raw_value: ParallelGaugeCacheValue | None = await self.internal_usage_cache.async_get_cache( @@ -4301,11 +4353,43 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): keys.append(op["key"]) args.extend([op["increment_value"], ttl_value]) + if self._defer_token_increment_script(keys, args, group_operations): + continue await self.token_increment_script( keys=keys, args=args, ) + def _defer_token_increment_script( + self, + keys: list[str], + args: list[int], + group_operations: list["RedisPipelineIncrementOperation"], + ) -> bool: + """Declared into the request's post-call pipeline instead of its own EVALSHA round trip; a failed + script falls back to the plain increment pipeline for its own group, as the direct path does.""" + redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache + script: Final = self.token_increment_script + batch: Final = None if redis_cache is None else active_post_call_redis_batch(redis_cache) + if batch is None or script is None: + return False + + async def fall_back(future: asyncio.Future[object]) -> None: + if future.cancelled() or future.exception() is None: + return + log_redis_failure( + verbose_proxy_logger, + logging.WARNING, + "TTL preservation failed, falling back to regular pipeline", + future.exception(), + ) + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( + increment_list=group_operations, + ) + + batch.script(TOKEN_INCREMENT_SCRIPT, script, keys, args).on_settled(fall_back) + return True + async def async_increment_tokens_with_ttl_preservation( self, pipeline_operations: list["RedisPipelineIncrementOperation"], @@ -4919,7 +5003,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): verbose_proxy_logger.debug("INSIDE parallel request limiter ASYNC SUCCESS LOGGING") stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) - await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span) + await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span, in_logging_callback=True) pipeline_operations: Final = self._build_success_event_pipeline_operations( kwargs=kwargs, @@ -5039,7 +5123,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = [] stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) - await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span) + await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span, in_logging_callback=True) # Skip the reservation refund if async_post_call_failure_hook # already released it (proxy-level rejection that also bubbles up @@ -5109,15 +5193,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) if pipeline_operations: - await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( - increment_list=pipeline_operations, - litellm_parent_otel_span=litellm_parent_otel_span, + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline_post_call( + pipeline_operations, parent_otel_span=litellm_parent_otel_span ) for project_operations in (itpm_operations, otpm_operations): if isinstance(project_operations, list): - await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( - increment_list=project_operations, - litellm_parent_otel_span=litellm_parent_otel_span, + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline_post_call( + project_operations, parent_otel_span=litellm_parent_otel_span ) elif project_operations: await self.async_increment_reservation_aware_tokens( diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 05995d22293..877dbfabe5d 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -285,6 +285,7 @@ class _ProxyDBLogger(CustomLogger): increment_spend_counters, proxy_logging_obj, update_cache, + update_cache_read_keys, ) verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback") @@ -378,6 +379,13 @@ class _ProxyDBLogger(CustomLogger): request_tags=tags, model_access_groups=model_access_groups, project_id=project_id, + update_cache_read_keys=update_cache_read_keys( + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + tags=tags, + response_cost=response_cost, + ), ) if not charged: return @@ -695,6 +703,7 @@ async def _update_database_and_spend_counters( request_tags: list[str] | None = None, model_access_groups: Sequence[str] | None = None, project_id: str | None = None, + update_cache_read_keys: Sequence[str] = (), ) -> bool: """The reservation is reconciled before the spend is persisted, from its own read. One spend counter batch then spans the database write and the counter update, so the post-call counters are read with a single MGET after the @@ -736,6 +745,7 @@ async def _update_database_and_spend_counters( request_tags=request_tags, model_access_groups=model_access_groups, project_id=project_id, + update_cache_read_keys=update_cache_read_keys, ) @@ -756,7 +766,10 @@ async def _update_database_and_spend_counters_in_batch( request_tags: list[str] | None, model_access_groups: Sequence[str] | None, project_id: str | None, + update_cache_read_keys: Sequence[str], ) -> bool: + from litellm.proxy.proxy_server import arm_update_cache_read + try: charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database( token=user_api_key, @@ -788,6 +801,7 @@ async def _update_database_and_spend_counters_in_batch( await _release_budget_reservation(budget_reservation=budget_reservation) return False + await arm_update_cache_read(update_cache_read_keys) try: await increment_spend_counters( token=user_api_key, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3f5e01ed0bc..a994293a4d8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -269,6 +269,12 @@ import litellm._redis from litellm import Router from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger from litellm.caching.caching import DualCache, RedisCache +from litellm.caching.dual_cache import DeclaredBatchRead +from litellm.caching.redis_batch import ( + active_post_call_redis_batch, + active_request_redis_batches, + drain_post_call_redis_batches, +) from litellm.caching.redis_cache import RedisCircuitBreakerOpenError, is_redis_timeout_failure from litellm.caching.redis_cluster_cache import RedisClusterCache from litellm.constants import ( @@ -1112,6 +1118,7 @@ async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = N verbose_proxy_logger.debug("Disconnecting from Prisma") await prisma_client.disconnect() + await drain_post_call_redis_batches() if litellm.cache is not None: await litellm.cache.disconnect() @@ -3715,6 +3722,8 @@ async def _invalidate_spend_counter(counter_key: str): async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> None: + if _defer_spend_counter_increments(pending): + return try: await increment_spend_counters_pipeline(pending=pending) except Exception as e: @@ -3723,6 +3732,41 @@ async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncremen raise +def _defer_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> bool: + """Post-call increments ride the request's post-call pipeline with the other counters. Each counter's + new value lands in memory when the pipeline settles; a failed one is invalidated so no reader trusts a + counter whose increment may not have applied, as ``increment_spend_counters_pipeline`` does.""" + redis_cache: Final = spend_counter_cache.redis_cache + if redis_cache is None or not pending: + return False + batch: Final = active_post_call_redis_batch(redis_cache) + if batch is None: + return False + ttl: Final = redis_cache.get_ttl() + for item in pending: + batch.increment(item.counter_key, item.increment, ttl).on_settled(_settle_spend_counter_increment(item)) + return True + + +def _settle_spend_counter_increment(item: PendingSpendIncrement) -> Callable[[asyncio.Future[float]], Awaitable[None]]: + async def settle(future: asyncio.Future[float]) -> None: + if not future.cancelled() and future.exception() is None: + current_value: Final = float(future.result()) + spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value) + record_spend_counter_value(item.counter_key, current_value) + return + if future.cancelled(): + if spend_counter_cache.in_memory_cache.get_cache(key=item.counter_key) is not None: + spend_counter_cache.in_memory_cache.increment_cache(key=item.counter_key, value=item.increment) + return + verbose_proxy_logger.warning( + "Spend counter %s increment did not land in the post-call pipeline; invalidating it", item.counter_key + ) + await _invalidate_spend_counter(counter_key=item.counter_key) + + return settle + + async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]: """One INCRBYFLOAT+EXPIRE pipeline for every pending counter, returning each counter's new value in order; on failure every counter is invalidated before the error propagates, so no caller can read a half-applied batch.""" @@ -3762,7 +3806,7 @@ async def run_spend_counter_pipeline(pending: Sequence[PendingSpendIncrement]) - return tuple(float(current_value) for current_value in results or ()) -def _update_cache_read_keys( +def update_cache_read_keys( user_id: str | None, end_user_id: str | None, team_id: str | None, @@ -3778,13 +3822,45 @@ def _update_cache_read_keys( return user_keys + end_user_keys + team_keys + tag_keys -async def _read_update_cache_values(keys: Sequence[str], parent_otel_span: Span | None) -> Mapping[str, object]: +_UPDATE_CACHE_PREFETCH_SLOT: Final = "update_cache_read" + + +async def arm_update_cache_read(keys: Sequence[str], cache: DualCache | None = None) -> None: + """Declares the ``update_cache`` read on the request pipeline once the spend is persisted, so it rides the same + round trip as the post-call spend counter read instead of its own.""" + request: Final = active_request_redis_batches() + target: Final = user_api_key_cache if cache is None else cache + if request is None or target.redis_cache is None or not keys: + return + request.prefetched[_UPDATE_CACHE_PREFETCH_SLOT] = await target.declare_batch_get( + keys, request.batch(target.redis_cache) + ) + + +async def _take_armed_update_cache_read(keys: Sequence[str], cache: DualCache) -> Mapping[str, object] | None: + request: Final = active_request_redis_batches() + if request is None: + return None + armed: Final = request.prefetched.pop(_UPDATE_CACHE_PREFETCH_SLOT, None) + if not isinstance(armed, DeclaredBatchRead) or armed.keys != tuple(keys): + return None + values: Final = await cache.async_resolve_batch_get(armed) + return MappingProxyType({key: value for key, value in zip(keys, values) if value is not None}) + + +async def _read_update_cache_values( + keys: Sequence[str], parent_otel_span: Span | None, cache: DualCache | None = None +) -> Mapping[str, object]: """One batched read for every object ``update_cache`` refreshes; a failed read leaves them all untouched, exactly as a failed per-object GET left that object untouched.""" if not keys: return MappingProxyType({}) + target: Final = user_api_key_cache if cache is None else cache try: - values: Final = await user_api_key_cache.async_batch_get_cache( + armed: Final = await _take_armed_update_cache_read(keys, target) + if armed is not None: + return armed + values: Final = await target.async_batch_get_cache( keys=list(keys), parent_otel_span=parent_otel_span, throttle_redis=False ) except Exception as e: @@ -3817,7 +3893,7 @@ async def update_cache( values_to_update_in_cache: Final[list[tuple[str, object]]] = [] cached_values: Final = await _read_update_cache_values( - keys=_update_cache_read_keys( + keys=update_cache_read_keys( user_id=user_id, end_user_id=end_user_id, team_id=team_id, tags=tags, response_cost=response_cost ), parent_otel_span=parent_otel_span, diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 909b47833cb..6e21d5d1f1f 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -304,7 +304,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): # update cache parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) ## TPM - await self.router_cache.async_increment_cache( + await self.router_cache.async_increment_cache_post_call( key=tpm_key, value=total_tokens, ttl=self.routing_args.ttl, diff --git a/tests/unit/caching/test_request_redis_batch_post_call.py b/tests/unit/caching/test_request_redis_batch_post_call.py new file mode 100644 index 00000000000..2b5d3b3dbbb --- /dev/null +++ b/tests/unit/caching/test_request_redis_batch_post_call.py @@ -0,0 +1,661 @@ +"""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).""" + +from __future__ import annotations + +import asyncio +import datetime +import hashlib +import json +from collections.abc import Awaitable, Callable, Mapping, Sequence +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm.caching.caching import Cache +from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_batch import ( + active_post_call_redis_batch, + active_request_redis_batches, + drain_post_call_redis_batches, + flush_post_call_redis_batches, + request_redis_batch_scope, +) +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + PARALLEL_RELEASE_SCRIPT, + TOKEN_INCREMENT_SCRIPT, + ParallelSlotAcquisition, + RequestRateLimiterStash, + _PROXY_MaxParallelRequestsHandler_v3, +) +from litellm.proxy.spend_tracking.spend_counter_batch import PendingSpendIncrement +from litellm.proxy.utils import InternalUsageCache +from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2 +from litellm.types.caching import RedisPipelineIncrementOperation +from litellm.types.utils import ModelResponse + +from .test_redis_batch import FakeClient, FakeRedisCache + + +async def _script_outside_the_pipeline(keys: Sequence[str], args: Sequence[object]) -> object: + raise AssertionError("post-call scripts must ride the post-call pipeline") + + +class PostCallFakeRedisCache(FakeRedisCache): + """Records the direct (non-pipelined) writes an owner falls back to.""" + + def async_register_script(self, script: str) -> Callable[..., Awaitable[object]]: + return _script_outside_the_pipeline + + async def async_increment_pipeline( + self, increment_list: list[RedisPipelineIncrementOperation], **kwargs: object + ) -> list[float]: + return [await self.async_increment(op["key"], op["increment_value"]) for op in increment_list] + + async def async_delete_cache(self, key: str, **kwargs: object) -> None: # pyright: ignore[reportIncompatibleMethodOverride] # the fake drops RedisCache's unused kwargs + self.alone.append(("DEL", key)) + self.store.pop(key, None) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.alone.append(("SET", key, dict(kwargs))) + self.store[key] = value + + +def sha_of(script: str) -> str: + return hashlib.sha1(script.encode()).hexdigest() # noqa: S324 + + +def _ok_replies(command: tuple[object, ...]) -> object: + match command[0]: + case "INCRBYFLOAT": + return b"7.5" + case "EXPIRE": + return 1 + case "SET": + return True + case "EVALSHA": + return [3, 0] + case "MGET": + return [json.dumps({"spend": 1.0}) for _ in command[1:]] + raise AssertionError(command) + + +async def _run_ready_callbacks(client: FakeClient) -> None: + for _ in range(20): + if client.pipelines: + return + await asyncio.sleep(0) + + +def _names(client: FakeClient, index: int = 0) -> list[str]: + return [command[0] for command in client.pipelines[index].commands] + + +def _limiter(redis_cache: FakeRedisCache) -> _PROXY_MaxParallelRequestsHandler_v3: + dual_cache = DualCache() + dual_cache.attach_redis_cache(redis_cache) + return _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache=dual_cache)) + + +def _slot_stash(slot_id: str, *counter_keys: str) -> RequestRateLimiterStash: + return RequestRateLimiterStash(parallel_slot=ParallelSlotAcquisition(slot_id=slot_id, counter_keys=list(counter_keys))) + + +def _token_ops(*keys: str) -> list[RedisPipelineIncrementOperation]: + return [RedisPipelineIncrementOperation(key=key, increment_value=10, ttl=60) for key in keys] + + +def _response_cache(redis_cache: FakeRedisCache) -> Cache: + cache = Cache(type="local") + cache.type = "redis" # pyright: ignore[reportAttributeAccessIssue] # the fake stands in for the Redis backend + cache.cache = redis_cache + return cache + + +def _tpm_router(redis_cache: FakeRedisCache) -> tuple[LowestTPMLoggingHandler_v2, DualCache]: + router_cache = DualCache() + router_cache.attach_redis_cache(redis_cache) + return LowestTPMLoggingHandler_v2(router_cache=router_cache, routing_args={"ttl": 60}), router_cache + + +def _tpm_kwargs() -> Mapping[str, object]: + return { + "standard_logging_object": { + "model_group": "gpt", + "model_id": "dep-a", + "hidden_params": {"litellm_model_name": "openai/gpt-4o-mini"}, + "total_tokens": 42, + }, + "litellm_params": {"metadata": {}}, + } + + +@pytest.mark.asyncio +async def test_every_post_call_owner_rides_one_pipeline_that_goes_out_when_the_callbacks_are_done(): + 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( + _slot_stash("slot-1", "{api_key:k1}:parallel"), None, in_logging_callback=True + ) + assert client.pipelines == [] # nothing goes out while the callbacks are still declaring + await flush_post_call_redis_batches() + + assert len(client.pipelines) == 1 + assert _names(client) == ["SET", "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 == [] + assert ( + await router_cache.in_memory_cache.async_get_cache( + next(k for k in router_cache.in_memory_cache.cache_dict if ":tpm:" in k) + ) + == 42 + ) + + +@pytest.mark.asyncio +async def test_the_response_cache_write_is_the_same_set_the_direct_path_issues(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + response_cache = _response_cache(redis_cache) + kwargs = {"messages": [{"role": "user", "content": "hi"}], "model": "gpt", "ttl": 120} + + with request_redis_batch_scope(): + await response_cache.async_add_cache({"id": "resp"}, **kwargs) + 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"} + + +@pytest.mark.asyncio +async def test_a_chat_response_written_through_the_handler_dual_cache_lands_in_memory_and_rides_the_pipeline(): + 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} + + 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 == [] + await flush_post_call_redis_batches() + + (command,) = client.pipelines[0].commands + assert (command[0], command[1], command[3]) == ("SET", cache_key, 120) + + +@pytest.mark.asyncio +async def test_a_failed_operation_fails_only_its_owner_and_the_owner_applies_its_own_fallback(): + def replies(command: tuple[object, ...]) -> object: + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:tokens": + return Exception("ERR Lua") + return _ok_replies(command) + + 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() + + assert len(client.pipelines) == 1 + # the failed group falls back to the plain increment (memory + Redis), the healthy group does not + assert redis_cache.alone == [("INCRBYFLOAT", "{api_key:k1}:tokens", 10)] + assert await limiter.internal_usage_cache.dual_cache.in_memory_cache.async_get_cache("{api_key:k1}:tokens") == 10 + assert await limiter.internal_usage_cache.dual_cache.in_memory_cache.async_get_cache("{team:t1}:tokens") is None + + +@pytest.mark.asyncio +async def test_a_failed_slot_release_script_releases_the_slot_in_memory(): + def replies(command: tuple[object, ...]) -> object: + if command[0] == "EVALSHA": + return Exception("ERR Lua") + return _ok_replies(command) + + redis_cache = PostCallFakeRedisCache(FakeClient(replies)) + limiter = _limiter(redis_cache) + memory = limiter.internal_usage_cache.dual_cache.in_memory_cache + await memory.async_set_cache("{api_key:k1}:parallel", {"slot-1": 1.0, "slot-2": 1.0}) + + with request_redis_batch_scope(): + await limiter._release_stashed_parallel_slot( + _slot_stash("slot-1", "{api_key:k1}:parallel"), None, in_logging_callback=True + ) + await flush_post_call_redis_batches() + + assert await memory.async_get_cache("{api_key:k1}:parallel") == {"slot-2": 1.0} + + +class DirectScriptFakeRedisCache(PostCallFakeRedisCache): + """Records the release script a pre-response caller runs outside the pipeline.""" + + def async_register_script(self, script: str) -> Callable[..., Awaitable[object]]: + async def run(keys: Sequence[str], args: Sequence[object]) -> object: + self.alone.append(("EVALSHA", tuple(keys), tuple(args))) + return [0 for _ in keys] + + return run + + +@pytest.mark.asyncio +async def test_a_slot_released_before_the_response_reaches_redis_at_once_not_on_the_pipeline(): + client = FakeClient(_ok_replies) + redis_cache = DirectScriptFakeRedisCache(client) + limiter = _limiter(redis_cache) + memory = limiter.internal_usage_cache.dual_cache.in_memory_cache + await memory.async_set_cache("{api_key:k1}:parallel", {"slot-1": 1.0}) + + with request_redis_batch_scope(): + await limiter._release_stashed_parallel_slot(_slot_stash("slot-1", "{api_key:k1}:parallel"), None) + assert redis_cache.alone == [("EVALSHA", ("{api_key:k1}:parallel",), ("slot-1",))] + assert await memory.async_get_cache("{api_key:k1}:parallel") == 0 + await flush_post_call_redis_batches() + + 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: + if command[0] == "EVALSHA": + return [2] + return _ok_replies(command) + + redis_cache = PostCallFakeRedisCache(FakeClient(replies)) + limiter = _limiter(redis_cache) + memory = limiter.internal_usage_cache.dual_cache.in_memory_cache + await memory.async_set_cache("{api_key:k1}:parallel", {"slot-1": 1.0, "slot-2": 1.0, "slot-3": 1.0}) + + with request_redis_batch_scope(): + await limiter._release_stashed_parallel_slot( + _slot_stash("slot-1", "{api_key:k1}:parallel"), None, in_logging_callback=True + ) + assert await memory.async_get_cache("{api_key:k1}:parallel") == {"slot-2": 1.0, "slot-3": 1.0} + await memory.async_set_cache("{api_key:k1}:parallel", {"slot-2": 1.0, "slot-3": 1.0, "slot-4": 1.0}) + await flush_post_call_redis_batches() + + assert await memory.async_get_cache("{api_key:k1}:parallel") == {"slot-2": 1.0, "slot-3": 1.0, "slot-4": 1.0} + + +@pytest.mark.asyncio +async def test_failure_refunds_ride_the_post_call_pipeline_and_count_in_memory_at_once(): + client = FakeClient(_ok_replies) + dual_cache = DualCache() + dual_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + refund = [RedisPipelineIncrementOperation(key="{api_key:k1}:tokens", increment_value=-500, ttl=60)] + + with request_redis_batch_scope(): + await dual_cache.async_increment_cache_pipeline_post_call(refund) + assert await dual_cache.in_memory_cache.async_get_cache("{api_key:k1}:tokens") == -500 + assert client.pipelines == [] + await flush_post_call_redis_batches() + + assert client.pipelines[0].commands[0] == ("INCRBYFLOAT", "{api_key:k1}:tokens", -500) + + +@pytest.mark.asyncio +async def test_outside_a_request_scope_owners_write_directly_as_before(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + dual_cache = DualCache() + dual_cache.attach_redis_cache(redis_cache) + response_cache = _response_cache(redis_cache) + + await dual_cache.async_increment_cache_post_call("dep:tpm", 42, ttl=60) + await response_cache.async_add_cache({"id": "resp"}, messages=[{"role": "user", "content": "hi"}], model="gpt") + + assert client.pipelines == [] + assert redis_cache.alone[0] == ("INCRBYFLOAT", "dep:tpm", 42) + assert active_post_call_redis_batch(redis_cache) is None + + +@pytest.mark.asyncio +async def test_a_set_with_options_keeps_the_direct_path(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + response_cache = _response_cache(redis_cache) + + with request_redis_batch_scope(): + await response_cache.async_add_cache( + {"id": "r"}, messages=[{"role": "user", "content": "hi"}], model="gpt", nx=True + ) + await flush_post_call_redis_batches() + + assert client.pipelines == [] + (direct_set,) = redis_cache.alone + assert direct_set[0] == "SET" and direct_set[2]["nx"] is True + + +@pytest.mark.asyncio +async def test_two_backends_get_one_post_call_pipeline_each(): + a_client, b_client = FakeClient(_ok_replies), FakeClient(_ok_replies) + a, b = DualCache(), DualCache() + a.attach_redis_cache(PostCallFakeRedisCache(a_client)) + b.attach_redis_cache(PostCallFakeRedisCache(b_client)) + + with request_redis_batch_scope(): + await a.async_increment_cache_post_call("x", 1, ttl=None) + await b.async_increment_cache_post_call("y", 1, ttl=None) + await a.async_increment_cache_post_call("z", 1, ttl=None) + await flush_post_call_redis_batches() + + assert len(a_client.pipelines) == 1 and len(b_client.pipelines) == 1 + 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) + dual_cache = DualCache() + dual_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + + with request_redis_batch_scope(post_call_deadline=60) as request: + await dual_cache.async_increment_cache_post_call("x", 1, ttl=None) + await request.flush_all() + assert client.pipelines == [] + + await drain_post_call_redis_batches() + assert len(client.pipelines) == 1 and _names(client) == ["INCRBYFLOAT"] + + await drain_post_call_redis_batches() + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_post_call_batch_nobody_closes_goes_out_at_the_deadline(monkeypatch: pytest.MonkeyPatch): + client = FakeClient(_ok_replies) + dual_cache = DualCache() + dual_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + + loop = asyncio.get_running_loop() + armed_at = loop.time() + + with request_redis_batch_scope(post_call_deadline=60) as request: + await dual_cache.async_increment_cache_post_call("x", 1, ttl=None) + await request.flush_all() + await _run_ready_callbacks(client) + assert client.pipelines == [], "the request boundary drains the immediate batch, not the post-call one" + + monkeypatch.setattr(loop, "time", lambda: armed_at + 61) + await _run_ready_callbacks(client) + + assert len(client.pipelines) == 1 and _names(client) == ["INCRBYFLOAT"] + + +@pytest.mark.asyncio +async def test_the_success_handler_closes_the_post_call_batch_after_the_last_callback(monkeypatch): + client = FakeClient(_ok_replies) + dual_cache = DualCache() + dual_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + pipelines_seen_by_callbacks: list[int] = [] + + class Counter(CustomLogger): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + await dual_cache.async_increment_cache_post_call("counted", 1, ttl=None) + pipelines_seen_by_callbacks.append(len(client.pipelines)) + + monkeypatch.setattr(litellm, "_async_success_callback", []) + logging_obj = LitellmLogging( + model="test-model", + messages=[], + stream=False, + call_type="completion", + start_time=datetime.datetime.now(), + litellm_call_id="post-call", + function_id="post-call", + dynamic_async_success_callbacks=[Counter(), Counter()], + ) + logging_obj.update_environment_variables(litellm_params={"metadata": {}}, optional_params={}) + payload = { + "id": "post-call", + "call_type": "completion", + "metadata": {}, + "model_group": "test-model", + "model_parameters": {}, + } + + with request_redis_batch_scope(): + await logging_obj.async_success_handler(result=ModelResponse(), standard_logging_object=payload) + + assert pipelines_seen_by_callbacks == [0, 0] + assert len(client.pipelines) == 1 and _names(client) == ["INCRBYFLOAT", "INCRBYFLOAT"] + + +@pytest.mark.asyncio +async def test_spend_counter_increments_ride_the_pipeline_and_settle_into_memory(monkeypatch): + from litellm.proxy import proxy_server + + client = FakeClient(_ok_replies) + spend_cache = DualCache() + spend_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + monkeypatch.setattr(proxy_server, "spend_counter_cache", spend_cache) + pending = [PendingSpendIncrement("spend:key:k1", 0.5), PendingSpendIncrement("spend:team:t1", 0.5)] + + with request_redis_batch_scope(): + await proxy_server._apply_spend_counter_increments(pending) + assert client.pipelines == [] + await flush_post_call_redis_batches() + + assert [c for c in client.pipelines[0].commands if c[0] == "INCRBYFLOAT"] == [ + ("INCRBYFLOAT", "spend:key:k1", 0.5), + ("INCRBYFLOAT", "spend:team:t1", 0.5), + ] + assert spend_cache.in_memory_cache.get_cache("spend:key:k1") == 7.5 + + +@pytest.mark.asyncio +async def test_a_spend_counter_whose_increment_failed_is_invalidated_not_trusted(monkeypatch): + from litellm.proxy import proxy_server + + def replies(command: tuple[object, ...]) -> object: + if command[0] == "INCRBYFLOAT" and command[1] == "spend:key:k1": + return Exception("OOM") + return _ok_replies(command) + + redis_cache = PostCallFakeRedisCache(FakeClient(replies)) + spend_cache = DualCache() + spend_cache.attach_redis_cache(redis_cache) + spend_cache.in_memory_cache.set_cache("spend:key:k1", 3.0) + spend_cache.in_memory_cache.set_cache("spend:team:t1", 3.0) + monkeypatch.setattr(proxy_server, "spend_counter_cache", spend_cache) + + with request_redis_batch_scope(): + await proxy_server._apply_spend_counter_increments( + [PendingSpendIncrement("spend:key:k1", 0.5), PendingSpendIncrement("spend:team:t1", 0.5)] + ) + await flush_post_call_redis_batches() + + assert spend_cache.in_memory_cache.get_cache("spend:key:k1") is None + assert redis_cache.alone == [("DEL", "spend:key:k1")] + assert spend_cache.in_memory_cache.get_cache("spend:team:t1") == 7.5 + + +@pytest.mark.asyncio +async def test_a_cancelled_post_call_flush_keeps_the_shared_spend_counter_and_counts_the_spend_locally(monkeypatch): + from litellm.proxy import proxy_server + + redis_cache = PostCallFakeRedisCache( + FakeClient(_ok_replies, fail=asyncio.CancelledError()) # pyright: ignore[reportArgumentType] # a cancel raised mid-pipeline + ) + spend_cache = DualCache() + spend_cache.attach_redis_cache(redis_cache) + spend_cache.in_memory_cache.set_cache("spend:key:k1", 3.0) + monkeypatch.setattr(proxy_server, "spend_counter_cache", spend_cache) + + with request_redis_batch_scope(): + await proxy_server._apply_spend_counter_increments( + [PendingSpendIncrement("spend:key:k1", 0.5), PendingSpendIncrement("spend:team:t1", 0.5)] + ) + with pytest.raises(asyncio.CancelledError): + await flush_post_call_redis_batches() + + assert redis_cache.alone == [], "a cancel says nothing about the shared counter, so Redis keeps it" + assert spend_cache.in_memory_cache.get_cache("spend:key:k1") == 3.5, "the local copy counts the cancelled spend" + assert spend_cache.in_memory_cache.get_cache("spend:team:t1") is None, "an absent local copy is not seeded" + + +@pytest.mark.asyncio +async def test_the_update_cache_read_armed_before_accounting_rides_the_pipeline_of_the_reconcile_read(): + from litellm.proxy.proxy_server import _read_update_cache_values, arm_update_cache_read + + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + cache = DualCache() + cache.attach_redis_cache(redis_cache) + keys = ["user-1", "team_id:t1"] + + with request_redis_batch_scope() as request: + await arm_update_cache_read(keys, cache=cache) + assert client.pipelines == [] + await request.batch(redis_cache).mget(["spend:key:k1"]) # the spend reconcile read of the same request + values = await _read_update_cache_values(keys, None, cache=cache) + + assert len(client.pipelines) == 1 + assert client.pipelines[0].commands == [("MGET", "user-1", "team_id:t1"), ("MGET", "spend:key:k1")] + assert values == {"user-1": {"spend": 1.0}, "team_id:t1": {"spend": 1.0}} + assert redis_cache.alone == [] + assert active_request_redis_batches() is None + + +@pytest.mark.asyncio +async def test_an_update_cache_read_armed_for_other_keys_is_ignored_and_the_read_happens_as_before(): + from litellm.proxy.proxy_server import _read_update_cache_values, arm_update_cache_read + + redis_cache = PostCallFakeRedisCache(FakeClient(_ok_replies)) + redis_cache.store["team_id:t1"] = {"spend": 2.0} + cache = DualCache() + cache.attach_redis_cache(redis_cache) + + with request_redis_batch_scope(): + await arm_update_cache_read(["user-1"], cache=cache) + values = await _read_update_cache_values(["team_id:t1"], None, cache=cache) + + assert values == {"team_id:t1": {"spend": 2.0}} + assert ("MGET", ("team_id:t1",)) in redis_cache.alone + + +@pytest.mark.asyncio +async def test_the_update_cache_read_sees_a_cached_spend_written_while_the_spend_was_persisted(monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy.hooks.proxy_track_cost_callback import _update_database_and_spend_counters + + cached_user_spend = {"user-1": 1.0} + + def replies(command: tuple[object, ...]) -> object: + if command[0] == "MGET": + return [ + json.dumps({"spend": cached_user_spend[key]}) if key in cached_user_spend else b"0.5" + for key in command[1:] + ] + return _ok_replies(command) + + client = FakeClient(replies) + redis_cache = PostCallFakeRedisCache(client) + spend_cache = DualCache() + spend_cache.attach_redis_cache(redis_cache) + user_cache = DualCache() + user_cache.attach_redis_cache(redis_cache) + monkeypatch.setattr(proxy_server, "spend_counter_cache", spend_cache) + monkeypatch.setattr(proxy_server, "user_api_key_cache", user_cache) + + async def _read_on_the_request_pipeline_then_a_concurrent_callback_writes_the_user(**kwargs: object) -> bool: + request = active_request_redis_batches() + assert request is not None + await request.batch(redis_cache).mget(["key-object"]) + cached_user_spend["user-1"] = 5.0 + return True + + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( + side_effect=_read_on_the_request_pipeline_then_a_concurrent_callback_writes_the_user + ) + reservation = { + "reserved_cost": 0.5, + "entries": [ + { + "counter_key": "spend:key:k1", + "entity_type": "Key", + "entity_id": "k1", + "reserved_cost": 0.5, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with request_redis_batch_scope(): + charged = await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=proxy_server.increment_spend_counters, + user_api_key="k1", + user_id="user-1", + end_user_id=None, + team_id=None, + org_id=None, + kwargs={}, + completion_response=None, + start_time=datetime.datetime.now(), + end_time=datetime.datetime.now(), + response_cost=0.2, + budget_reservation=reservation, + update_cache_read_keys=("user-1",), + ) + values = await proxy_server._read_update_cache_values(("user-1",), None) + + assert charged is True + assert values == {"user-1": {"spend": 5.0}}, client.pipelines