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 <yassin@berri.ai>
Co-authored-by: Chenglun Hu <chenglunhu@gmail.com>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 11:37:43 -05:00 • committed by GitHub
parent aeec8703a8
commit 1e88d9955d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 127 additions and 9 deletions

View file

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

View file

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

View file

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