From c80de961a0b005395581b0926c0dab2a8330ec66 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 21 Jul 2026 21:44:46 -0700 Subject: [PATCH] fix(proxy): reserve the largest explicit token cap when max_tokens and max_completion_tokens conflict --- .../hooks/parallel_request_limiter_v3.py | 5 +++- .../spend_tracking/budget_reservation.py | 13 ++++---- .../io_token_rate_limit_check.py | 18 ++++++----- tests/e2e/load/test_session_anomaly.py | 4 +-- .../hooks/test_parallel_request_limiter_v3.py | 25 ++++++++++++++++ .../proxy/test_budget_reservation.py | 30 +++++++++++++++++++ .../test_router/test_io_token_rate_limits.py | 9 +++++- 7 files changed, 88 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 22ea9fe176a..60845e4fe97 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -517,7 +517,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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 = max( + (int(v) for v in (data.get("max_tokens"), data.get("max_completion_tokens")) if v), + default=None, + ) match (explicit_max_tokens, input_text): case (mt, _) if mt is not None: diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 80fd8a1594e..b74976a01d4 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -1215,11 +1215,14 @@ def _estimate_output_tokens( if _is_input_only_route(route=route): return 0 - requested: Optional[int] = None - for key in ("max_completion_tokens", "max_tokens", "max_output_tokens"): - requested = _to_int(request_body.get(key)) - if requested is not None: - break + requested: Optional[int] = max( + ( + value + for key in ("max_completion_tokens", "max_tokens", "max_output_tokens") + if (value := _to_int(request_body.get(key))) is not None + ), + default=None, + ) # Clamp at min(requested-or-default, model_max-or-default). Two purposes: # (1) Without an explicit cap we still need a finite reservation so the diff --git a/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py b/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py index 803fdc4b353..e63d4e09d38 100644 --- a/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py +++ b/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py @@ -135,13 +135,17 @@ def _resolve_max_tokens(request_kwargs: Optional[dict[str, Any]], deployment: di if request_kwargs: # An explicit max_tokens=0 must be honored, not treated as absent and # replaced by the model default. - explicit = request_kwargs.get("max_tokens") - if explicit is None: - explicit = request_kwargs.get("max_completion_tokens") - if explicit is None: - explicit = request_kwargs.get("max_output_tokens") - if explicit is not None: - return max(0, int(explicit)) + explicit_caps = [ + int(value) + for value in ( + request_kwargs.get("max_tokens"), + request_kwargs.get("max_completion_tokens"), + request_kwargs.get("max_output_tokens"), + ) + if value is not None + ] + if explicit_caps: + return max(0, max(explicit_caps)) model_name = (deployment.get("litellm_params") or {}).get("model") if model_name: diff --git a/tests/e2e/load/test_session_anomaly.py b/tests/e2e/load/test_session_anomaly.py index 80f8aff3ba4..15ea55cf37a 100644 --- a/tests/e2e/load/test_session_anomaly.py +++ b/tests/e2e/load/test_session_anomaly.py @@ -62,7 +62,7 @@ class TestSummarizePlannedTurns: class TestRetried: def test_transient_failures_then_success_returns_the_success(self) -> None: - outcome = Success(data=SessionMessagesResponse()) + outcome = Success(status_code=200, data=SessionMessagesResponse()) calls = iter( (NetworkError(message="overloaded"), NetworkError(message="overloaded"), outcome) ) @@ -88,7 +88,7 @@ class TestRetried: raise AssertionError("slept after a successful attempt") result = retried( - lambda: Success(data=SessionMessagesResponse()), + lambda: Success(status_code=200, data=SessionMessagesResponse()), attempts=3, sleep=sleep_means_retry, ) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index c76e1a60afd..c9a263767a6 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -4652,3 +4652,28 @@ async def test_streaming_mirror_matches_non_streaming_header_shape(monkeypatch): f" non_streaming={_rl_only(non_stream_headers)}" ) assert "x-ratelimit-model_per_key-remaining-requests" in stream_slp_headers + + +def test_estimate_tokens_reserves_larger_cap_when_max_tokens_conflicts(): + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()), + ) + messages = [{"role": "user", "content": "hi"}] + + conflicting_small_first = handler._estimate_tokens_for_request( + data={"messages": messages, "max_tokens": 1, "max_completion_tokens": 5000}, + ) + conflicting_large_first = handler._estimate_tokens_for_request( + data={"messages": messages, "max_tokens": 5000, "max_completion_tokens": 1}, + ) + only_max_completion_tokens = handler._estimate_tokens_for_request( + data={"messages": messages, "max_completion_tokens": 5000}, + ) + only_max_tokens = handler._estimate_tokens_for_request( + data={"messages": messages, "max_tokens": 5000}, + ) + + assert conflicting_small_first == only_max_completion_tokens + assert conflicting_small_first == only_max_tokens + assert conflicting_large_first == conflicting_small_first + assert conflicting_small_first > 5000 diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 1db76aed61d..1f7fc7560b4 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -2484,3 +2484,33 @@ async def test_streaming_slow_path_processes_and_yields_chunk(spend_counter_stat assert received == [{"content": "hi"}] streaming_logging_obj.async_post_call_streaming_hook.assert_awaited_once() + + +def test_estimate_output_tokens_conflicting_caps_reserve_larger(): + from litellm.proxy.spend_tracking.budget_reservation import _estimate_output_tokens + + model_info = {"max_output_tokens": 200000} + assert ( + _estimate_output_tokens( + request_body={"max_completion_tokens": 1, "max_tokens": 5000}, + route="/chat/completions", + model_info=model_info, + ) + == 5000 + ) + assert ( + _estimate_output_tokens( + request_body={"max_completion_tokens": 5000, "max_tokens": 1}, + route="/chat/completions", + model_info=model_info, + ) + == 5000 + ) + assert ( + _estimate_output_tokens( + request_body={"max_completion_tokens": 7}, + route="/chat/completions", + model_info=model_info, + ) + == 7 + ) diff --git a/tests/test_litellm/test_router/test_io_token_rate_limits.py b/tests/test_litellm/test_router/test_io_token_rate_limits.py index a5a68271111..e1fbc0ee0f2 100644 --- a/tests/test_litellm/test_router/test_io_token_rate_limits.py +++ b/tests/test_litellm/test_router/test_io_token_rate_limits.py @@ -51,10 +51,17 @@ class TestIOTokenRateLimitHelpers: deployment = {"litellm_params": {"model": "openai/gpt-4o-mini"}} # An explicit max_tokens=0 is honored, not replaced by the model default. assert _resolve_max_tokens({"max_tokens": 0}, deployment) == 0 - # max_completion_tokens is the fallback only when max_tokens is absent. + # Any of the explicit cap aliases is honored on its own; when several + # are present the reservation uses the largest one. assert _resolve_max_tokens({"max_completion_tokens": 12}, deployment) == 12 assert _resolve_max_tokens({"max_output_tokens": 9}, deployment) == 9 + def test_resolve_max_tokens_conflicting_caps_reserve_larger(self): + deployment = {"litellm_params": {"model": "openai/gpt-4o-mini"}} + assert _resolve_max_tokens({"max_tokens": 1, "max_completion_tokens": 5000}, deployment) == 5000 + assert _resolve_max_tokens({"max_tokens": 5000, "max_completion_tokens": 1}, deployment) == 5000 + assert _resolve_max_tokens({"max_tokens": 0, "max_completion_tokens": 12}, deployment) == 12 + def test_build_io_token_rate_limit_headers(self): headers = build_io_token_rate_limit_headers( itpm_limit=200,