From ffb15f946f586c102bf0359b2fa9ac46b340b658 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 16:42:13 -0700 Subject: [PATCH] perf(proxy): one request-scoped Redis pipeline for auth, spend, rate-limit and routing reads (#43407) RedisBatch: one pipeline per Redis backend for independently declared operations (MGET, GET, Lua scripts, INCRBYFLOAT, SET, DEL), a future per operation so each owner keeps its own fallback, Redis Cluster hash-slot fallback. A request-scoped batch middleware shares that pipeline across the auth identity reads and write-back, the spend counter MGET, the rate limiter Lua groups and the routing read. A rate-limit denial stands when another pipelined group fails; every pipelined group is refunded on rejection; local cooldowns win over the prefetch. The routing prefetch failure log line strips request line breaks (CodeQL py/log-injection) Resolves LIT-8882 Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/redis_batch.py | 422 ++++++++++++++ litellm/proxy/auth/auth_object_prefetch.py | 23 +- litellm/proxy/common_request_processing.py | 3 + .../hooks/parallel_request_limiter_v3.py | 146 ++++- .../redis_request_batch_middleware.py | 25 + litellm/proxy/proxy_server.py | 2 + .../spend_tracking/spend_counter_batch.py | 37 +- litellm/router.py | 25 +- litellm/router_utils/routing_read_batch.py | 114 +++- .../router_code_coverage.py | 1 + .../hooks/test_parallel_request_limiter_v3.py | 33 ++ tests/unit/caching/test_redis_batch.py | 299 ++++++++++ .../test_request_redis_batch_pre_call.py | 530 ++++++++++++++++++ tests/unit/test_router/test_router.py | 34 ++ 14 files changed, 1667 insertions(+), 27 deletions(-) create mode 100644 litellm/caching/redis_batch.py create mode 100644 litellm/proxy/middleware/redis_request_batch_middleware.py create mode 100644 tests/unit/caching/test_redis_batch.py create mode 100644 tests/unit/caching/test_request_redis_batch_pre_call.py diff --git a/litellm/caching/redis_batch.py b/litellm/caching/redis_batch.py new file mode 100644 index 00000000000..8bd9554ec55 --- /dev/null +++ b/litellm/caching/redis_batch.py @@ -0,0 +1,422 @@ +"""One Redis pipeline for several independent operations, each with its own result and its own failure. + +A ``RedisBatch`` collects MGETs, Lua scripts and increments declared by unrelated callers and sends them +in one ``pipeline(transaction=False)`` round trip. Every declaration returns an awaitable; awaiting one +flushes whatever has been declared so far, so callers keep their existing ``await`` shape and their own +error handling while sharing the wire. Redis Cluster clients run each operation on its own, as before: +a cluster pipeline is per node anyway and the existing per-operation paths already group by slot. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +import logging +import time +from collections.abc import Awaitable, Callable, Generator, Mapping, Sequence +from contextvars import ContextVar, Token +from dataclasses import dataclass, field +from datetime import timedelta +from types import MappingProxyType, TracebackType +from typing import Final, Generic, Protocol, TypeVar + +from litellm._logging import verbose_logger +from litellm.caching.redis_cache import ( + RedisCache, + _run_under_circuit_breaker, # pyright: ignore[reportPrivateUsage] # same health signal as every RedisCache method + log_redis_failure, +) +from litellm.caching.redis_cluster_cache import RedisClusterCache +from litellm.types.services import ServiceTypes + +_T = TypeVar("_T") +_ScriptArg = str | bytes | int | float + + +class RegisteredScript(Protocol): + def __call__(self, keys: Sequence[str], args: Sequence[_ScriptArg]) -> Awaitable[object]: ... + + +class _RedisPipeline(Protocol): + def mget(self, keys: Sequence[str]) -> object: ... + def evalsha(self, sha: str, numkeys: int, *keys_and_args: _ScriptArg) -> object: ... + def incrbyfloat(self, name: str, amount: float) -> object: ... + def expire(self, name: str, time: timedelta) -> object: ... + def set(self, name: str, value: str, ex: timedelta | None = None) -> object: ... + async def execute(self, raise_on_error: bool = True) -> list[object]: ... + + +class _Op(Generic[_T]): + """One declared operation: how many pipeline replies it consumes, how to turn them into a result, and + how to run on its own when the batch cannot pipeline (cluster client, or a reply the pipeline cannot + settle, like NOSCRIPT).""" + + __slots__ = ("future",) + + def __init__(self) -> None: + self.future: Final[asyncio.Future[_T]] = asyncio.get_running_loop().create_future() + self.future.add_done_callback(_mark_retrieved) + + def enqueue(self, pipe: _RedisPipeline) -> int: + raise NotImplementedError + + def resolve(self, replies: Sequence[object]) -> _T: + raise NotImplementedError + + async def run_alone(self) -> _T: + raise NotImplementedError + + def settle(self, replies: Sequence[object]) -> Awaitable[None] | None: + """Resolve from pipeline replies; return a coroutine when the op has to be retried on its own.""" + failure: Final = next((reply for reply in replies if isinstance(reply, Exception)), None) + if failure is None: + try: + self.future.set_result(self.resolve(replies)) + except Exception as e: # noqa: BLE001 # a reply this op cannot decode fails this op alone + self.future.set_exception(e) + return None + if _is_missing_script(failure): + return self._settle_alone() + self.future.set_exception(failure) + return None + + async def _settle_alone(self) -> None: + try: + self.future.set_result(await self.run_alone()) + except Exception as e: # noqa: BLE001 # the declaring caller owns the failure of its own operation + self.future.set_exception(e) + + +def _is_missing_script(failure: Exception) -> bool: + """Imported lazily: this module is reachable from a base ``import litellm`` while redis is not a base dependency.""" + from redis.exceptions import NoScriptError + + return isinstance(failure, NoScriptError) + + +def _mark_retrieved(future: asyncio.Future[object]) -> None: + """A caller that stops awaiting (cancelled request) must not leave an 'exception never retrieved' log.""" + if not future.cancelled(): + future.exception() + + +class _MGet(_Op[Mapping[str, object]]): + __slots__ = ("_keys", "_redis_cache") + + def __init__(self, redis_cache: RedisCache, keys: Sequence[str]) -> None: + super().__init__() + self._redis_cache: Final = redis_cache + self._keys: Final[tuple[str, ...]] = tuple(dict.fromkeys(keys)) + + def enqueue(self, pipe: _RedisPipeline) -> int: + pipe.mget(tuple(self._redis_cache.check_and_fix_namespace(key=key) for key in self._keys)) + return 1 + + def resolve(self, replies: Sequence[object]) -> Mapping[str, object]: + values: Final = replies[0] + if not isinstance(values, (list, tuple)): + raise TypeError(f"MGET reply is not a list: {type(values).__name__}") + return MappingProxyType( + {key: self._redis_cache._get_cache_logic(value) for key, value in zip(self._keys, values)} # pyright: ignore[reportPrivateUsage, reportUnknownMemberType, reportUnknownArgumentType] # shared decode with async_batch_get_cache + ) + + async def run_alone(self) -> Mapping[str, object]: + found: Mapping[str, object] = await self._redis_cache.async_batch_get_cache(key_list=list(self._keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API # mutable-ok: the cache API takes a list + if any(key not in found for key in self._keys): + raise ConnectionError("batch get did not return every key") + return found + + +class _Script(_Op[object]): + __slots__ = ("_args", "_keys", "_redis_cache", "_run", "_sha") + + def __init__( + self, + redis_cache: RedisCache, + source: str, + run: RegisteredScript, + keys: Sequence[str], + args: Sequence[_ScriptArg], + ) -> None: + super().__init__() + self._redis_cache: Final = redis_cache + self._sha: Final = hashlib.sha1(source.encode()).hexdigest() # noqa: S324 # EVALSHA identifies scripts by SHA-1 + self._run: Final = run + self._keys: Final[tuple[str, ...]] = tuple(keys) + self._args: Final[tuple[_ScriptArg, ...]] = tuple(args) + + def enqueue(self, pipe: _RedisPipeline) -> int: + namespaced: Final = tuple(self._redis_cache.check_and_fix_namespace(key=key) for key in self._keys) + pipe.evalsha(self._sha, len(namespaced), *namespaced, *self._args) + return 1 + + def resolve(self, replies: Sequence[object]) -> object: + return replies[0] + + async def run_alone(self) -> object: + return await self._run(keys=self._keys, args=self._args) + + +class _Increment(_Op[float]): + __slots__ = ("_key", "_redis_cache", "_ttl", "_value") + + def __init__(self, redis_cache: RedisCache, key: str, value: float, ttl: int | None) -> None: + super().__init__() + self._redis_cache: Final = redis_cache + self._key: Final = key + self._value: Final = value + self._ttl: Final = ttl + + def enqueue(self, pipe: _RedisPipeline) -> int: + name: Final = self._redis_cache.check_and_fix_namespace(key=self._key) + pipe.incrbyfloat(name, self._value) + if self._ttl is None: + return 1 + pipe.expire(name, timedelta(seconds=self._ttl)) + return 2 + + def resolve(self, replies: Sequence[object]) -> float: + reply: Final = replies[0] + if not isinstance(reply, (int, float, str, bytes)): + raise TypeError(f"INCRBYFLOAT reply is not numeric: {type(reply).__name__}") + return float(reply) + + async def run_alone(self) -> float: + value: object = await self._redis_cache.async_increment(key=self._key, value=self._value, ttl=self._ttl) # pyright: ignore[reportUnknownMemberType] # untyped cache API + if not isinstance(value, (int, float)): + raise TypeError(f"increment did not return a number: {type(value).__name__}") + return float(value) + + +class _Set(_Op[None]): + """SET with the cache's TTL rules, same encoding as ``async_set_cache_pipeline_with_ttls``.""" + + __slots__ = ("_key", "_redis_cache", "_ttl", "_value") + + def __init__(self, redis_cache: RedisCache, key: str, value: object, ttl: float | None) -> None: + super().__init__() + self._redis_cache: Final = redis_cache + self._key: Final = key + self._value: Final = value + self._ttl: Final = ttl + + def enqueue(self, pipe: _RedisPipeline) -> int: + ttl: Final = self._redis_cache.get_ttl(ttl=self._ttl) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API + pipe.set( + self._redis_cache.check_and_fix_namespace(key=self._key), + json.dumps(self._value), + ex=None if ttl is None else timedelta(seconds=ttl), + ) + return 1 + + def resolve(self, replies: Sequence[object]) -> None: + return None + + async def run_alone(self) -> None: + await self._redis_cache.async_set_cache_pipeline_with_ttls(((self._key, self._value, self._ttl),)) + + +class BatchResult(Generic[_T]): + """Awaitable handle for one declared operation; awaiting it flushes the batch it belongs to.""" + + __slots__ = ("_batch", "_op") + + def __init__(self, batch: RedisBatch, op: _Op[_T]) -> None: + self._batch: Final = batch + self._op: Final = op + + def __await__(self) -> Generator[object, None, _T]: + return self._wait().__await__() + + async def _wait(self) -> _T: + if not self._op.future.done(): + await self._batch.flush() + return self._op.future.result() + + @property + def done(self) -> bool: + return self._op.future.done() + + +@dataclass(slots=True) +class RedisBatch: + """Operations declared here go out in one pipeline the next time any of them is awaited or ``flush`` runs.""" + + redis_cache: RedisCache + name: str = "redis_batch" + _pending: list[_Op[object]] = field(default_factory=list) # mutable-ok: drained by flush + _flush_hooks: list[Callable[[], None]] = field(default_factory=list) # mutable-ok: append-only registry + _lock: asyncio.Lock = field(default_factory=asyncio.Lock) + flushes: int = 0 + + def mget(self, keys: Sequence[str]) -> BatchResult[Mapping[str, object]]: + return self._declare(_MGet(self.redis_cache, keys)) + + def script( + self, source: str, run: RegisteredScript, keys: Sequence[str], args: Sequence[_ScriptArg] + ) -> BatchResult[object]: + return self._declare(_Script(self.redis_cache, source, run, keys, args)) + + def increment(self, key: str, value: float, ttl: int | None = None) -> BatchResult[float]: + return self._declare(_Increment(self.redis_cache, key, value, ttl)) + + def set(self, key: str, value: object, ttl: float | None = None) -> BatchResult[None]: + return self._declare(_Set(self.redis_cache, key, value, ttl)) + + def add_flush_hook(self, hook: Callable[[], None]) -> None: + """Called at the start of every flush so lazily bound readers can declare their keys into the same trip.""" + self._flush_hooks.append(hook) + + @property + def pending(self) -> int: + return len(self._pending) + + def _declare(self, op: _Op[_T]) -> BatchResult[_T]: + self._pending.append(op) # pyright: ignore[reportArgumentType] # heterogeneous ops share the flush loop + return BatchResult(self, op) + + async def flush(self) -> None: + async with self._lock: + for hook in self._flush_hooks: + hook() + ops: Final = tuple(self._pending) + self._pending.clear() + if not ops: + return + self.flushes += 1 + try: + if isinstance(self.redis_cache, RedisClusterCache): + await asyncio.gather(*(op._settle_alone() for op in ops)) # pyright: ignore[reportPrivateUsage] # batch owns its ops + else: + await self._flush_pipeline(ops) + finally: + for op in ops: + if not op.future.done(): + op.future.cancel() + + async def _flush_pipeline(self, ops: Sequence[_Op[object]]) -> None: + start_time: Final = time.time() + widths: list[int] = [] # mutable-ok: filled while enqueuing + + async def run() -> list[object]: + client: Final = self.redis_cache.init_async_client() + async with client.pipeline(transaction=False) as pipe: + widths.extend(op.enqueue(pipe) for op in ops) + return await pipe.execute(raise_on_error=False) + + try: + replies: Final = await _run_under_circuit_breaker(self.redis_cache._circuit_breaker, self.name, run) # pyright: ignore[reportPrivateUsage] # same breaker as the cache's own methods + except Exception as e: # noqa: BLE001 # each declaring caller applies its own Redis fallback + log_redis_failure(verbose_logger, logging.WARNING, f"{self.name}: pipeline of {len(ops)} ops failed", e) + asyncio.create_task( + self.redis_cache.service_logger_obj.async_service_failure_hook( + service=ServiceTypes.REDIS, + duration=time.time() - start_time, + error=e, + call_type=f"{self.name}[{len(ops)}]", + start_time=start_time, + end_time=time.time(), + ) + ) + for op in ops: + op.future.set_exception(e) + return + asyncio.create_task( + self.redis_cache.service_logger_obj.async_service_success_hook( + service=ServiceTypes.REDIS, + duration=time.time() - start_time, + call_type=f"{self.name}[{len(ops)}]", + start_time=start_time, + end_time=time.time(), + ) + ) + retries: list[Awaitable[None]] = [] # mutable-ok: collected while slicing replies + offset = 0 + for op, width in zip(ops, widths): + retry = op.settle(replies[offset : offset + width]) + offset += width + if retry is not None: + retries.append(retry) + if retries: + await asyncio.gather(*retries) + + +def _backend_key(redis_cache: RedisCache) -> object: + """Two ``RedisCache`` instances built from the same connection settings and namespace talk to the same server + under the same key prefix, so the proxy's cache and the router's cache share one pipeline (the router gets its + port as a string, hence the ``str`` comparison); a cache whose settings cannot be compared (a test double) gets + its own.""" + try: + settings: Final = tuple(sorted((str(k), str(v)) for k, v in redis_cache.redis_kwargs.items() if v is not None)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType, reportUnknownArgumentType] # untyped cache API + except AttributeError: + return ("instance", id(redis_cache)) + return (type(redis_cache), redis_cache.namespace, settings) + + +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.""" + + __slots__ = ("_batches", "prefetched") + + def __init__(self) -> None: + self._batches: Final[dict[object, RedisBatch]] = {} # mutable-ok: lazily filled per backend + # 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 + + def batch(self, redis_cache: RedisCache) -> RedisBatch: + key: Final = _backend_key(redis_cache) + batch = self._batches.get(key) + if batch is None: + batch = RedisBatch(redis_cache, name="request_redis_batch") + self._batches[key] = batch + return batch + + 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)) + + @property + def batches(self) -> tuple[RedisBatch, ...]: + return tuple(self._batches.values()) + + +_active_request_batches: Final[ContextVar[RequestRedisBatches | None]] = ContextVar( + "request_redis_batches", default=None +) + + +def active_request_redis_batch(redis_cache: RedisCache) -> RedisBatch | None: + """The request's 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.batch(redis_cache) + + +def active_request_redis_batches() -> RequestRedisBatches | None: + return _active_request_batches.get() + + +class request_redis_batch_scope: + """Redis reads declared inside share one pipeline per backend; nested scopes join the outer one.""" + + __slots__ = ("_token",) + + def __init__(self) -> None: + self._token: Token[RequestRedisBatches | None] | None = None + + def __enter__(self) -> RequestRedisBatches: + outer: Final = _active_request_batches.get() + if outer is not None: + return outer + batches: Final = RequestRedisBatches() + self._token = _active_request_batches.set(batches) + return batches + + def __exit__( + self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: TracebackType | None + ) -> None: + if self._token is not None: + _active_request_batches.reset(self._token) diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py index 52e26e885c9..ce55190aa02 100644 --- a/litellm/proxy/auth/auth_object_prefetch.py +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -13,6 +13,7 @@ from typing import Final, Literal, Protocol, TypeAlias from pydantic import BaseModel, TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger +from litellm.caching.redis_batch import active_request_redis_batch from litellm.caching.redis_cache import RedisCache from litellm.constants import DEFAULT_IN_MEMORY_TTL from litellm.models.organization import LiteLLM_OrganizationTable @@ -218,11 +219,23 @@ def _set_in_memory(memory: _InMemoryCache, cache_key: str, value: object, ttl: f memory.set_cache(key=cache_key, value=value, ttl=ttl) +async def _read_redis_rows(keys: list[str], redis_cache: RedisCache) -> Mapping[str, object]: + """On the request pipeline when one is open; a failed pipeline reads as a miss, like ``async_batch_get_cache``.""" + batch: Final = active_request_redis_batch(redis_cache) + if batch is None: + return await redis_cache.async_batch_get_cache(key_list=keys) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API + try: + return await batch.mget(keys) + except Exception as e: # noqa: BLE001 # the DB fill below takes over, as it does after a failed MGET today + verbose_proxy_logger.debug("auth prefetch Redis read failed, filling from the database: %s", e) + return MappingProxyType({}) + + async def _fill_from_redis(entries: Sequence[_CacheEntry], redis_cache: RedisCache, memory: _InMemoryCache) -> None: if not entries: return found: Final = _RowValues.validate_python( - await redis_cache.async_batch_get_cache(key_list=sorted(entry.cache_key for entry in entries)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped cache API + await _read_redis_rows(sorted(entry.cache_key for entry in entries), redis_cache) ) for entry, value in ((entry, found.get(entry.cache_key)) for entry in entries): if value is not None: @@ -267,8 +280,14 @@ async def _write_back(entries: Sequence[tuple[_CacheEntry, BaseModel]], cache: U memory: Final[_InMemoryCache] = cache.in_memory_cache for cache_key, payload, ttl in payloads: _set_in_memory(memory, cache_key, payload, cache.default_in_memory_ttl if ttl is None else ttl) - if cache.redis_cache is not None: + if cache.redis_cache is None: + return + batch: Final = active_request_redis_batch(cache.redis_cache) + if batch is None: await cache.redis_cache.async_set_cache_pipeline_with_ttls(payloads) + return + for cache_key, payload, ttl in payloads: # rides the request's next round trip; the scope drains leftovers + batch.set(cache_key, payload, ttl) async def _fill_from_db( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 97de2488b8d..10724e9e7e6 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2199,6 +2199,9 @@ class ProxyBaseLLMRequestProcessing: if self._tags_before_guardrails is None: self._tags_before_guardrails = frozenset(get_tags_from_request_body(request_body=self.data)) + prefetch_model = self.data.get("model") + if llm_router is not None and isinstance(prefetch_model, str): + llm_router.arm_routing_read_prefetch(prefetch_model, self.data) self.data = await proxy_logging_obj.pre_call_hook( user_api_key_dict=user_api_key_dict, data=self.data, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index e4b782d5ff3..589aa7da7b5 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -33,6 +33,7 @@ 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_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 @@ -474,6 +475,19 @@ CacheCounterValue: TypeAlias = int | float | str | bytes CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None] + +def _as_counter_values(reply: object) -> list[CacheCounterValue]: + """A Lua reply read back off the pipeline is the same array the script returns when called directly.""" + if not isinstance(reply, (list, tuple)): + raise TypeError(f"rate limiter script reply is not a list: {type(reply).__name__}") + values: Final[list[CacheCounterValue]] = [] # mutable-ok: each element is narrowed before it is kept + for value in reply: # pyright: ignore[reportUnknownVariableType] # raw Redis reply + if not isinstance(value, (int, float, str, bytes)): + raise TypeError(f"rate limiter script reply holds {type(value).__name__}") # pyright: ignore[reportUnknownArgumentType] # raw Redis reply + values.append(value) + return values + + ReservationWindowIdentity: TypeAlias = tuple[str, str, Literal["redis", "local"]] ParallelGaugeCacheValue: TypeAlias = dict[str, object] | int | float | str | bytes @@ -1323,6 +1337,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): crc: Final = binascii.crc_hqx(key.encode("utf-8"), 0) return crc % REDIS_CLUSTER_SLOTS + def _pipeline_scripts( + self, + source: str, + run: RegisteredScript, + calls: Sequence[tuple[Sequence[str], Sequence[int]]], + ) -> tuple[BatchResult[object] | None, ...]: + """Declare one Lua call per group on the request's Redis batch, so all groups share one round trip + with whatever else the request declared (the routing read). Returns ``None`` per call when no batch + is open, and the caller runs the script directly as before.""" + redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache + batch: Final = None if redis_cache is None else active_request_redis_batch(redis_cache) + if batch is None: + return (None,) * len(calls) + return tuple(batch.script(source, run, keys, args) for keys, args in calls) + def _group_keys_by_hash_tag(self, keys: list[str]) -> dict[str, list[str]]: """ Group keys by their Redis hash tag to ensure cluster compatibility. @@ -1404,7 +1433,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) return await self._batch_get_counter_values(keys=keys, parent_otel_span=parent_otel_span, local_only=True) - def _reject_if_rate_limit_unverifiable(self, failed_operation: str, error: Exception) -> None: + def _reject_if_rate_limit_unverifiable(self, failed_operation: str, error: BaseException) -> None: if not self._fail_closed_resolver(): return log_redis_failure( @@ -1436,12 +1465,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): key_groups: Final = list(self._group_keys_by_hash_tag(keys_to_fetch).items()) all_cache_values: Final[list[CacheCounterValue | None]] = [] + args: Final = (now_int, self.window_size) + pipelined: Final = self._pipeline_scripts( + BATCH_RATE_LIMITER_SCRIPT, + self.batch_rate_limiter_script, + tuple((group_keys, args) for _tag, group_keys in key_groups), + ) - for index, (hash_tag, group_keys) in enumerate(key_groups): + for index, ((hash_tag, group_keys), group_result) in enumerate(zip(key_groups, pipelined)): try: - group_cache_values: CacheCounterValues = await self.batch_rate_limiter_script( - keys=group_keys, - args=[now_int, self.window_size], # Use integer timestamp + group_cache_values: CacheCounterValues = ( + await self.batch_rate_limiter_script(keys=group_keys, args=args) + if group_result is None + else _as_counter_values(await group_result) ) all_cache_values.extend(group_cache_values) except Exception as e: @@ -1450,6 +1486,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): await self._refund_counter_increments( self._counter_refunds_from_batch_values(applied_keys, all_cache_values) ) + await self._refund_later_pipelined_groups(key_groups[index + 1 :], pipelined[index + 1 :]) self._reject_if_rate_limit_unverifiable("batch_rate_limiter_script", e) log_redis_failure( verbose_proxy_logger, logging.WARNING, f"Redis Lua script failed for hash tag {hash_tag}", e @@ -1464,6 +1501,22 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return all_cache_values + async def _refund_later_pipelined_groups( + self, + key_groups: Sequence[tuple[str, list[str]]], + pipelined: Sequence[BatchResult[object] | None], + ) -> None: + """Groups declared on the request batch ran in the same round trip as the one that failed, so their + increments landed even though the loop never read them.""" + for (_tag, group_keys), group_result in zip(key_groups, pipelined): + if group_result is None: + continue + try: + group_values = _as_counter_values(await group_result) + except Exception: # noqa: BLE001 # a group that failed in Redis incremented nothing to refund + continue + await self._refund_counter_increments(self._counter_refunds_from_batch_values(group_keys, group_values)) + async def should_rate_limit( self, descriptors: Sequence[RateLimitDescriptor], @@ -2061,7 +2114,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): reservation_windows: Final[set[ReservationWindowIdentity]] = set() # mutable-ok: filled by the group loop raw: list[CacheCounterValue] - for _idx, (keys, args, meta) in enumerate(descriptor_groups): + pipelined: Final = self._pipeline_scripts( + CHECK_AND_INCREMENT_BY_N_SCRIPT, + self.check_and_increment_by_n_script, # pyright: ignore[reportArgumentType] # sole caller guards it is not None + tuple((keys, args) for keys, args, _meta in descriptor_groups), + ) + batched: Final = tuple(result for result in pipelined if result is not None) + if len(batched) == len(descriptor_groups): + return await self._settle_pipelined_descriptor_groups(descriptor_groups, batched, parent_otel_span) + + for keys, args, meta in descriptor_groups: try: raw = await self.check_and_increment_by_n_script( # pyright: ignore[reportOptionalCall] # sole caller guards it is not None keys=keys, @@ -2105,6 +2167,76 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): reservation_windows=frozenset(reservation_windows), ) + async def _settle_pipelined_descriptor_groups( + self, + descriptor_groups: list[DescriptorAtomicGroup], + results: Sequence[BatchResult[object]], + parent_otel_span: Span | None, + ) -> RateLimitResponse: + """Every group's Lua call left in one pipeline, so each group has already checked and incremented on + its own before any result is read. A failed or over-limit group therefore refunds every group that + incremented, after it as well as before it, where the one-at-a-time loop only unwinds the groups it ran. + A Redis denial stands even when another group failed: the in-memory fallback only replaces a verdict + Redis never gave.""" + replies: Final = await asyncio.gather(*results, return_exceptions=True) + responses: Final = tuple( + self._pipelined_group_response(reply, meta) + for reply, (_keys, _args, meta) in zip(replies, descriptor_groups) + ) + applied: Final[list[tuple[CounterRefund, ...]]] = [] # mutable-ok: filled by the group loop + statuses: Final[list[RateLimitStatus]] = [] # mutable-ok: filled by the group loop + reservation_windows: Final[set[ReservationWindowIdentity]] = set() # mutable-ok: filled by the group loop + for reply, response, (_keys, _args, meta) in zip(replies, responses, descriptor_groups): + if isinstance(response, BaseException) or response["overall_code"] != "OK": + continue + applied.append(self._counter_refunds_from_atomic_response(_as_counter_values(reply), meta)) + statuses.extend(response["statuses"]) + reservation_windows.update(response.get("reservation_windows", frozenset())) + + over_limit: Final = next( + (r for r in responses if not isinstance(r, BaseException) and r["overall_code"] == "OVER_LIMIT"), None + ) + if over_limit is not None: + await self._refund_applied_descriptor_groups(applied) + return over_limit + failure: Final = next((r for r in responses if isinstance(r, BaseException)), None) + if failure is not None: + await self._refund_applied_descriptor_groups(applied) + self._reject_if_rate_limit_unverifiable("check_and_increment_by_n_script", failure) + log_redis_failure( + verbose_proxy_logger, + logging.ERROR, + f"atomic_check_and_increment_by_n: Redis Lua execution failed ({type(failure).__name__}). Refunding " + f"{len(applied)} pipelined descriptors and falling back to in-memory enforcement, counters will " + f"diverge from Redis until window expires (window_size={self.window_size}s)", + failure, + ) + flat_meta: Final = tuple( + itertools.chain.from_iterable(group_meta for _k, _a, group_meta in descriptor_groups) + ) + async with self._check_and_increment_lock: + return await self._atomic_check_and_increment_in_memory( + per_counter_meta=flat_meta, + parent_otel_span=parent_otel_span, + ) + if len(responses) == 1 and not isinstance(responses[0], BaseException): + return responses[0] + return RateLimitResponse( + overall_code="OK", + statuses=statuses, + reservation_windows=frozenset(reservation_windows), + ) + + def _pipelined_group_response( + self, reply: object, per_counter_meta: list[AtomicCounterMeta] + ) -> RateLimitResponse | BaseException: + if isinstance(reply, BaseException): + return reply + try: + return self._build_atomic_response(_as_counter_values(reply), per_counter_meta) + except Exception as e: # noqa: BLE001 # a reply this group cannot read is that group's Lua failure + return e + async def _refund_applied_descriptor_groups( self, applied: Sequence[Sequence[CounterRefund]], @@ -2233,7 +2365,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _atomic_check_and_increment_in_memory( self, - per_counter_meta: list[AtomicCounterMeta], + per_counter_meta: Sequence[AtomicCounterMeta], parent_otel_span: Span | None = None, ) -> RateLimitResponse: """In-memory all-or-nothing check-and-increment. Caller holds lock. diff --git a/litellm/proxy/middleware/redis_request_batch_middleware.py b/litellm/proxy/middleware/redis_request_batch_middleware.py new file mode 100644 index 00000000000..bfb5f79a174 --- /dev/null +++ b/litellm/proxy/middleware/redis_request_batch_middleware.py @@ -0,0 +1,25 @@ +from typing import Final + +from starlette.types import ASGIApp, Receive, Scope, Send + +from litellm.caching.redis_batch import request_redis_batch_scope + +_REQUEST_SCOPES: Final = frozenset({"http", "websocket"}) + + +class RedisRequestBatchMiddleware: + """Opens the request's Redis batch scope so auth, admission and routing reads issued anywhere in the + request (dependencies, the endpoint, tasks it spawns) share one pipeline per Redis backend.""" + + def __init__(self, app: ASGIApp) -> None: + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] not in _REQUEST_SCOPES: + await self.app(scope, receive, send) + return + with request_redis_batch_scope() as batches: + try: + await self.app(scope, receive, send) + finally: + await batches.flush_all() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b4a497ea1e4..3f5e01ed0bc 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -681,6 +681,7 @@ from litellm.proxy.middleware.billable_request_metrics_middleware import ( from litellm.proxy.middleware.budget_reservation_release_middleware import ( BudgetReservationReleaseMiddleware, ) +from litellm.proxy.middleware.redis_request_batch_middleware import RedisRequestBatchMiddleware from litellm.proxy.plugin_routes import ( register_plugins_from_config, ) @@ -2417,6 +2418,7 @@ app.add_middleware( sink_factory=lambda: gateway_request_accumulator if prisma_client is not None else None, ) app.add_middleware(BudgetReservationReleaseMiddleware, release=release_unbound_budget_reservation) +app.add_middleware(RedisRequestBatchMiddleware) app.add_middleware(InFlightRequestsMiddleware) app.add_middleware(SecurityHeadersMiddleware) diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py index a6694895a27..ae24331c236 100644 --- a/litellm/proxy/spend_tracking/spend_counter_batch.py +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -10,6 +10,7 @@ from typing import Final from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger +from litellm.caching.redis_batch import BatchResult, RedisBatch, active_request_redis_batch from litellm.caching.redis_cache import RedisCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import ( @@ -30,9 +31,13 @@ class PendingSpendIncrement: class SpendCounterBatch: """Bound counters are read with one MGET on first use; counters bound later join the next MGET. ``async_batch_get_cache`` maps a clean miss to ``None`` and drops keys only when Redis failed, so an absent - key means "read it yourself" and a present ``None`` is an authoritative miss.""" + key means "read it yourself" and a present ``None`` is an authoritative miss. - __slots__ = ("_fetched", "_keys", "_loaded", "_lock", "_open", "_redis_cache") + Inside a ``request_redis_batch_scope`` the MGET rides the request's pipeline instead: the batch's flush + hook declares whatever is bound but unread, so whoever flushes first (the auth object prefetch, usually) + carries the spend counters in the same round trip.""" + + __slots__ = ("_fetched", "_inflight", "_keys", "_loaded", "_lock", "_open", "_redis_cache", "_request_batch") def __init__(self, redis_cache: RedisCache) -> None: self._redis_cache: Final = redis_cache @@ -41,6 +46,10 @@ class SpendCounterBatch: self._keys: frozenset[str] = frozenset() self._fetched: frozenset[str] = frozenset() self._loaded: Mapping[str, float | None] = _NO_VALUES + self._inflight: Final[list[BatchResult[Mapping[str, object]]]] = [] # mutable-ok: drained by _load + self._request_batch: Final[RedisBatch | None] = active_request_redis_batch(redis_cache) + if self._request_batch is not None: + self._request_batch.add_flush_hook(self._declare_pending) @property def counter_keys(self) -> frozenset[str]: @@ -85,6 +94,10 @@ class SpendCounterBatch: async def _load(self) -> Mapping[str, float | None]: async with self._lock: + if self._request_batch is not None: + self._declare_pending() + await self._collect_inflight() + return self._loaded pending: Final = self._keys - self._fetched if pending: self._fetched = self._fetched | pending @@ -92,6 +105,26 @@ class SpendCounterBatch: self._loaded = MappingProxyType({**fetched, **self._loaded}) return self._loaded + def _declare_pending(self) -> None: + """Flush hook: put every bound-but-unread counter on the request pipeline that is about to go out.""" + if self._request_batch is None or not self._open: + return + pending: Final = self._keys - self._fetched + if pending: + self._fetched = self._fetched | pending + self._inflight.append(self._request_batch.mget(sorted(pending))) + + async def _collect_inflight(self) -> None: + results: Final = tuple(self._inflight) + self._inflight.clear() + for result in results: + try: + fetched: Mapping[str, float | None] = _CounterValues.validate_python(await result) + except Exception as e: # noqa: BLE001 # per-key reads take over and apply their own Redis fallback + verbose_proxy_logger.debug("spend counter batch read failed, falling back to per-key reads: %s", e) + continue + self._loaded = MappingProxyType({**fetched, **self._loaded}) + async def _fetch(self, keys: frozenset[str]) -> Mapping[str, float | None]: try: return _CounterValues.validate_python( diff --git a/litellm/router.py b/litellm/router.py index c18dfea1a36..8aaf58d5a3e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -259,7 +259,7 @@ from litellm.router_utils.routing_groups import ( parse_routing_groups, validate_routing_strategy, ) -from litellm.router_utils.routing_read_batch import RoutingReadBatch +from litellm.router_utils.routing_read_batch import RoutingPrefetch, RoutingReadBatch from litellm.scheduler import FlowItem, Scheduler from litellm.types.litellm_params import RoutingStrategyName from litellm.types.llms.openai import ( @@ -539,6 +539,10 @@ def _is_retriable_anthropic_status(status_code: int) -> bool: return status_code == 429 or status_code >= 500 +def _without_line_breaks(value: object) -> str: + return str(value).replace("\r", "").replace("\n", "") + + def _anthropic_stream_error_is_gateway_verdict(chunk: object) -> bool: """AgenticAnthropicStreamingIterator's own retrieval-failure frame is the gateway's verdict, not a provider failure: another deployment would rerun the same failed hook, so it reaches the client instead of falling back.""" @@ -1730,6 +1734,25 @@ class Router: normalized for normalized in map(self._normalize_strategy, configured) if normalized is not None ) + def arm_routing_read_prefetch(self, model: str, request_kwargs: dict[str, object] | None = None) -> None: + """Declare the cooldown read (and, for usage-based routing, the usage read) that + `async_get_available_deployment` will make for `model` on the request's Redis batch, so admission's + flush carries it. A miss (alias, no batch) costs nothing: routing then reads as it always has.""" + try: + strategy, selector = self._get_routing_context(model, request_kwargs) + usage_selector: Final = ( + selector + if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2) + else None + ) + deployments: Final = self.get_model_list(model_name=model) + if deployments: + RoutingPrefetch.arm(self, usage_selector, deployments) + except Exception as e: # noqa: BLE001 # a prefetch is an optimisation, never a reason to fail the request + verbose_router_logger.debug( + "routing read prefetch not armed for %s: %s", _without_line_breaks(model), _without_line_breaks(e) + ) + def _get_routing_context( self, model: str, request_kwargs: dict | None = None ) -> tuple[str | None, RouterStrategySelector | None]: diff --git a/litellm/router_utils/routing_read_batch.py b/litellm/router_utils/routing_read_batch.py index e9410d0e586..4039d7b1508 100644 --- a/litellm/router_utils/routing_read_batch.py +++ b/litellm/router_utils/routing_read_batch.py @@ -8,10 +8,15 @@ different objects. `RoutingReadBatch` fetches both key sets in one the usage slice to the strategy, so selection does not read again. """ +import itertools +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_router_logger from litellm.caching.dual_cache import DualCache +from litellm.caching.redis_batch import BatchResult, active_request_redis_batches from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage from litellm.router_utils.cooldown_cache import CooldownCache @@ -27,16 +32,68 @@ else: Span = Any +_PREFETCH_SLOT: Final = "routing_read" + + +@dataclass(frozen=True, slots=True) +class RoutingPrefetch: + """The cooldown and usage keys of a model group, declared on the request's Redis batch before admission + flushes it, so the routing read rides the same round trip as the rate limiter's Lua calls.""" + + keys: frozenset[str] + result: BatchResult[Mapping[str, object]] + + @staticmethod + def arm( + litellm_router_instance: LitellmRouter, + usage_selector: LowestTPMLoggingHandler_v2 | None, + deployments: list, + ) -> None: + request: Final = active_request_redis_batches() + redis_cache: Final = litellm_router_instance.cache.redis_cache + if request is None or redis_cache is None or _PREFETCH_SLOT in request.prefetched: + return + cooldown_keys: Final = tuple( + CooldownCache.get_cooldown_cache_key(model_id) for model_id in litellm_router_instance.get_model_ids() + ) + usage_keys: Final = ( + () if usage_selector is None else tuple(itertools.chain(*usage_selector.usage_counter_keys(deployments))) + ) + keys: Final = (*cooldown_keys, *usage_keys) + request.prefetched[_PREFETCH_SLOT] = RoutingPrefetch( + keys=frozenset(keys), result=request.batch(redis_cache).mget(keys) + ) + + @staticmethod + def armed() -> bool: + request: Final = active_request_redis_batches() + return request is not None and _PREFETCH_SLOT in request.prefetched + + @staticmethod + def take(needed: Sequence[str]) -> "RoutingPrefetch | None": + """The armed prefetch when it covers every key this read needs; taken once, so a retry reads fresh.""" + request: Final = active_request_redis_batches() + if request is None: + return None + armed: Final = request.prefetched.pop(_PREFETCH_SLOT, None) + if isinstance(armed, RoutingPrefetch) and armed.keys.issuperset(needed): + return armed + return None + + class RoutingReadBatch: - def __init__(self, usage_selector: LowestTPMLoggingHandler_v2) -> None: + def __init__(self, usage_selector: LowestTPMLoggingHandler_v2 | None) -> None: self.usage_selector: Final = usage_selector self.prefetched_usage: PrefetchedUsage | None = None @staticmethod def for_strategy(strategy: str | None, selector: object) -> "RoutingReadBatch | None": + """Usage-based routing reads its counters with the cooldown state; every other strategy reads only the + cooldown state, and only through this batch when the request armed a prefetch for it. Otherwise the + router's plain cooldown read stays in charge.""" if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2): return RoutingReadBatch(usage_selector=selector) - return None + return RoutingReadBatch(usage_selector=None) if RoutingPrefetch.armed() else None async def async_get_cooldown_deployments( self, @@ -50,23 +107,50 @@ class RoutingReadBatch: """ model_ids: Final = litellm_router_instance.get_model_ids() cooldown_keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids] - tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments) - usage_keys: Final = tpm_keys + rpm_keys - - cooldown_results, usage_values = await DualCache.async_batch_get_cache_shared( - [ - (litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys), - (self.usage_selector.router_cache, usage_keys), - ], - parent_otel_span=parent_otel_span, - ) - self.prefetched_usage = PrefetchedUsage( - keys=frozenset(usage_keys), - values=None if usage_values is None else dict(zip(usage_keys, usage_values)), + reads: Final[list[tuple[DualCache, list[str]]]] = [ # mutable-ok: the usage read is appended below + (litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys) + ] + usage_keys: list[str] = [] # mutable-ok: DualCache batch reads take a list + if self.usage_selector is not None: + tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments) + usage_keys = tpm_keys + rpm_keys + reads.append((self.usage_selector.router_cache, usage_keys)) + results: Final = await self._read_prefetched(reads) or await DualCache.async_batch_get_cache_shared( + reads, parent_otel_span=parent_otel_span ) + cooldown_results: Final = results[0] + if self.usage_selector is not None: + usage_values: Final = results[1] + self.prefetched_usage = PrefetchedUsage( + keys=frozenset(usage_keys), + values=None if usage_values is None else MappingProxyType(dict(zip(usage_keys, usage_values))), + ) cooldown_models: Final = litellm_router_instance.cooldown_cache.active_cooldowns_from_results( model_ids, cooldown_results ) verbose_router_logger.debug("retrieve cooldown models: %s", cooldown_models) return [model_id for model_id, _ in cooldown_models] + + @staticmethod + async def _read_prefetched( + reads: list[tuple[DualCache, list[str]]], + ) -> list[list[object | None] | None] | None: + """Serve the reads from the request's armed `RoutingPrefetch`, backfilling each cache's memory tier as + its own batch read would. None when nothing usable was armed or the prefetch failed.""" + prefetch: Final = RoutingPrefetch.take(tuple(itertools.chain.from_iterable(keys for _, keys in reads))) + if prefetch is None: + return None + try: + values: Final = await prefetch.result + except Exception as e: # noqa: BLE001 # the shared read below applies the caches' own Redis fallback + verbose_router_logger.debug("routing prefetch failed, reading again: %s", e) + return None + results: Final[list[list[object | None] | None]] = [] # mutable-ok: filled per read below + for cache, keys in reads: + pending = await cache._prepare_batch_get(keys, local_only=True) # pyright: ignore[reportPrivateUsage] # same two-step read as async_batch_get_cache_shared + missed = { # mutable-ok: _apply_batch_get takes a dict + key: values.get(key) for key, local in zip(keys, pending.result) if local is None + } + results.append(await cache._apply_batch_get(pending, missed)) # pyright: ignore[reportPrivateUsage] # same two-step read as async_batch_get_cache_shared + return results diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 7c247ae3303..df149f6c56a 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -90,6 +90,7 @@ ignored_function_names = [ "_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py "_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py "_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name) + "arm_routing_read_prefetch", # Tested in tests/unit/caching/test_request_redis_batch_pre_call.py (file lacks "router" in name) ] diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 9aff2636c42..3be8501f5da 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -6974,6 +6974,39 @@ async def test_batch_increment_refunds_counters_already_applied_when_a_later_clu assert redis.increments == [] +@pytest.mark.parametrize("fail_closed", [True, False], ids=["fail_closed", "fail_open"]) +@pytest.mark.asyncio +async def test_batch_increment_refunds_pipelined_groups_declared_after_the_one_that_failed(fail_closed): + from unittest.mock import patch + + redis = _ScriptedRedis() + handler = _handler_with_redis(redis, fail_closed=fail_closed) + now = int(time.time()) + groups = {"a": ["{a}:window", "{a}:requests"], "b": ["{b}:window", "{b}:requests"]} + loop = asyncio.get_running_loop() + failed_group = loop.create_future() + failed_group.set_exception(ConnectionError("Error 61 connecting to 127.0.0.1:6379. Connection refused.")) + landed_group = loop.create_future() + landed_group.set_result([now, 1]) + + with ( + patch.object(handler, "_group_keys_by_hash_tag", return_value=groups), + patch.object(handler, "_pipeline_scripts", return_value=[failed_group, landed_group]), + ): + if fail_closed: + with pytest.raises(HTTPException) as exc: + await handler._execute_redis_batch_rate_limiter_script( + keys_to_fetch=[*groups["a"], *groups["b"]], now_int=now + ) + assert exc.value.status_code == 503 + else: + await handler._execute_redis_batch_rate_limiter_script( + keys_to_fetch=[*groups["a"], *groups["b"]], now_int=now + ) + + assert redis.guarded_increments == ([(groups["b"], [str(now), -1, 0])] if fail_closed else []) + + @pytest.mark.parametrize( "limits, request_data, counter_scope", [ diff --git a/tests/unit/caching/test_redis_batch.py b/tests/unit/caching/test_redis_batch.py new file mode 100644 index 00000000000..9433aeac524 --- /dev/null +++ b/tests/unit/caching/test_redis_batch.py @@ -0,0 +1,299 @@ +"""RedisBatch: independent operations share one pipeline, each keeps its own result and failure.""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +from collections.abc import Callable, Sequence +from datetime import timedelta +from typing import Any + +import pytest +from redis.exceptions import NoScriptError + +from litellm._service_logger import ServiceLogging +from litellm.caching.redis_batch import ( + RedisBatch, + active_request_redis_batch, + request_redis_batch_scope, +) +from litellm.caching.redis_cache import RedisCache, RedisCircuitBreaker +from litellm.caching.redis_cluster_cache import RedisClusterCache + +SCRIPT = "return redis.call('GET', KEYS[1])" +SHA = hashlib.sha1(SCRIPT.encode()).hexdigest() # noqa: S324 + + +class FakePipeline: + def __init__(self, reply_for: Callable[[tuple[object, ...]], object], fail: Exception | None) -> None: + self.commands: list[tuple[Any, ...]] = [] + self.reply_for = reply_for + self.fail = fail + self.executed = False + + async def __aenter__(self) -> FakePipeline: + return self + + async def __aexit__(self, *exc: object) -> None: + return None + + def mget(self, keys: Sequence[str]) -> FakePipeline: + self.commands.append(("MGET", *keys)) + return self + + def evalsha(self, sha: str, numkeys: int, *keys_and_args: object) -> FakePipeline: + self.commands.append(("EVALSHA", sha, numkeys, *keys_and_args)) + return self + + def incrbyfloat(self, name: str, amount: float) -> FakePipeline: + self.commands.append(("INCRBYFLOAT", name, amount)) + return self + + def expire(self, name: str, time: timedelta) -> FakePipeline: + self.commands.append(("EXPIRE", name, int(time.total_seconds()))) + return self + + def set(self, name: str, value: str, ex: timedelta | None = None) -> FakePipeline: + self.commands.append(("SET", name, value, None if ex is None else int(ex.total_seconds()))) + return self + + async def execute(self, raise_on_error: bool = True) -> list[Any]: + assert raise_on_error is False + self.executed = True + if self.fail is not None: + raise self.fail + return [self.reply_for(command) for command in self.commands] + + +class FakeClient: + def __init__(self, reply_for: Callable[[tuple[object, ...]], object], fail: Exception | None = None) -> None: + self.pipelines: list[FakePipeline] = [] + self.reply_for = reply_for + self.fail = fail + + def pipeline(self, transaction: bool = True) -> FakePipeline: + assert transaction is False + pipe = FakePipeline(self.reply_for, self.fail) + self.pipelines.append(pipe) + return pipe + + +class FakeRedisCache(RedisCache): + def __init__(self, client: FakeClient, namespace: str | None = None) -> None: # super().__init__ needs a server + self.client = client + self.namespace = namespace + self._circuit_breaker = RedisCircuitBreaker(failure_threshold=5, recovery_timeout=30) + self.service_logger_obj = ServiceLogging() + self.default_ttl = None + self.alone: list[tuple[str, Any]] = [] + self.store: dict[str, Any] = {} + + def init_async_client(self) -> FakeClient: # pyright: ignore[reportIncompatibleMethodOverride] # fake client, no server + return self.client + + async def async_batch_get_cache(self, key_list: Sequence[str], **kwargs: object) -> dict[str, Any]: # pyright: ignore[reportIncompatibleMethodOverride] # records the direct read + self.alone.append(("MGET", tuple(key_list))) + return {key: self.store.get(key) for key in key_list} + + async def async_increment(self, key: str, value: float, ttl: int | None = None, **kwargs: object) -> float: # pyright: ignore[reportIncompatibleMethodOverride] # records the direct write + self.alone.append(("INCRBYFLOAT", key, value)) + self.store[key] = float(self.store.get(key, 0.0)) + value + return self.store[key] + + async def async_set_cache_pipeline_with_ttls(self, cache_list: Sequence[tuple[str, object, float | None]]) -> None: + self.alone.append(("SET_PIPELINE", tuple(cache_list))) + for key, value, _ttl in cache_list: + self.store[key] = value + + +class FakeClusterCache(RedisClusterCache, FakeRedisCache): + def __init__(self, client: FakeClient) -> None: # super().__init__ needs a server + FakeRedisCache.__init__(self, client) + + +def replies(command: tuple[Any, ...]) -> Any: + match command[0]: + case "MGET": + return [json.dumps({"k": key}) if key.endswith("hit") else None for key in command[1:]] + case "EVALSHA": + return [1, 2] + case "INCRBYFLOAT": + return b"3.5" + case "EXPIRE": + return 1 + case "SET": + return True + raise AssertionError(command) + + +def make(fail: Exception | None = None, namespace: str | None = None) -> tuple[FakeRedisCache, FakeClient]: + client = FakeClient(replies, fail) + return FakeRedisCache(client, namespace), client + + +async def run_alone_script(keys: Sequence[str], args: Sequence[Any]) -> object: + return ["alone", *keys, *args] + + +@pytest.mark.asyncio +async def test_one_pipeline_carries_every_declared_operation_and_awaiting_one_flushes_all() -> None: + cache, client = make(namespace="ns") + batch = RedisBatch(cache) + got = batch.mget(["a:hit", "b", "a:hit"]) + script = batch.script(SCRIPT, run_alone_script, ["w"], [7, "x"]) + incr = batch.increment("cnt", 2.5, ttl=60) + plain = batch.increment("cnt2", 1) + assert client.pipelines == [] + + assert await got == {"a:hit": {"k": "ns:a:hit"}, "b": None} + assert script.done and incr.done and plain.done + assert await script == [1, 2] + assert await incr == 3.5 + assert await plain == 3.5 + assert batch.flushes == 1 + assert [pipe.commands for pipe in client.pipelines] == [ + [ + ("MGET", "ns:a:hit", "ns:b"), + ("EVALSHA", SHA, 1, "ns:w", 7, "x"), + ("INCRBYFLOAT", "ns:cnt", 2.5), + ("EXPIRE", "ns:cnt", 60), + ("INCRBYFLOAT", "ns:cnt2", 1), + ] + ] + assert cache.alone == [] + + +@pytest.mark.asyncio +async def test_operations_declared_after_a_flush_go_out_in_the_next_pipeline() -> None: + cache, client = make() + batch = RedisBatch(cache) + await batch.mget(["a"]) + later = batch.increment("cnt", 1) + assert not later.done + assert await later == 3.5 + assert batch.flushes == 2 + assert [pipe.commands for pipe in client.pipelines] == [[("MGET", "a")], [("INCRBYFLOAT", "cnt", 1)]] + + +@pytest.mark.asyncio +async def test_a_failing_reply_fails_only_its_own_operation() -> None: + def reply_for(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA": + return ValueError("script blew up") + return replies(command) + + client = FakeClient(reply_for) + cache = FakeRedisCache(client) + batch = RedisBatch(cache) + got = batch.mget(["a:hit"]) + script = batch.script(SCRIPT, run_alone_script, ["w"], []) + assert await got == {"a:hit": {"k": "a:hit"}} + with pytest.raises(ValueError, match="script blew up"): + await script + assert cache.alone == [] + + +@pytest.mark.asyncio +async def test_a_reply_an_operation_cannot_decode_fails_only_that_operation() -> None: + def reply_for(command: tuple[Any, ...]) -> Any: + if command[0] == "MGET": + return "not-a-list" + return replies(command) + + client = FakeClient(reply_for) + cache = FakeRedisCache(client) + batch = RedisBatch(cache) + got = batch.mget(["a:hit"]) + written = batch.set("w", {"k": 1}) + script = batch.script(SCRIPT, run_alone_script, ["w"], []) + with pytest.raises(TypeError, match="MGET reply is not a list"): + await got + assert await written is None + assert await script == [1, 2] + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_pipeline_failure_fails_every_operation_and_trips_the_breaker() -> None: + cache, _client = make(fail=ConnectionError("redis down")) + batch = RedisBatch(cache) + got = batch.mget(["a"]) + incr = batch.increment("cnt", 1) + with pytest.raises(ConnectionError): + await got + with pytest.raises(ConnectionError): + await incr + assert cache._circuit_breaker._failure_count == 1 # pyright: ignore[reportPrivateUsage] + + +@pytest.mark.asyncio +async def test_noscript_reply_reruns_that_script_through_the_registered_executor() -> None: + def reply_for(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA": + return NoScriptError("NOSCRIPT No matching script") + return replies(command) + + client = FakeClient(reply_for) + cache = FakeRedisCache(client) + batch = RedisBatch(cache) + script = batch.script(SCRIPT, run_alone_script, ["w"], [1]) + incr = batch.increment("cnt", 1) + assert await script == ["alone", "w", 1] + assert await incr == 3.5 + assert batch.flushes == 1 + + +@pytest.mark.asyncio +async def test_cluster_cache_runs_each_operation_on_its_own_path() -> None: + client = FakeClient(replies) + cache = FakeClusterCache(client) + cache.store["a"] = 4 + batch = RedisBatch(cache) + got = batch.mget(["a", "b"]) + incr = batch.increment("cnt", 2) + assert await got == {"a": 4, "b": None} + assert await incr == 2.0 + assert client.pipelines == [] + assert cache.alone == [("MGET", ("a", "b")), ("INCRBYFLOAT", "cnt", 2)] + + +@pytest.mark.asyncio +async def test_flush_hook_lets_a_lazy_reader_join_the_pipeline_that_is_going_out() -> None: + cache, client = make() + batch = RedisBatch(cache) + joined: list[Any] = [] + batch.add_flush_hook(lambda: joined.append(batch.mget(["late"]))) + await batch.mget(["early"]) + assert len(joined) == 1 and joined[0].done + assert await joined[0] == {"late": None} + assert [pipe.commands for pipe in client.pipelines] == [[("MGET", "early"), ("MGET", "late")]] + + +@pytest.mark.asyncio +async def test_concurrent_awaiters_share_one_flush() -> None: + cache, client = make() + batch = RedisBatch(cache) + first = batch.mget(["a"]) + second = batch.mget(["b"]) + results = await asyncio.gather(first._wait(), second._wait()) # pyright: ignore[reportPrivateUsage] + assert results == [{"a": None}, {"b": None}] + assert batch.flushes == 1 + assert len(client.pipelines) == 1 + + +def test_request_scope_hands_out_one_batch_per_backend_and_nests() -> None: + cache_a, _ = make() + cache_b, _ = make() + assert active_request_redis_batch(cache_a) is None + with request_redis_batch_scope() as batches: + first = active_request_redis_batch(cache_a) + assert first is not None + assert active_request_redis_batch(cache_a) is first + assert active_request_redis_batch(cache_b) is not first + with request_redis_batch_scope() as inner: + assert inner is batches + assert active_request_redis_batch(cache_a) is first + assert active_request_redis_batch(cache_a) is first + assert len(batches.batches) == 2 + assert active_request_redis_batch(cache_a) is None diff --git a/tests/unit/caching/test_request_redis_batch_pre_call.py b/tests/unit/caching/test_request_redis_batch_pre_call.py new file mode 100644 index 00000000000..c0834974f26 --- /dev/null +++ b/tests/unit/caching/test_request_redis_batch_pre_call.py @@ -0,0 +1,530 @@ +"""One Redis pipeline per backend for the pre-call reads a request makes: rate limiter Lua groups, the +router's cooldown and usage read, auth identity and spend counters all join the request batch.""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +from typing import Any, Final +from unittest.mock import AsyncMock + +import pytest + +from litellm import Router +from litellm.caching.dual_cache import DualCache +from litellm.caching.redis_batch import active_request_redis_batches, request_redis_batch_scope +from litellm.proxy._types import LiteLLM_UserTable +from litellm.proxy.auth.auth_object_prefetch import _CacheEntry, _write_back +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + CHECK_AND_INCREMENT_BY_N_SCRIPT, + RateLimitDescriptor, + RateLimitUnverifiableError, + _PROXY_MaxParallelRequestsHandler_v3, +) +from litellm.proxy.utils import InternalUsageCache +from litellm.router_utils.cooldown_cache import CooldownCache +from litellm.router_utils.routing_read_batch import RoutingPrefetch + +from .test_redis_batch import FakeClient, FakeRedisCache + +_MODEL_GROUP = "claude" +_FAR_FUTURE = 4_102_444_800.0 # 2100-01-01, a cooldown stamped then is still active + + +def sha_of(script: str) -> str: + return hashlib.sha1(script.encode()).hexdigest() # noqa: S324 + + +def _limiter(redis_cache: FakeRedisCache, fail_closed: bool = False) -> _PROXY_MaxParallelRequestsHandler_v3: + dual_cache = DualCache() + limiter = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=InternalUsageCache(dual_cache=dual_cache), + fail_closed_resolver=lambda: fail_closed, + ) + dual_cache.attach_redis_cache(redis_cache) # after init: the fake has no server to register scripts on + limiter.check_and_increment_by_n_script = AsyncMock( + side_effect=AssertionError("descriptor groups must ride the request pipeline") + ) + limiter.window_guarded_token_increment_script = AsyncMock(return_value=[1, 0]) + return limiter + + +def _descriptor(key: str, value: str, rpm: int) -> RateLimitDescriptor: + return {"key": key, "value": value, "rate_limit": {"requests_per_unit": rpm}} + + +def _refunds(limiter: _PROXY_MaxParallelRequestsHandler_v3) -> list[tuple[str, float]]: + refund_script = limiter.window_guarded_token_increment_script + assert isinstance(refund_script, AsyncMock) + return [(call.kwargs["keys"][1], call.kwargs["args"][1]) for call in refund_script.await_args_list] + + +def _lua_ok_replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA": + return [0, 1, 1700000000] # OK: one counter, new_counter=1, window_start + if command[0] == "MGET": + return [None for _ in command[1:]] + if command[0] == "SET": + return True + raise AssertionError(command) + + +@pytest.mark.asyncio +async def test_descriptor_lua_calls_share_one_pipeline_and_each_keeps_its_result(): + client = FakeClient(_lua_ok_replies) + limiter = _limiter(FakeRedisCache(client)) + descriptors = [ + _descriptor("api_key", "k1", 10), + _descriptor("model_per_key", "k1:gpt", 5), + _descriptor("team", "t1", 20), + ] + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=descriptors, + increments=[{"requests": 1}, {"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OK" + assert [s["descriptor_key"] for s in response["statuses"]] == ["api_key", "model_per_key", "team"] + assert len(client.pipelines) == 1 + evalshas = [c for c in client.pipelines[0].commands if c[0] == "EVALSHA"] + assert len(evalshas) == 3 + assert {c[1] for c in evalshas} == {sha_of(CHECK_AND_INCREMENT_BY_N_SCRIPT)} + assert [c[3] for c in evalshas] == ["{api_key:k1}:window", "{model_per_key:k1:gpt}:window", "{team:t1}:window"] + + +@pytest.mark.asyncio +async def test_an_over_limit_descriptor_in_the_pipeline_refunds_the_groups_that_were_applied(): + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA" and command[3] == "{team:t1}:window": + return [1, 1, 21, 20] # OVER_LIMIT on its first counter + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OVER_LIMIT" + assert response["statuses"][0]["descriptor_key"] == "team" + assert _refunds(limiter) == [("{api_key:k1}:requests", -1.0)] + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_an_over_limit_descriptor_also_refunds_the_groups_the_pipeline_incremented_after_it(): + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:window": + return [1, 1, 11, 10] # OVER_LIMIT on the first group; the later groups already incremented + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[ + _descriptor("api_key", "k1", 10), + _descriptor("team", "t1", 20), + _descriptor("model_per_key", "k1:gpt", 5), + ], + increments=[{"requests": 1}, {"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OVER_LIMIT" + assert response["statuses"][0]["descriptor_key"] == "api_key" + assert _refunds(limiter) == [("{team:t1}:requests", -1.0), ("{model_per_key:k1:gpt}:requests", -1.0)] + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_redis_denial_stands_when_another_pipelined_group_fails(): + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:window": + return [1, 1, 11, 10] # OVER_LIMIT + if command[0] == "EVALSHA" and command[3] == "{team:t1}:window": + return ValueError("script blew up") + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[ + _descriptor("api_key", "k1", 10), + _descriptor("team", "t1", 20), + _descriptor("model_per_key", "k1:gpt", 5), + ], + increments=[{"requests": 1}, {"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OVER_LIMIT" # not the in-memory fallback's verdict + assert response["statuses"][0]["descriptor_key"] == "api_key" + assert _refunds(limiter) == [("{model_per_key:k1:gpt}:requests", -1.0)] + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_one_failed_lua_group_refunds_the_other_pipelined_groups_and_falls_back_to_in_memory(): + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:window": + return ValueError("script blew up") + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OK" + assert len(response["statuses"]) == 2 # in-memory enforcement covered both descriptors + assert _refunds(limiter) == [("{team:t1}:requests", -1.0)] + assert len(client.pipelines) == 1 + + +@pytest.mark.parametrize( + "client, refunded", + [ + ( + FakeClient( + lambda command: ( + ValueError("script blew up") + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:window" + else _lua_ok_replies(command) + ) + ), + [("{team:t1}:requests", -1.0)], + ), + (FakeClient(_lua_ok_replies, fail=ConnectionError("redis down")), []), + ], + ids=["one_group_failed", "pipeline_failed"], +) +@pytest.mark.asyncio +async def test_fail_closed_rejects_when_a_pipelined_lua_group_cannot_be_verified( + client: FakeClient, refunded: list[tuple[str, float]] +): + limiter = _limiter(FakeRedisCache(client), fail_closed=True) + + with request_redis_batch_scope(), pytest.raises(RateLimitUnverifiableError) as exc: + await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert exc.value.status_code == 503 + assert _refunds(limiter) == refunded + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_pipeline_failure_refunds_nothing_and_falls_back_to_in_memory_enforcement(): + client = FakeClient(_lua_ok_replies, fail=ConnectionError("redis down")) + limiter = _limiter(FakeRedisCache(client)) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OK" + assert len(response["statuses"]) == 2 + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_without_a_request_scope_descriptor_groups_run_the_script_directly_as_before(): + client = FakeClient(_lua_ok_replies) + limiter = _limiter(FakeRedisCache(client)) + limiter.check_and_increment_by_n_script = AsyncMock(return_value=[0, 1, 1700000000]) + + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OK" + assert limiter.check_and_increment_by_n_script.await_count == 2 + assert client.pipelines == [] + + +def _deployment(deployment_id: str) -> dict: + return { + "model_name": _MODEL_GROUP, + "litellm_params": {"model": "anthropic/claude-x", "api_key": "test", "mock_response": "pong"}, + "model_info": {"id": deployment_id}, + } + + +def _router(redis_cache: FakeRedisCache, routing_strategy: str = "usage-based-routing-v2") -> Router: + router = Router(model_list=[_deployment("dep-a"), _deployment("dep-b")], routing_strategy=routing_strategy) + router._update_redis_cache(cache=redis_cache) + return router + + +@pytest.mark.asyncio +async def test_armed_routing_read_rides_the_admission_pipeline_and_routing_issues_no_read_of_its_own(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert len(client.pipelines) == 1 + commands = client.pipelines[0].commands + assert [c[0] for c in commands] == ["MGET", "EVALSHA", "EVALSHA"] + mget_keys = set(commands[0][1:]) + assert {CooldownCache.get_cooldown_cache_key("dep-a"), CooldownCache.get_cooldown_cache_key("dep-b")} <= mget_keys + assert any(":tpm:" in key for key in mget_keys) and any(":rpm:" in key for key in mget_keys) + assert redis_cache.alone == [] + + +@pytest.mark.asyncio +async def test_a_cooldown_recorded_locally_after_the_prefetch_left_still_excludes_its_deployment(): + expired = {"exception_received": "429", "status_code": "429", "timestamp": 0.0, "cooldown_time": 60} + + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "MGET": # Redis holds a stale cooldown for dep-b and nothing for dep-a + return [ + json.dumps(expired) if key == CooldownCache.get_cooldown_cache_key("dep-b") else None + for key in command[1:] + ] + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + cooldown_store = router.cooldown_cache.cooldown_store + assert cooldown_store.in_memory_cache is not None + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + cooldown_store.in_memory_cache.set_cache( + CooldownCache.get_cooldown_cache_key("dep-a"), + {"exception_received": "429", "status_code": "429", "timestamp": _FAR_FUTURE, "cooldown_time": 60}, + ) + picks = { + ( + await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + )["model_info"]["id"] + for _ in range(5) + } + + assert picks == {"dep-b"} + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_prefetch_that_does_not_cover_the_routing_keys_is_ignored_and_routing_reads_itself(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + + with request_redis_batch_scope() as request: + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + armed = request.prefetched["routing_read"] + assert isinstance(armed, RoutingPrefetch) + request.prefetched["routing_read"] = RoutingPrefetch(keys=frozenset({"other"}), result=armed.result) + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + assert request.prefetched == {} + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert len(redis_cache.alone) == 1 # the shared cooldown+usage read, one round trip as in P1 + + +@pytest.mark.asyncio +async def test_a_failed_prefetch_falls_back_to_the_shared_read(): + client = FakeClient(_lua_ok_replies, fail=ConnectionError("redis down")) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert len(redis_cache.alone) == 1 + + +@pytest.mark.asyncio +async def test_arming_outside_a_request_scope_is_a_no_op(): + redis_cache = FakeRedisCache(FakeClient(_lua_ok_replies)) + router = _router(redis_cache) + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + assert active_request_redis_batches() is None + + +@pytest.mark.asyncio +async def test_simple_shuffle_prefetches_only_its_cooldown_read_into_the_admission_pipeline(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache, routing_strategy="simple-shuffle") + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10)], + increments=[{"requests": 1}], + ) + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert len(client.pipelines) == 1 + commands = client.pipelines[0].commands + assert [c[0] for c in commands] == ["MGET", "EVALSHA"] + assert set(commands[0][1:]) == { + CooldownCache.get_cooldown_cache_key("dep-a"), + CooldownCache.get_cooldown_cache_key("dep-b"), + } + assert redis_cache.alone == [] + + shuffle = Router(model_list=[_deployment("dep-a")], routing_strategy="simple-shuffle") + shuffle._update_redis_cache(cache=redis_cache) + with request_redis_batch_scope() as request: + shuffle.arm_routing_read_prefetch(_MODEL_GROUP, {}) + armed = request.prefetched["routing_read"] + assert isinstance(armed, RoutingPrefetch) + assert armed.keys == {CooldownCache.get_cooldown_cache_key("dep-a")} # no usage counters for shuffle + + +@pytest.mark.asyncio +async def test_two_backends_flush_concurrently_one_pipeline_each(): + a_client, b_client = FakeClient(_lua_ok_replies), FakeClient(_lua_ok_replies) + a, b = FakeRedisCache(a_client), FakeRedisCache(b_client) + with request_redis_batch_scope() as request: + ra = request.batch(a).mget(["x", "y"]) + rb = request.batch(b).mget(["x"]) + await asyncio.gather(ra, rb) + assert len(a_client.pipelines) == 1 and len(b_client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_single_lua_group_rides_the_pipeline_with_the_armed_routing_read(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10)], + increments=[{"requests": 1}], + ) + await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + + assert len(client.pipelines) == 1 + assert [c[0] for c in client.pipelines[0].commands] == ["MGET", "EVALSHA"] + assert redis_cache.alone == [] + + +class _SameServerCache(FakeRedisCache): + def __init__(self, client: FakeClient, namespace: str | None = None, **redis_kwargs: object) -> None: + super().__init__(client, namespace) + self.redis_kwargs = redis_kwargs + + +@pytest.mark.asyncio +async def test_caches_built_from_the_same_connection_settings_share_the_request_pipeline(): + client = FakeClient(_lua_ok_replies) + proxy_cache = _SameServerCache(client, host="r", port=6379, db=0) + router_cache = _SameServerCache(FakeClient(_lua_ok_replies), port="6379", host="r", db=0, password=None) + other_cache = _SameServerCache(FakeClient(_lua_ok_replies), host="r", port=6380, db=0) + with request_redis_batch_scope() as request: + assert request.batch(proxy_cache) is request.batch(router_cache) + assert request.batch(proxy_cache) is not request.batch(other_cache) + a = request.batch(proxy_cache).mget(["a"]) + b = request.batch(router_cache).mget(["b"]) + await asyncio.gather(a, b) + assert len(client.pipelines) == 1 + assert [c[0] for c in client.pipelines[0].commands] == ["MGET", "MGET"] + + +@pytest.mark.asyncio +async def test_caches_on_one_server_with_different_namespaces_keep_their_own_key_prefix(): + proxy_client, router_client = FakeClient(_lua_ok_replies), FakeClient(_lua_ok_replies) + proxy_cache = _SameServerCache(proxy_client, namespace="proxy", host="r", port=6379, db=0) + router_cache = _SameServerCache(router_client, namespace="router", host="r", port=6379, db=0) + with request_redis_batch_scope() as request: + await asyncio.gather(request.batch(proxy_cache).mget(["a"]), request.batch(router_cache).mget(["b"])) + sent: Final = tuple( + tuple(command for pipe in client.pipelines for command in pipe.commands) + for client in (proxy_client, router_client) + ) + assert sent == ((("MGET", "proxy:a"),), (("MGET", "router:b"),)), "each cache reads under its own namespace" + + +def _user_entry() -> tuple[_CacheEntry, LiteLLM_UserTable]: + entry = _CacheEntry("user-1", "user_row", LiteLLM_UserTable, 42) + return entry, LiteLLM_UserTable(user_id="user-1", max_budget=None, spend=0.0) + + +@pytest.mark.asyncio +async def test_auth_write_back_rides_the_next_round_trip_and_the_scope_drains_what_nobody_awaited(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + cache = UserApiKeyCache(redis_cache=redis_cache) + with request_redis_batch_scope() as request: + await _write_back([_user_entry()], cache) + assert client.pipelines == [] # not sent yet: the SET waits for the next round trip + await request.batch(redis_cache).mget(["spend:key:k1"]) + assert len(client.pipelines) == 1 + kinds = [c[0] for c in client.pipelines[0].commands] + assert kinds == ["MGET", "SET"] or kinds == ["SET", "MGET"] + set_command = next(c for c in client.pipelines[0].commands if c[0] == "SET") + assert set_command[1] == "user-1" and set_command[3] == 42 + assert json.loads(set_command[2])["user_id"] == "user-1" + assert cache.in_memory_cache.get_cache("user-1") is not None + + await _write_back([_user_entry()], cache) + assert len(client.pipelines) == 1 + await request.flush_all() + assert len(client.pipelines) == 2 + assert [c[0] for c in client.pipelines[1].commands] == ["SET"] + + +@pytest.mark.asyncio +async def test_auth_write_back_outside_a_scope_writes_through_as_before(): + redis_cache = FakeRedisCache(FakeClient(_lua_ok_replies)) + cache = UserApiKeyCache(redis_cache=redis_cache) + await _write_back([_user_entry()], cache) + assert [(op[0], [(key, ttl) for key, _value, ttl in op[1]]) for op in redis_cache.alone] == [ + ("SET_PIPELINE", [("user-1", 42)]) + ] diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 0aedfce3598..e4e65f8904c 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -47,6 +47,7 @@ from litellm.router import ( _anthropic_stream_should_decline_fallback, _is_retriable_anthropic_status, _responses_stream_holds_event, + _without_line_breaks, ) from litellm.router_strategy import simple_shuffle from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit @@ -18803,3 +18804,36 @@ async def test_a_guardrail_verdict_is_neither_retried_nor_fallen_back(verdict: E await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}]) assert [c.kwargs["metadata"]["model_group"] for c in mock_acompletion.call_args_list] == ["primary"] + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("gpt-4\r\nERROR forged entry\n", "gpt-4ERROR forged entry"), + (RuntimeError("no deployments\r\nfor gpt-4"), "no deploymentsfor gpt-4"), + ("gpt-4", "gpt-4"), + ], +) +def test_without_line_breaks_drops_every_cr_and_lf_from_the_logged_value(value: object, expected: str) -> None: + assert _without_line_breaks(value) == expected + + +def test_a_failed_routing_read_prefetch_logs_the_request_model_without_its_line_breaks(monkeypatch, caplog) -> None: + router = litellm.Router( + model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "openai/gpt-4", "api_key": "k"}}] + ) + forged_model: Final = "gpt-4\r\nERROR forged entry\n" + + def fail_lookup(model_name: str | None = None, team_id: str | None = None) -> None: + raise RuntimeError(f"no deployments for {model_name}") + + monkeypatch.setattr(router, "get_model_list", fail_lookup) + caplog.clear() + + with caplog.at_level(logging.DEBUG, logger="LiteLLM Router"): + router.arm_routing_read_prefetch(forged_model, {}) + + messages: Final = [r.getMessage() for r in caplog.records if "routing read prefetch not armed" in r.getMessage()] + assert messages == [ + "routing read prefetch not armed for gpt-4ERROR forged entry: no deployments for gpt-4ERROR forged entry" + ]