mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
aeec8703a8
commit
1e88d9955d
3 changed files with 127 additions and 9 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue