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 511eb5bbb89..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 @@ -1157,18 +1157,19 @@ async def test_async_increment_tokens_fallback_behavior(): # Redis Cluster Compatibility Tests -def test_group_keys_by_hash_tag_regular_redis(): +def test_group_keys_by_hash_tag(): """ - Test that keys are correctly grouped for regular Redis (non-cluster). + Test that keys are correctly grouped by Redis hash tag for cluster compatibility. - For regular Redis, all keys should be grouped together under a single group. + 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 + # Test keys with different hash tags that would cause cluster slot conflicts test_keys = [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", @@ -1180,77 +1181,32 @@ def test_group_keys_by_hash_tag_regular_redis(): "no_hash_tag_key" ] - # Group the keys (should be single group for regular Redis) + # Group the keys groups = handler._group_keys_by_hash_tag(test_keys) - # 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 = [ + # 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", - ] - - # 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) - ) + "{user:user-456}:requests" + ], + "{team:team-789}": [ + "{team:team-789}:window", + "{team:team-789}:tokens" + ], + "no_hash_tag": ["no_hash_tag_key"] + } - # Test basic key - slot1 = handler.keyslot_for_redis_cluster("user:1000") - assert 0 <= slot1 < 16384, "Slot should be in valid range" + assert len(groups) == 4, f"Expected 4 groups, got {len(groups)}" - # 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" + 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 @@ -1261,76 +1217,69 @@ 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, patch + from unittest.mock import AsyncMock local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - # 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" + # 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 slot. + by grouping operations by hash tag. This ensures token increments work correctly in cluster environments. """ from typing import List - from unittest.mock import AsyncMock, patch + from unittest.mock import AsyncMock from litellm.types.caching import RedisPipelineIncrementOperation @@ -1339,55 +1288,52 @@ async def test_execute_token_increment_script_cluster_compatibility(): internal_usage_cache=InternalUsageCache(local_cache) ) - # 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" + # 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 } - 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" + ] + + # 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]