From 7e0ede8da3435258673cd217bbaf64d2ddb2074a Mon Sep 17 00:00:00 2001 From: Harshit Jain Date: Thu, 8 Jan 2026 13:57:15 +0530 Subject: [PATCH] fix: prevent concurrent requests from bypassing TPM rate limits (#18730) --- .../hooks/parallel_request_limiter_v3.py | 617 +++++++++++++++++- .../hooks/test_tpm_concurrent_bypass_fix.py | 259 ++++++++ 2 files changed, 864 insertions(+), 12 deletions(-) create mode 100644 tests/test_litellm/proxy/hooks/test_tpm_concurrent_bypass_fix.py diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index c416527990e..de3be1e714b 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -107,10 +107,66 @@ end return results """ +# Lua script for atomic TPM token reservation +# This script atomically: +# 1. Increments the token counter by the reservation amount +# 2. Checks if the new value exceeds the limit +# 3. Returns both the new value and whether it's over limit +# +# This eliminates the race condition where another request could slip in +# between increment and limit check. +# +# KEYS: [key1, key2, ...] - Token counter keys +# ARGV: [window_size, tokens_to_reserve, limit1, limit2, ...] +# - window_size: TTL for the keys +# - tokens_to_reserve: Amount to reserve (same for all keys) +# - limit1, limit2, ...: TPM limit for each key +# +# Returns: [new_value1, over_limit1, new_value2, over_limit2, ...] +# where over_limit is 1 if exceeded, 0 otherwise +TPM_RESERVATION_SCRIPT = """ +local results = {} +local window_size = tonumber(ARGV[1]) +local tokens_to_reserve = tonumber(ARGV[2]) + +for i = 1, #KEYS do + local key = KEYS[i] + local limit = tonumber(ARGV[i + 2]) -- Limits start at ARGV[3] + + -- Atomically increment and get new value + local new_value = redis.call('INCRBY', key, tokens_to_reserve) + + -- Set TTL if not already set + local current_ttl = redis.call('TTL', key) + if current_ttl == -1 or current_ttl == -2 then + redis.call('EXPIRE', key, window_size) + end + + -- Check if over limit + local over_limit = 0 + if new_value > limit then + over_limit = 1 + end + + table.insert(results, new_value) + table.insert(results, over_limit) +end + +return results +""" + # Redis cluster slot count REDIS_CLUSTER_SLOTS = 16384 REDIS_NODE_HASHTAG_NAME = "all_keys" +# TPM Token Reservation Constants +# When max_tokens is not specified in the request, use this default for estimation +DEFAULT_MAX_TOKENS_ESTIMATE = 256 +# Fallback: approximate characters per token for rough estimation +DEFAULT_CHARS_PER_TOKEN = 4 +# Metadata key for storing reserved tokens in request data +TPM_RESERVED_TOKENS_KEY = "_litellm_tpm_reserved_tokens" + class RateLimitDescriptorRateLimitObject(TypedDict, total=False): requests_per_unit: Optional[int] @@ -167,7 +223,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.token_increment_script = None self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) - + # Batch rate limiter (lazy loaded) self._batch_rate_limiter: Optional[Any] = None @@ -193,6 +249,110 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """Return the current time for rate limiting calculations.""" return self._time_provider() + def _estimate_tokens_for_request( + self, + data: dict, + model: Optional[str] = None, + ) -> int: + """ + Estimate total tokens that will be used by this request. + Used for token reservation in pre-call hook to prevent concurrent bypass. + + This follows the AWS Bedrock approach: + estimated_tokens = input_tokens + max_tokens + + Handles multiple request types: + - Chat completions (messages) + - Completions (prompt) + - Embeddings (input) + + If max_tokens is not specified, we use a conservative estimate based on + input size to avoid over-reserving while still providing protection. + + Args: + data: Request data containing messages, max_tokens, etc. + model: Model name for accurate token counting + + Returns: + Estimated total tokens for this request + """ + estimated_input_tokens = 0 + + # Try to count input tokens from different request types + messages = data.get("messages") + prompt = data.get("prompt") + input_text = data.get("input") # For embeddings + + if messages: + # Chat completions + try: + from litellm import token_counter + + estimated_input_tokens = token_counter( + model=model or "", + messages=messages, + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Token counting failed, using fallback estimation: {str(e)}" + ) + # Fallback: rough estimate based on character count + total_chars = sum( + len(str(m.get("content", ""))) + for m in messages + if isinstance(m, dict) + ) + estimated_input_tokens = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) + + elif prompt: + # Completions API + if isinstance(prompt, str): + estimated_input_tokens = max(1, len(prompt) // DEFAULT_CHARS_PER_TOKEN) + elif isinstance(prompt, list): + # List of prompts + total_chars = sum(len(str(p)) for p in prompt) + estimated_input_tokens = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) + + elif input_text: + # Embeddings API + if isinstance(input_text, str): + estimated_input_tokens = max( + 1, len(input_text) // DEFAULT_CHARS_PER_TOKEN + ) + elif isinstance(input_text, list): + total_chars = sum(len(str(i)) for i in input_text) + estimated_input_tokens = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) + + # Get max_tokens from request + explicit_max_tokens = data.get("max_tokens") or data.get( + "max_completion_tokens" + ) + + if explicit_max_tokens is not None: + # User specified max_tokens - trust their estimate + max_tokens_estimate = int(explicit_max_tokens) + elif input_text: + # Embeddings don't have output tokens + max_tokens_estimate = 0 + else: + # No max_tokens specified - use conservative estimate + # Use input tokens as a baseline, with a minimum floor + # This balances protection against over-reservation + max_tokens_estimate = max( + estimated_input_tokens, # At least as many as input + DEFAULT_MAX_TOKENS_ESTIMATE // 4, # Minimum 64 tokens + ) + + total_estimated = estimated_input_tokens + max_tokens_estimate + + verbose_proxy_logger.debug( + f"TPM reservation estimate: input={estimated_input_tokens}, " + f"max_tokens={max_tokens_estimate} (explicit={explicit_max_tokens is not None}), " + f"total={total_estimated}" + ) + + return total_estimated + def _is_redis_cluster(self) -> bool: """ Check if the dual cache is using Redis cluster. @@ -588,6 +748,261 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) return rate_limit_response + async def reserve_tpm_tokens( + self, + descriptors: List[RateLimitDescriptor], + estimated_tokens: int, + parent_otel_span: Optional[Span] = None, + ) -> RateLimitResponse: + """ + Reserve estimated tokens for TPM rate limiting BEFORE the request is processed. + + This prevents the concurrent bypass bug by using an atomic Lua script that + increments the token counter AND checks the limit in a single Redis operation. + No other request can slip in between increment and check. + + Args: + descriptors: Rate limit descriptors containing TPM limits + estimated_tokens: Estimated tokens to reserve (input + max_tokens) + parent_otel_span: Optional OpenTelemetry span for tracing + + Returns: + RateLimitResponse indicating if reservation succeeded or exceeded limit + """ + # Collect TPM keys and their limits + tpm_keys: List[str] = [] + tpm_limits: List[int] = [] + descriptor_keys: List[str] = [] + + for descriptor in descriptors: + rate_limit = descriptor.get("rate_limit") or {} + tokens_limit = rate_limit.get("tokens_per_unit") + + if tokens_limit is not None: + tpm_key = self.create_rate_limit_keys( + descriptor["key"], descriptor["value"], "tokens" + ) + tpm_keys.append(tpm_key) + tpm_limits.append(tokens_limit) + descriptor_keys.append(descriptor["key"]) + + if not tpm_keys: + # No TPM limits configured, nothing to reserve + return RateLimitResponse(overall_code="OK", statuses=[]) + + # Try to use atomic Lua script via Redis + redis_cache = self._get_redis_cache() + + if redis_cache is not None: + # Use atomic Lua script for Redis + results = await self._execute_atomic_tpm_reservation( + redis_cache=redis_cache, + keys=tpm_keys, + limits=tpm_limits, + tokens_to_reserve=estimated_tokens, + ) + else: + # Fallback to in-memory cache (less concurrent-safe but works) + results = await self._execute_inmemory_tpm_reservation( + keys=tpm_keys, + limits=tpm_limits, + tokens_to_reserve=estimated_tokens, + ) + + # Build response from results + statuses: List[RateLimitStatus] = [] + overall_code = "OK" + + for i, descriptor_key in enumerate(descriptor_keys): + if i < len(results): + new_value, over_limit = results[i] + limit = tpm_limits[i] + + if over_limit: + overall_code = "OVER_LIMIT" + statuses.append( + RateLimitStatus( + code="OVER_LIMIT", + current_limit=limit, + limit_remaining=limit - new_value, + rate_limit_type="tokens", + descriptor_key=descriptor_key, + ) + ) + else: + statuses.append( + RateLimitStatus( + code="OK", + current_limit=limit, + limit_remaining=limit - new_value, + rate_limit_type="tokens", + descriptor_key=descriptor_key, + ) + ) + + return RateLimitResponse(overall_code=overall_code, statuses=statuses) + + def _get_redis_cache(self): + """Get the Redis cache instance if available.""" + try: + dual_cache = self.internal_usage_cache.dual_cache + if ( + hasattr(dual_cache, "redis_cache") + and dual_cache.redis_cache is not None + ): + return dual_cache.redis_cache + except Exception: + pass + return None + + async def _execute_atomic_tpm_reservation( + self, + redis_cache, + keys: List[str], + limits: List[int], + tokens_to_reserve: int, + ) -> List[tuple]: + """ + Execute atomic TPM reservation using Lua script. + + Returns list of (new_value, over_limit) tuples. + """ + try: + # Build ARGV: [window_size, tokens_to_reserve, limit1, limit2, ...] + argv = [self.window_size, tokens_to_reserve] + limits + + # Execute the Lua script + client = redis_cache.redis_client + if client is None: + # Fallback if no client + return await self._execute_inmemory_tpm_reservation( + keys, limits, tokens_to_reserve + ) + + # Execute script + raw_results = await client.eval( + TPM_RESERVATION_SCRIPT, + len(keys), + *keys, + *argv, + ) + + # Parse results: [new_value1, over_limit1, new_value2, over_limit2, ...] + results = [] + for i in range(0, len(raw_results), 2): + new_value = int(raw_results[i]) + over_limit = bool(int(raw_results[i + 1])) + results.append((new_value, over_limit)) + + verbose_proxy_logger.debug( + f"Atomic TPM reservation: reserved={tokens_to_reserve}, results={results}" + ) + + return results + + except Exception as e: + verbose_proxy_logger.warning( + f"Atomic TPM reservation failed, falling back to non-atomic: {e}" + ) + # Fallback to non-atomic approach + return await self._execute_inmemory_tpm_reservation( + keys, limits, tokens_to_reserve + ) + + async def _execute_inmemory_tpm_reservation( + self, + keys: List[str], + limits: List[int], + tokens_to_reserve: int, + ) -> List[tuple]: + """ + Fallback in-memory TPM reservation (less concurrent-safe). + + Returns list of (new_value, over_limit) tuples. + """ + from litellm.types.caching import RedisPipelineIncrementOperation + + # Create increment operations + pipeline_operations = [ + RedisPipelineIncrementOperation( + key=key, + increment_value=tokens_to_reserve, + ttl=self.window_size, + ) + for key in keys + ] + + # Execute increments + await self.async_increment_tokens_with_ttl_preservation( + pipeline_operations=pipeline_operations, + parent_otel_span=None, + ) + + # Fetch new values + new_values = await self.internal_usage_cache.async_batch_get_cache( + keys=keys, + parent_otel_span=None, + local_only=False, + ) + + # Build results + results = [] + for i, limit in enumerate(limits): + current_value = 0 + if new_values and i < len(new_values) and new_values[i] is not None: + try: + current_value = int(float(new_values[i])) + except (ValueError, TypeError): + current_value = 0 + + over_limit = current_value > limit + results.append((current_value, over_limit)) + + return results + + async def release_tpm_tokens( + self, + descriptors: List[RateLimitDescriptor], + tokens_to_release: int, + parent_otel_span: Optional[Span] = None, + ) -> None: + """ + Release reserved tokens (decrement TPM counter). + + Called when a request fails before completion to refund reserved tokens. + + Args: + descriptors: Rate limit descriptors containing TPM limits + tokens_to_release: Number of tokens to release (negative increment) + parent_otel_span: Optional OpenTelemetry span for tracing + """ + from litellm.types.caching import RedisPipelineIncrementOperation + + pipeline_operations: List[RedisPipelineIncrementOperation] = [] + + for descriptor in descriptors: + rate_limit = descriptor.get("rate_limit") or {} + tokens_limit = rate_limit.get("tokens_per_unit") + + if tokens_limit is not None: + tpm_key = self.create_rate_limit_keys( + descriptor["key"], descriptor["value"], "tokens" + ) + # Negative increment = decrement + pipeline_operations.append( + RedisPipelineIncrementOperation( + key=tpm_key, + increment_value=-tokens_to_release, + ttl=self.window_size, + ) + ) + + if pipeline_operations: + await self.async_increment_tokens_with_ttl_preservation( + pipeline_operations=pipeline_operations, + parent_otel_span=parent_otel_span, + ) + def create_organization_rate_limit_descriptor( self, user_api_key_dict: UserAPIKeyAuth, requested_model: Optional[str] = None ) -> List[RateLimitDescriptor]: @@ -1013,7 +1428,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) # Fail safe: enforce limits if we can't check return True - + def get_rate_limiter_for_call_type(self, call_type: str) -> Optional[Any]: """Get the rate limiter for the call type.""" if call_type == "acreate_batch": @@ -1095,9 +1510,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): now = self._get_current_time().timestamp() reset_time = now + self.window_size - reset_time_formatted = datetime.fromtimestamp( - reset_time - ).strftime("%Y-%m-%d %H:%M:%S UTC") + reset_time_formatted = datetime.fromtimestamp(reset_time).strftime( + "%Y-%m-%d %H:%M:%S UTC" + ) remaining_display = max(0, status["limit_remaining"]) rate_limit_type = status["rate_limit_type"] @@ -1137,7 +1552,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Check if the call type has a specific rate limiter # eg. for Batch APIs we need to use the batch rate limiter to read the input file and count the tokens and requests ######################################################### - call_type_specific_rate_limiter = self.get_rate_limiter_for_call_type(call_type=call_type) + call_type_specific_rate_limiter = self.get_rate_limiter_for_call_type( + call_type=call_type + ) if call_type_specific_rate_limiter: return await call_type_specific_rate_limiter.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -1191,6 +1608,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) # Only check rate limits if we have descriptors with actual limits if descriptors: + # First, check RPM and max_parallel_requests limits response = await self.should_rate_limit( descriptors=descriptors, parent_otel_span=user_api_key_dict.parent_otel_span, @@ -1205,6 +1623,56 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # add descriptors to request headers data["litellm_proxy_rate_limit_response"] = response + ######################################################### + # TPM Token Reservation (prevents concurrent bypass) + # Reserve estimated tokens BEFORE the request is processed + # This is the fix for GitHub Issue #18730 + ######################################################### + + # Check if any descriptor has TPM limits + has_tpm_limits = any( + (d.get("rate_limit") or {}).get("tokens_per_unit") is not None + for d in descriptors + ) + + # Check if request has content we can estimate tokens for + # Supports: chat (messages), completions (prompt), embeddings (input) + has_estimable_content = bool( + data.get("messages") or data.get("prompt") or data.get("input") + ) + + if has_tpm_limits and has_estimable_content: + # Estimate tokens for this request + estimated_tokens = self._estimate_tokens_for_request( + data=data, + model=requested_model, + ) + + # Only reserve if the estimate is meaningful (> 0) + if estimated_tokens > 0: + # Reserve tokens atomically using Lua script (if Redis available) + tpm_response = await self.reserve_tpm_tokens( + descriptors=descriptors, + estimated_tokens=estimated_tokens, + parent_otel_span=user_api_key_dict.parent_otel_span, + ) + + if tpm_response["overall_code"] == "OVER_LIMIT": + # TPM limit exceeded, raise 429 + self._handle_rate_limit_error( + response=tpm_response, + descriptors=descriptors, + ) + else: + # Store reservation info for adjustment in success/failure callbacks + data[TPM_RESERVED_TOKENS_KEY] = estimated_tokens + # Store descriptors for use in callbacks + data["_litellm_rate_limit_descriptors"] = descriptors + + verbose_proxy_logger.debug( + f"TPM tokens reserved: {estimated_tokens} for model {requested_model}" + ) + def _create_pipeline_operations( self, key: str, @@ -1233,7 +1701,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return pipeline_operations - def _get_total_tokens_from_usage(self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"]) -> int: + def _get_total_tokens_from_usage( + self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"] + ) -> int: # Get total tokens from response total_tokens = 0 # spot fix for /responses api @@ -1336,6 +1806,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def get_rate_limit_type(self) -> Literal["output", "input", "total"]: from litellm.proxy.proxy_server import general_settings + specified_rate_limit_type = general_settings.get( "token_rate_limit_type", "total" ) @@ -1381,9 +1852,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): user_api_key_organization_id = standard_logging_metadata.get( "user_api_key_org_id" ) - user_api_key_end_user_id = kwargs.get("user") or standard_logging_metadata.get( - "user_api_key_end_user_id" - ) + user_api_key_end_user_id = kwargs.get( + "user" + ) or standard_logging_metadata.get("user_api_key_end_user_id") model_group = get_model_group_from_litellm_kwargs(kwargs) # Get total tokens from response @@ -1393,7 +1864,28 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): response_obj, BaseLiteLLMOpenAIResponseObject ): _usage = getattr(response_obj, "usage", None) - total_tokens = self._get_total_tokens_from_usage(usage=_usage, rate_limit_type=rate_limit_type) + total_tokens = self._get_total_tokens_from_usage( + usage=_usage, rate_limit_type=rate_limit_type + ) + + ######################################################### + # TPM Token Adjustment (fix for GitHub Issue #18730) + # If tokens were reserved upfront, we need to adjust based on + # the difference between actual and reserved tokens. + # If actual > reserved: we need to add the difference + # If actual < reserved: we need to subtract (refund) the difference + ######################################################### + reserved_tokens = standard_logging_metadata.get(TPM_RESERVED_TOKENS_KEY, 0) + + if reserved_tokens > 0: + # Tokens were reserved upfront - calculate adjustment + token_adjustment = total_tokens - reserved_tokens + verbose_proxy_logger.debug( + f"TPM token adjustment: reserved={reserved_tokens}, " + f"actual={total_tokens}, adjustment={token_adjustment}" + ) + # Use adjustment instead of full total_tokens + total_tokens = token_adjustment # Create pipeline operations for TPM increments pipeline_operations: List[RedisPipelineIncrementOperation] = [] @@ -1544,18 +2036,119 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) + ######################################################### + # TPM Token Release (fix for GitHub Issue #18730) + # If tokens were reserved upfront, we need to release them + # when the request fails (refund the reservation) + ######################################################### + reserved_tokens = standard_logging_metadata.get(TPM_RESERVED_TOKENS_KEY, 0) + + if reserved_tokens > 0: + verbose_proxy_logger.debug( + f"Releasing reserved TPM tokens on failure: {reserved_tokens}" + ) + + # Get user identifiers for releasing tokens across all rate limit keys + user_api_key_user_id = standard_logging_metadata.get( + "user_api_key_user_id" + ) + user_api_key_team_id = standard_logging_metadata.get( + "user_api_key_team_id" + ) + user_api_key_organization_id = standard_logging_metadata.get( + "user_api_key_org_id" + ) + user_api_key_end_user_id = kwargs.get( + "user" + ) or standard_logging_metadata.get("user_api_key_end_user_id") + + # Release tokens for API key + if user_api_key: + tpm_key = self.create_rate_limit_keys( + key="api_key", + value=user_api_key, + rate_limit_type="tokens", + ) + pipeline_operations.append( + RedisPipelineIncrementOperation( + key=tpm_key, + increment_value=-reserved_tokens, + ttl=self.window_size, + ) + ) + + # Release tokens for user + if user_api_key_user_id: + tpm_key = self.create_rate_limit_keys( + key="user", + value=user_api_key_user_id, + rate_limit_type="tokens", + ) + pipeline_operations.append( + RedisPipelineIncrementOperation( + key=tpm_key, + increment_value=-reserved_tokens, + ttl=self.window_size, + ) + ) + + # Release tokens for team + if user_api_key_team_id: + tpm_key = self.create_rate_limit_keys( + key="team", + value=user_api_key_team_id, + rate_limit_type="tokens", + ) + pipeline_operations.append( + RedisPipelineIncrementOperation( + key=tpm_key, + increment_value=-reserved_tokens, + ttl=self.window_size, + ) + ) + + # Release tokens for organization + if user_api_key_organization_id: + tpm_key = self.create_rate_limit_keys( + key="organization", + value=user_api_key_organization_id, + rate_limit_type="tokens", + ) + pipeline_operations.append( + RedisPipelineIncrementOperation( + key=tpm_key, + increment_value=-reserved_tokens, + ttl=self.window_size, + ) + ) + + # Release tokens for end user + if user_api_key_end_user_id: + tpm_key = self.create_rate_limit_keys( + key="end_user", + value=user_api_key_end_user_id, + rate_limit_type="tokens", + ) + pipeline_operations.append( + RedisPipelineIncrementOperation( + key=tpm_key, + increment_value=-reserved_tokens, + ttl=self.window_size, + ) + ) + # 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, ) + except Exception as e: verbose_proxy_logger.exception( f"Error in rate limit failure event: {str(e)}" ) - async def async_post_call_success_hook( self, data: dict, user_api_key_dict: UserAPIKeyAuth, response ): diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent_bypass_fix.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent_bypass_fix.py new file mode 100644 index 00000000000..be77485a114 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent_bypass_fix.py @@ -0,0 +1,259 @@ +""" +Unit Test for TPM Rate Limit Concurrent Bypass Fix +=================================================== + +This test directly verifies that the token reservation mechanism +prevents concurrent requests from bypassing TPM limits. + +It does NOT require a running proxy, database, or Redis - it tests +the core logic in isolation. +""" + +import asyncio +import pytest +from datetime import datetime +from typing import Dict, Any + +from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler, + TPM_RESERVED_TOKENS_KEY, +) +from litellm.proxy.utils import InternalUsageCache, hash_token + + +class TestTPMConcurrentBypassFix: + """Test suite for verifying the TPM concurrent bypass fix.""" + + @pytest.fixture + def rate_limiter(self): + """Create a rate limiter instance for testing.""" + cache = DualCache() + handler = RateLimitHandler(internal_usage_cache=InternalUsageCache(cache)) + return handler, cache + + @pytest.mark.asyncio + async def test_token_reservation_prevents_concurrent_bypass(self, rate_limiter): + """ + Test that token reservation prevents multiple concurrent requests + from bypassing the TPM limit. + + Scenario: + - TPM limit: 100 tokens + - 5 concurrent requests, each estimated to use ~50 tokens + - Without fix: All 5 would pass (check then update) + - With fix: Only 2 should pass (reserve then check) + """ + handler, cache = rate_limiter + + # Create API key auth with TPM limit of 100 + api_key = hash_token("sk-test-key") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + tpm_limit=100, # Low limit to trigger rate limiting + ) + + # Track token reservations + reservations = [] + + # Mock the token increment to track what's being reserved + original_increment = handler.async_increment_tokens_with_ttl_preservation + + async def track_increment(pipeline_operations, **kwargs): + for op in pipeline_operations: + if "tokens" in op["key"]: + reservations.append( + { + "key": op["key"], + "increment": op["increment_value"], + "timestamp": datetime.now(), + } + ) + # Still call the original to update cache + await original_increment(pipeline_operations, **kwargs) + + handler.async_increment_tokens_with_ttl_preservation = track_increment + + # Create request data with messages (to trigger token estimation) + request_data = { + "model": "gpt-3.5-turbo", + "messages": [ + { + "role": "user", + "content": "Hello, this is a test message for concurrent bypass testing.", + } + ], + "max_tokens": 50, # Request ~50 output tokens + } + + # Fire 5 concurrent pre_call_hook requests + async def make_request(request_id: int) -> Dict[str, Any]: + data = request_data.copy() + try: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + return { + "request_id": request_id, + "success": True, + "reserved_tokens": data.get(TPM_RESERVED_TOKENS_KEY, 0), + } + except Exception as e: + return { + "request_id": request_id, + "success": False, + "error": str(e), + "status_code": getattr(e, "status_code", None), + } + + # Fire all requests concurrently + tasks = [make_request(i) for i in range(5)] + results = await asyncio.gather(*tasks) + + # Analyze results + successful = [r for r in results if r["success"]] + failed = [r for r in results if not r["success"]] + rate_limited = [r for r in failed if r.get("status_code") == 429] + + # Results: successful, rate_limited, reservations tracked + + # With the fix in place, we should see some rate limited requests + # because tokens are reserved upfront + assert len(rate_limited) > 0, ( + f"Expected some rate limited requests, but all {len(successful)} succeeded. " + f"This suggests the concurrent bypass bug still exists." + ) + + # Total tokens reserved should not exceed limit by much + _total_reserved = sum(r.get("reserved_tokens", 0) for r in successful) + # Total tokens reserved by successful requests is in _total_reserved + + @pytest.mark.asyncio + async def test_token_adjustment_on_success(self, rate_limiter): + """ + Test that after a successful request, tokens are adjusted based on + actual usage vs reserved. + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-test-adjust") + + # Create mock kwargs for success event + mock_kwargs = { + "standard_logging_object": { + "metadata": { + "user_api_key_hash": api_key, + TPM_RESERVED_TOKENS_KEY: 100, # Reserved 100 tokens + } + }, + "model": "gpt-3.5-turbo", + } + + # Create mock response with actual usage + from litellm.types.utils import ModelResponse, Usage + + mock_response = ModelResponse( + id="test", + object="chat.completion", + created=int(datetime.now().timestamp()), + model="gpt-3.5-turbo", + usage=Usage(prompt_tokens=20, completion_tokens=30, total_tokens=50), + choices=[], + ) + + # Track increments + increments = [] + + async def mock_increment(increment_list, **kwargs): + for op in increment_list: + increments.append( + { + "key": op["key"], + "increment": op["increment_value"], + } + ) + + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) + + # Call success event + await handler.async_log_success_event( + kwargs=mock_kwargs, + response_obj=mock_response, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Find the token adjustment + token_adjustments = [i for i in increments if "tokens" in i["key"]] + + # Token adjustments are in token_adjustments list + + # The adjustment should be actual - reserved = 50 - 100 = -50 + # (Refunding 50 tokens) + assert any( + i["increment"] == -50 for i in token_adjustments + ), f"Expected a -50 token adjustment (refund), but got: {token_adjustments}" + + @pytest.mark.asyncio + async def test_token_release_on_failure(self, rate_limiter): + """ + Test that when a request fails, all reserved tokens are released. + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-test-fail") + + # Create mock kwargs for failure event + mock_kwargs = { + "standard_logging_object": { + "metadata": { + "user_api_key_hash": api_key, + TPM_RESERVED_TOKENS_KEY: 100, # Reserved 100 tokens + } + }, + } + + # Track increments + increments = [] + + async def mock_increment(increment_list, **kwargs): + for op in increment_list: + increments.append( + { + "key": op["key"], + "increment": op["increment_value"], + } + ) + + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) + + # Call failure event + await handler.async_log_failure_event( + kwargs=mock_kwargs, + response_obj=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Find the token releases + token_releases = [i for i in increments if "tokens" in i["key"]] + + # Token releases are in token_releases list + + # Should release all reserved tokens (-100) + assert any( + i["increment"] == -100 for i in token_releases + ), f"Expected all reserved tokens to be released (-100), but got: {token_releases}" + + +if __name__ == "__main__": + # Run tests + pytest.main([__file__, "-v", "-s"])