mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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:
parent
bcc7fe74db
commit
36cc98254c
3 changed files with 298 additions and 25 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue