diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d40f518d4b8..822ab969c43 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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=( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index a34dc99e472..2c5fe4fbeb5 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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"], diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f81a3166a28..200208fb27c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d201ac4bc88..34643a43d1a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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) 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 48f980086fd..6c99c6a4627 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 @@ -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 diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index bdaca9ffc2d..2bd003423b1 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -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, } ] diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 04173ced776..8bfd3a541d5 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index b26f5e25b6f..55f75a2a646 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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.