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 <yassin@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-09-29 17:26:05 -07:00 • committed by GitHub
parent c129ea4fc9
commit 9525452d37
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 1091 additions and 36 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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