fix(proxy): reserve the largest explicit token cap when max_tokens and max_completion_tokens conflict

This commit is contained in:
mateo-berri 2026-07-21 21:44:46 -07:00
parent 6375923f65
commit c80de961a0
7 changed files with 88 additions and 16 deletions

View file

@ -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:

View file

@ -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

View file

@ -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:

View file

@ -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,
)

View file

@ -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

View file

@ -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
)

View file

@ -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,