From 1e88d9955dcfcb760b639da62e99d2034a3c06c6 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 7 Oct 2026 11:37:43 -0500 Subject: [PATCH] fix(proxy): opt-in Redis hash-tag grouping for v3 rate limiter (#45085) Add litellm_settings.force_redis_hash_tag_grouping so Redis endpoints that enforce cluster slot rules behind a standalone protocol (Redis Enterprise clustering policy) group multi-key scripts by slot like RedisClusterCache, and return grouped batch values in the caller's key order. Co-authored-by: yassin Co-authored-by: Chenglun Hu Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/__init__.py | 1 + .../hooks/parallel_request_limiter_v3.py | 31 +++++- .../hooks/test_parallel_request_limiter_v3.py | 104 +++++++++++++++++- 3 files changed, 127 insertions(+), 9 deletions(-) diff --git a/litellm/__init__.py b/litellm/__init__.py index 208c33d2d25..c0dc61e2911 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -506,6 +506,7 @@ prometheus_metrics_ttl_seconds: Optional[float] = None prometheus_metrics_cleanup_interval_seconds: Optional[float] = 60.0 disable_add_prefix_to_prompt: bool = False # used by anthropic, to disable adding prefix to prompt disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. +force_redis_hash_tag_grouping: bool = False public_mcp_servers: Optional[List[str]] = None public_mcp_hub_strict_whitelist: bool = True public_model_groups: Optional[List[str]] = None diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 882c4d126cf..3430a4a8863 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -161,6 +161,12 @@ def _fail_closed_rate_limit_enforcement_from_general_settings() -> bool: return fail_closed_rate_limit_enforcement_enabled(general_settings) +def _force_redis_hash_tag_grouping_from_litellm_settings() -> bool: + from litellm import force_redis_hash_tag_grouping + + return force_redis_hash_tag_grouping is True + + def _sibling_counter_keys(window_key: str) -> tuple[str, str]: prefix: Final = window_key.removesuffix(":window") return f"{prefix}:requests", f"{prefix}:tokens" @@ -483,6 +489,22 @@ CacheCounterValue: TypeAlias = int | float | str | bytes CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None] +def _values_in_caller_order( + keys: Sequence[str], + key_groups: Sequence[tuple[str, list[str]]], + grouped_values: CacheCounterValues, +) -> list[CacheCounterValue | None]: + tag_by_key: Final = dict( + itertools.chain.from_iterable(((key, tag) for key in group_keys) for tag, group_keys in key_groups) + ) + caller_positions: Final = itertools.chain.from_iterable( + tuple(position for position, key in enumerate(keys) if tag_by_key[key] == tag) + for tag, _group_keys in key_groups + ) + value_by_position: Final = dict(zip(caller_positions, grouped_values)) + return [value_by_position.get(position) for position in range(len(keys))] + + def _as_counter_values(reply: object) -> list[CacheCounterValue]: """A Lua reply read back off the pipeline is the same array the script returns when called directly.""" if not isinstance(reply, (list, tuple)): @@ -766,12 +788,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): tag_rate_limit_resolver: TagRateLimitResolver = resolve_tag_rate_limits_from_db, model_group_resolver: Callable[[str], str | None] = _resolve_model_group_alias_via_proxy_router, fail_closed_resolver: Callable[[], bool] = _fail_closed_rate_limit_enforcement_from_general_settings, + force_hash_tag_grouping_resolver: Callable[[], bool] = _force_redis_hash_tag_grouping_from_litellm_settings, ): self.internal_usage_cache = internal_usage_cache self._time_provider = time_provider or datetime.now self._tag_rate_limit_resolver = tag_rate_limit_resolver self._model_group_resolver = model_group_resolver self._fail_closed_resolver = fail_closed_resolver + self._force_hash_tag_grouping_resolver = force_hash_tag_grouping_resolver 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 @@ -1367,12 +1391,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): Group keys by their Redis hash tag to ensure cluster compatibility. For Redis clusters, uses slot calculation to group keys that belong to the same slot. - For regular Redis, no grouping is needed - all keys can be processed together. + For regular Redis, grouping is skipped unless forced by configuration. """ groups: Final[dict[str, list[str]]] = {} - # Use slot calculation for Redis clusters only - if self._is_redis_cluster(): + if self._is_redis_cluster() or self._force_hash_tag_grouping_resolver(): for key in keys: slot = self.keyslot_for_redis_cluster(key) slot_key = f"slot_{slot}" @@ -1509,7 +1532,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) all_cache_values.extend(group_cache_values) - return all_cache_values + return _values_in_caller_order(keys_to_fetch, key_groups, all_cache_values) async def _refund_later_pipelined_groups( self, diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index 783e544eb6c..d9a9ecf807f 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -1890,7 +1890,7 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): Exception( "EVALSHA - all keys must map to the same key slot" ), # First group fails - [1234, 1, 1234, 2], # Second group succeeds + [1234, 2], # Second group succeeds ] handler.batch_rate_limiter_script = mock_script @@ -1910,8 +1910,7 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): keys_to_fetch=test_keys, now_int=1234 ) - # Verify results: 2 from fallback + 4 from successful script = 6 total - assert len(results) == 6, f"Expected 6 results, got {len(results)}" + assert results == [1234, 1, 1234, 2] # Verify script was called twice (once per slot group) assert mock_script.call_count == 2 @@ -6738,16 +6737,111 @@ class _ScriptedRedis: return run -def _handler_with_redis(redis, fail_closed: bool | None = None): +def _handler_with_redis( + redis, + fail_closed: bool | None = None, + force_hash_tag_grouping: bool | None = None, +): internal_usage_cache = InternalUsageCache(DualCache(redis_cache=redis)) # pyright: ignore[reportArgumentType] # duck-typed Redis double - if fail_closed is None: + if fail_closed is None and force_hash_tag_grouping is None: return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=internal_usage_cache) + if fail_closed is None: + return _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=internal_usage_cache, + force_hash_tag_grouping_resolver=lambda: force_hash_tag_grouping, + ) + if force_hash_tag_grouping is None: + return _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=internal_usage_cache, + fail_closed_resolver=lambda: fail_closed, + ) return _PROXY_MaxParallelRequestsHandler( internal_usage_cache=internal_usage_cache, fail_closed_resolver=lambda: fail_closed, + force_hash_tag_grouping_resolver=lambda: force_hash_tag_grouping, ) +@pytest.mark.parametrize("force_hash_tag_grouping", [True, False], ids=["forced", "default"]) +@pytest.mark.asyncio +async def test_redis_batch_rate_limiter_respects_hash_tag_grouping_opt_in(force_hash_tag_grouping: bool): + redis: Final = _ScriptedRedis() + handler: Final = _handler_with_redis(redis, force_hash_tag_grouping=force_hash_tag_grouping) + keys: Final = [ + "{api_key:sk-abc}:window", + "{api_key:sk-abc}:requests", + "{team:t1}:window", + "{team:t1}:requests", + "{end_user:u28551}:window", + "{end_user:u28551}:requests", + ] + + await handler._execute_redis_batch_rate_limiter_script(keys_to_fetch=keys, now_int=1234) + + expected_calls: Final = ( + [ + [ + "{api_key:sk-abc}:window", + "{api_key:sk-abc}:requests", + "{end_user:u28551}:window", + "{end_user:u28551}:requests", + ], + ["{team:t1}:window", "{team:t1}:requests"], + ] + if force_hash_tag_grouping + else [keys] + ) + assert redis.batch_call_keys == expected_calls + + +@pytest.mark.asyncio +async def test_redis_batch_rate_limiter_returns_values_in_caller_order(): + redis: Final = _ScriptedRedis() + handler: Final = _handler_with_redis(redis, force_hash_tag_grouping=True) + now: Final = 1234 + keys_to_fetch: Final = [ + "{api_key:sk-abc}:window", + "{api_key:sk-abc}:requests", + "{team:t1}:window", + "{team:t1}:requests", + "{end_user:u28551}:window", + "{end_user:u28551}:requests", + ] + + results: Final = await handler._execute_redis_batch_rate_limiter_script( + keys_to_fetch=keys_to_fetch, now_int=now + ) + + assert results == [now, 1, now, 2, now, 1] + + +def test_default_hash_tag_grouping_resolver_uses_litellm_setting(monkeypatch: pytest.MonkeyPatch): + redis: Final = _ScriptedRedis() + handler: Final = _handler_with_redis(redis) + keys: Final = [ + "{api_key:sk-abc}:window", + "{api_key:sk-abc}:requests", + "{team:t1}:window", + "{team:t1}:requests", + "{end_user:u28551}:window", + "{end_user:u28551}:requests", + ] + + monkeypatch.setattr(litellm, "force_redis_hash_tag_grouping", True) + assert list(handler._group_keys_by_hash_tag(keys).values()) == [ + [ + "{api_key:sk-abc}:window", + "{api_key:sk-abc}:requests", + "{end_user:u28551}:window", + "{end_user:u28551}:requests", + ], + ["{team:t1}:window", "{team:t1}:requests"], + ] + + monkeypatch.setattr(litellm, "force_redis_hash_tag_grouping", False) + assert handler._group_keys_by_hash_tag(keys) == {"all_keys": keys} + + async def _admit(handler, auth, data=None): await handler.async_pre_call_hook( user_api_key_dict=auth,