fix(proxy): count Responses images safely for TPM limits

This commit is contained in:
Devin AI 2026-07-11 03:17:37 +00:00
parent ee6e8077ae
commit f920a43efa
2 changed files with 96 additions and 23 deletions

View file

@ -23,6 +23,8 @@ from typing import (
cast,
)
from pydantic import TypeAdapter
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE
@ -30,6 +32,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
from litellm.litellm_core_utils.token_counter import token_counter
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
get_key_tag_rpm_limit,
@ -42,11 +45,18 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import (
map_v3_rate_limit_type,
)
from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject
from litellm.types.llms.openai import (
AllMessageValues,
BaseLiteLLMOpenAIResponseObject,
)
from litellm.types.utils import (
CallTypes,
EmbeddingResponse,
Message,
ModelResponse,
TextCompletionResponse,
Usage,
@ -377,6 +387,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
data: dict,
model: Optional[str] = None,
min_configured_tpm_limit: Optional[int] = None,
call_type: Optional[str] = None,
) -> int:
"""
Estimate total tokens this request will consume so we can reserve them
@ -393,33 +404,55 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
messages = data.get("messages")
prompt = data.get("prompt")
input_text = data.get("input") # embeddings
input_text = data.get("input")
is_responses_request = call_type in ("responses", "aresponses")
match (messages, prompt, input_text):
case (messages, _, _) if messages:
total_chars = len(get_str_from_messages(messages))
case (_, str() as p, _):
total_chars = len(p)
case (_, list() as p, _):
total_chars = sum(len(str(item)) for item in p)
case (_, _, str() as t):
total_chars = len(t)
case (_, _, list() as t):
total_chars = sum(len(str(item)) for item in t)
case _:
total_chars = 0
if is_responses_request and input_text is not None:
response_messages = TypeAdapter(List[Union[AllMessageValues, Message]]).validate_python(
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=input_text,
responses_api_request=data,
)
)
estimated_input_tokens = token_counter(
model=model or "",
messages=response_messages,
tools=data.get("tools"),
tool_choice=data.get("tool_choice"),
use_default_image_token_count=True,
)
else:
match (messages, prompt, input_text):
case (messages, _, _) if messages:
total_chars = len(get_str_from_messages(messages))
case (_, str() as p, _):
total_chars = len(p)
case (_, list() as p, _):
total_chars = sum(len(str(item)) for item in p)
case (_, _, str() as t):
total_chars = len(t)
case (_, _, list() as t):
total_chars = sum(len(str(item)) for item in t)
case _:
total_chars = 0
estimated_input_tokens = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) if total_chars > 0 else 0
estimated_input_tokens = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) if total_chars > 0 else 0
explicit_max_tokens = data.get("max_tokens") or data.get("max_completion_tokens")
explicit_max_tokens = next(
(
data.get(key)
for key in ("max_tokens", "max_completion_tokens", "max_output_tokens")
if data.get(key) is not None
),
None,
)
match (explicit_max_tokens, input_text):
case (mt, _) if mt is not None:
max_tokens_estimate = int(mt)
case (_, embeddings_input) if embeddings_input:
# Embeddings have no output tokens
case (_, embeddings_input) if embeddings_input and not is_responses_request:
max_tokens_estimate = 0
case _ if total_chars == 0:
case _ if estimated_input_tokens == 0:
# Fully contentless request (no messages, prompt, or input).
# Don't apply the conservative output-budget floor here — it
# would over-reserve and could push small TPM limits into a
@ -2067,11 +2100,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
capped_floor = self._no_max_tokens_output_floor(min_configured_tpm_limit)
baseline_floor = DEFAULT_MAX_TOKENS_ESTIMATE // _TPM_FLOOR_FRACTION
has_explicit_max_tokens = (
data.get("max_tokens") is not None or data.get("max_completion_tokens") is not None
data.get("max_tokens") is not None
or data.get("max_completion_tokens") is not None
or data.get("max_output_tokens") is not None
)
is_embedding = data.get("input") is not None
is_responses_request = call_type in ("responses", "aresponses")
is_embedding = data.get("input") is not None and not is_responses_request
if capped_floor < baseline_floor and not has_explicit_max_tokens and not is_embedding:
data["max_tokens"] = capped_floor
max_tokens_key = "max_output_tokens" if is_responses_request else "max_tokens"
data[max_tokens_key] = capped_floor
# Floor at 1 token so contentless requests (/responses,
# tool-call continuations, empty messages) still flow
@ -2085,6 +2122,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
data=data,
model=requested_model,
min_configured_tpm_limit=min_configured_tpm_limit,
call_type=call_type,
),
1,
)

View file

@ -1047,6 +1047,41 @@ async def test_estimate_tokens_floor_unchanged_when_kwarg_omitted(rate_limiter):
assert estimate == 1 + 1024
@pytest.mark.asyncio
async def test_responses_base64_image_does_not_exhaust_tpm(rate_limiter):
handler, cache = rate_limiter
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token("sk-responses-image"),
tpm_limit=50_000,
)
data = {
"model": "gpt-4o",
"input": [
{
"role": "user",
"content": [
{"type": "input_text", "text": "Describe this image"},
{
"type": "input_image",
"image_url": "data:image/png;base64,"
+ ("A" * 2_500_000),
},
],
}
],
}
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="aresponses",
)
reserved_tokens = (data.get("metadata") or {})[TPM_RESERVED_TOKENS_KEY]
assert 0 < reserved_tokens < user_api_key_dict.tpm_limit
@pytest.mark.asyncio
async def test_small_tpm_cap_admits_no_max_tokens_request(rate_limiter):
"""