Merge pull request #14237 from yeahyung/fix/tpm_limit_bug

Fixes #14204 TPM Rate Limit Bug
This commit is contained in:
Krish Dholakia 2025-09-04 07:29:47 -07:00 • committed by GitHub
commit a8cce5fbf3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 292 additions and 3 deletions

View file

@ -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:

View file

@ -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"