perf(proxy): bound event-loop blocking from oversized requests (#31497)

Skip token counting in Router._pre_call_checks when no deployment in the
group declares max_input_tokens, and skip the full-body surrogate-repair
regex in _read_request_body above a configurable size, raising the existing
400 immediately.

Resolves LIT-3541
This commit is contained in:
Yassin Kortam 2026-06-27 22:06:48 +03:00 • committed by GitHub
parent 64d8d7f8cb
commit 437acc9b09
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 164 additions and 18 deletions

View file

@ -88,6 +88,12 @@ MAX_IMAGE_URL_DOWNLOAD_SIZE_MB = float(os.getenv("MAX_IMAGE_URL_DOWNLOAD_SIZE_MB
MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB = int(
os.getenv("MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB", 1024)
) # 1MB = 1024KB
# Surrogate-repair fallback in _read_request_body runs two full-body re.sub passes
# that block the event loop on multi-MB malformed bodies. Skip the repair above this
# size and raise the existing 400 immediately. Set to 0 to disable the cap.
MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB = get_env_int(
"MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 1
)
SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD = int(
os.getenv("SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD", 1000)
) # Minimum number of requests to consider "reasonable traffic". Used for single-deployment cooldown logic.

View file

@ -6,6 +6,7 @@ import orjson
from fastapi import Request, UploadFile, status
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB
from litellm.proxy._types import ProxyException
from litellm.proxy.common_utils.callback_utils import (
get_metadata_variable_name_from_kwargs,
@ -95,6 +96,23 @@ async def _read_request_body(request: Optional[Request]) -> Dict:
try:
parsed_body = orjson.loads(body)
except orjson.JSONDecodeError as e:
# The surrogate-repair fallback below runs two full-body re.sub
# passes, which block the event loop on multi-MB malformed bodies.
# Above the configured size, skip the repair and raise the 400 now.
repair_limit_bytes = (
MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB * 1024 * 1024
)
if repair_limit_bytes > 0 and len(body) > repair_limit_bytes:
verbose_proxy_logger.error(
f"Invalid JSON payload received: {str(e)}"
)
raise ProxyException(
message=f"Invalid JSON payload: {str(e)}",
type="invalid_request_error",
param="request_body",
code=status.HTTP_400_BAD_REQUEST,
)
# First try the standard json module which is more forgiving
# First decode bytes to string if needed
body_str = body.decode("utf-8") if isinstance(body, bytes) else body

View file

@ -10740,15 +10740,10 @@ class Router:
invalid_model_indices = set() # Use set for O(1) membership checks
try:
input_tokens = litellm.token_counter(messages=messages)
except Exception as e:
verbose_router_logger.error(
"litellm.router.py::_pre_call_checks: failed to count tokens. Returning initial list of deployments. Got - {}".format(
str(e)
)
)
return _returned_deployments
# Token counting (tiktoken) is the dominant on-loop cost for large prompts.
# Only count when a deployment actually declares max_input_tokens, and count
# at most once; for model groups with no context-window limit it is skipped.
input_tokens: Optional[int] = None
_context_window_error = False
_potential_error_str = ""
@ -10781,20 +10776,29 @@ class Router:
)
_deployment_model = base_model or _litellm_params.get("model", None)
if (
isinstance(model_info, dict)
and model_info.get("max_input_tokens", None) is not None
):
if (
isinstance(model_info["max_input_tokens"], int)
and input_tokens > model_info["max_input_tokens"]
):
max_input_tokens = (
model_info.get("max_input_tokens")
if isinstance(model_info, dict)
else None
)
if isinstance(max_input_tokens, int):
if input_tokens is None:
try:
input_tokens = litellm.token_counter(messages=messages)
except Exception as e:
verbose_router_logger.error(
"litellm.router.py::_pre_call_checks: failed to count tokens. Returning initial list of deployments. Got - {}".format(
str(e)
)
)
return _returned_deployments
if input_tokens > max_input_tokens:
invalid_model_indices.add(idx)
_context_window_error = True
_potential_error_str += (
"Model={}, Max Input Tokens={}, Got={}".format(
_deployment_model,
model_info["max_input_tokens"],
max_input_tokens,
input_tokens,
)
)

View file

@ -447,6 +447,56 @@ async def test_json_parsing_error_handling():
assert result["tools"][0]["type"] == "mcp"
def _make_json_request(body: bytes) -> MagicMock:
mock_request = MagicMock()
mock_request.body = AsyncMock(return_value=body)
mock_request.headers = {"content-type": "application/json"}
mock_request.scope = {}
return mock_request
@pytest.mark.asyncio
async def test_surrogate_repair_skipped_above_size_limit(monkeypatch):
"""
The surrogate-repair fallback runs two full-body re.sub passes that block the
event loop on multi-MB malformed bodies. Above MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB
the repair must be skipped and the existing 400 raised immediately, while bodies
at or below the limit still get repaired.
`\\ud83d` is a lone high-surrogate escape: orjson rejects it, the json fallback
accepts it, so a body containing it is only salvaged when the repair path runs.
"""
import litellm.proxy.common_utils.http_parsing_utils as http_parsing_utils
# Cap the repair at ~100 bytes so the test stays fast and independent of the default.
monkeypatch.setattr(
http_parsing_utils, "MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 100 / (1024 * 1024)
)
small_body = b'{"model":"gpt-4o","x":"\\ud83d"}'
assert len(small_body) <= 100
repaired = await _read_request_body(_make_json_request(small_body))
assert repaired["model"] == "gpt-4o"
padding = "a" * 200
large_body = (
b'{"model":"gpt-4o","pad":"' + padding.encode() + b'","x":"\\ud83d"}'
)
assert len(large_body) > 100
with pytest.raises(ProxyException) as exc_info:
await _read_request_body(_make_json_request(large_body))
assert exc_info.value.code == "400"
assert "Invalid JSON payload" in exc_info.value.message
# Disabling the cap (0) restores repair for the same large body, proving the cap
# — not the malformed content — is what short-circuits the repair.
monkeypatch.setattr(
http_parsing_utils, "MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 0
)
repaired_large = await _read_request_body(_make_json_request(large_body))
assert repaired_large["model"] == "gpt-4o"
@pytest.mark.asyncio
async def test_get_form_data():
"""

View file

@ -2670,6 +2670,74 @@ def test_should_include_deployment():
)
def test_pre_call_checks_skips_token_count_without_max_input_tokens(monkeypatch):
"""
tiktoken token counting is the dominant on-loop cost for large prompts. When no
deployment in the group declares max_input_tokens, the count is never consumed, so
_pre_call_checks must not run it at all.
"""
router = litellm.Router(
model_list=[
{"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
],
enable_pre_call_checks=True,
)
monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {})
calls = []
monkeypatch.setattr(
litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000
)
deployments = [
{"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
{"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d2"}},
]
result = router._pre_call_checks(
model="m",
healthy_deployments=deployments,
messages=[{"role": "user", "content": "hi"}],
)
assert calls == []
assert len(result) == 2
def test_pre_call_checks_counts_once_and_filters_on_max_input_tokens(monkeypatch):
"""
When a deployment declares max_input_tokens the count must still run, be performed
at most once across the group (memoized), and filter deployments whose limit is
exceeded.
"""
router = litellm.Router(
model_list=[
{"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
],
enable_pre_call_checks=True,
)
monkeypatch.setattr(
router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5}
)
calls = []
monkeypatch.setattr(
litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000
)
deployments = [
{"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
{"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d2"}},
]
with pytest.raises(litellm.ContextWindowExceededError):
router._pre_call_checks(
model="m",
healthy_deployments=deployments,
messages=[{"role": "user", "content": "hi"}],
)
assert calls == [1]
def test_get_deployment_model_info_base_model_flow():
"""Test that get_deployment_model_info correctly handles the base model flow"""
from unittest.mock import patch