Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
422 changes: 422 additions & 0 deletions litellm/caching/redis_batch.py

Large diffs are not rendered by default.

23 changes: 21 additions & 2 deletions litellm/proxy/auth/auth_object_prefetch.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand Down
3 changes: 3 additions & 0 deletions litellm/proxy/common_request_processing.py
Original file line number Diff line number Diff line change
Expand Up @@ -2199,6 +2199,9 @@ async def common_processing_pre_call_logic(

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,
Expand Down
146 changes: 139 additions & 7 deletions litellm/proxy/hooks/parallel_request_limiter_v3.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@

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
Expand Down Expand Up @@ -474,6 +475,19 @@ def _sibling_counter_keys(window_key: str) -> tuple[str, str]:

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
Expand Down Expand Up @@ -1323,6 +1337,21 @@ def keyslot_for_redis_cluster(self, key: str) -> int:
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.
Expand Down Expand Up @@ -1404,7 +1433,7 @@ async def _read_counter_values_without_incrementing(
)
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(
Expand Down Expand Up @@ -1436,12 +1465,19 @@ async def _execute_redis_batch_rate_limiter_script(

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:
Expand All @@ -1450,6 +1486,7 @@ async def _execute_redis_batch_rate_limiter_script(
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
Expand All @@ -1464,6 +1501,22 @@ async def _execute_redis_batch_rate_limiter_script(

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],
Expand Down Expand Up @@ -2061,7 +2114,16 @@ async def _atomic_lua_per_descriptor(
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),
)
Comment thread
greptile-apps[bot] marked this conversation as resolved.
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,
Expand Down Expand Up @@ -2105,6 +2167,76 @@ async def _atomic_lua_per_descriptor(
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:
Comment thread
greptile-apps[bot] marked this conversation as resolved.
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,
)
Comment thread
cursor[bot] marked this conversation as resolved.
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]],
Expand Down Expand Up @@ -2233,7 +2365,7 @@ def _build_atomic_response(

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.
Expand Down
25 changes: 25 additions & 0 deletions litellm/proxy/middleware/redis_request_batch_middleware.py
Original file line number Diff line number Diff line change
@@ -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()
2 changes: 2 additions & 0 deletions litellm/proxy/proxy_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -681,6 +681,7 @@ def generate_feedback_box():
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,
)
Expand Down Expand Up @@ -2417,6 +2418,7 @@ def _restructure_ui_html_files(ui_root: str) -> None:
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)

Expand Down
Loading
Loading