From 9e3010daa4342c731e5e929869a5227e38b9a92e Mon Sep 17 00:00:00 2001 From: yeahyung Date: Thu, 4 Sep 2025 16:41:23 +0900 Subject: [PATCH 1/2] (#14204) increase token usage with TTL preservation --- .../hooks/parallel_request_limiter_v3.py | 94 ++++++++++++++++++- 1 file changed, 91 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b04d14bcc8b..b3840761d2a 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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: From 7829de294816214c9582526e90fc60c139410774 Mon Sep 17 00:00:00 2001 From: yeahyung Date: Thu, 4 Sep 2025 16:41:29 +0900 Subject: [PATCH 2/2] (#14204) add test code --- .../hooks/test_parallel_request_limiter_v3.py | 201 ++++++++++++++++++ 1 file changed, 201 insertions(+) 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 694a49159c0..da4218a9547 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 @@ -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"