fix(proxy): preserve decision request bodies under token limits (#43920)

This commit is contained in:
shrey-berri 2026-10-01 09:23:33 -07:00 • committed by GitHub
parent c2ae483782
commit d96477abce
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 123 additions and 12 deletions

View file

@ -75,6 +75,7 @@ from litellm.router_utils.add_retry_fallback_headers import (
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
from litellm.types.utils import (
CallTypes,
EmbeddingResponse,
@ -978,6 +979,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
min_configured_limit: int | None,
call_type: str | None,
configured_output_tokens: int | None = None,
endpoint_type: EndpointType = EndpointType.GENERIC,
) -> None:
"""Hard-cap generation length when the request has no explicit cap.
@ -1006,6 +1008,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
capped_floor >= baseline_floor
or _PROXY_MaxParallelRequestsHandler_v3._has_explicit_output_cap(data, call_type)
or is_embedding
or endpoint_type == EndpointType.DECISIONS
):
return
effective_cap: Final = max(capped_floor, configured_output_tokens or 0)
@ -3757,6 +3760,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
tpm_reservation_scopes: Sequence[tuple[str, str]],
tpm_reservation_amount: int,
call_type: str | None = None,
endpoint_type: EndpointType = EndpointType.GENERIC,
) -> None:
"""
Reserve project-scoped ITPM/OTPM tokens (Bedrock Mantle-style
@ -3810,6 +3814,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
data=data,
min_configured_limit=min_configured_otpm_limit,
call_type=call_type,
endpoint_type=endpoint_type,
)
io_response, itpm_reserved, otpm_reserved = await self.reserve_io_tokens(
@ -3984,6 +3989,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
cache: DualCache,
data: dict,
call_type: str,
endpoint_type: EndpointType = EndpointType.GENERIC,
):
"""
Pre-call hook to check rate limits before making the API call.
@ -4115,6 +4121,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
min_configured_limit=min_configured_tpm_limit,
call_type=call_type,
configured_output_tokens=configured_output_tokens,
endpoint_type=endpoint_type,
)
# Floor at 1 token so contentless requests (/responses,
@ -4201,6 +4208,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
tpm_reservation_scopes=tpm_reservation_scopes,
tpm_reservation_amount=tpm_reservation_amount,
call_type=call_type,
endpoint_type=endpoint_type,
)
def _create_pipeline_operations(

View file

@ -377,8 +377,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
return return_headers
@staticmethod
def get_endpoint_type(url: str) -> EndpointType:
def get_endpoint_type(url: str, custom_llm_provider: str | None = None) -> EndpointType:
parsed_url: Final = urlparse(url)
if custom_llm_provider == "typesafe" and parsed_url.path.removesuffix("/").endswith("/v1/systemone"):
return EndpointType.DECISIONS
if (
("generateContent") in url
or ("streamGenerateContent") in url
@ -1093,7 +1095,9 @@ async def pass_through_request(
requested_query_params: dict | None = query_params or dict(request.query_params) or None
endpoint_type: Final[EndpointType] = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url))
endpoint_type: Final[EndpointType] = HttpPassThroughEndpointHelpers.get_endpoint_type(
str(url), custom_llm_provider
)
# SigV4-signed callers (e.g. Bedrock) attach the exact bytes that were
# signed via request.state; we must send those instead of re-encoding the
@ -1180,6 +1184,7 @@ async def pass_through_request(
user_api_key_dict=user_api_key_dict,
data=_parsed_body,
call_type="pass_through_endpoint",
endpoint_type=endpoint_type,
)
resolved_timeout: Final = resolve_pass_through_request_timeout(timeout)
async_client_obj: Final = get_async_httpx_client(

View file

@ -245,6 +245,7 @@ from litellm.types.mcp import (
MCPPreCallRequestObject,
MCPPreCallResponseObject,
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult
from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams
from litellm.utils import (
@ -2353,6 +2354,7 @@ class ProxyLogging:
call_type: CallTypesLiteral,
guardrails_only: bool = False,
skip_guardrails: bool = False,
endpoint_type: EndpointType = EndpointType.GENERIC,
) -> None:
pass
@ -2364,6 +2366,7 @@ class ProxyLogging:
call_type: CallTypesLiteral,
guardrails_only: bool = False,
skip_guardrails: bool = False,
endpoint_type: EndpointType = EndpointType.GENERIC,
) -> dict:
pass
@ -2374,6 +2377,7 @@ class ProxyLogging:
call_type: CallTypesLiteral,
guardrails_only: bool = False,
skip_guardrails: bool = False,
endpoint_type: EndpointType = EndpointType.GENERIC,
) -> dict | None:
"""
Allows users to modify/reject the incoming request to the proxy, without having to deal with parsing Request body.
@ -2512,11 +2516,21 @@ class ProxyLogging:
if call_type in MCP_GUARDRAIL_CALL_TYPES and user_api_key_dict is None:
continue
response: Exception | str | Mapping[str, object] | None = await _callback.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=self.call_details["user_api_key_cache"],
data=data,
call_type=call_type,
response: Exception | str | Mapping[str, object] | None = (
await _callback.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=self.call_details["user_api_key_cache"],
data=data,
call_type=call_type,
endpoint_type=endpoint_type,
)
if isinstance(_callback, _PROXY_MaxParallelRequestsHandler_v3)
else await _callback.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=self.call_details["user_api_key_cache"],
data=data,
call_type=call_type,
)
)
if response is not None:
data = await self.process_pre_call_hook_response(

View file

@ -30,6 +30,7 @@ class EndpointType(str, Enum):
OPENAI = "openai"
TINYFISH = "tinyfish"
GENERIC = "generic"
DECISIONS = "decisions"
class PassthroughStandardLoggingPayload(TypedDict, total=False):

View file

@ -7,7 +7,7 @@ import os
import traceback
from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Mapping
from types import MappingProxyType, SimpleNamespace
from typing import Final
from typing import Final, Literal
from unittest import mock
from unittest.mock import AsyncMock, MagicMock, Mock, patch
from urllib.parse import parse_qs
@ -7292,6 +7292,82 @@ class TestTypeSafePassthroughRoute:
assert sent.headers["authorization"] == "Bearer typesafe-test-key"
assert json.loads(sent.content or b"{}") == (body or {})
@pytest.mark.parametrize(
"provider, endpoint, is_decision_request",
(
("typesafe", "systemone", True),
("typesafe", "systemone/", True),
("typesafe", "systemone?trace=1", True),
("typesafe", "systemone/?trace=1", True),
("typesafe", "systemone/other", False),
("typesafe", "systemone/other/", False),
("typesafe", "systemone-other", False),
("typesafe", "chat/completions?next=/typesafe/v1/systemone", False),
("openrouter", "systemone", False),
("openrouter", "systemone/", False),
("openrouter", "chat/completions", False),
),
)
@pytest.mark.parametrize("quota_scope", ("key", "project_output"))
@pytest.mark.parametrize("token_limit", (0, 1000))
def test_token_limits_preserve_decisions_cap_generation_and_enforce_quota(
self,
client: TestClient,
monkeypatch: pytest.MonkeyPatch,
provider: Literal["typesafe", "openrouter"],
endpoint: str,
is_decision_request: bool,
quota_scope: Literal["key", "project_output"],
token_limit: int,
) -> None:
from litellm.caching.caching import DualCache
from litellm.proxy import proxy_server
from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
get_request_stash,
)
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
cache: Final = DualCache()
limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache))
monkeypatch.setattr(litellm, "callbacks", list((limiter, _PROXY_CacheControlCheck())))
monkeypatch.setattr(proxy_server, "proxy_logging_obj", ProxyLogging(user_api_key_cache=cache))
monkeypatch.setenv("OPENROUTER_API_KEY", "openrouter-test-key")
monkeypatch.setenv("OPENROUTER_API_BASE", "https://typesafe.example/base")
model: Final = "jev-latest" if provider == "typesafe" else "test-generative-model"
auth: Final = UserAPIKeyAuth(
api_key="sk-limited",
tpm_limit=token_limit if quota_scope == "key" else None,
project_id="test-project" if quota_scope == "project_output" else None,
project_metadata={"model_otpm_limit": {model: token_limit}} if quota_scope == "project_output" else {},
)
monkeypatch.setitem(proxy_server.app.dependency_overrides, user_api_key_auth, lambda: auth)
body: Final = (
{
"model": model,
"state": "A request for help",
"questions": {"urgent": {"type": "noul", "instructions": "Is this urgent?"}},
}
if is_decision_request
else {"model": model, "messages": [{"role": "user", "content": "Hello"}]}
)
def upstream_response(request: httpx.Request) -> httpx.Response:
expected_body: Final = body if is_decision_request else {**body, "max_tokens": token_limit // 4}
assert json.loads(request.content) == expected_body
stash: Final = get_request_stash()
assert stash is not None
assert (stash.reserved_tokens if quota_scope == "key" else stash.otpm_reserved_tokens) > 0
return httpx.Response(200, json={"model": model})
with respx.mock(assert_all_called=False) as upstream:
route: Final = upstream.post(f"https://typesafe.example/base/v1/{endpoint}").mock(side_effect=upstream_response)
response: Final = client.post(f"/{provider}/v1/{endpoint}", json=body)
assert response.status_code == (429 if token_limit == 0 else 200), response.text
assert route.call_count == (0 if token_limit == 0 else 1)
@pytest.mark.asyncio
async def test_forwards_target_auth_headers_provider_and_query(self, monkeypatch):
monkeypatch.setenv("TYPESAFE_API_KEY", "typesafe-test-key")

View file

@ -49,6 +49,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.types import utils as types_utils
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
EndpointType,
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
@ -1689,7 +1690,7 @@ async def test_pass_through_request_streamed_response_is_owned_by_the_caller():
cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler)))
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type, endpoint_type: data)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
mock_proxy_logging.get_proxy_hook = MagicMock(return_value=MagicMock())
@ -2591,7 +2592,9 @@ async def _run_pass_through_and_capture_wire_url(
mock_request.body = AsyncMock(return_value=b"")
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
mock_proxy_logging.pre_call_hook = AsyncMock(
side_effect=lambda user_api_key_dict, data, call_type, endpoint_type=None: data
)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
mock_proxy_logging.get_proxy_hook = MagicMock(return_value=managed_files_hook)
@ -4889,7 +4892,9 @@ async def test_pass_through_request_mid_stream_upstream_drop_fires_failure_hook(
cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler)))
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
mock_proxy_logging.pre_call_hook = AsyncMock(
side_effect=lambda user_api_key_dict, data, call_type, endpoint_type=None: data
)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_proxy_logging.get_proxy_hook = MagicMock(return_value=None)
@ -7093,7 +7098,9 @@ async def _drive_passthrough_request_and_capture_logging(
captured_data: dict = {} # mutable-ok: the pre-call hook records the request data into it
async def capture_pre_call_hook(user_api_key_dict, data, call_type):
async def capture_pre_call_hook(
user_api_key_dict, data, call_type, endpoint_type: EndpointType = EndpointType.GENERIC
):
captured_data.update(data)
if on_pre_call is not None:
on_pre_call(data.get("litellm_logging_obj"))