mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
9e18526887
commit
0175c7da1c
9 changed files with 602 additions and 34 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")))
|
||||
|
|
|
|||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue