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 <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-29 16:42:13 -07:00 • committed by GitHub
parent e7460f1cff
commit ffb15f946f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1667 additions and 27 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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