mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
c129ea4fc9
commit
9525452d37
9 changed files with 1091 additions and 36 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
661
tests/unit/caching/test_request_redis_batch_post_call.py
Normal file
661
tests/unit/caching/test_request_redis_batch_post_call.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue