mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #14237 from yeahyung/fix/tpm_limit_bug
Fixes #14204 TPM Rate Limit Bug
This commit is contained in:
commit
a8cce5fbf3
2 changed files with 292 additions and 3 deletions
|
|
@ -68,6 +68,32 @@ end
|
|||
return results
|
||||
"""
|
||||
|
||||
TOKEN_INCREMENT_SCRIPT = """
|
||||
local results = {}
|
||||
|
||||
-- Process each key/increment_value/ttl triplet
|
||||
for i = 1, #KEYS do
|
||||
local key = KEYS[i]
|
||||
local increment_value = tonumber(ARGV[i * 2 - 1])
|
||||
local ttl_seconds = tonumber(ARGV[i * 2])
|
||||
|
||||
-- Increment the value
|
||||
local new_value = redis.call('INCRBYFLOAT', key, increment_value)
|
||||
|
||||
-- Handle TTL: only set expire if ttl_seconds > 0 and key has no current TTL
|
||||
-- ttl_seconds can be 0 (no TTL) or positive (set TTL)
|
||||
if ttl_seconds and ttl_seconds > 0 then
|
||||
local current_ttl = redis.call('TTL', key)
|
||||
if current_ttl == -1 then
|
||||
redis.call('EXPIRE', key, ttl_seconds)
|
||||
end
|
||||
end
|
||||
|
||||
table.insert(results, new_value)
|
||||
end
|
||||
|
||||
return results
|
||||
"""
|
||||
|
||||
class RateLimitDescriptorRateLimitObject(TypedDict, total=False):
|
||||
requests_per_unit: Optional[int]
|
||||
|
|
@ -109,8 +135,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
BATCH_RATE_LIMITER_SCRIPT
|
||||
)
|
||||
)
|
||||
self.token_increment_script = (
|
||||
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
|
||||
TOKEN_INCREMENT_SCRIPT
|
||||
)
|
||||
)
|
||||
else:
|
||||
self.batch_rate_limiter_script = None
|
||||
self.token_increment_script = None
|
||||
|
||||
self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60))
|
||||
|
||||
|
|
@ -567,6 +599,62 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
return pipeline_operations
|
||||
|
||||
async def async_increment_tokens_with_ttl_preservation(
|
||||
self,
|
||||
pipeline_operations: List["RedisPipelineIncrementOperation"],
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Increment token counters using Lua script to preserve existing TTL.
|
||||
This prevents TTL reset on every token increment.
|
||||
"""
|
||||
if not pipeline_operations:
|
||||
return
|
||||
|
||||
# Check if script is available
|
||||
if self.token_increment_script is None:
|
||||
verbose_proxy_logger.debug("TTL preservation script not available, using regular pipeline")
|
||||
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
|
||||
increment_list=pipeline_operations,
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Successfully executed TTL-preserving increment for {len(pipeline_operations)} keys"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
f"TTL preservation failed, falling back to regular pipeline: {str(e)}"
|
||||
)
|
||||
# Fallback to regular pipeline on error
|
||||
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
|
||||
increment_list=pipeline_operations,
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
def get_rate_limit_type(self) -> Literal["output", "input", "total"]:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
|
|
@ -713,9 +801,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
# Execute all increments in a single pipeline
|
||||
if pipeline_operations:
|
||||
await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
|
||||
increment_list=pipeline_operations,
|
||||
litellm_parent_otel_span=litellm_parent_otel_span,
|
||||
await self.async_increment_tokens_with_ttl_preservation(
|
||||
pipeline_operations=pipeline_operations,
|
||||
parent_otel_span=litellm_parent_otel_span,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -934,3 +934,204 @@ async def test_team_member_rate_limits_v3():
|
|||
assert team_member_descriptor["value"] == f"{_team_id}:{_user_id}", "Team member value should combine team_id and user_id"
|
||||
assert team_member_descriptor["rate_limit"]["requests_per_unit"] == 10, "Team member RPM limit should be set"
|
||||
assert team_member_descriptor["rate_limit"]["tokens_per_unit"] == 1000, "Team member TPM limit should be set"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_increment_tokens_with_ttl_preservation():
|
||||
"""
|
||||
Test TTL preservation functionality for token increment operations.
|
||||
|
||||
This test verifies that:
|
||||
1. Keys are created with proper TTL on first increment
|
||||
2. TTL is preserved on subsequent increments (not reset)
|
||||
3. Both TTL and non-TTL operations work correctly in the same call
|
||||
|
||||
Environment variables required:
|
||||
- REDIS_HOST: Redis server hostname
|
||||
- REDIS_PORT: Redis server port
|
||||
- REDIS_PASSWORD: Redis password (optional)
|
||||
|
||||
Test scenario:
|
||||
1. First call: Create keys with TTL=60s and TTL=None
|
||||
2. Wait 2 seconds
|
||||
3. Second call: Increment same keys
|
||||
4. Verify TTL decreased but wasn't reset to 60s
|
||||
"""
|
||||
import os
|
||||
import time
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
||||
# Skip test if Redis environment variables are not set
|
||||
redis_host = os.getenv("REDIS_HOST")
|
||||
redis_port = os.getenv("REDIS_PORT")
|
||||
redis_password = os.getenv("REDIS_PASSWORD")
|
||||
|
||||
if not redis_host or not redis_port:
|
||||
pytest.skip("Redis environment variables (REDIS_HOST, REDIS_PORT) not set")
|
||||
|
||||
# Setup Redis cache
|
||||
redis_cache = RedisCache(
|
||||
host=redis_host,
|
||||
port=int(redis_port),
|
||||
password=redis_password,
|
||||
)
|
||||
|
||||
local_cache = DualCache(redis_cache=redis_cache)
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
|
||||
# Verify Redis connection is working
|
||||
try:
|
||||
await redis_cache.ping()
|
||||
except Exception as e:
|
||||
pytest.skip(f"Redis connection failed: {str(e)}")
|
||||
|
||||
# Test keys
|
||||
test_key_with_ttl = "test_ttl_preservation:with_ttl"
|
||||
test_key_without_ttl = "test_ttl_preservation:without_ttl"
|
||||
|
||||
try:
|
||||
# Clean up any existing test keys
|
||||
try:
|
||||
await redis_cache.async_delete_cache(test_key_with_ttl)
|
||||
await redis_cache.async_delete_cache(test_key_without_ttl)
|
||||
except Exception:
|
||||
# Keys might not exist, ignore cleanup errors
|
||||
pass
|
||||
|
||||
# First increment: Create operations with mixed TTL scenarios
|
||||
pipeline_operations_first = [
|
||||
RedisPipelineIncrementOperation(
|
||||
key=test_key_with_ttl,
|
||||
increment_value=10.0,
|
||||
ttl=60
|
||||
),
|
||||
RedisPipelineIncrementOperation(
|
||||
key=test_key_without_ttl,
|
||||
increment_value=5.0,
|
||||
ttl=None # No TTL
|
||||
)
|
||||
]
|
||||
|
||||
# Execute first increment
|
||||
await parallel_request_handler.async_increment_tokens_with_ttl_preservation(
|
||||
pipeline_operations=pipeline_operations_first
|
||||
)
|
||||
|
||||
# Verify keys exist and check initial TTL
|
||||
ttl_after_first = await redis_cache.async_get_ttl(test_key_with_ttl)
|
||||
value_after_first_with_ttl = await redis_cache.async_get_cache(test_key_with_ttl)
|
||||
value_after_first_without_ttl = await redis_cache.async_get_cache(test_key_without_ttl)
|
||||
|
||||
assert value_after_first_with_ttl == 10.0, "First increment should set value to 10.0"
|
||||
assert value_after_first_without_ttl == 5.0, "First increment should set value to 5.0"
|
||||
assert ttl_after_first is not None and ttl_after_first > 0, "Key with TTL should have positive TTL after first increment"
|
||||
assert ttl_after_first <= 60, "TTL should not exceed the set value"
|
||||
|
||||
# Check TTL for key without TTL (should be None, meaning no expiry)
|
||||
ttl_no_ttl_key = await redis_cache.async_get_ttl(test_key_without_ttl)
|
||||
assert ttl_no_ttl_key is None, "Key without TTL should have no expiry (None from async_get_ttl)"
|
||||
|
||||
# Wait a moment to ensure TTL decreases
|
||||
await asyncio.sleep(2)
|
||||
|
||||
# Second increment: Same operations to test TTL preservation
|
||||
pipeline_operations_second = [
|
||||
RedisPipelineIncrementOperation(
|
||||
key=test_key_with_ttl,
|
||||
increment_value=15.0,
|
||||
ttl=60 # Same TTL value
|
||||
),
|
||||
RedisPipelineIncrementOperation(
|
||||
key=test_key_without_ttl,
|
||||
increment_value=7.0,
|
||||
ttl=None # No TTL
|
||||
)
|
||||
]
|
||||
|
||||
# Execute second increment
|
||||
await parallel_request_handler.async_increment_tokens_with_ttl_preservation(
|
||||
pipeline_operations=pipeline_operations_second
|
||||
)
|
||||
|
||||
# Verify TTL preservation and value updates
|
||||
ttl_after_second = await redis_cache.async_get_ttl(test_key_with_ttl)
|
||||
value_after_second_with_ttl = await redis_cache.async_get_cache(test_key_with_ttl)
|
||||
value_after_second_without_ttl = await redis_cache.async_get_cache(test_key_without_ttl)
|
||||
|
||||
assert value_after_second_with_ttl == 25.0, "Second increment should update value to 25.0"
|
||||
assert value_after_second_without_ttl == 12.0, "Second increment should update value to 12.0"
|
||||
|
||||
# Critical test: TTL should be preserved (not reset to 60)
|
||||
assert ttl_after_second is not None, "TTL should still exist"
|
||||
assert ttl_after_second < ttl_after_first, "TTL should have decreased (not been reset)"
|
||||
assert ttl_after_second > 0, "TTL should still be positive"
|
||||
|
||||
# TTL should not be close to the original 60 seconds (proving it wasn't reset)
|
||||
assert ttl_after_second < 59, "TTL should be significantly less than original, proving preservation"
|
||||
|
||||
# Key without TTL should still have no expiry
|
||||
ttl_no_ttl_key_after_second = await redis_cache.async_get_ttl(test_key_without_ttl)
|
||||
assert ttl_no_ttl_key_after_second is None, "Key without TTL should still have no expiry"
|
||||
|
||||
finally:
|
||||
# Clean up test keys
|
||||
try:
|
||||
await redis_cache.async_delete_cache(test_key_with_ttl)
|
||||
await redis_cache.async_delete_cache(test_key_without_ttl)
|
||||
except Exception:
|
||||
# Ignore cleanup errors
|
||||
pass
|
||||
|
||||
# Properly close Redis connections to prevent warnings
|
||||
try:
|
||||
await redis_cache.disconnect()
|
||||
except Exception:
|
||||
# Ignore disconnect errors
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_increment_tokens_fallback_behavior():
|
||||
"""
|
||||
Test fallback behavior when Lua script is not available.
|
||||
"""
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
|
||||
# Mock the token_increment_script to None to simulate unavailable script
|
||||
parallel_request_handler.token_increment_script = None
|
||||
|
||||
# Mock the fallback method
|
||||
fallback_called = False
|
||||
original_method = parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline
|
||||
|
||||
async def mock_fallback(*args, **kwargs):
|
||||
nonlocal fallback_called
|
||||
fallback_called = True
|
||||
return await original_method(*args, **kwargs)
|
||||
|
||||
parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_fallback
|
||||
|
||||
# Test operations
|
||||
pipeline_operations = [
|
||||
RedisPipelineIncrementOperation(
|
||||
key="test_fallback_key",
|
||||
increment_value=10.0,
|
||||
ttl=60
|
||||
)
|
||||
]
|
||||
|
||||
# Execute increment
|
||||
await parallel_request_handler.async_increment_tokens_with_ttl_preservation(
|
||||
pipeline_operations=pipeline_operations
|
||||
)
|
||||
|
||||
# Verify fallback was called
|
||||
assert fallback_called, "Fallback method should be called when Lua script is not available"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue