feat(router): resolve max_tokens to the tier model's ceiling on auto-routed requests (#40209)

A client behind an auto-router sends one max_tokens for every tier, so a value
sized for the smallest tier starves a bigger tier's thinking budget and a value
sized for the biggest is rejected by the smallest. After the complexity router
picks a tier, its per-tier litellm_params now carry max_tokens set to the
smallest max_output_tokens across that tier model's deployments (model_info,
then the cost map), applied the same way a per-tier reasoning_effort already
is, on every routing exit including plan mode, the empty-ask default and the
classifier fallback. The router seam collapses whichever ceiling alias a tier
carries onto the surface's own name, so one tier max_tokens reaches chat,
/v1/messages and /v1/responses alike, drops the caller's other carriers of the
same setting before the merge, and stamps the caller's original once so a
fallback into a group no tier owns gets it back instead of a ceiling sized for
the tier that failed.

Proxy-level reservations were sized from the caller's cap before routing, so a
raised cap left them short. Both owners now re-validate at the deployment hook:
the v3 limiter tops up its combined-TPM and project-OTPM reservations to the
final cap or writes the admitted cap back, and the budget limiter re-estimates
on the chosen deployment and grows the reservation or writes the admitted cap
back. An auto-router alias also reserves budget at its priciest tier model now
instead of pricing to zero.

An explicit per-tier max_tokens, max_completion_tokens or max_output_tokens
still wins, and max_tokens_from_tier_model: false forwards the caller's value
unchanged.
This commit is contained in:
tin-berri 2026-09-08 13:31:28 -07:00 • committed by GitHub
parent 9e18526887
commit 0175c7da1c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 602 additions and 34 deletions

View file

@ -1461,6 +1461,8 @@ LITELLM_METADATA_FIELD: Final = "litellm_metadata"
OLD_LITELLM_METADATA_FIELD: Final = "metadata"
RETURN_RAW_MODEL_NAME_METADATA_KEY: Final = "_complexity_router_return_raw_model_name"
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY: Final = "_session_deployment_affinity_ttl"
OUTPUT_TOKEN_CEILING_PARAMS: Final = frozenset({"max_tokens", "max_completion_tokens", "max_output_tokens"})
CLIENT_OUTPUT_CEILING_METADATA_KEY: Final = "_client_output_ceiling"
CONSUMED_REQUEST_TAGS_METADATA_KEY: Final = "_consumed_request_tags"
INTERNAL_CALL_ORIGIN_METADATA_KEY: Final = "internal_call_origin"
SESSION_ID_GENERATED_METADATA_KEY: Final = "litellm_session_id_generated"

View file

@ -10,6 +10,7 @@ import litellm
from litellm import get_secret
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
CLIENT_OUTPUT_CEILING_METADATA_KEY,
CONSUMED_REQUEST_TAGS_METADATA_KEY,
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
@ -507,6 +508,7 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset(
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
CONSUMED_REQUEST_TAGS_METADATA_KEY,
CLIENT_OUTPUT_CEILING_METADATA_KEY,
"disable_global_guardrails",
"disable_global_guardrail",
"opted_out_global_guardrails",

View file

@ -18,6 +18,7 @@ from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._service_logger import ServiceLogging
from litellm._uuid import uuid
from litellm.constants import (
CLIENT_OUTPUT_CEILING_METADATA_KEY,
CONSUMED_REQUEST_TAGS_METADATA_KEY,
INTERNAL_CALL_ORIGIN_METADATA_KEY,
LITELLM_PROXY_MASTER_KEY_ALIAS,
@ -325,7 +326,9 @@ _CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logg
# ``attempted_fallbacks`` and ``original_model_group`` are written by the router
# and read by spend logs as fact; a client value has no legitimate meaning and no
# key or team setting keeps it, so the strip is never gated.
_ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset({"attempted_fallbacks", "original_model_group"})
_ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset(
{"attempted_fallbacks", "original_model_group", CLIENT_OUTPUT_CEILING_METADATA_KEY}
)
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override"
# Request fields whose value, when URL-valued, becomes the outbound destination

View file

@ -21,7 +21,16 @@ import time
import traceback
import weakref
from collections import defaultdict
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Iterator, Mapping, Sequence
from collections.abc import (
AsyncGenerator,
AsyncIterator,
Callable,
Generator,
Iterator,
Mapping,
MutableMapping,
Sequence,
)
from functools import lru_cache, partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
@ -45,12 +54,14 @@ from litellm.caching.caching import (
RedisClusterCache,
)
from litellm.constants import (
CLIENT_OUTPUT_CEILING_METADATA_KEY,
CONSUMED_REQUEST_TAGS_METADATA_KEY,
DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS,
DEFAULT_HEALTH_CHECK_INTERVAL,
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
DEFAULT_MAX_LRU_CACHE_SIZE,
INTERNAL_CALL_ORIGIN_METADATA_KEY,
OUTPUT_TOKEN_CEILING_PARAMS,
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
)
@ -648,6 +659,18 @@ class FallbackAwareStreamWrapper(CustomStreamWrapper):
self.fallback_headers_adopted = True
def as_output_cap(value: object) -> int | None:
"""A client-sent output cap coerced to an int: ints, floats and numeric strings, never bools
or negatives."""
if isinstance(value, bool) or not isinstance(value, (int, float, str)):
return None
try:
cap: Final = int(float(value))
except (ValueError, OverflowError):
return None
return cap if cap >= 0 else None
class Router:
model_names: set = set()
cache_responses: bool | None = False
@ -12606,7 +12629,83 @@ class Router:
request_kwargs.pop(carrier, None)
@staticmethod
def _drop_client_effort_carriers_a_tier_pin_supersedes(
def _tier_ceiling_under_the_surface_name(
tier_litellm_params: Mapping[str, object], responses_call: bool
) -> Mapping[str, object]:
"""``max_tokens``, ``max_completion_tokens`` and ``max_output_tokens`` are one
ceiling under three names, and each surface reads exactly one of them: the
Responses bridge builds its internal ``max_tokens`` from ``max_output_tokens``
and would overwrite the tier's, chat and /v1/messages never read
``max_output_tokens``, and litellm already renames ``max_tokens`` to
``max_completion_tokens`` for the OpenAI models that require it. Collapse
whatever the tier carries onto the surface's own name, preferring a value the
operator already wrote under that name."""
surface_key: Final = "max_output_tokens" if responses_call else "max_tokens"
carried: Final = tuple(
key
for key in (surface_key, "max_tokens", "max_completion_tokens", "max_output_tokens")
if key in tier_litellm_params
)
if not carried:
return tier_litellm_params
return MappingProxyType(
{
**{k: v for k, v in tier_litellm_params.items() if k not in OUTPUT_TOKEN_CEILING_PARAMS},
surface_key: tier_litellm_params[carried[0]],
}
)
def _pin_tier_params_onto_request(
self,
model: str,
tier_litellm_params: Mapping[str, object] | None,
request_kwargs: dict,
responses_call: bool,
) -> bool:
"""Apply a routing strategy's per-tier litellm_params on top of the request and report
whether they pinned an output ceiling, so the caller can hand the request its own ceiling
back on a routing pass that pins none."""
if not tier_litellm_params:
return False
accepted_tier_params: Final = self._tier_params_the_target_accepts(model, tier_litellm_params, request_kwargs)
surface_tier_params: Final = self._tier_ceiling_under_the_surface_name(
accepted_tier_params, responses_call=responses_call
)
self._drop_client_carriers_a_tier_pin_supersedes(request_kwargs, surface_tier_params)
request_kwargs.update(surface_tier_params)
return not OUTPUT_TOKEN_CEILING_PARAMS.isdisjoint(surface_tier_params)
@staticmethod
def _restore_client_ceiling_no_tier_pins(request_kwargs: MutableMapping[str, object]) -> None:
"""A model-group fallback re-enters routing with the kwargs an earlier auto-router pass
already rewrote, so a ceiling sized for that pass's tier would ride onto a group no tier
chose. When this pass pins none, hand the request back exactly the carriers the caller
sent, which the first pinning pass stamped. The stamp lives in a metadata bucket a
caller can also write, so the proxy strips the key at ingestion and this read takes
nothing but the three ceiling carriers as integers: no other key ever reaches kwargs."""
stamped: Final = next(
(
bucket.get(CLIENT_OUTPUT_CEILING_METADATA_KEY)
for bucket in (request_kwargs.get("metadata"), request_kwargs.get("litellm_metadata"))
if isinstance(bucket, dict) and CLIENT_OUTPUT_CEILING_METADATA_KEY in bucket
),
None,
)
if not isinstance(stamped, dict):
return
callers_ceiling: Final = MappingProxyType(
{
carrier: cap
for carrier, value in stamped.items()
if carrier in OUTPUT_TOKEN_CEILING_PARAMS and (cap := as_output_cap(value)) is not None
}
)
for carrier in OUTPUT_TOKEN_CEILING_PARAMS:
request_kwargs.pop(carrier, None)
request_kwargs.update(callers_ceiling)
@staticmethod
def _drop_client_carriers_a_tier_pin_supersedes(
request_kwargs: dict[str, object],
tier_litellm_params: Mapping[str, object],
) -> None:
@ -12616,7 +12715,22 @@ class Router:
the ``reasoning_effort`` alias, so a pinned effort only reaches the wire
if the client's other encodings are removed before the merge. Non-effort
fields a carrier also holds (``output_config.format``,
``reasoning.summary``) are kept."""
``reasoning.summary``) are kept. An output ceiling has the same shape:
``max_tokens``, ``max_completion_tokens`` and ``max_output_tokens`` are
one setting under three names, and a provider handed two of them either
rejects the request or picks one by iteration order."""
if not OUTPUT_TOKEN_CEILING_PARAMS.isdisjoint(tier_litellm_params):
_, metadata_bucket = get_or_create_metadata_bucket(request_kwargs)
metadata_bucket.setdefault(
CLIENT_OUTPUT_CEILING_METADATA_KEY,
{
carrier: request_kwargs[carrier]
for carrier in OUTPUT_TOKEN_CEILING_PARAMS
if carrier in request_kwargs
},
)
for carrier in OUTPUT_TOKEN_CEILING_PARAMS:
request_kwargs.pop(carrier, None)
if "reasoning_effort" not in tier_litellm_params:
return
request_kwargs.pop("thinking", None)
@ -12657,6 +12771,7 @@ class Router:
# Execute Pre-Routing Hooks
# this hook can modify the model, messages before the routing decision is made
#########################################################
responses_call: Final = input is not None and messages is None
pre_routing_hook_response: Final = await self.async_pre_routing_hook(
model=model,
request_kwargs=request_kwargs,
@ -12668,12 +12783,14 @@ class Router:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
record_pre_routing_selection(request_kwargs, model)
if pre_routing_hook_response.litellm_params:
accepted_tier_params: Final = self._tier_params_the_target_accepts(
model, pre_routing_hook_response.litellm_params, request_kwargs
)
self._drop_client_effort_carriers_a_tier_pin_supersedes(request_kwargs, accepted_tier_params)
request_kwargs.update(accepted_tier_params)
tier_pins_ceiling: Final = self._pin_tier_params_onto_request(
model=model,
tier_litellm_params=pre_routing_hook_response.litellm_params if pre_routing_hook_response else None,
request_kwargs=request_kwargs,
responses_call=responses_call,
)
if not tier_pins_ceiling:
self._restore_client_ceiling_no_tier_pins(request_kwargs)
#########################################################
# Resolve the strategy and logger AFTER the pre-routing hook, since
@ -12773,6 +12890,7 @@ class Router:
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs)
# 1. Execute pre-routing hook
responses_call: Final = input is not None and messages is None
pre_routing_hook_response: Final = await self.async_pre_routing_hook(
model=model,
request_kwargs=request_kwargs,
@ -12784,12 +12902,14 @@ class Router:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
record_pre_routing_selection(request_kwargs, model)
if pre_routing_hook_response.litellm_params:
accepted_tier_params: Final = self._tier_params_the_target_accepts(
model, pre_routing_hook_response.litellm_params, request_kwargs
)
self._drop_client_effort_carriers_a_tier_pin_supersedes(request_kwargs, accepted_tier_params)
request_kwargs.update(accepted_tier_params)
tier_pins_ceiling: Final = self._pin_tier_params_onto_request(
model=model,
tier_litellm_params=pre_routing_hook_response.litellm_params if pre_routing_hook_response else None,
request_kwargs=request_kwargs,
responses_call=responses_call,
)
if not tier_pins_ceiling:
self._restore_client_ceiling_no_tier_pins(request_kwargs)
# 2. Get healthy deployments
healthy_deployments: Final = await self.async_get_healthy_deployments(

View file

@ -31,6 +31,7 @@ from litellm._logging import verbose_router_logger
from litellm.constants import (
EMPTY_MAPPING,
INTERNAL_CALL_ORIGIN_METADATA_KEY,
OUTPUT_TOKEN_CEILING_PARAMS,
RETURN_RAW_MODEL_NAME_METADATA_KEY,
SESSION_ID_GENERATED_METADATA_KEY,
)
@ -2110,11 +2111,15 @@ class ComplexityRouter(CustomLogger):
raise ValueError(f"No model configured for tier {tier_key} and no default_model set")
def _litellm_params_for_model(self, tier: ComplexityTier | str | None, model: str) -> Mapping[str, object]:
if tier is None:
return MappingProxyType({})
entries: Final = self.config.tier_model_configs.get(_tier_name(tier), ())
entries: Final = self.config.tier_model_configs.get(_tier_name(tier), ()) if tier is not None else ()
entry: Final = next((candidate for candidate in entries if candidate.model_name == model), None)
return entry.litellm_params if entry is not None else MappingProxyType({})
explicit: Final = entry.litellm_params if entry is not None else MappingProxyType({})
if not self.config.max_tokens_from_tier_model or not OUTPUT_TOKEN_CEILING_PARAMS.isdisjoint(explicit):
return explicit
ceiling: Final = self._group_output_ceiling(model)
if ceiling is None:
return explicit
return MappingProxyType({**explicit, "max_tokens": ceiling})
@staticmethod
def _pick_from_tier_value(model: str | Sequence[str], tier_key: str) -> str:
@ -2449,12 +2454,15 @@ class ComplexityRouter(CustomLogger):
return name if self.config.has_custom_tiers else ComplexityTier(name)
def _deployment_window(self, group: str, deployment: Mapping[str, object]) -> int | None:
return self._deployment_limit(group, deployment, "max_input_tokens")
def _deployment_limit(
self, group: str, deployment: Mapping[str, object], key: Literal["max_input_tokens", "max_output_tokens"]
) -> int | None:
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
deployment_model_info: Final = deployment.get("model_info")
declared: Final = (
deployment_model_info.get("max_input_tokens") if isinstance(deployment_model_info, Mapping) else None
)
declared: Final = deployment_model_info.get(key) if isinstance(deployment_model_info, Mapping) else None
if isinstance(declared, int):
return declared
litellm_params: Final = deployment.get("litellm_params")
@ -2471,18 +2479,34 @@ class ComplexityRouter(CustomLogger):
deployment=cast(dict, deployment), # cast-ok: router deployments are plain dicts
received_model_name=group,
)
window: Final = model_info.get("max_input_tokens")
limit: Final = model_info.get(key)
except Exception: # noqa: BLE001 # best-effort: an unmappable deployment must not hide the others
return None
return window if isinstance(window, int) else None
return limit if isinstance(limit, int) else None
def _group_deployments(self, group: str) -> Sequence[Mapping[str, object]]:
list_models: Final = getattr(self.litellm_router_instance, "get_model_list", None)
deployments: Final = list_models(model_name=group) if callable(list_models) else None
return tuple(deployments) if isinstance(deployments, list) else ()
def _group_output_ceiling(self, group: str) -> int | None:
"""Smallest max_output_tokens across the group's deployments, or None when any deployment
declares none: the core router picks within the group without a fit check, and a ceiling
above an unmapped member's real limit is a provider 400 on that member."""
deployments: Final = self._group_deployments(group)
ceilings: Final = tuple(
ceiling
for deployment in deployments
if (ceiling := self._deployment_limit(group, deployment, "max_output_tokens")) is not None
)
return min(ceilings) if ceilings and len(ceilings) == len(deployments) else None
def _group_window_facts(self, group: str) -> tuple[int | None, bool]:
"""(smallest declared context window across the group's deployments, whether any deployment
declares none). The core router picks a deployment within the group without a fit check, so
the group is only as safe as its smallest member."""
list_models: Final = getattr(self.litellm_router_instance, "get_model_list", None)
deployments: Final = list_models(model_name=group) if callable(list_models) else None
if not isinstance(deployments, list) or not deployments:
deployments: Final = self._group_deployments(group)
if not deployments:
return (None, True)
windows: Final = tuple(
window for deployment in deployments if (window := self._deployment_window(group, deployment)) is not None
@ -3531,6 +3555,7 @@ class ComplexityRouter(CustomLogger):
ComplexityTier.MEDIUM, messages, resolved_messages, request_kwargs
)
fallback_tier: Final = None if default_model_first else ComplexityTier.MEDIUM
default_tier_params: Final = self._litellm_params_for_model(fallback_tier, routed_model)
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
@ -3539,7 +3564,9 @@ class ComplexityRouter(CustomLogger):
cause="default_fallback",
tier=fallback_tier,
conversation_continuing=conversation_continuing,
tier_litellm_params=default_tier_params,
),
litellm_params=default_tier_params,
)
ask: Final = user_message or ""
@ -3566,6 +3593,7 @@ class ComplexityRouter(CustomLogger):
_tier_name(plan_floor),
routed_model,
)
plan_tier_params: Final = self._litellm_params_for_model(plan_floor, routed_model)
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
@ -3577,7 +3605,9 @@ class ComplexityRouter(CustomLogger):
matched_keyword=plan_mode_sentinel,
escalation_keyword=escalation_keyword,
escalated=False,
tier_litellm_params=plan_tier_params,
),
litellm_params=plan_tier_params,
)
override: Final = await self._resolve_keyword_tier_override(ask, request_kwargs)
@ -3670,6 +3700,7 @@ class ComplexityRouter(CustomLogger):
outcome.signals,
fallback_model,
)
fallback_tier_params: Final = self._litellm_params_for_model(None, fallback_model)
return PreRoutingHookResponse(
model=fallback_model,
messages=messages if has_original_messages else None,
@ -3680,7 +3711,9 @@ class ComplexityRouter(CustomLogger):
signals=outcome.signals,
escalation_keyword=escalation_keyword,
escalated=False,
tier_litellm_params=fallback_tier_params,
),
litellm_params=fallback_tier_params,
)
if self.config.adaptive:
# hard_floor rather than a hard pick, and passed whenever the sentinel is present

View file

@ -1089,6 +1089,20 @@ class ComplexityRouterConfig(BaseModel):
"wording the built-ins don't cover, or after a client release changes its strings."
),
)
max_tokens_from_tier_model: bool = Field(
default=True,
description=(
"Set max_tokens on every routed request to the output ceiling of the tier model it "
"lands on, replacing whatever the caller sent. A caller behind an auto-router cannot "
"pick one value that fits every tier: the smallest tier's ceiling starves a bigger "
"tier's thinking budget, and a bigger tier's ceiling is rejected by the smallest. The "
"ceiling is the smallest max_output_tokens across the tier model's deployments, read "
"from each deployment's model_info and then the model cost map; a tier model with a "
"deployment whose ceiling is unknown keeps the caller's value. A max_tokens, "
"max_completion_tokens or max_output_tokens in the tier's own litellm_params still "
"wins. Set false to forward the caller's value unchanged."
),
)
route_housekeeping_to_cheapest_tier: bool = Field(
default=True,
description=(

View file

@ -7272,7 +7272,12 @@ def _reserved_stamp_key(key_metadata: dict | None = None) -> UserAPIKeyAuth:
)
_PLANTED_STAMPS = {"attempted_fallbacks": 99, "original_model_group": "spoofed-group", "client_key": "client_value"}
_PLANTED_STAMPS = {
"attempted_fallbacks": 99,
"original_model_group": "spoofed-group",
"_client_output_ceiling": {"api_base": "https://attacker.example"},
"client_key": "client_value",
}
@pytest.mark.asyncio
@ -7301,6 +7306,7 @@ async def test_add_litellm_data_to_request_strips_router_reserved_stamps_from_bo
assert "litellm_metadata" not in updated
assert "attempted_fallbacks" not in updated["metadata"]
assert "original_model_group" not in updated["metadata"]
assert "_client_output_ceiling" not in updated["metadata"]
assert updated["metadata"]["client_key"] == "client_value"

View file

@ -16,6 +16,7 @@ from pydantic import ValidationError
import litellm
from litellm import Router
from litellm.integrations.custom_logger import CustomLogger
from litellm.router_utils.auto_router_model_naming import (
CUSTOMIZATION_CAPABILITY,
GATED_AUTO_ROUTER_CAPABILITIES,
@ -24,7 +25,12 @@ from litellm.router_utils.auto_router_model_naming import (
)
from litellm._logging import verbose_router_logger
from litellm.caching.dual_cache import DualCache
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY, SESSION_ID_GENERATED_METADATA_KEY
from litellm.constants import (
OUTPUT_TOKEN_CEILING_PARAMS,
RETURN_RAW_MODEL_NAME_METADATA_KEY,
SESSION_ID_GENERATED_METADATA_KEY,
)
from litellm.router import as_output_cap
from litellm.router_strategy.complexity_router.complexity_router import (
_CLASSIFICATION_CURRENT_MESSAGE_ONLY,
_CLASSIFICATION_WITH_CONVERSATION,
@ -3415,11 +3421,11 @@ class TestRouterPreRoutingAliasOverrides:
def test_drop_client_effort_carriers_helper_edge_shapes(self):
no_pin: Dict = {"thinking": {"type": "adaptive"}}
Router._drop_client_effort_carriers_a_tier_pin_supersedes(no_pin, {"temperature": 0.1})
Router._drop_client_carriers_a_tier_pin_supersedes(no_pin, {"temperature": 0.1})
assert no_pin == {"thinking": {"type": "adaptive"}}
non_dict_carriers: Dict = {"output_config": "max", "reasoning": 3}
Router._drop_client_effort_carriers_a_tier_pin_supersedes(non_dict_carriers, {"reasoning_effort": "low"})
Router._drop_client_carriers_a_tier_pin_supersedes(non_dict_carriers, {"reasoning_effort": "low"})
assert non_dict_carriers == {"output_config": "max", "reasoning": 3}
effort_only: Dict = {"output_config": {"effort": "max"}, "reasoning": {"effort": "high"}}
@ -3873,11 +3879,11 @@ class TestRouterPreRoutingSharedAliasName:
}
@staticmethod
async def _routed_call_kwargs(router: Router, **request_params) -> dict:
async def _routed_call_kwargs(router: Router, prompt: str = "hi", **request_params) -> dict:
mock_acompletion = AsyncMock(return_value=litellm.ModelResponse(choices=[{"message": {"content": "hi"}}]))
with patch.object(litellm, "acompletion", mock_acompletion):
await router.acompletion(
model="smart-router", messages=[{"role": "user", "content": "hi"}], **request_params
model="smart-router", messages=[{"role": "user", "content": prompt}], **request_params
)
return mock_acompletion.call_args.kwargs
@ -13278,3 +13284,379 @@ class TestClassifierVision:
def test_max_images_must_be_positive(self):
with pytest.raises(ValidationError):
ClassifierLLMConfig(model="clf", vision={"enabled": True, "max_images": 0})
class TestMaxTokensFromTierModel:
"""The auto-router replaces the caller's output ceiling with the tier model's own, so one
client-side value no longer starves a bigger tier or gets rejected by a smaller one."""
COMPLEX_PROMPT: Final = (
"Design a distributed rate limiter with Redis, sharding and failover. Analyze the consistency "
"tradeoffs and implement the algorithm step by step with tests."
)
SMALL: Final = {
"model_name": "small",
"litellm_params": {"model": "anthropic/claude-haiku-4-5", "api_key": "k"},
"model_info": {"max_output_tokens": 8192},
}
@staticmethod
def _router(
tier_litellm_params: dict | None = None,
max_tokens_from_tier_model: bool | None = None,
simple_deployments: list[dict] | None = None,
extra_config: dict | None = None,
) -> Router:
simple_tier: dict = {"model_name": "small"}
if tier_litellm_params:
simple_tier["litellm_params"] = tier_litellm_params
config: dict = {
"tiers": {"SIMPLE": simple_tier, "MEDIUM": "big", "COMPLEX": "big", "REASONING": "big"},
**(extra_config or {}),
}
if max_tokens_from_tier_model is not None:
config["max_tokens_from_tier_model"] = max_tokens_from_tier_model
return Router(
model_list=[
{
"model_name": "smart-router",
"litellm_params": {"model": "auto_router/complexity_router", "complexity_router_config": config},
},
*(simple_deployments or [TestMaxTokensFromTierModel.SMALL]),
{
"model_name": "big",
"litellm_params": {"model": "anthropic/claude-sonnet-5", "api_key": "k"},
"model_info": {"max_output_tokens": 64000},
},
]
)
@staticmethod
async def _routed(router: Router, prompt: str = "hi", **request_kwargs) -> dict:
"""Drive the real routing entry point and return the request kwargs it leaves behind."""
deployment = await router.async_get_available_deployment(
model="smart-router", request_kwargs=request_kwargs, messages=[{"role": "user", "content": prompt}]
)
return {"model": deployment["litellm_params"]["model"], **request_kwargs}
@staticmethod
async def _routed_responses(router: Router, prompt: str = "hi", **request_kwargs) -> dict:
"""The Responses surface hands the router `input` both as the prompt argument and inside the
request kwargs, so the hook sees the same shape the real call carries."""
routed: dict = {"input": prompt, **request_kwargs}
deployment = await router.async_get_available_deployment(
model="smart-router", request_kwargs=routed, input=prompt
)
return {"model": deployment["litellm_params"]["model"], **routed}
@pytest.mark.asyncio
async def test_client_ceiling_is_replaced_by_the_routed_tier_models_ceiling(self):
router = self._router()
simple = await self._routed(router, max_tokens=8192)
complex_ = await self._routed(router, self.COMPLEX_PROMPT, max_tokens=8192)
assert (simple["model"], simple["max_tokens"]) == ("anthropic/claude-haiku-4-5", 8192)
assert (complex_["model"], complex_["max_tokens"]) == ("anthropic/claude-sonnet-5", 64000)
assert "max_output_tokens" not in complex_
@pytest.mark.asyncio
async def test_every_client_carrier_of_the_ceiling_is_replaced(self):
sent = await self._routed(self._router(), self.COMPLEX_PROMPT, max_completion_tokens=8192)
assert sent["max_tokens"] == 64000
assert "max_completion_tokens" not in sent
@pytest.mark.asyncio
async def test_responses_surface_gets_the_ceiling_under_its_own_name(self):
sent = await self._routed_responses(self._router(), self.COMPLEX_PROMPT, max_output_tokens=8192)
assert (sent["model"], sent["max_output_tokens"]) == ("anthropic/claude-sonnet-5", 64000)
assert "max_tokens" not in sent
@pytest.mark.asyncio
@pytest.mark.parametrize(
"tier_params, responses_call",
[
({"max_tokens": 4321}, False),
({"max_tokens": 4321}, True),
({"max_completion_tokens": 4321}, False),
({"max_completion_tokens": 4321}, True),
({"max_output_tokens": 4321}, False),
],
)
async def test_operators_own_tier_ceiling_wins_under_the_surface_name(self, tier_params, responses_call):
router = self._router(tier_litellm_params=tier_params)
if responses_call:
sent = await self._routed_responses(router, max_output_tokens=8192)
else:
sent = await self._routed(router, max_tokens=8192)
surface_key = "max_output_tokens" if responses_call else "max_tokens"
assert sent[surface_key] == 4321
assert not (OUTPUT_TOKEN_CEILING_PARAMS - {surface_key}) & sent.keys()
@pytest.mark.asyncio
async def test_opting_out_forwards_the_client_value_unchanged(self):
sent = await self._routed(self._router(max_tokens_from_tier_model=False), self.COMPLEX_PROMPT, max_tokens=8192)
assert sent["max_tokens"] == 8192
@pytest.mark.asyncio
async def test_a_tier_model_with_an_unknown_ceiling_keeps_the_client_value(self):
unmapped: dict = {"model_name": "small", "litellm_params": {"model": "openai/not-in-any-map", "api_key": "k"}}
sent = await self._routed(self._router(simple_deployments=[self.SMALL, unmapped]), max_tokens=4000)
assert sent["max_tokens"] == 4000
@pytest.mark.asyncio
async def test_a_multi_deployment_tier_model_uses_its_smallest_ceiling(self):
smaller: dict = {
**self.SMALL,
"litellm_params": {**self.SMALL["litellm_params"], "api_key": "k2"},
"model_info": {"max_output_tokens": 4096},
}
sent = await self._routed(self._router(simple_deployments=[self.SMALL, smaller]), max_tokens=100000)
assert sent["max_tokens"] == 4096
@pytest.mark.asyncio
async def test_ceiling_falls_back_to_the_cost_map(self, monkeypatch):
monkeypatch.setitem(
litellm.model_cost,
"auto-cap-probe-model",
{"litellm_provider": "openai", "mode": "chat", "max_output_tokens": 4242, "max_input_tokens": 100000},
)
mapped_only: dict = {
"model_name": "small",
"litellm_params": {"model": "openai/auto-cap-probe-model", "api_key": "k"},
}
sent = await self._routed(self._router(simple_deployments=[mapped_only]), max_tokens=8192)
assert sent["max_tokens"] == 4242
@pytest.mark.asyncio
@pytest.mark.parametrize("client_kwargs", [{}, {"max_tokens": 0}], ids=["omitted", "zero"])
async def test_omitted_and_zero_are_replaced_like_any_other_value(self, client_kwargs):
sent = await self._routed(self._router(), self.COMPLEX_PROMPT, **client_kwargs)
assert sent["max_tokens"] == 64000
@pytest.mark.parametrize(
"tier_params, responses_call, expected",
[
({"max_tokens": 1, "temperature": 0.2}, False, {"max_tokens": 1, "temperature": 0.2}),
({"max_tokens": 1}, True, {"max_output_tokens": 1}),
({"max_completion_tokens": 2}, False, {"max_tokens": 2}),
({"max_completion_tokens": 2}, True, {"max_output_tokens": 2}),
({"max_output_tokens": 3}, False, {"max_tokens": 3}),
({"max_output_tokens": 3}, True, {"max_output_tokens": 3}),
({"max_tokens": 1, "max_completion_tokens": 2, "max_output_tokens": 3}, False, {"max_tokens": 1}),
({"max_tokens": 1, "max_completion_tokens": 2, "max_output_tokens": 3}, True, {"max_output_tokens": 3}),
({"max_completion_tokens": 2, "max_output_tokens": 3}, False, {"max_tokens": 2}),
({"reasoning_effort": "low"}, True, {"reasoning_effort": "low"}),
],
)
def test_every_tier_alias_collapses_onto_the_surface_key(self, tier_params, responses_call, expected):
assert dict(Router._tier_ceiling_under_the_surface_name(tier_params, responses_call=responses_call)) == expected
@pytest.mark.asyncio
async def test_the_default_fallback_exit_carries_the_ceiling(self):
routed: dict = {"max_tokens": 8192}
deployment = await self._router().async_get_available_deployment(
model="smart-router", request_kwargs=routed, messages=[{"role": "system", "content": "be nice"}]
)
assert routed["metadata"]["routing_decision"]["cause"] == "default_fallback"
assert (deployment["litellm_params"]["model"], routed["max_tokens"]) == ("anthropic/claude-sonnet-5", 64000)
@pytest.mark.asyncio
async def test_the_plan_mode_exit_carries_the_ceiling(self):
routed: dict = {"max_tokens": 8192}
deployment = await self._router(
extra_config={"plan_mode_min_tier": "REASONING"}
).async_get_available_deployment(
model="smart-router",
request_kwargs=routed,
messages=[
{"role": "user", "content": "plan the refactor"},
{"role": "system", "content": "Plan mode is active"},
],
)
assert routed["metadata"]["routing_decision"]["cause"] == "plan_mode"
assert (deployment["litellm_params"]["model"], routed["max_tokens"]) == ("anthropic/claude-sonnet-5", 64000)
@pytest.mark.asyncio
async def test_a_default_model_landing_with_no_tier_still_gets_its_ceiling(self):
strategy = ComplexityRouter(
model_name="smart-router",
litellm_router_instance=self._router(),
complexity_router_config={"tiers": {"SIMPLE": "small"}, "default_model": "big"},
)
assert dict(strategy._litellm_params_for_model(None, "big")) == {"max_tokens": 64000}
@pytest.mark.asyncio
async def test_a_fallback_into_a_plain_group_gets_the_callers_ceiling_back(self):
"""A model-group fallback re-enters routing with the same kwargs; a Sonnet-sized ceiling
must not ride onto the plain group the caller configured as the fallback."""
big: dict = {
"model_name": "big",
"litellm_params": {
"model": "anthropic/claude-sonnet-5",
"api_key": "k",
"mock_response": "litellm.InternalServerError",
},
"model_info": {"max_output_tokens": 64000},
}
plain: dict = {
"model_name": "plain",
"litellm_params": {"model": "anthropic/claude-haiku-4-5", "api_key": "k", "mock_response": "ok"},
}
router = Router(
model_list=[
{
"model_name": "smart-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {
"tiers": {"SIMPLE": "big", "MEDIUM": "big", "COMPLEX": "big", "REASONING": "big"}
},
},
},
big,
plain,
],
fallbacks=[{"smart-router": ["plain"]}],
num_retries=0,
)
recorder = _OutputCeilingRecorder()
litellm.callbacks.append(recorder)
try:
await router.acompletion(
model="smart-router", messages=[{"role": "user", "content": self.COMPLEX_PROMPT}], max_tokens=8192
)
finally:
litellm.callbacks.remove(recorder)
assert recorder.seen == [("claude-sonnet-5", 64000), ("claude-haiku-4-5", 8192)]
@pytest.mark.asyncio
async def test_a_caller_seeded_stamp_cannot_inject_kwargs_on_a_plain_group(self):
"""The stamp sits in a metadata bucket a caller can write; a planted one must yield
nothing but integer ceiling carriers, never a redirected api_base or credential."""
planted: dict = {
"api_base": "https://attacker.example",
"api_key": "stolen",
"max_tokens": "not-an-int",
"max_completion_tokens": True,
"max_output_tokens": 321,
}
routed: dict = {"max_tokens": 8192, "metadata": {"_client_output_ceiling": planted}}
await self._router().async_get_available_deployment(
model="big", request_kwargs=routed, messages=[{"role": "user", "content": "hi"}]
)
assert {k: v for k, v in routed.items() if k not in ("metadata", "model_info")} == {"max_output_tokens": 321}
@pytest.mark.asyncio
async def test_the_pass_through_routing_entry_point_pins_and_restores_the_same_way(self):
pass_through: dict = {**self.SMALL["litellm_params"], "use_in_pass_through": True}
small: dict = {**self.SMALL, "litellm_params": pass_through}
plain: dict = {**small, "model_name": "plain"}
router = self._router(simple_deployments=[small, plain])
for deployment in router.model_list:
deployment["litellm_params"]["use_in_pass_through"] = True
routed: dict = {"max_tokens": 8192}
deployment = await router.async_get_available_deployment_for_pass_through(
model="smart-router", request_kwargs=routed, messages=[{"role": "user", "content": self.COMPLEX_PROMPT}]
)
pinned = routed["max_tokens"]
await router.async_get_available_deployment_for_pass_through(
model="plain", request_kwargs=routed, messages=[{"role": "user", "content": "hi"}]
)
assert (deployment["litellm_params"]["model"], pinned, routed["max_tokens"]) == (
"anthropic/claude-sonnet-5",
64000,
8192,
)
@pytest.mark.asyncio
async def test_the_classifier_fallback_exit_carries_the_ceiling(self):
router = self._router(
extra_config={
"classifier_type": "llm",
"classifier_llm_config": {"model": "no-such-classifier", "timeout_ms": 400},
"classifier_fallback": "default_model",
"default_model": "big",
}
)
routed: dict = {"max_tokens": 8192}
deployment = await router.async_get_available_deployment(
model="smart-router", request_kwargs=routed, messages=[{"role": "user", "content": "hi"}]
)
assert routed["metadata"]["routing_decision"]["cause"] == "default_model_fallback"
assert (deployment["litellm_params"]["model"], routed["max_tokens"]) == ("anthropic/claude-sonnet-5", 64000)
@pytest.mark.parametrize(
"value, expected",
[(8192, 8192), ("8192", 8192), (100.9, 100), (0, 0), (-1, None), (True, None), ("x", None), (None, None)],
)
def test_a_client_cap_is_read_as_an_integer_or_ignored(self, value, expected):
assert as_output_cap(value) == expected
def test_restoring_the_callers_ceiling_reads_the_stamp_and_replaces_every_carrier(self):
stamped: dict = {"max_output_tokens": 500, "metadata": {"_client_output_ceiling": {"max_tokens": 8192}}}
Router._restore_client_ceiling_no_tier_pins(stamped)
assert {k: v for k, v in stamped.items() if k != "metadata"} == {"max_tokens": 8192}
coerced: dict = {
"max_tokens": 64000,
"metadata": {"_client_output_ceiling": {"max_tokens": "8192", "max_completion_tokens": 100.0}},
}
Router._restore_client_ceiling_no_tier_pins(coerced)
assert {k: v for k, v in coerced.items() if k != "metadata"} == {
"max_tokens": 8192,
"max_completion_tokens": 100,
}
unstamped: dict = {"max_tokens": 64000, "metadata": {}}
Router._restore_client_ceiling_no_tier_pins(unstamped)
assert unstamped["max_tokens"] == 64000
@pytest.mark.asyncio
async def test_pinning_stamps_the_callers_carriers_once(self):
router = self._router()
request_kwargs: dict = {"max_completion_tokens": 8192}
first = router._pin_tier_params_onto_request(
model="big", tier_litellm_params={"max_tokens": 64000}, request_kwargs=request_kwargs, responses_call=False
)
second = router._pin_tier_params_onto_request(
model="big", tier_litellm_params={"max_tokens": 32000}, request_kwargs=request_kwargs, responses_call=False
)
none = router._pin_tier_params_onto_request(
model="big", tier_litellm_params=None, request_kwargs=request_kwargs, responses_call=False
)
assert (first, second, none) == (True, True, False)
assert request_kwargs["max_tokens"] == 32000
assert request_kwargs["metadata"]["_client_output_ceiling"] == {"max_completion_tokens": 8192}
class _OutputCeilingRecorder(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.seen: list[tuple[str, int | None]] = []
def log_pre_api_call(self, model, messages, kwargs):
self.seen.append((model, kwargs.get("optional_params", {}).get("max_tokens")))

View file

@ -34996,6 +34996,12 @@ export interface components {
* @default 0.5
*/
match_threshold: number;
/**
* Max Tokens From Tier Model
* @description Set max_tokens on every routed request to the output ceiling of the tier model it lands on, replacing whatever the caller sent. A caller behind an auto-router cannot pick one value that fits every tier: the smallest tier's ceiling starves a bigger tier's thinking budget, and a bigger tier's ceiling is rejected by the smallest. The ceiling is the smallest max_output_tokens across the tier model's deployments, read from each deployment's model_info and then the model cost map; a tier model with a deployment whose ceiling is unknown keeps the caller's value. A max_tokens, max_completion_tokens or max_output_tokens in the tier's own litellm_params still wins. Set false to forward the caller's value unchanged.
* @default true
*/
max_tokens_from_tier_model: boolean;
/**
* Modality Pin Override
* @description Let modality_routing replace a kept session-affinity pin on the turns that carry an image. Without this, a session pinned to a text-only model fails every image turn with a provider 400, since the pin is exempt from the modality gate. When enabled, such a turn routes to a capable model for that request only and the stored pin is left untouched, so the next text turn replays the session's own model; the override is reported as cause modality_pin_override and is never itself pinned. Inert unless modality_routing is also enabled.