[Fix] Parallel Request Limiter v3 - ensure Lua scripts can execute on redis cluster (#14968)

* use hashtag slots for rate limiting logic

* redis_startup_nodes fix

* test_execute_redis_batch_rate_limiter_script_cluster_compatibility

* async_increment_tokens_with_ttl_preservation
This commit is contained in:
Ishaan Jaff 2025-09-26 18:15:57 -07:00 • committed by GitHub
parent bcc7fe74db
commit 36cc98254c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 298 additions and 25 deletions

View file

@ -292,6 +292,71 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return RateLimitResponse(overall_code=overall_code, statuses=statuses)
def _group_keys_by_hash_tag(self, keys: List[str]) -> Dict[str, List[str]]:
"""
Group keys by their Redis hash tag to ensure cluster compatibility.
Keys with the same hash tag will be processed together.
"""
groups = {}
for key in keys:
# Extract hash tag from key like "{api_key:sk-123}:requests"
if "{" in key and "}" in key:
start = key.find("{")
end = key.find("}", start)
hash_tag = key[start:end+1]
else:
# Fallback for keys without hash tags
hash_tag = "no_hash_tag"
if hash_tag not in groups:
groups[hash_tag] = []
groups[hash_tag].append(key)
return groups
async def _execute_redis_batch_rate_limiter_script(
self,
keys_to_fetch: List[str],
now_int: int,
) -> List[Any]:
"""
Execute Redis operations grouped by hash tag for cluster compatibility.
Args:
keys_to_fetch: List[str] - List of keys to fetch
now_int: int - Current timestamp
Returns:
List[Any] - List of cache values
"""
if self.batch_rate_limiter_script is None:
return []
key_groups = self._group_keys_by_hash_tag(keys_to_fetch)
all_cache_values = []
for hash_tag, group_keys in key_groups.items():
try:
group_cache_values = await self.batch_rate_limiter_script(
keys=group_keys,
args=[now_int, self.window_size], # Use integer timestamp
)
all_cache_values.extend(group_cache_values)
except Exception as e:
verbose_proxy_logger.warning(
f"Redis Lua script failed for hash tag {hash_tag}: {str(e)}"
)
# Fallback to in-memory cache for this group
group_cache_values = await self.in_memory_cache_sliding_window(
keys=group_keys,
now_int=now_int,
window_size=self.window_size,
)
all_cache_values.extend(group_cache_values)
return all_cache_values
async def should_rate_limit(
self,
descriptors: List[RateLimitDescriptor],
@ -374,9 +439,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
## IF under limit, check Redis
if self.batch_rate_limiter_script is not None:
cache_values = await self.batch_rate_limiter_script(
keys=keys_to_fetch,
args=[now_int, self.window_size], # Use integer timestamp
# Group keys by hash tag for Redis cluster compatibility
cache_values = await self._execute_redis_batch_rate_limiter_script(
keys_to_fetch=keys_to_fetch,
now_int=now_int,
)
# update in-memory cache with new values
@ -627,6 +693,44 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
return pipeline_operations
async def _execute_token_increment_script(
self,
pipeline_operations: List["RedisPipelineIncrementOperation"],
) -> None:
"""
Execute token increment script grouped by hash tag for cluster compatibility.
"""
if self.token_increment_script is None:
return
# Group operations by hash tag for Redis cluster compatibility
operation_keys = [op["key"] for op in pipeline_operations]
key_groups = self._group_keys_by_hash_tag(operation_keys)
for _hash_tag, group_keys in key_groups.items():
# Get operations for this hash tag group
group_operations = [op for op in pipeline_operations if op["key"] in group_keys]
keys = []
args = []
for op in group_operations:
# Convert None TTL to 0 for Lua script
ttl_value = op["ttl"] if op["ttl"] is not None else 0
verbose_proxy_logger.debug(
f"Executing TTL-preserving increment for key={op['key']}, "
f"increment={op['increment_value']}, ttl={ttl_value}"
)
keys.append(op["key"])
args.extend([op["increment_value"], ttl_value])
await self.token_increment_script(
keys=keys,
args=args,
)
async def async_increment_tokens_with_ttl_preservation(
self,
@ -652,26 +756,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return
try:
# Use Lua script for all operations
keys = []
args = []
for op in pipeline_operations:
# Convert None TTL to 0 for Lua script
ttl_value = op["ttl"] if op["ttl"] is not None else 0
verbose_proxy_logger.debug(
f"Executing TTL-preserving increment for key={op['key']}, "
f"increment={op['increment_value']}, ttl={ttl_value}"
)
keys.append(op["key"])
args.extend([op["increment_value"], ttl_value])
await self.token_increment_script(
keys=keys,
args=args,
)
await self._execute_token_increment_script(pipeline_operations)
verbose_proxy_logger.debug(
f"Successfully executed TTL-preserving increment for {len(pipeline_operations)} keys"
)
@ -708,8 +794,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
_get_parent_otel_span_from_kwargs,
)
from litellm.proxy.common_utils.callback_utils import (
get_metadata_variable_name_from_kwargs,
get_model_group_from_litellm_kwargs,
get_metadata_variable_name_from_kwargs
)
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import ModelResponse, Usage

View file

@ -43,4 +43,8 @@ litellm_settings:
turn_off_message_logging: true
datadog_llm_observability_params:
turn_off_message_logging: true
# proxy_config.yaml
cache: True
cache_params:
type: redis
redis_startup_nodes: [{"host": "127.0.0.1", "port": "7000"}, {"host": "127.0.0.1", "port": "7001"}, {"host": "127.0.0.1", "port": "7002"}, {"host": "127.0.0.1", "port": "7003"}]

View file

@ -1154,3 +1154,186 @@ async def test_async_increment_tokens_fallback_behavior():
# Verify fallback was called
assert fallback_called, "Fallback method should be called when Lua script is not available"
# Redis Cluster Compatibility Tests
def test_group_keys_by_hash_tag():
"""
Test that keys are correctly grouped by Redis hash tag for cluster compatibility.
This ensures that keys with the same hash tag (e.g., {api_key:sk-123}) are grouped
together so they can be processed in the same Redis cluster slot.
"""
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
# Test keys with different hash tags that would cause cluster slot conflicts
test_keys = [
"{api_key:sk-123}:window",
"{api_key:sk-123}:requests",
"{api_key:sk-123}:tokens",
"{user:user-456}:window",
"{user:user-456}:requests",
"{team:team-789}:window",
"{team:team-789}:tokens",
"no_hash_tag_key"
]
# Group the keys
groups = handler._group_keys_by_hash_tag(test_keys)
# Verify correct grouping
expected_groups = {
"{api_key:sk-123}": [
"{api_key:sk-123}:window",
"{api_key:sk-123}:requests",
"{api_key:sk-123}:tokens"
],
"{user:user-456}": [
"{user:user-456}:window",
"{user:user-456}:requests"
],
"{team:team-789}": [
"{team:team-789}:window",
"{team:team-789}:tokens"
],
"no_hash_tag": ["no_hash_tag_key"]
}
assert len(groups) == 4, f"Expected 4 groups, got {len(groups)}"
for expected_tag, expected_keys in expected_groups.items():
assert expected_tag in groups, f"Missing group {expected_tag}"
assert set(groups[expected_tag]) == set(expected_keys), f"Group {expected_tag} keys mismatch"
@pytest.mark.asyncio
async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility():
"""
Test that the Redis batch rate limiter script execution handles cluster compatibility
by grouping keys and falling back gracefully on errors.
This simulates the Redis cluster error scenario and verifies fallback behavior.
"""
from unittest.mock import AsyncMock
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
# Mock script that simulates Redis cluster slot conflict
mock_script = AsyncMock()
mock_script.side_effect = [
Exception("EVALSHA - all keys must map to the same key slot"), # First group fails
[1234, 1, 1234, 2] # Second group succeeds
]
handler.batch_rate_limiter_script = mock_script
# Mock in-memory fallback (returns 2 values for 2 keys: window_start, counter)
handler.in_memory_cache_sliding_window = AsyncMock(return_value=[1234, 1])
# Test keys from different hash tags (would fail in cluster without grouping)
test_keys = [
"{api_key:sk-123}:window",
"{api_key:sk-123}:requests",
"{user:user-456}:window",
"{user:user-456}:requests"
]
# Execute the method
results = await handler._execute_redis_batch_rate_limiter_script(
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)}"
# Verify script was called twice (once per hash tag group)
assert mock_script.call_count == 2
# Verify fallback was called for the failed group
handler.in_memory_cache_sliding_window.assert_called_once()
# Verify the calls were made with grouped keys
call_args_list = mock_script.call_args_list
# First call should have api_key group keys
first_call_keys = call_args_list[0][1]['keys']
assert all(key.startswith("{api_key:sk-123}") for key in first_call_keys)
# Second call should have user group keys
second_call_keys = call_args_list[1][1]['keys']
assert all(key.startswith("{user:user-456}") for key in second_call_keys)
@pytest.mark.asyncio
async def test_execute_token_increment_script_cluster_compatibility():
"""
Test that token increment script execution handles Redis cluster compatibility
by grouping operations by hash tag.
This ensures token increments work correctly in cluster environments.
"""
from typing import List
from unittest.mock import AsyncMock
from litellm.types.caching import RedisPipelineIncrementOperation
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
# Mock script
mock_script = AsyncMock()
handler.token_increment_script = mock_script
# Create pipeline operations with different hash tags
pipeline_operations: List[RedisPipelineIncrementOperation] = [
{
"key": "{api_key:sk-123}:tokens",
"increment_value": 100,
"ttl": 60
},
{
"key": "{api_key:sk-123}:max_parallel_requests",
"increment_value": -1,
"ttl": 60
},
{
"key": "{user:user-456}:tokens",
"increment_value": 50,
"ttl": 60
}
]
# Execute the method
await handler._execute_token_increment_script(pipeline_operations)
# Verify script was called twice (once per hash tag group)
assert mock_script.call_count == 2
call_args_list = mock_script.call_args_list
# Verify first call has api_key operations
first_call_keys = call_args_list[0][1]['keys']
assert len(first_call_keys) == 2
assert all(key.startswith("{api_key:sk-123}") for key in first_call_keys)
# Verify second call has user operations
second_call_keys = call_args_list[1][1]['keys']
assert len(second_call_keys) == 1
assert second_call_keys[0] == "{user:user-456}:tokens"
# Verify args are correctly mapped
first_call_args = call_args_list[0][1]['args']
assert len(first_call_args) == 4 # 2 operations * 2 args each (increment_value, ttl)
assert first_call_args == [100, 60, -1, 60] # increment_value, ttl for each operation
second_call_args = call_args_list[1][1]['args']
assert len(second_call_args) == 2 # 1 operation * 2 args
assert second_call_args == [50, 60]