feat(router): support complexity tier generation params

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
tin 2026-08-06 04:05:38 +00:00
parent d26ef670e2
commit fe1e42b949
7 changed files with 291 additions and 20 deletions

View file

@ -11204,6 +11204,11 @@ class Router:
# actual outbound LLM call downstream by litellm.types.utils.all_litellm_params,
# not here.
if pre_routing_hook_response is not None:
for key, value in (
pre_routing_hook_response.params or {}
).items(): # mutable-ok: preserve request kwargs identity
if value is not None:
request_kwargs.setdefault(key, value)
alias_index: Final = self.model_name_to_deployment_indices.get(model, [])
if alias_index:
alias_litellm_params: Final = self.model_list[alias_index[0]].get("litellm_params", {})

View file

@ -81,6 +81,37 @@ Where the names show up depends on your classifier. Under the default heuristic
Spend logs keep `routing_decision.tier` canonical so rows from before and after a rename stay comparable, and gain `routing_decision.tier_label` on the tiers you renamed.
### Per-tier generation parameters
Each tier can use an object instead of a model string. The object keeps the model
or model pool alongside generation parameters, so similar deployments do not need
near-duplicate entries:
```yaml
tiers:
SIMPLE: {model: gpt-5-mini, reasoning_effort: low}
MEDIUM: {model: gpt-5-mini, reasoning_effort: high}
COMPLEX: gpt-5
REASONING: [gpt-5, o3]
```
An object can also contain a pool:
```yaml
tiers:
COMPLEX:
model: [gpt-5, o3]
thinking: {type: enabled}
```
Caller-supplied request parameters take precedence over tier parameters, which
take precedence over alias `litellm_params` defaults. Parameters unsupported by
the selected model follow LiteLLM's existing `drop_params` behavior. Unknown
parameter names produce a configuration warning but remain available for
provider-specific parameters. `thinking` is the supported structured parameter;
`thinking_budget` is not a top-level LiteLLM parameter and therefore produces a
warning, as intended.
### Full Configuration
```yaml

View file

@ -46,6 +46,7 @@ from .config import (
TIER_SEVERITY_ORDER,
ComplexityRouterConfig,
ComplexityTier,
TierTarget,
)
if TYPE_CHECKING:
@ -1151,7 +1152,9 @@ class ComplexityRouter(CustomLogger):
raise ValueError(f"No model configured for tier {tier_key} and no default_model set")
@staticmethod
def _pick_from_tier_value(model: str | list[str], tier_key: str) -> str:
def _pick_from_tier_value(model: str | list[str] | TierTarget, tier_key: str) -> str:
if isinstance(model, TierTarget):
model = model.model
if isinstance(model, str):
return model
if not model:
@ -1159,7 +1162,26 @@ class ComplexityRouter(CustomLogger):
return random.choice(model)
def _tier_pools(self) -> dict[str, list[str]]:
return {tier: (models if isinstance(models, list) else [models]) for tier, models in self.config.tiers.items()}
return { # mutable-ok: router consumers require mutable tier pool mappings
tier: (
(target.model if isinstance(target.model, list) else [target.model])
if isinstance(target, TierTarget)
else models
if isinstance(models, list)
else [models]
)
for tier, models in self.config.tiers.items()
for target in (models,)
}
def _tier_params(self, tier: ComplexityTier | str) -> dict[str, object]:
tier_key: Final = tier.value if isinstance(tier, ComplexityTier) else tier
target: Final = self.config.tiers.get(tier_key)
return target.params if isinstance(target, TierTarget) else {} # mutable-ok: response schema requires a dict
def _params_for_model(self, model: str) -> dict[str, object]:
tier: Final = self._tier_for_model(model)
return self._tier_params(tier) if tier is not None else {} # mutable-ok: response schema requires a dict
async def _pick_model_for_tier(
self,
@ -1704,6 +1726,7 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
params=self._params_for_model(routed_model),
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause=cause,
@ -1789,6 +1812,9 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
params=self._tier_params(ComplexityTier.MEDIUM)
if not self.config.default_model or self.config.plugins
else {}, # mutable-ok: response schema requires a dict
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause="default_fallback",
@ -1817,6 +1843,7 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
params=self._tier_params(routed_tier),
routing_decision=self._build_routing_decision(
routed_model=routed_model,
conversation_continuing=conversation_continuing,
@ -1857,6 +1884,7 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=fallback_model,
messages=messages if has_original_messages else None,
params={}, # mutable-ok: response schema requires a dict
routing_decision=self._build_routing_decision(
routed_model=fallback_model,
conversation_continuing=conversation_continuing,
@ -1910,6 +1938,7 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
params=self._params_for_model(routed_model) if self.config.adaptive else self._tier_params(tier),
routing_decision=self._build_routing_decision(
routed_model=routed_model,
conversation_continuing=conversation_continuing,

View file

@ -6,11 +6,14 @@ All values are configurable via proxy config.yaml.
"""
from enum import Enum
from typing import Final, Literal
from typing import Final, Literal, cast
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from litellm._logging import verbose_router_logger
from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS
from litellm.types.router import AdaptiveRouterWeights, RoutingPlugin
from litellm.types.utils import all_litellm_params
class ComplexityTier(str, Enum):
@ -239,6 +242,33 @@ DEFAULT_TIER_MODELS: Final[dict[str, str]] = {
}
class TierTarget(BaseModel):
model: str | list[str]
model_config = ConfigDict(extra="allow")
@field_validator("model", mode="before")
@classmethod
def _validate_model(cls, value: object) -> object:
if isinstance(value, str):
if not value.strip():
raise ValueError("model must be a non-empty string")
return value
if isinstance(value, list):
if not value:
raise ValueError("model pool must be non-empty")
if any(not isinstance(item, str) for item in value):
raise ValueError("model pool entries must be strings")
if any(not item.strip() for item in value):
raise ValueError("model pool entries must be non-empty strings")
return value
raise ValueError("model must be a string or a list of strings")
@property
def params(self) -> dict[str, object]:
return cast(dict[str, object], self.__pydantic_extra__ or {}) # cast-ok: pydantic owns extra params
class ClassifierLLMConfig(BaseModel):
"""Configuration for the LLM-based complexity classifier."""
@ -279,7 +309,7 @@ class ComplexityRouterConfig(BaseModel):
"""Configuration for the ComplexityRouter."""
# string = pin; list = random pick when adaptive=False, soft-floor home pool when adaptive=True
tiers: dict[str, str | list[str]] = Field(
tiers: dict[str, str | list[str] | TierTarget] = Field(
default_factory=lambda: DEFAULT_TIER_MODELS.copy(),
description=(
"Mapping of complexity tiers to a model or model pool. "
@ -514,14 +544,21 @@ class ComplexityRouterConfig(BaseModel):
def _coerce_tier_values(cls, value: object) -> object:
if not isinstance(value, dict):
return value
coerced: Final[dict[str, object]] = {}
for key, item in value.items():
if isinstance(item, str):
coerced[key] = item
elif isinstance(item, (list, tuple)):
coerced[key] = list(item)
else:
coerced[key] = item
coerced: Final[dict[str, object]] = { # mutable-ok: pydantic requires a mutable input mapping
key: list(item)
if isinstance(item, tuple)
else TierTarget.model_validate(item)
if isinstance(item, dict)
else item
for key, item in cast(dict[str, object], value).items() # cast-ok: validator narrowed the mapping
}
invalid: Final = next(
((key, item) for key, item in coerced.items() if not isinstance(item, (str, list, TierTarget))),
None,
)
if invalid is not None:
key, _ = invalid
raise ValueError(f"tier {key!r} must be a string, list of strings, or object with a model")
return coerced
@field_validator("escalation_keywords")
@ -537,14 +574,57 @@ class ComplexityRouterConfig(BaseModel):
raise ValueError("classifier_llm_config is required when classifier_type is 'llm'")
return self
@model_validator(mode="after")
def _warn_unknown_tier_params(self) -> "ComplexityRouterConfig":
known: Final = frozenset(OPENAI_CHAT_COMPLETION_PARAMS) | frozenset(all_litellm_params)
for tier, target in self.tiers.items():
if isinstance(target, TierTarget):
for key in target.params:
if key not in known:
verbose_router_logger.warning(
"ComplexityRouter tier %s has an unrecognized parameter key: %s",
tier,
key,
)
return self
@model_validator(mode="after")
def _validate_tier_models(self) -> "ComplexityRouterConfig":
for tier, target in self.tiers.items():
if isinstance(target, TierTarget):
continue
if isinstance(target, str) and not target.strip():
raise ValueError(f"tier {tier!r} model must be a non-empty string")
if not target:
raise ValueError(f"tier {tier!r} model pool must be non-empty")
return self
@model_validator(mode="after")
def _validate_adaptive_pools(self) -> "ComplexityRouterConfig":
if not self.adaptive:
return self
normalized = {tier: (models if isinstance(models, list) else [models]) for tier, models in self.tiers.items()}
if not any(normalized.values()):
normalized: Final[
dict[str, str | list[str] | TierTarget]
] = { # mutable-ok: pydantic requires normalized tier mappings
tier: (
target.model_copy(
update={"model": [*target.model] if isinstance(target.model, list) else [target.model]}
)
if isinstance(target, TierTarget)
else models
if isinstance(models, list)
else [models]
)
for tier, models in self.tiers.items()
for target in (models,)
}
if not any(target.model if isinstance(target, TierTarget) else target for target in normalized.values()):
raise ValueError("adaptive=True requires at least one non-empty tier pool")
empty: Final = [tier for tier, models in normalized.items() if not models]
empty: Final = [
tier
for tier, target in normalized.items()
if not (target.model if isinstance(target, TierTarget) else target)
]
if empty:
raise ValueError(f"adaptive=True tier pools must be non-empty; empty tiers: {empty}")
self.tiers = normalized

View file

@ -819,6 +819,7 @@ class PreRoutingHookResponse(BaseModel):
model: str
messages: list[dict[str, Any]] | None
params: dict[str, Any] | None = None
routing_decision: StandardLoggingRoutingDecision | None = None

View file

@ -37,6 +37,7 @@ from litellm.router_strategy.complexity_router.config import (
DEFAULT_TECHNICAL_KEYWORDS,
ComplexityRouterConfig,
ComplexityTier,
TierTarget,
)
from litellm.types.router import (
Deployment,
@ -156,6 +157,126 @@ class TestComplexityRouterInit:
metadata = request_kwargs.get("metadata", {})
assert metadata.get(RETURN_RAW_MODEL_NAME_METADATA_KEY, False) is return_raw_model_name
@pytest.mark.parametrize(
"tier_value",
[
"",
[],
{"reasoning_effort": "low"},
{"model": 3},
],
)
def test_tier_target_validation_errors_are_clear(self, tier_value):
with pytest.raises(ValidationError, match="tier|model|pool"):
ComplexityRouterConfig(tiers={"SIMPLE": tier_value})
def test_tier_target_accepts_pool_and_extra_params(self):
config = ComplexityRouterConfig(
tiers={"SIMPLE": {"model": ["gpt-5", "o3"], "reasoning_effort": "high"}}
)
target = config.tiers["SIMPLE"]
assert isinstance(target, TierTarget)
assert target.model == ["gpt-5", "o3"]
assert target.params == {"reasoning_effort": "high"}
def test_unknown_tier_param_warns_and_is_preserved(self, caplog):
with caplog.at_level(logging.WARNING, logger=verbose_router_logger.name):
config = ComplexityRouterConfig(tiers={"SIMPLE": {"model": "gpt-5", "thinking_level": "high"}})
assert "thinking_level" in caplog.text
assert isinstance(config.tiers["SIMPLE"], TierTarget)
assert config.tiers["SIMPLE"].params["thinking_level"] == "high"
@pytest.mark.asyncio
async def test_tier_params_are_applied_with_request_and_alias_precedence(self):
router = Router(
model_list=[
{
"model_name": "smart-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"temperature": 0.1,
"complexity_router_config": {
"tiers": {
"SIMPLE": {"model": "cheap", "temperature": 0.2},
"MEDIUM": "cheap",
"COMPLEX": "cheap",
"REASONING": "cheap",
},
"keyword_tier_rules": [{"keywords": ["force"], "tier": "SIMPLE"}],
"session_affinity": False,
},
},
},
{"model_name": "cheap", "litellm_params": {"model": "openai/gpt-4o-mini"}},
]
)
request_kwargs = {"temperature": 0.3}
response = await router.async_pre_routing_hook(
model="smart-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "force this request"}],
)
assert response is not None
assert request_kwargs["temperature"] == 0.3
request_kwargs = {}
await router.async_pre_routing_hook(
model="smart-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "force this request"}],
)
assert request_kwargs["temperature"] == 0.2
@pytest.mark.asyncio
async def test_adaptive_cross_tier_model_uses_model_pool_tier_params(self, mock_router_instance):
router = ComplexityRouter(
model_name="hybrid",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"adaptive": True,
"adaptive_eligible": "all",
"tiers": {
"SIMPLE": {"model": "cheap", "reasoning_effort": "low"},
"MEDIUM": {"model": "premium", "reasoning_effort": "high"},
},
},
)
router.adaptive_router = MagicMock()
with patch.object(router, "_soft_floor_pick", return_value="premium"):
response = await router.async_pre_routing_hook(
model="hybrid",
request_kwargs={},
messages=[{"role": "user", "content": "hello"}],
)
assert response is not None
assert response.params == {"reasoning_effort": "high"}
@pytest.mark.asyncio
async def test_session_affinity_pin_keeps_pinned_tier_params(self, mock_router_instance):
mock_router_instance.cache = DualCache()
router = ComplexityRouter(
model_name="hybrid",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"session_affinity": True,
"tiers": {
"SIMPLE": {"model": "cheap", "reasoning_effort": "low"},
"MEDIUM": "mid",
},
},
)
request_kwargs = {"metadata": {"session_id": "session-params"}}
cache_key = router._get_session_affinity_cache_key("session-params", request_kwargs)
await mock_router_instance.cache.async_set_cache(key=cache_key, value="cheap")
response = await router.async_pre_routing_hook(
model="hybrid",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hello"}],
)
assert response is not None
assert response.params == {"reasoning_effort": "low"}
class TestTokenScoring:
"""Test token count scoring."""
@ -750,10 +871,7 @@ class TestSingletonMutation:
def test_default_config_not_mutated(self, mock_router_instance):
"""Test that creating routers without config doesn't mutate defaults."""
from litellm.router_strategy.complexity_router.config import (
DEFAULT_CLASSIFIER_CONTEXT_WINDOW_SIZE,
ComplexityRouterConfig,
)
from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig
# Get original default
original_default = ComplexityRouterConfig().default_model

View file

@ -31737,7 +31737,7 @@ export interface components {
* @description Mapping of complexity tiers to a model or model pool. A list is randomly picked from when adaptive=False, and used as a soft-floor home pool when adaptive=True
*/
tiers?: {
[key: string]: string | string[];
[key: string]: string | string[] | components["schemas"]["TierTarget"];
};
/**
* Token Thresholds
@ -33259,6 +33259,13 @@ export interface components {
[key: string]: unknown;
};
};
/** TierTarget */
TierTarget: {
/** Model */
model: string | string[];
} & {
[key: string]: unknown;
};
/**
* TokenCountDetailsResponse
* @description Response structure for token count details with modality breakdown.