fix: prevent concurrent requests from bypassing TPM rate limits (#18730)

This commit is contained in:
Harshit Jain 2026-01-08 13:57:15 +05:30
parent e285e2b91d
commit 7e0ede8da3
No known key found for this signature in database
GPG key ID: 36C392CD4415B4CF
2 changed files with 864 additions and 12 deletions

View file

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

View file

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