From 43146036e12ac1aae12cd543690b369709a51faf Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Tue, 15 Sep 2026 13:18:55 +0200 Subject: [PATCH] fix(proxy): refuse rate_limit_remote_replicas without a local Redis, clear lint gates --- .../hooks/parallel_request_limiter_v3.py | 93 +++++--------- litellm/proxy/proxy_server.py | 34 ++---- litellm/proxy/utils.py | 6 +- .../hooks/test_parallel_request_limiter_v3.py | 115 ++++++++++++++---- 4 files changed, 137 insertions(+), 111 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 2c5fe4fbeb5..87200905c80 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 200208fb27c..bba6a412e44 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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)}, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 34643a43d1a..ea1017e58a9 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 ## diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 6c99c6a4627..af91874f76e 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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 == ()