fix(proxy): refuse rate_limit_remote_replicas without a local Redis, clear lint gates

This commit is contained in:
michelligabriele 2026-09-15 13:18:55 +02:00
parent 3b4b8b8ef3
commit 43146036e1
No known key found for this signature in database
4 changed files with 137 additions and 111 deletions

View file

@ -398,30 +398,21 @@ 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).
# (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
if isinstance(value, bool) or not isinstance(value, int | float | str | bytes):
return 0
text: Final = value.decode("utf-8", errors="ignore") if isinstance(value, bytes) else value
try:
return int(float(text))
except (ValueError, OverflowError):
return 0
def _merged_counter(local_value: CacheCounterValue | None, remote_offset: int) -> CacheCounterValue | None:
@ -430,7 +421,7 @@ def _merged_counter(local_value: CacheCounterValue | None, remote_offset: int) -
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.
hands back the full limit, dropping the remote usage on the floor.
"""
if local_value is None:
return remote_offset or None
@ -519,8 +510,7 @@ 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.
# Subtracted from the limit sent to Lua, added back when reporting.
remote_offset: ReadOnly[int]
@ -653,11 +643,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
):
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.
# `RedisClusterCache` overrides the batch read with `mget_nonatomic`, so a cluster
# replica is slot-safe here whatever mode the local primary runs in.
self.remote_replica_caches = tuple(remote_replica_caches)
if self.remote_replica_caches and self.internal_usage_cache.dual_cache.redis_cache is None:
raise ValueError(
"general_settings.rate_limit_remote_replicas is set, but this region keeps its rate-limit "
"counters in memory only, so the other regions can never read them and the shared limit is "
"enforced in one direction. Give this region a Redis of its own under "
"general_settings.coordination_redis (or litellm_settings.cache with type: redis), or remove "
"general_settings.rate_limit_remote_replicas."
)
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
@ -1179,17 +1175,8 @@ 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 {}
"""`RedisCache` logs and swallows its own failures, so an unreachable replica comes back empty."""
return await replica.async_batch_get_cache(key_list=list(keys))
async def _remote_counter_offsets(
self,
@ -1198,16 +1185,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
) -> 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.
only while its copy of that counter's window is still current. A replica
that reads back empty contributes nothing, so a region whose replica link
is down falls back to the per-region enforcement it has today.
"""
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)
)
@ -1236,11 +1221,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
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.
Each probe carries its own descriptor's window size, not `self.window_size`,
because a descriptor may override it with `rate_limit.window_size`.
"""
offsets: Final = await self._remote_counter_offsets(
probes=tuple(
@ -1486,9 +1468,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
window_size=self.window_size,
)
# 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.
# Must stay after the in-memory write-back above, which stores the local value only.
merged_values: Final = (
await self._merge_remote_counters(keys_to_fetch, cache_values, key_metadata, now_int)
if self.remote_replica_caches
@ -1858,9 +1838,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
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.
# The first pass exists only to learn which counter keys are in play; it is pure,
# so with no offsets to apply the groups it built are reused as-is.
remote_offsets: Final = await self._remote_counter_offsets(
probes=tuple(
(meta["window_key"], meta["counter_key"], meta["window_size"])
@ -1927,7 +1906,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
`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
`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"]
@ -1963,8 +1942,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
keys.extend([window_key, counter_key])
# 4-tuple matches the Lua ARGV layout:
# [limit, increment, 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.
# A remote region that already burned the quota drives the limit negative, which blocks.
args.extend([int(limit_value) - remote_offset, inc_amount, ttl_seconds, window_size_seconds])
meta.append(
{
@ -2095,9 +2073,7 @@ 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.
# `raw[3]` is the reduced limit the script was given, so report `meta` instead.
descriptor_index: Final = int(raw[1]) - 1
current_counter: Final = int(raw[2])
meta = per_counter_meta[descriptor_index]
@ -2178,9 +2154,6 @@ 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"] > effective_limit

View file

@ -4446,11 +4446,8 @@ def _build_redis_usage_cache(
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.
`allow_env_cluster_fallback=False` suppresses the REDIS_CLUSTER_NODES fallback, which names
THIS pod's local cluster. Pass it whenever the target is not the local coordination Redis.
"""
startup_nodes = redis_params.get("startup_nodes")
if startup_nodes is None and allow_env_cluster_fallback:
@ -5083,10 +5080,9 @@ class ProxyConfig:
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.
`general_settings.rate_limit_remote_replicas`. These are replicas of another region's
Redis, so they are deliberately not attached to the usage cache, the pod lock manager,
or anything else the proxy writes to.
"""
raw_replicas: Final = (config.get("general_settings") or {}).get("rate_limit_remote_replicas")
if raw_replicas is None:
@ -5105,16 +5101,11 @@ class ProxyConfig:
"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.
# litellm._redis applies REDIS_CLUSTER_NODES to any client with no explicit `startup_nodes`,
# so a host-only replica entry would silently connect to THIS region's cluster and the
# limiter would add this region's counters to themselves.
if get_secret_str("REDIS_CLUSTER_NODES") is not None:
env_ambiguous: Final = [
params for params in replica_params if not params.startup_nodes
]
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 "
@ -5126,11 +5117,8 @@ class ProxyConfig:
"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.
# REDIS_SOCKET_TIMEOUT instead of RedisCache's 5.0s default, so a degraded replica cannot
# add seconds to every request. A per-entry socket_timeout still wins.
replicas: Final = tuple(
_build_redis_usage_cache(
{"socket_timeout": REDIS_SOCKET_TIMEOUT, **params.model_dump(exclude_none=True)},

View file

@ -990,8 +990,7 @@ 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.
# Injected into any hook that names `remote_replica_caches`, see _add_proxy_hooks.
self.rate_limit_remote_replica_caches: tuple[RedisCache, ...] = ()
# Guard flags to prevent duplicate background tasks
@ -1010,8 +1009,7 @@ class ProxyLogging:
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.
# Must be set before _init_litellm_callbacks below, which constructs the hooks that read it.
self.rate_limit_remote_replica_caches = rate_limit_remote_replica_caches
## UPDATE INTERNAL USAGE CACHE ##

View file

@ -23,6 +23,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
PARALLEL_REQUEST_SLOT_TTL_SECONDS,
ParallelSlotAcquisition,
RequestRateLimiterStash,
_as_counter,
_request_stash,
get_or_create_request_stash,
get_request_stash,
@ -6546,9 +6547,21 @@ class _ReplicaRedis:
return {key: self.snapshot[key] for key in key_list if key in self.snapshot}
class _UnreachableReplicaRedis:
class _LuaFailurePrimaryRedis:
"""
Primary Redis whose Lua calls fail. That is how production reaches the in-memory
enforcement fallback, since a limiter with no primary Redis at all is now refused
whenever remote replicas are configured.
"""
def async_register_script(self, script: str):
async def failing(keys, args):
raise ConnectionError("primary Redis Lua unavailable")
return failing
async def async_batch_get_cache(self, key_list, parent_otel_span=None):
raise ConnectionError("replica unreachable")
return {}
class _ScriptedPrimaryRedis:
@ -6605,9 +6618,13 @@ class _ScriptedPrimaryRedis:
def _rpm_handler(replicas, time_controller, rpm_limit: int):
"""A limiter with no primary Redis, so the windowed check takes the in-memory path."""
"""A limiter whose primary Redis fails every Lua call, so the windowed check lands on
the in-memory fallback."""
primary = _LuaFailurePrimaryRedis()
return _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache()),
internal_usage_cache=InternalUsageCache(
DualCache(redis_cache=primary) # pyright: ignore[reportArgumentType] # duck-typed Redis double
),
time_provider=time_controller.now,
remote_replica_caches=replicas,
), [
@ -6724,25 +6741,31 @@ async def test_in_memory_atomic_fallback_still_counts_remote_usage(time_controll
@pytest.mark.asyncio
async def test_unreachable_replica_fails_open_and_warns(time_controller, caplog):
handler, descriptors = _rpm_handler(
[_UnreachableReplicaRedis()], time_controller, 100
async def test_a_replica_that_reads_back_empty_fails_open(time_controller):
"""RedisCache logs and swallows its own read failures and hands back an empty mapping,
so a replica outage reaches the limiter as an empty read. It must degrade to per-region
enforcement rather than reject traffic the other region cannot vouch for."""
handler, descriptors = _rpm_handler([_ReplicaRedis({})], time_controller, 100)
response = await handler.should_rate_limit(descriptors=descriptors)
assert response["overall_code"] == "OK"
assert response["statuses"][0]["limit_remaining"] == 99, (
"nothing was readable on the replica, so no remote term is charged"
)
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_a_replica_read_missing_the_window_key_contributes_nothing(time_controller):
"""A read that returns the counter but not its window cannot say which window that
counter belongs to, so charging it could bill usage from an already rolled-over window."""
_window_key, counter_key = _rpm_keys()
handler, descriptors = _rpm_handler([_ReplicaRedis({counter_key: 90})], 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
@ -6895,12 +6918,10 @@ async def test_reserve_tpm_without_replicas_allows_a_request_that_fits_locally(
@pytest.mark.asyncio
async def test_replica_failure_on_the_reservation_path_enforces_local_limits_only(
async def test_an_empty_replica_read_on_the_reservation_path_enforces_local_limits_only(
time_controller,
):
handler, descriptors = _tpm_handler(
[_UnreachableReplicaRedis()], time_controller, 1000
)
handler, descriptors = _tpm_handler([_ReplicaRedis({})], time_controller, 1000)
response = await handler.reserve_tpm_tokens(
descriptors=descriptors, estimated_tokens=200
@ -6908,3 +6929,49 @@ async def test_replica_failure_on_the_reservation_path_enforces_local_limits_onl
assert response["overall_code"] == "OK"
assert response["statuses"][0]["limit_remaining"] == 800
@pytest.mark.parametrize(
"value, expected",
[
(12, 12),
(12.9, 12),
(True, 0),
(False, 0),
("12", 12),
("12.0", 12),
(b"12", 12),
(b"12.0", 12),
("", 0),
("garbage", 0),
(b"garbage", 0),
(None, 0),
({"counter": 1}, 0),
("inf", 0),
(float("inf"), 0),
(float("nan"), 0),
],
)
def test_as_counter_coerces_whatever_redis_hands_back(value, expected):
"""Counters come back as ints from a Lua call, as strings or bytes from a raw read, and
as None for a key that does not exist. A bool is never a counter, so it reads as 0."""
assert _as_counter(value) == expected
def test_remote_replicas_without_a_local_redis_are_refused():
"""Other regions read this region's usage off this region's Redis. With counters in
memory only there is nothing for them to read, so the shared limit would be enforced
in one direction and the operator would never be told."""
with pytest.raises(ValueError, match="rate_limit_remote_replicas"):
_PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache()),
remote_replica_caches=[_ReplicaRedis({})],
)
def test_no_remote_replicas_still_runs_without_a_local_redis():
"""The refusal is scoped to the cross-region feature; single-region in-memory
enforcement is untouched."""
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
assert handler.remote_replica_caches == ()