From d96477abce06f7591be19b2429db13ac3e955aa0 Mon Sep 17 00:00:00 2001 From: shrey-berri Date: Thu, 1 Oct 2026 09:23:33 -0700 Subject: [PATCH] fix(proxy): preserve decision request bodies under token limits (#43920) --- .../hooks/parallel_request_limiter_v3.py | 8 ++ .../pass_through_endpoints.py | 9 ++- litellm/proxy/utils.py | 24 ++++-- .../pass_through_endpoints.py | 1 + .../test_llm_pass_through_endpoints.py | 78 ++++++++++++++++++- .../test_pass_through_endpoints.py | 15 +++- 6 files changed, 123 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 36896bd6d44..c8fe49afcc9 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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( diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 2e7f9c4a41c..d15a14bb4b2 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index c0e36e6e172..0af4a5cce8a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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( diff --git a/litellm/types/passthrough_endpoints/pass_through_endpoints.py b/litellm/types/passthrough_endpoints/pass_through_endpoints.py index 619001a5791..bdfe99403a5 100644 --- a/litellm/types/passthrough_endpoints/pass_through_endpoints.py +++ b/litellm/types/passthrough_endpoints/pass_through_endpoints.py @@ -30,6 +30,7 @@ class EndpointType(str, Enum): OPENAI = "openai" TINYFISH = "tinyfish" GENERIC = "generic" + DECISIONS = "decisions" class PassthroughStandardLoggingPayload(TypedDict, total=False): diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 4577d578263..fb903043799 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -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") diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 81ccc66942a..793db970dd5 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -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"))