diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b6b82e4b376..03ddfa1f7a5 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index b35c7e95a16..5445ba3e2b2 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -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"}] \ No newline at end of file diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index e0d47f5c179..7ebed1b5991 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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]