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 7ebed1b5991..511eb5bbb89 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 @@ -1157,19 +1157,18 @@ async def test_async_increment_tokens_fallback_behavior(): # Redis Cluster Compatibility Tests -def test_group_keys_by_hash_tag(): +def test_group_keys_by_hash_tag_regular_redis(): """ - Test that keys are correctly grouped by Redis hash tag for cluster compatibility. + Test that keys are correctly grouped for regular Redis (non-cluster). - 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. + For regular Redis, all keys should be grouped together under a single group. """ 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 with different hash tags test_keys = [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", @@ -1181,32 +1180,77 @@ def test_group_keys_by_hash_tag(): "no_hash_tag_key" ] - # Group the keys + # Group the keys (should be single group for regular Redis) groups = handler._group_keys_by_hash_tag(test_keys) - # Verify correct grouping - expected_groups = { - "{api_key:sk-123}": [ + # Verify all keys are in single group for regular Redis + assert len(groups) == 1, f"Expected 1 group for regular Redis, got {len(groups)}" + assert "all_keys" in groups, "Expected 'all_keys' group for regular Redis" + assert set(groups["all_keys"]) == set(test_keys), "All keys should be in single group" + + +def test_group_keys_by_hash_tag_redis_cluster(): + """ + Test that keys are correctly grouped by Redis cluster slots when using Redis cluster. + + This ensures that keys are grouped by their slot number for cluster compatibility. + """ + from unittest.mock import patch + + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + # Mock _is_redis_cluster to return True + with patch.object(handler, '_is_redis_cluster', return_value=True): + # Test keys with different hash tags + test_keys = [ "{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"] - } + "{user:user-456}:requests", + ] + + # Group the keys (should be grouped by slot for Redis cluster) + groups = handler._group_keys_by_hash_tag(test_keys) + + # Verify keys are grouped by slot + assert len(groups) >= 1, "Should have at least 1 slot group" + + # All group keys should start with "slot_" + for group_key in groups.keys(): + assert group_key.startswith("slot_"), f"Group key {group_key} should start with 'slot_'" + + # Verify all original keys are present across groups + all_grouped_keys = [] + for group_keys in groups.values(): + all_grouped_keys.extend(group_keys) + assert set(all_grouped_keys) == set(test_keys), "All keys should be present in groups" + + +def test_keyslot_for_redis_cluster(): + """ + Test the keyslot calculation for Redis cluster. + """ + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) - assert len(groups) == 4, f"Expected 4 groups, got {len(groups)}" + # Test basic key + slot1 = handler.keyslot_for_redis_cluster("user:1000") + assert 0 <= slot1 < 16384, "Slot should be in valid range" - 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" + # Test key with hash tag + slot2 = handler.keyslot_for_redis_cluster("foo{bar}baz") + slot3 = handler.keyslot_for_redis_cluster("{bar}") + assert slot2 == slot3, "Keys with same hash tag should have same slot" + + # Test keys with same hash tag should have same slot + slot4 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:requests") + slot5 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:window") + assert slot4 == slot5, "Keys with same hash tag should have same slot" @pytest.mark.asyncio @@ -1217,69 +1261,76 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): This simulates the Redis cluster error scenario and verifies fallback behavior. """ - from unittest.mock import AsyncMock + from unittest.mock import AsyncMock, patch 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) + # Mock _is_redis_cluster to return True for this test + with patch.object(handler, '_is_redis_cluster', return_value=True): + # 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 slot 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 + + # Both calls should have keys, but we can't predict exact grouping without knowing slots + # Just verify that keys were grouped and calls were made + assert len(call_args_list) == 2, "Should have made 2 script calls" + + # Verify all keys were processed + all_processed_keys = [] + for call_args in call_args_list: + all_processed_keys.extend(call_args[1]['keys']) + + # Should have processed all keys (some might be duplicated due to fallback) + unique_processed_keys = set(all_processed_keys) + assert len(unique_processed_keys) >= 2, "Should have processed at least some 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. + by grouping operations by slot. This ensures token increments work correctly in cluster environments. """ from typing import List - from unittest.mock import AsyncMock + from unittest.mock import AsyncMock, patch from litellm.types.caching import RedisPipelineIncrementOperation @@ -1288,52 +1339,55 @@ async def test_execute_token_increment_script_cluster_compatibility(): 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 + # Mock _is_redis_cluster to return True for this test + with patch.object(handler, '_is_redis_cluster', return_value=True): + # 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 (at least once, possibly more depending on slot grouping) + assert mock_script.call_count >= 1, "Script should be called at least once" + + call_args_list = mock_script.call_args_list + + # Verify all operations were processed + all_processed_keys = [] + for call_args in call_args_list: + all_processed_keys.extend(call_args[1]['keys']) + + # Should have processed all 3 keys + expected_keys = { + "{api_key:sk-123}:tokens", + "{api_key:sk-123}:max_parallel_requests", + "{user:user-456}:tokens" } - ] - - # 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] + assert set(all_processed_keys) == expected_keys, "All operation keys should be processed" + + # Verify args structure is correct for each call + for call_args in call_args_list: + keys = call_args[1]['keys'] + args = call_args[1]['args'] + # Each key should have 2 args (increment_value, ttl) + assert len(args) == len(keys) * 2, f"Each key should have 2 args, got {len(args)} args for {len(keys)} keys"