mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
feat(proxy): enforce rate limits across regions via read-only replicas
This commit is contained in:
parent
3ed6c19b8d
commit
3b4b8b8ef3
8 changed files with 852 additions and 26 deletions
|
|
@ -2531,6 +2531,15 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
"borrowing the `cache_params` Redis and over the REDIS_* env fallback"
|
||||
),
|
||||
)
|
||||
rate_limit_remote_replicas: tuple[CoordinationRedisParams, ...] | None = Field(
|
||||
None,
|
||||
description=(
|
||||
"read-only Redis replicas of OTHER regions' coordination Redis, used only to read "
|
||||
"rate-limit counters. The limiter adds each replica's counter to the local one before "
|
||||
"comparing against the limit, so an active-active deployment enforces one shared limit "
|
||||
"instead of one per region. Never written to, and never used for spend, locks, or caching"
|
||||
),
|
||||
)
|
||||
control_plane_url: str | None = Field(
|
||||
None,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ from typing_extensions import NotRequired, ReadOnly
|
|||
|
||||
from litellm import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
from litellm.caching.redis_cache import RedisCache, 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
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
|
|
@ -398,6 +398,44 @@ CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None]
|
|||
|
||||
ParallelGaugeCacheValue: TypeAlias = dict[str, object] | int | float | str | bytes
|
||||
|
||||
# One counter to look up on the remote-region replicas: (window key, counter key, window size).
|
||||
RemoteCounterProbe: TypeAlias = tuple[str, str, int]
|
||||
|
||||
# Returned whenever no remote replica is configured, so every call site can branch on
|
||||
# falsiness without allocating a dict per request.
|
||||
_NO_REMOTE_OFFSETS: Final[Mapping[str, int]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _as_counter(value: object) -> int:
|
||||
"""Coerce a counter read out of Redis (or a replica) to an int, treating anything unreadable as 0."""
|
||||
match value:
|
||||
case bool():
|
||||
return 0
|
||||
case int() | float():
|
||||
return int(value)
|
||||
case bytes():
|
||||
return _as_counter(value.decode("utf-8", errors="ignore"))
|
||||
case str():
|
||||
try:
|
||||
return int(float(value))
|
||||
except ValueError:
|
||||
return 0
|
||||
case _:
|
||||
return 0
|
||||
|
||||
|
||||
def _merged_counter(local_value: CacheCounterValue | None, remote_offset: int) -> CacheCounterValue | None:
|
||||
"""
|
||||
Add a remote-region counter to the local one.
|
||||
|
||||
A missing local counter with a non-zero remote term must report the remote value
|
||||
rather than None: `is_cache_list_over_limit` reads None as "no counter yet" and
|
||||
hands back the full limit, which would drop the remote usage on the floor.
|
||||
"""
|
||||
if local_value is None:
|
||||
return remote_offset or None
|
||||
return _as_counter(local_value) + remote_offset
|
||||
|
||||
|
||||
class _AsyncLuaScript(Protocol):
|
||||
"""A Lua script registered against the async Redis client, called with KEYS and ARGV."""
|
||||
|
|
@ -481,6 +519,9 @@ class AtomicCounterMeta(TypedDict):
|
|||
increment: int
|
||||
ttl: int
|
||||
window_size: int
|
||||
# Sum of this counter across the remote-region replicas. Subtracted from the limit
|
||||
# sent to Lua, and added back when reporting, so `current_limit` stays the configured one.
|
||||
remote_offset: ReadOnly[int]
|
||||
|
||||
|
||||
class AtomicCounterState(TypedDict):
|
||||
|
|
@ -608,9 +649,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
self,
|
||||
internal_usage_cache: InternalUsageCache,
|
||||
time_provider: Callable[[], datetime] | None = None,
|
||||
remote_replica_caches: Sequence[RedisCache] = (),
|
||||
):
|
||||
self.internal_usage_cache = internal_usage_cache
|
||||
self._time_provider = time_provider or datetime.now
|
||||
# Read-only replicas of other regions' coordination Redis. Deliberately NOT tied to
|
||||
# `_is_redis_cluster()`: that reports the mode of the primary, and a replica may
|
||||
# differ. `RedisClusterCache` overrides the batch read with `mget_nonatomic`, so a
|
||||
# cluster replica is slot-safe here without hash-tag grouping.
|
||||
self.remote_replica_caches = tuple(remote_replica_caches)
|
||||
if self.internal_usage_cache.dual_cache.redis_cache is not None:
|
||||
self.batch_rate_limiter_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
|
||||
BATCH_RATE_LIMITER_SCRIPT
|
||||
|
|
@ -1131,6 +1178,88 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
return RateLimitResponse(overall_code=overall_code, statuses=statuses)
|
||||
|
||||
async def _read_replica_counters(self, replica: RedisCache, keys: Sequence[str]) -> Mapping[str, object]:
|
||||
"""Read one remote replica's copy of these counters. An unreachable replica contributes nothing."""
|
||||
try:
|
||||
return await replica.async_batch_get_cache(key_list=list(keys))
|
||||
except Exception as e: # noqa: BLE001 # any replica failure must degrade to local-only enforcement
|
||||
log_redis_failure(
|
||||
verbose_proxy_logger,
|
||||
logging.WARNING,
|
||||
"rate_limit_remote_replicas: replica read failed, enforcing against local counters only",
|
||||
e,
|
||||
)
|
||||
return {}
|
||||
|
||||
async def _remote_counter_offsets(
|
||||
self,
|
||||
probes: Sequence[RemoteCounterProbe],
|
||||
now_int: int,
|
||||
) -> Mapping[str, int]:
|
||||
"""
|
||||
Sum each counter's value across the remote replicas, counting a replica
|
||||
only while its copy of that counter's window is still current. An
|
||||
unreachable replica contributes nothing (fail open): a region whose
|
||||
replica link is down falls back to the per-region enforcement it has
|
||||
today, rather than rejecting traffic the other region cannot vouch for.
|
||||
"""
|
||||
if not self.remote_replica_caches or not probes:
|
||||
return _NO_REMOTE_OFFSETS
|
||||
|
||||
keys: Final = [key for window_key, counter_key, _ in probes for key in (window_key, counter_key)]
|
||||
# One concurrent round trip per replica, not one after another.
|
||||
replica_reads: Final = await asyncio.gather(
|
||||
*(self._read_replica_counters(replica, keys) for replica in self.remote_replica_caches)
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
counter_key: sum(
|
||||
_as_counter(read.get(counter_key))
|
||||
for read in replica_reads
|
||||
if (window_value := read.get(window_key)) is not None
|
||||
and now_int - _as_counter(window_value) < window_size
|
||||
)
|
||||
for window_key, counter_key, window_size in probes
|
||||
}
|
||||
)
|
||||
|
||||
async def _merge_remote_counters(
|
||||
self,
|
||||
keys_to_fetch: Sequence[str],
|
||||
cache_values: CacheCounterValues,
|
||||
key_metadata: Mapping[str, WindowKeyMetadata],
|
||||
now_int: int,
|
||||
) -> CacheCounterValues:
|
||||
"""
|
||||
Add the remote-region counters to a (window, counter) pair list, leaving the
|
||||
window values untouched. The merged list is for the limit comparison only and
|
||||
must never be written back to the in-memory cache, or the next request's
|
||||
in-memory pre-check would count the remote term a second time.
|
||||
|
||||
Each probe carries its OWN descriptor's window size, not `self.window_size`:
|
||||
a descriptor may override it (`rate_limit.window_size`), and using the global
|
||||
one would mis-judge whether the replica's copy of that window is still current
|
||||
-- dropping remote usage for a longer window, and counting expired remote usage
|
||||
for a shorter one.
|
||||
"""
|
||||
offsets: Final = await self._remote_counter_offsets(
|
||||
probes=tuple(
|
||||
(
|
||||
keys_to_fetch[i],
|
||||
keys_to_fetch[i + 1],
|
||||
key_metadata[keys_to_fetch[i]]["window_size"],
|
||||
)
|
||||
for i in range(0, len(keys_to_fetch) - 1, 2)
|
||||
),
|
||||
now_int=now_int,
|
||||
)
|
||||
if not offsets:
|
||||
return cache_values
|
||||
return [
|
||||
value if index % 2 == 0 else _merged_counter(value, offsets.get(keys_to_fetch[index], 0))
|
||||
for index, value in enumerate(cache_values)
|
||||
]
|
||||
|
||||
def keyslot_for_redis_cluster(self, key: str) -> int:
|
||||
"""
|
||||
Compute the Redis Cluster slot for a given key.
|
||||
|
|
@ -1357,7 +1486,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
window_size=self.window_size,
|
||||
)
|
||||
|
||||
windowed_response = self.is_cache_list_over_limit(keys_to_fetch, cache_values, key_metadata)
|
||||
# Single merge point for all three branches above (read-only, Redis Lua, in-memory
|
||||
# fallback), so response headers become cross-region aware for free. It sits AFTER
|
||||
# the in-memory write-back, which keeps storing the local Redis result.
|
||||
merged_values: Final = (
|
||||
await self._merge_remote_counters(keys_to_fetch, cache_values, key_metadata, now_int)
|
||||
if self.remote_replica_caches
|
||||
else cache_values
|
||||
)
|
||||
windowed_response = self.is_cache_list_over_limit(keys_to_fetch, merged_values, key_metadata)
|
||||
if windowed_response["overall_code"] == "OVER_LIMIT":
|
||||
return windowed_response
|
||||
|
||||
|
|
@ -1716,18 +1853,28 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
# Build per-descriptor (keys, args, meta) groups. All keys within a
|
||||
# group share the descriptor's {key:value} hash tag, so a single Lua
|
||||
# call per group never triggers CROSSSLOT on Redis Cluster.
|
||||
descriptor_groups: Final[list[DescriptorAtomicGroup]] = []
|
||||
for descriptor, increment_amounts in zip(descriptors, increments):
|
||||
keys, args, meta = self._build_descriptor_atomic_payload(
|
||||
descriptor=descriptor,
|
||||
increment_amounts=increment_amounts,
|
||||
)
|
||||
if keys:
|
||||
descriptor_groups.append((keys, args, meta))
|
||||
descriptor_groups: Final = self._build_descriptor_groups(descriptors, increments, _NO_REMOTE_OFFSETS)
|
||||
|
||||
if not descriptor_groups:
|
||||
return RateLimitResponse(overall_code="OK", statuses=[])
|
||||
|
||||
# The first pass above exists only to learn which counter keys are in play;
|
||||
# it is pure, and with no replicas configured `_remote_counter_offsets` returns
|
||||
# the shared empty mapping and the groups are reused as-is.
|
||||
remote_offsets: Final = await self._remote_counter_offsets(
|
||||
probes=tuple(
|
||||
(meta["window_key"], meta["counter_key"], meta["window_size"])
|
||||
for _keys, _args, group_meta in descriptor_groups
|
||||
for meta in group_meta
|
||||
),
|
||||
now_int=int(self._get_current_time().timestamp()),
|
||||
)
|
||||
effective_groups: Final = (
|
||||
self._build_descriptor_groups(descriptors, increments, remote_offsets)
|
||||
if remote_offsets
|
||||
else descriptor_groups
|
||||
)
|
||||
|
||||
# Multi-process atomicity via Redis Lua, per descriptor for slot
|
||||
# co-location. Single-process atomicity falls back to the
|
||||
# asyncio.Lock + in-memory sliding window below — there are no
|
||||
|
|
@ -1735,12 +1882,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
# critical section for true cross-descriptor atomicity.
|
||||
if self.check_and_increment_by_n_script is not None:
|
||||
return await self._atomic_lua_per_descriptor(
|
||||
descriptor_groups=descriptor_groups,
|
||||
descriptor_groups=effective_groups,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
flat_meta: Final[list[AtomicCounterMeta]] = [
|
||||
m for _keys, _args, group_meta in descriptor_groups for m in group_meta
|
||||
m for _keys, _args, group_meta in effective_groups for m in group_meta
|
||||
]
|
||||
async with self._check_and_increment_lock:
|
||||
return await self._atomic_check_and_increment_in_memory(
|
||||
|
|
@ -1748,14 +1895,40 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
def _build_descriptor_groups(
|
||||
self,
|
||||
descriptors: Sequence[RateLimitDescriptor],
|
||||
increments: Sequence[Mapping[Literal["requests", "tokens"], int]],
|
||||
remote_offsets: Mapping[str, int],
|
||||
) -> tuple[DescriptorAtomicGroup, ...]:
|
||||
"""Build one (KEYS, ARGV, meta) group per descriptor that has at least one enforced counter."""
|
||||
return tuple(
|
||||
group
|
||||
for descriptor, increment_amounts in zip(descriptors, increments)
|
||||
if (
|
||||
group := self._build_descriptor_atomic_payload(
|
||||
descriptor=descriptor,
|
||||
increment_amounts=increment_amounts,
|
||||
remote_offsets=remote_offsets,
|
||||
)
|
||||
)[0]
|
||||
)
|
||||
|
||||
def _build_descriptor_atomic_payload(
|
||||
self,
|
||||
descriptor: RateLimitDescriptor,
|
||||
increment_amounts: dict[Literal["requests", "tokens"], int],
|
||||
increment_amounts: Mapping[Literal["requests", "tokens"], int],
|
||||
remote_offsets: Mapping[str, int] = _NO_REMOTE_OFFSETS,
|
||||
) -> DescriptorAtomicGroup:
|
||||
"""
|
||||
Build (KEYS, ARGV, per-counter meta) for a single descriptor's Lua
|
||||
call. All keys returned share the descriptor's {key:value} hash tag.
|
||||
|
||||
`remote_offsets` carries each counter's usage in the other regions. It is
|
||||
subtracted from the limit sent to Lua rather than added to the counter,
|
||||
because `local + remote + increment > limit` is the same test as
|
||||
`local + increment > limit - remote` — so the shared Lua script needs no
|
||||
change. `meta` keeps the true configured limit for reporting.
|
||||
"""
|
||||
descriptor_key: Final = descriptor["key"]
|
||||
descriptor_value: Final = descriptor["value"]
|
||||
|
|
@ -1786,10 +1959,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
# custom-TTL descriptor doesn't reintroduce a silent expiry bug.
|
||||
ttl_seconds = int(window_size)
|
||||
window_size_seconds = int(window_size)
|
||||
remote_offset = remote_offsets.get(counter_key, 0)
|
||||
keys.extend([window_key, counter_key])
|
||||
# 4-tuple matches the Lua ARGV layout:
|
||||
# [limit, increment, ttl_seconds, window_size_seconds].
|
||||
args.extend([int(limit_value), inc_amount, ttl_seconds, window_size_seconds])
|
||||
# A remote region that has already burned the quota drives the limit
|
||||
# negative, which the script blocks on — the correct answer.
|
||||
args.extend([int(limit_value) - remote_offset, inc_amount, ttl_seconds, window_size_seconds])
|
||||
meta.append(
|
||||
{
|
||||
"descriptor_key": descriptor_key,
|
||||
|
|
@ -1801,13 +1977,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
"increment": inc_amount,
|
||||
"ttl": ttl_seconds,
|
||||
"window_size": window_size_seconds,
|
||||
"remote_offset": remote_offset,
|
||||
}
|
||||
)
|
||||
return keys, args, meta
|
||||
|
||||
async def _atomic_lua_per_descriptor(
|
||||
self,
|
||||
descriptor_groups: list[DescriptorAtomicGroup],
|
||||
descriptor_groups: Sequence[DescriptorAtomicGroup],
|
||||
parent_otel_span: Span | None = None,
|
||||
) -> RateLimitResponse:
|
||||
"""
|
||||
|
|
@ -1918,17 +2095,20 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
status_code: Final = int(raw[0])
|
||||
if status_code == 1:
|
||||
# Over limit: { 1, counter_index (1-based), current_counter, limit }
|
||||
# `raw[3]` is the limit the script was given, which is the configured limit
|
||||
# minus the remote-region usage; report the configured one and add the remote
|
||||
# term back onto the local counter. Identical to raw[3] at a zero offset.
|
||||
descriptor_index: Final = int(raw[1]) - 1
|
||||
current_counter: Final = int(raw[2])
|
||||
limit: Final = int(raw[3])
|
||||
meta = per_counter_meta[descriptor_index]
|
||||
merged_counter: Final = current_counter + meta["remote_offset"]
|
||||
return RateLimitResponse(
|
||||
overall_code="OVER_LIMIT",
|
||||
statuses=[
|
||||
RateLimitStatus(
|
||||
code="OVER_LIMIT",
|
||||
current_limit=limit,
|
||||
limit_remaining=max(0, limit - current_counter),
|
||||
current_limit=meta["current_limit"],
|
||||
limit_remaining=max(0, meta["current_limit"] - merged_counter),
|
||||
rate_limit_type=meta["rate_limit_type"],
|
||||
descriptor_key=meta["descriptor_key"],
|
||||
descriptor_value=meta["descriptor_value"],
|
||||
|
|
@ -1943,7 +2123,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
RateLimitStatus(
|
||||
code="OK",
|
||||
current_limit=meta["current_limit"],
|
||||
limit_remaining=max(0, meta["current_limit"] - int(new_counter)),
|
||||
limit_remaining=max(0, meta["current_limit"] - int(new_counter) - meta["remote_offset"]),
|
||||
rate_limit_type=meta["rate_limit_type"],
|
||||
descriptor_key=meta["descriptor_key"],
|
||||
descriptor_value=meta["descriptor_value"],
|
||||
|
|
@ -1998,10 +2178,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
)
|
||||
current_counter = 0 if window_expired else int(raw_counter or 0)
|
||||
# Same reduced-limit algebra as the Lua path: `local + remote + increment > limit`
|
||||
# is `local + increment > limit - remote`. Without this, a Lua failure would fall
|
||||
# back to enforcement that silently ignores the other regions.
|
||||
effective_limit = meta["current_limit"] - meta["remote_offset"]
|
||||
over_limit = (
|
||||
current_counter + meta["increment"] > meta["current_limit"]
|
||||
current_counter + meta["increment"] > effective_limit
|
||||
if meta["increment"] > 0
|
||||
else current_counter >= meta["current_limit"]
|
||||
else current_counter >= effective_limit
|
||||
)
|
||||
if over_limit:
|
||||
return RateLimitResponse(
|
||||
|
|
@ -2010,7 +2194,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
RateLimitStatus(
|
||||
code="OVER_LIMIT",
|
||||
current_limit=meta["current_limit"],
|
||||
limit_remaining=max(0, meta["current_limit"] - current_counter),
|
||||
limit_remaining=max(0, effective_limit - current_counter),
|
||||
rate_limit_type=meta["rate_limit_type"],
|
||||
descriptor_key=meta["descriptor_key"],
|
||||
descriptor_value=meta["descriptor_value"],
|
||||
|
|
@ -2048,7 +2232,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
RateLimitStatus(
|
||||
code="OK",
|
||||
current_limit=meta["current_limit"],
|
||||
limit_remaining=max(0, meta["current_limit"] - new_counter),
|
||||
limit_remaining=max(0, meta["current_limit"] - new_counter - meta["remote_offset"]),
|
||||
rate_limit_type=meta["rate_limit_type"],
|
||||
descriptor_key=meta["descriptor_key"],
|
||||
descriptor_value=meta["descriptor_value"],
|
||||
|
|
|
|||
|
|
@ -276,6 +276,7 @@ from litellm.constants import (
|
|||
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS,
|
||||
REALTIME_SESSION_FAILURE_LOGGED_KEY,
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
|
||||
REDIS_SOCKET_TIMEOUT,
|
||||
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG,
|
||||
USER_SPEND_ALERTS_JOB_ID,
|
||||
WEEKLY_SPEND_REPORT_JOB_ID,
|
||||
|
|
@ -1244,6 +1245,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
|
|||
llm_router=llm_router,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
redis_usage_cache=transaction_buffer_redis_cache,
|
||||
rate_limit_remote_replica_caches=rate_limit_remote_replica_caches,
|
||||
)
|
||||
|
||||
## V2 OTEL: publish the chosen V2 logger's TracerProvider as the OTel global.
|
||||
|
|
@ -2359,6 +2361,8 @@ cli_sso_session_cache: Final = DualCache(default_in_memory_ttl=CLI_SSO_SESSION_T
|
|||
model_max_budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=spend_counter_cache)
|
||||
litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter)
|
||||
redis_usage_cache: RedisCache | None = None # redis cache used for tracking spend, tpm/rpm limits
|
||||
# read-only replicas of OTHER regions' coordination Redis; read by the v3 rate limiter, never written to
|
||||
rate_limit_remote_replica_caches: tuple[RedisCache, ...] = ()
|
||||
polling_via_cache_enabled: Literal["all"] | list[str] | bool = False
|
||||
native_background_mode: list[str] = [] # Models that should use native provider background mode instead of polling
|
||||
polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache
|
||||
|
|
@ -4431,16 +4435,25 @@ def _resolve_coordination_redis_env_refs(raw_params: Mapping[str, object]) -> di
|
|||
}
|
||||
|
||||
|
||||
def _build_redis_usage_cache(redis_params: Mapping[str, object]) -> RedisCache:
|
||||
def _build_redis_usage_cache(
|
||||
redis_params: Mapping[str, object],
|
||||
allow_env_cluster_fallback: bool = True,
|
||||
) -> RedisCache:
|
||||
"""
|
||||
Builds the proxy's coordination Redis client from resolved connection
|
||||
params. Cluster-mode targets (explicit `startup_nodes` or the
|
||||
REDIS_CLUSTER_NODES env var) get a `RedisClusterCache`, so consumers that
|
||||
branch on cluster mode (e.g. the v3 rate limiter) take the cluster path;
|
||||
everything else (host/url/sentinel) gets a plain `RedisCache`.
|
||||
|
||||
`allow_env_cluster_fallback=False` suppresses the REDIS_CLUSTER_NODES
|
||||
fallback. The env var names THIS pod's local cluster, so borrowing it for a
|
||||
client that is meant to reach a different Redis (a remote region's replica)
|
||||
would silently connect to the local one instead — set it False whenever the
|
||||
caller's connection target is not the local coordination Redis.
|
||||
"""
|
||||
startup_nodes = redis_params.get("startup_nodes")
|
||||
if startup_nodes is None:
|
||||
if startup_nodes is None and allow_env_cluster_fallback:
|
||||
env_cluster_nodes: Final = get_secret_str("REDIS_CLUSTER_NODES")
|
||||
if env_cluster_nodes is not None:
|
||||
startup_nodes = json.loads(env_cluster_nodes)
|
||||
|
|
@ -5067,6 +5080,71 @@ class ProxyConfig:
|
|||
)
|
||||
return coordination_redis_cache
|
||||
|
||||
def _init_rate_limit_remote_replicas(self, config: Mapping[str, Any]) -> tuple[RedisCache, ...]:
|
||||
"""
|
||||
Builds the read-only remote-region replica clients from
|
||||
`general_settings.rate_limit_remote_replicas`. Deliberately does NOT
|
||||
call `_attach_redis_usage_cache`: these clients are replicas of another
|
||||
region's Redis and must never back spend counters, the pod lock
|
||||
manager, or any cache the proxy writes to.
|
||||
"""
|
||||
raw_replicas: Final = (config.get("general_settings") or {}).get("rate_limit_remote_replicas")
|
||||
if raw_replicas is None:
|
||||
return ()
|
||||
if not isinstance(raw_replicas, list):
|
||||
raise TypeError("general_settings.rate_limit_remote_replicas must be a list of Redis connection params")
|
||||
|
||||
replica_params: Final = tuple(
|
||||
CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(raw_params))
|
||||
for raw_params in raw_replicas
|
||||
)
|
||||
for params in replica_params:
|
||||
if not params.has_connection_target():
|
||||
raise ValueError(
|
||||
"every general_settings.rate_limit_remote_replicas entry needs a connection target: "
|
||||
"set one of host, url, startup_nodes, or sentinel_nodes"
|
||||
)
|
||||
|
||||
# REDIS_CLUSTER_NODES names THIS region's cluster, and litellm._redis applies it to any
|
||||
# client built without an explicit `startup_nodes` -- so a replica entry that only sets
|
||||
# `host` or `url` would silently connect to the local Redis instead. The limiter would then
|
||||
# add this region's counters to themselves and every limit would run at half its configured
|
||||
# value, with nothing in the logs to say so. Refuse to start instead: the operator either
|
||||
# names the replica's own cluster nodes, or the env var has no business being set here.
|
||||
if get_secret_str("REDIS_CLUSTER_NODES") is not None:
|
||||
env_ambiguous: Final = [
|
||||
params for params in replica_params if not params.startup_nodes
|
||||
]
|
||||
if env_ambiguous:
|
||||
raise ValueError(
|
||||
"REDIS_CLUSTER_NODES is set, which would point every "
|
||||
"general_settings.rate_limit_remote_replicas entry at THIS region's cluster "
|
||||
"instead of the remote region's, silently enforcing half of every configured "
|
||||
f"limit. Give each of the {len(env_ambiguous)} affected replica entr"
|
||||
f"{'y' if len(env_ambiguous) == 1 else 'ies'} its own `startup_nodes`, or unset "
|
||||
"REDIS_CLUSTER_NODES and name the local cluster under "
|
||||
"general_settings.coordination_redis.startup_nodes instead."
|
||||
)
|
||||
|
||||
# REDIS_SOCKET_TIMEOUT (0.1s) instead of RedisCache's 5.0s default, so a degraded
|
||||
# replica cannot add seconds to every request; a per-entry socket_timeout wins.
|
||||
#
|
||||
# allow_env_cluster_fallback=False keeps the proxy-level builder from doing the same
|
||||
# substitution for an entry that reached here legitimately.
|
||||
replicas: Final = tuple(
|
||||
_build_redis_usage_cache(
|
||||
{"socket_timeout": REDIS_SOCKET_TIMEOUT, **params.model_dump(exclude_none=True)},
|
||||
allow_env_cluster_fallback=False,
|
||||
)
|
||||
for params in replica_params
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"rate_limit_remote_replicas: reading rate-limit counters from %s remote replica(s); "
|
||||
"limits are enforced against the sum of local and remote counters.",
|
||||
len(replicas),
|
||||
)
|
||||
return replicas
|
||||
|
||||
@staticmethod
|
||||
async def _init_coordination_redis_env_fallback(litellm_settings: Mapping[str, object]) -> RedisCache | None:
|
||||
"""
|
||||
|
|
@ -5439,6 +5517,10 @@ class ProxyConfig:
|
|||
if coordination_redis_cache is not None:
|
||||
_set_redis_usage_cache(coordination_redis_cache)
|
||||
|
||||
## Read-only replicas of other regions' coordination Redis, read by the v3 rate limiter
|
||||
global rate_limit_remote_replica_caches
|
||||
rate_limit_remote_replica_caches = self._init_rate_limit_remote_replicas(config=config)
|
||||
|
||||
## Callback settings
|
||||
callback_settings: Final = config.get("callback_settings", {})
|
||||
if callback_settings:
|
||||
|
|
@ -9327,12 +9409,17 @@ class ProxyStartupEvent:
|
|||
llm_router: Router | None,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
redis_usage_cache: RedisCache | None,
|
||||
rate_limit_remote_replica_caches: tuple[RedisCache, ...] = (),
|
||||
):
|
||||
"""Initialize logging and alerting on startup"""
|
||||
## COST TRACKING ##
|
||||
cost_tracking()
|
||||
|
||||
proxy_logging_obj.startup_event(llm_router=llm_router, redis_usage_cache=redis_usage_cache)
|
||||
proxy_logging_obj.startup_event(
|
||||
llm_router=llm_router,
|
||||
redis_usage_cache=redis_usage_cache,
|
||||
rate_limit_remote_replica_caches=rate_limit_remote_replica_caches,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _warn_if_mock_testing_params_enabled(general_settings: dict) -> None:
|
||||
|
|
|
|||
|
|
@ -990,6 +990,9 @@ class ProxyLogging:
|
|||
self.service_logging_obj = ServiceLogging()
|
||||
self.db_spend_update_writer = DBSpendUpdateWriter()
|
||||
self.proxy_hook_mapping: dict[str, CustomLogger] = {}
|
||||
# Read-only replicas of other regions' coordination Redis, handed to hooks that ask for
|
||||
# them by name in _add_proxy_hooks. Empty unless the proxy configured them at startup.
|
||||
self.rate_limit_remote_replica_caches: tuple[RedisCache, ...] = ()
|
||||
|
||||
# Guard flags to prevent duplicate background tasks
|
||||
self.daily_report_started: bool = False
|
||||
|
|
@ -1000,11 +1003,17 @@ class ProxyLogging:
|
|||
self,
|
||||
llm_router: Router | None,
|
||||
redis_usage_cache: RedisCache | None,
|
||||
rate_limit_remote_replica_caches: tuple[RedisCache, ...] = (),
|
||||
):
|
||||
"""Initialize logging and alerting on proxy startup"""
|
||||
## UPDATE SLACK ALERTING ##
|
||||
self.slack_alerting_instance.update_values(llm_router=llm_router)
|
||||
|
||||
## REMOTE-REGION REPLICAS ##
|
||||
# Set before _init_litellm_callbacks below, which constructs the proxy hooks
|
||||
# that read this off the instance.
|
||||
self.rate_limit_remote_replica_caches = rate_limit_remote_replica_caches
|
||||
|
||||
## UPDATE INTERNAL USAGE CACHE ##
|
||||
self.update_values(
|
||||
redis_cache=redis_usage_cache
|
||||
|
|
@ -1130,6 +1139,8 @@ class ProxyLogging:
|
|||
passed_in_args["internal_usage_cache"] = self.internal_usage_cache
|
||||
if "prisma_client" in expected_args:
|
||||
passed_in_args["prisma_client"] = prisma_client
|
||||
if "remote_replica_caches" in expected_args:
|
||||
passed_in_args["remote_replica_caches"] = self.rate_limit_remote_replica_caches
|
||||
proxy_hook_obj = cast(CustomLogger, proxy_hook(**passed_in_args))
|
||||
litellm.logging_callback_manager.add_litellm_callback(proxy_hook_obj)
|
||||
|
||||
|
|
|
|||
|
|
@ -6529,3 +6529,382 @@ async def test_request_capacity_rejection_keeps_existing_redis_mirror():
|
|||
pytest.fail("rejection released another request's mirrored slot")
|
||||
assert exc.value.status_code == 429
|
||||
assert await cache.async_get_cache(counter_key, local_only=True) == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cross-region rate limiting: read-only replicas of other regions' Redis
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _ReplicaRedis:
|
||||
"""Read-only replica stand-in: answers the batch read from a fixed snapshot."""
|
||||
|
||||
def __init__(self, snapshot):
|
||||
self.snapshot = snapshot
|
||||
|
||||
async def async_batch_get_cache(self, key_list, parent_otel_span=None):
|
||||
return {key: self.snapshot[key] for key in key_list if key in self.snapshot}
|
||||
|
||||
|
||||
class _UnreachableReplicaRedis:
|
||||
async def async_batch_get_cache(self, key_list, parent_otel_span=None):
|
||||
raise ConnectionError("replica unreachable")
|
||||
|
||||
|
||||
class _ScriptedPrimaryRedis:
|
||||
"""
|
||||
Primary-Redis stand-in that runs the documented check-and-increment-by-N
|
||||
contract against a dict, so the TPM reservation path can be driven end to end.
|
||||
"""
|
||||
|
||||
def __init__(self, now_int: int):
|
||||
self.now_int = now_int
|
||||
self.counters: Dict[str, int] = {}
|
||||
self.windows: Dict[str, int] = {}
|
||||
|
||||
def async_register_script(self, script: str):
|
||||
if "Atomic check-and-increment-by-N" not in script:
|
||||
|
||||
async def unsupported(keys, args):
|
||||
raise AssertionError(
|
||||
"this double only scripts the check-and-increment-by-N contract"
|
||||
)
|
||||
|
||||
return unsupported
|
||||
|
||||
async def check_and_increment(keys, args):
|
||||
state = []
|
||||
for index in range(len(keys) // 2):
|
||||
window_key = keys[index * 2]
|
||||
counter_key = keys[index * 2 + 1]
|
||||
limit, increment, _ttl, window_size = args[index * 4 : index * 4 + 4]
|
||||
window_start = self.windows.get(window_key)
|
||||
expired = (
|
||||
window_start is None
|
||||
or self.now_int - window_start >= int(window_size)
|
||||
)
|
||||
current = 0 if expired else self.counters.get(counter_key, 0)
|
||||
blocked = (
|
||||
current + int(increment) > int(limit)
|
||||
if int(increment) > 0
|
||||
else current >= int(limit)
|
||||
)
|
||||
if blocked:
|
||||
return [1, index + 1, current, int(limit)]
|
||||
state.append((window_key, counter_key, expired, current, int(increment)))
|
||||
|
||||
results: List[int] = [0]
|
||||
for window_key, counter_key, expired, current, increment in state:
|
||||
if expired:
|
||||
self.windows[window_key] = self.now_int
|
||||
self.counters[counter_key] = (0 if expired else current) + increment
|
||||
results.extend([self.counters[counter_key], self.windows[window_key]])
|
||||
return results
|
||||
|
||||
return check_and_increment
|
||||
|
||||
|
||||
def _rpm_handler(replicas, time_controller, rpm_limit: int):
|
||||
"""A limiter with no primary Redis, so the windowed check takes the in-memory path."""
|
||||
return _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache()),
|
||||
time_provider=time_controller.now,
|
||||
remote_replica_caches=replicas,
|
||||
), [
|
||||
{
|
||||
"key": "api_key",
|
||||
"value": "sk-cross-region",
|
||||
"rate_limit": {"requests_per_unit": rpm_limit},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def _rpm_keys():
|
||||
return "{api_key:sk-cross-region}:window", "{api_key:sk-cross-region}:requests"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remote_replica_counters_are_summed_into_the_rpm_decision(time_controller):
|
||||
window_key, counter_key = _rpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
handler, descriptors = _rpm_handler(
|
||||
[_ReplicaRedis({window_key: now_int, counter_key: 100})], time_controller, 100
|
||||
)
|
||||
|
||||
response = await handler.should_rate_limit(descriptors=descriptors)
|
||||
|
||||
assert response["overall_code"] == "OVER_LIMIT", (
|
||||
"one local request plus 100 already spent in the other region is 101 against "
|
||||
"a limit of 100 — neither region sees that on its own counter"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_remote_replicas_leaves_the_rpm_decision_unchanged(time_controller):
|
||||
handler, descriptors = _rpm_handler([], time_controller, 100)
|
||||
|
||||
response = await handler.should_rate_limit(descriptors=descriptors)
|
||||
|
||||
assert response["overall_code"] == "OK"
|
||||
assert response["statuses"][0]["limit_remaining"] == 99
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_expired_remote_window_is_not_counted(time_controller):
|
||||
window_key, counter_key = _rpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
replica = _ReplicaRedis({})
|
||||
handler, descriptors = _rpm_handler([replica], time_controller, 100)
|
||||
replica.snapshot = {
|
||||
window_key: now_int - handler.window_size - 1,
|
||||
counter_key: 999,
|
||||
}
|
||||
|
||||
response = await handler.should_rate_limit(descriptors=descriptors)
|
||||
|
||||
assert response["overall_code"] == "OK", (
|
||||
"a remote counter whose window has already rolled over is spent quota from a "
|
||||
"previous window and must not be charged against this one"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remote_window_currency_uses_the_descriptors_own_window_size(time_controller):
|
||||
"""A descriptor may override the window (`rate_limit.window_size`). Judging the
|
||||
replica's copy against the limiter-wide default instead would keep charging remote
|
||||
usage from a window that has already rolled over on a short-window descriptor."""
|
||||
window_key, counter_key = _rpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
replica = _ReplicaRedis({})
|
||||
handler, descriptors = _rpm_handler([replica], time_controller, 100)
|
||||
descriptors[0]["rate_limit"]["window_size"] = 10
|
||||
assert handler.window_size > 10, "the default window has to be the longer one for this to bite"
|
||||
# Expired for this descriptor's 10s window, still current under the 60s default.
|
||||
replica.snapshot = {window_key: now_int - 20, counter_key: 999}
|
||||
|
||||
response = await handler.should_rate_limit(descriptors=descriptors)
|
||||
|
||||
assert response["overall_code"] == "OK", (
|
||||
"the remote window rolled over 20s ago on a 10s window, so that usage is spent "
|
||||
"quota from a previous window and must not be charged against this one"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_in_memory_atomic_fallback_still_counts_remote_usage(time_controller):
|
||||
"""The reservation path falls back to in-memory enforcement whenever the Lua call
|
||||
fails. Dropping the remote offset there would let a Redis blip silently downgrade
|
||||
a region to per-region limits, which is the bug this feature exists to fix."""
|
||||
window_key = "{api_key:sk-cross-region}:window"
|
||||
counter_key = "{api_key:sk-cross-region}:tokens"
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
# No primary Redis on this handler, so atomic_check_and_increment_by_n takes the
|
||||
# same in-memory branch the post-Lua-failure fallback lands on.
|
||||
handler, _ = _rpm_handler([_ReplicaRedis({window_key: now_int, counter_key: 60})], time_controller, 100)
|
||||
|
||||
response = await handler.atomic_check_and_increment_by_n(
|
||||
descriptors=[
|
||||
{
|
||||
"key": "api_key",
|
||||
"value": "sk-cross-region",
|
||||
"rate_limit": {"tokens_per_unit": 100},
|
||||
}
|
||||
],
|
||||
increments=[{"tokens": 50}],
|
||||
)
|
||||
|
||||
assert response["overall_code"] == "OVER_LIMIT", (
|
||||
"50 tokens locally on top of 60 already spent in the other region is 110 "
|
||||
"against a limit of 100"
|
||||
)
|
||||
assert response["statuses"][0]["current_limit"] == 100, (
|
||||
"the configured limit is what the customer set; the remote offset is an "
|
||||
"implementation detail and must not leak into what we report"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unreachable_replica_fails_open_and_warns(time_controller, caplog):
|
||||
handler, descriptors = _rpm_handler(
|
||||
[_UnreachableReplicaRedis()], time_controller, 100
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
response = await handler.should_rate_limit(descriptors=descriptors)
|
||||
|
||||
assert response["overall_code"] == "OK", (
|
||||
"a replica outage must degrade to per-region enforcement, not reject traffic "
|
||||
"the other region cannot vouch for"
|
||||
)
|
||||
warnings = [
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if record.levelno >= logging.WARNING
|
||||
]
|
||||
assert len(warnings) == 1
|
||||
assert "rate_limit_remote_replicas: replica read failed" in warnings[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_two_remote_replicas_are_summed(time_controller):
|
||||
window_key, counter_key = _rpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
handler, descriptors = _rpm_handler(
|
||||
[
|
||||
_ReplicaRedis({window_key: now_int, counter_key: 40}),
|
||||
_ReplicaRedis({window_key: now_int, counter_key: 40}),
|
||||
],
|
||||
time_controller,
|
||||
75,
|
||||
)
|
||||
|
||||
response = await handler.should_rate_limit(descriptors=descriptors)
|
||||
|
||||
assert response["overall_code"] == "OVER_LIMIT", "40 + 40 + 1 local exceeds 75"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_in_memory_counter_excludes_remote_counters_after_a_merged_check(
|
||||
time_controller,
|
||||
):
|
||||
window_key, counter_key = _rpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
handler, descriptors = _rpm_handler(
|
||||
[_ReplicaRedis({window_key: now_int, counter_key: 10})], time_controller, 100
|
||||
)
|
||||
|
||||
await handler.should_rate_limit(descriptors=descriptors)
|
||||
|
||||
assert (
|
||||
await handler.internal_usage_cache.async_get_cache(
|
||||
key=counter_key, litellm_parent_otel_span=None, local_only=True
|
||||
)
|
||||
== 1
|
||||
), (
|
||||
"the merged value is for the comparison only; writing it back would make the "
|
||||
"next request's in-memory pre-check count the remote term a second time"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_limit_remaining_reflects_remote_counters(time_controller):
|
||||
window_key, counter_key = _rpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
handler, descriptors = _rpm_handler(
|
||||
[_ReplicaRedis({window_key: now_int, counter_key: 30})], time_controller, 100
|
||||
)
|
||||
|
||||
response = await handler.should_rate_limit(descriptors=descriptors)
|
||||
|
||||
assert response["overall_code"] == "OK"
|
||||
assert response["statuses"][0]["limit_remaining"] == 69, (
|
||||
"100 - (1 local + 30 remote); the response headers are what the customer's "
|
||||
"clients back off on, so they have to be cross-region aware too"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_only_check_merges_remote_counters(time_controller):
|
||||
window_key, counter_key = _rpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
handler, descriptors = _rpm_handler(
|
||||
[_ReplicaRedis({window_key: now_int, counter_key: 60})], time_controller, 100
|
||||
)
|
||||
|
||||
response = await handler.should_rate_limit(descriptors=descriptors, read_only=True)
|
||||
|
||||
assert response["overall_code"] == "OK"
|
||||
assert response["statuses"][0]["limit_remaining"] == 40
|
||||
|
||||
|
||||
def _tpm_handler(replicas, time_controller, tpm_limit: int):
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
primary = _ScriptedPrimaryRedis(now_int=now_int)
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(
|
||||
DualCache(redis_cache=primary) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
),
|
||||
time_provider=time_controller.now,
|
||||
remote_replica_caches=replicas,
|
||||
)
|
||||
descriptors = [
|
||||
{
|
||||
"key": "api_key",
|
||||
"value": "sk-cross-region",
|
||||
"rate_limit": {"tokens_per_unit": tpm_limit},
|
||||
}
|
||||
]
|
||||
return handler, descriptors
|
||||
|
||||
|
||||
def _tpm_keys():
|
||||
return "{api_key:sk-cross-region}:window", "{api_key:sk-cross-region}:tokens"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reserve_tpm_tokens_counts_remote_usage_against_the_limit(time_controller):
|
||||
window_key, counter_key = _tpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
handler, descriptors = _tpm_handler(
|
||||
[_ReplicaRedis({window_key: now_int, counter_key: 900})], time_controller, 1000
|
||||
)
|
||||
|
||||
response = await handler.reserve_tpm_tokens(
|
||||
descriptors=descriptors, estimated_tokens=200
|
||||
)
|
||||
|
||||
assert response["overall_code"] == "OVER_LIMIT", (
|
||||
"900 already reserved in the other region leaves 100 of headroom, so a "
|
||||
"200-token estimate must not be admitted"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reserve_tpm_over_limit_reports_the_configured_limit(time_controller):
|
||||
window_key, counter_key = _tpm_keys()
|
||||
now_int = int(time_controller.now().timestamp())
|
||||
handler, descriptors = _tpm_handler(
|
||||
[_ReplicaRedis({window_key: now_int, counter_key: 900})], time_controller, 1000
|
||||
)
|
||||
|
||||
response = await handler.reserve_tpm_tokens(
|
||||
descriptors=descriptors, estimated_tokens=200
|
||||
)
|
||||
|
||||
status = response["statuses"][0]
|
||||
assert status["current_limit"] == 1000, (
|
||||
"the limit handed to Lua is reduced by the remote usage, but the customer "
|
||||
"must be told the limit they configured"
|
||||
)
|
||||
assert status["limit_remaining"] == 100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reserve_tpm_without_replicas_allows_a_request_that_fits_locally(
|
||||
time_controller,
|
||||
):
|
||||
handler, descriptors = _tpm_handler([], time_controller, 1000)
|
||||
|
||||
response = await handler.reserve_tpm_tokens(
|
||||
descriptors=descriptors, estimated_tokens=200
|
||||
)
|
||||
|
||||
assert response["overall_code"] == "OK"
|
||||
assert response["statuses"][0]["current_limit"] == 1000
|
||||
assert response["statuses"][0]["limit_remaining"] == 800
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_replica_failure_on_the_reservation_path_enforces_local_limits_only(
|
||||
time_controller,
|
||||
):
|
||||
handler, descriptors = _tpm_handler(
|
||||
[_UnreachableReplicaRedis()], time_controller, 1000
|
||||
)
|
||||
|
||||
response = await handler.reserve_tpm_tokens(
|
||||
descriptors=descriptors, estimated_tokens=200
|
||||
)
|
||||
|
||||
assert response["overall_code"] == "OK"
|
||||
assert response["statuses"][0]["limit_remaining"] == 800
|
||||
|
|
|
|||
|
|
@ -1723,6 +1723,7 @@ async def test_atomic_lua_response_carries_redis_window_identity(rate_limiter):
|
|||
"current_limit": 100,
|
||||
"rate_limit_type": "tokens",
|
||||
"counter_key": counter_key,
|
||||
"remote_offset": 0,
|
||||
}
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ import litellm
|
|||
import litellm.proxy.proxy_server as proxy_server_module
|
||||
from litellm.caching.caching import RedisCache
|
||||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||||
from litellm.constants import REDIS_SOCKET_TIMEOUT
|
||||
from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy._types import LitellmUserRoles, TokenCountRequest, UserAPIKeyAuth
|
||||
|
|
@ -11937,6 +11938,155 @@ def test_init_coordination_redis_absent_leaves_usage_cache_unset():
|
|||
assert spend_redis is None
|
||||
|
||||
|
||||
def _run_init_rate_limit_remote_replicas(config, env=None):
|
||||
"""Run ProxyConfig._init_rate_limit_remote_replicas against a stubbed module
|
||||
state, returning (replicas, spend_counter redis, config-cache redis)."""
|
||||
fresh_spend_cache = DualCache()
|
||||
fresh_config_cache = types.SimpleNamespace(redis_cache=None)
|
||||
|
||||
with (
|
||||
_patched_coordination_redis_module_state(spend_cache=fresh_spend_cache, config_cache=fresh_config_cache),
|
||||
mock.patch.dict(os.environ, env or {}, clear=False),
|
||||
):
|
||||
built = proxy_server_module.ProxyConfig()._init_rate_limit_remote_replicas(config=config)
|
||||
return (
|
||||
built,
|
||||
fresh_spend_cache.redis_cache,
|
||||
fresh_config_cache.redis_cache,
|
||||
)
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_builds_one_client_per_entry():
|
||||
"""Each entry names one other region's Redis, so an active-active pair with
|
||||
three regions gets two replica clients, each with its own connection target."""
|
||||
replicas, _, _ = _run_init_rate_limit_remote_replicas(
|
||||
config={
|
||||
"general_settings": {
|
||||
"rate_limit_remote_replicas": [
|
||||
{"host": "us-west-replica", "port": 6379},
|
||||
{"host": "eu-central-replica", "port": 6380},
|
||||
]
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
assert [replica.init_kwargs["host"] for replica in replicas] == [
|
||||
"us-west-replica",
|
||||
"eu-central-replica",
|
||||
]
|
||||
assert [replica.init_kwargs["port"] for replica in replicas] == [6379, 6380]
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_builds_a_cluster_client_for_startup_nodes():
|
||||
"""The customer's Redis is cluster-mode; the replica client has to be a cluster
|
||||
client so the counter read uses the slot-safe non-atomic MGET."""
|
||||
replicas, _, _ = _run_init_rate_limit_remote_replicas(
|
||||
config={
|
||||
"general_settings": {
|
||||
"rate_limit_remote_replicas": [{"startup_nodes": [{"host": "replica-node-1", "port": 7000}]}]
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
assert isinstance(replicas[0], _EnvBuiltClusterCache)
|
||||
assert replicas[0].init_kwargs["startup_nodes"] == [{"host": "replica-node-1", "port": 7000}]
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_refuse_ambiguous_cluster_env():
|
||||
"""REDIS_CLUSTER_NODES names THIS region's cluster, and litellm._redis applies it to any
|
||||
client built without an explicit startup_nodes -- deeper than _build_redis_usage_cache, so
|
||||
suppressing the fallback there is not enough on its own. A host-only replica entry would
|
||||
therefore connect to the local Redis and the limiter would add this region's counters to
|
||||
themselves, halving every limit silently. Startup must refuse instead."""
|
||||
with pytest.raises(ValueError, match="REDIS_CLUSTER_NODES is set"):
|
||||
_run_init_rate_limit_remote_replicas(
|
||||
config={"general_settings": {"rate_limit_remote_replicas": [{"host": "us-west-replica"}]}},
|
||||
env={"REDIS_CLUSTER_NODES": '[{"host": "local-cluster-node", "port": 7000}]'},
|
||||
)
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_allow_cluster_env_when_entry_names_its_nodes():
|
||||
"""An entry that names its own startup_nodes is unambiguous: litellm._redis prefers an
|
||||
explicit startup_nodes over the env var, so the client reaches the remote region's cluster.
|
||||
That is the supported way to run this feature on cluster-mode Redis."""
|
||||
replicas, _, _ = _run_init_rate_limit_remote_replicas(
|
||||
config={
|
||||
"general_settings": {
|
||||
"rate_limit_remote_replicas": [{"startup_nodes": [{"host": "replica-node-1", "port": 7000}]}]
|
||||
}
|
||||
},
|
||||
env={"REDIS_CLUSTER_NODES": '[{"host": "local-cluster-node", "port": 7000}]'},
|
||||
)
|
||||
|
||||
assert replicas[0].init_kwargs["startup_nodes"] == [{"host": "replica-node-1", "port": 7000}]
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_resolves_os_environ_references():
|
||||
"""os.environ/ values must be resolved per entry, as in coordination_redis."""
|
||||
replicas, _, _ = _run_init_rate_limit_remote_replicas(
|
||||
config={"general_settings": {"rate_limit_remote_replicas": [{"host": "os.environ/WEST_REPLICA_HOST"}]}},
|
||||
env={"WEST_REPLICA_HOST": "resolved-replica-host"},
|
||||
)
|
||||
|
||||
assert replicas[0].init_kwargs["host"] == "resolved-replica-host"
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_defaults_a_short_socket_timeout():
|
||||
"""A degraded replica must not add RedisCache's 5s default to every request;
|
||||
an entry that sets its own socket_timeout still wins."""
|
||||
replicas, _, _ = _run_init_rate_limit_remote_replicas(
|
||||
config={
|
||||
"general_settings": {
|
||||
"rate_limit_remote_replicas": [
|
||||
{"host": "default-timeout-replica"},
|
||||
{"host": "tuned-replica", "socket_timeout": 0.5},
|
||||
]
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
assert replicas[0].init_kwargs["socket_timeout"] == REDIS_SOCKET_TIMEOUT
|
||||
assert replicas[1].init_kwargs["socket_timeout"] == 0.5
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_absent_builds_nothing():
|
||||
"""The feature is opt-in: without the field the limiter gets no replicas and
|
||||
enforcement is exactly what it is today."""
|
||||
replicas, _, _ = _run_init_rate_limit_remote_replicas(config={"general_settings": {}})
|
||||
|
||||
assert replicas == ()
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_without_a_connection_target_raises():
|
||||
"""An entry with no host, url, startup_nodes, or sentinel_nodes is a config
|
||||
error and must fail startup loudly rather than silently skip a region."""
|
||||
with pytest.raises(ValueError, match="connection target"):
|
||||
_run_init_rate_limit_remote_replicas(
|
||||
config={"general_settings": {"rate_limit_remote_replicas": [{"ssl": True}]}},
|
||||
)
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_non_list_raises():
|
||||
"""A single mapping instead of a list is a config error."""
|
||||
with pytest.raises(TypeError, match="must be a list"):
|
||||
_run_init_rate_limit_remote_replicas(
|
||||
config={"general_settings": {"rate_limit_remote_replicas": {"host": "west"}}},
|
||||
)
|
||||
|
||||
|
||||
def test_init_rate_limit_remote_replicas_are_not_attached_to_the_spend_caches():
|
||||
"""A replica is another region's Redis, read-only to us. Attaching it to the
|
||||
spend counter, config, or auth caches would have this region writing into it."""
|
||||
replicas, spend_redis, config_redis = _run_init_rate_limit_remote_replicas(
|
||||
config={"general_settings": {"rate_limit_remote_replicas": [{"host": "us-west-replica"}]}},
|
||||
)
|
||||
|
||||
assert len(replicas) == 1
|
||||
assert spend_redis is None
|
||||
assert config_redis is None
|
||||
assert proxy_server_module.redis_usage_cache is None
|
||||
|
||||
|
||||
def test_explicit_coordination_redis_takes_precedence_over_cache_backend():
|
||||
"""When both an explicit coordination_redis block and a plain-Redis
|
||||
response cache are configured, the explicit block must win; the cache
|
||||
|
|
|
|||
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -26279,6 +26279,11 @@ export interface components {
|
|||
* @default 30
|
||||
*/
|
||||
proxy_config_reload_interval_seconds: number;
|
||||
/**
|
||||
* Rate Limit Remote Replicas
|
||||
* @description read-only Redis replicas of OTHER regions' coordination Redis, used only to read rate-limit counters. The limiter adds each replica's counter to the local one before comparing against the limit, so an active-active deployment enforces one shared limit instead of one per region. Never written to, and never used for spend, locks, or caching
|
||||
*/
|
||||
rate_limit_remote_replicas?: components["schemas"]["CoordinationRedisParams"][] | null;
|
||||
/**
|
||||
* Reject Clientside Metadata Tags
|
||||
* @description When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue