diff --git a/litellm/constants.py b/litellm/constants.py index 8ef9523a60a..8a6a044ff7b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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" diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 39e74d2c8bd..d9d56778405 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -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", diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 3c186e19829..893174dab03 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 843b916d90f..cf62ce1eb2d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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( diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index ae1f82fc40b..f9b4f50c15c 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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 diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 37375d9727d..58f56d746c5 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -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=( diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 7070617ce3e..74799072984 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -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" diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index b0f2a65bb9b..565e77b0ae0 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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"))) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index d1431d8956f..b52f0922390 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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.