feat(proxy): enforce rate limits across regions via read-only replicas

This commit is contained in:
michelligabriele 2026-09-15 12:01:46 +02:00
parent 3ed6c19b8d
commit 3b4b8b8ef3
No known key found for this signature in database
8 changed files with 852 additions and 26 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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