mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
d26ef670e2
commit
fe1e42b949
7 changed files with 291 additions and 20 deletions
|
|
@ -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", {})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
9
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
9
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue