mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
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:
parent
e7460f1cff
commit
ffb15f946f
14 changed files with 1667 additions and 27 deletions
422
litellm/caching/redis_batch.py
Normal file
422
litellm/caching/redis_batch.py
Normal 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)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
25
litellm/proxy/middleware/redis_request_batch_middleware.py
Normal file
25
litellm/proxy/middleware/redis_request_batch_middleware.py
Normal 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()
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
299
tests/unit/caching/test_redis_batch.py
Normal file
299
tests/unit/caching/test_redis_batch.py
Normal 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
|
||||
530
tests/unit/caching/test_request_redis_batch_pre_call.py
Normal file
530
tests/unit/caching/test_request_redis_batch_pre_call.py
Normal 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)])
|
||||
]
|
||||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue