mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
test_keyslot_for_redis_cluster
This commit is contained in:
parent
67abd8880a
commit
52de33787b
1 changed files with 173 additions and 119 deletions
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue