fix(router): preserve legacy complexity tier pools

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
tin 2026-08-06 04:13:57 +00:00
parent fe1e42b949
commit e7bb85b430
5 changed files with 41 additions and 58 deletions

View file

@ -11204,11 +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)
tier_params: Final = pre_routing_hook_response.params
if tier_params is not None:
for key, value in tier_params.items():
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

@ -47,6 +47,7 @@ from .config import (
ComplexityRouterConfig,
ComplexityTier,
TierTarget,
_tier_pool,
)
if TYPE_CHECKING:
@ -1152,7 +1153,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] | TierTarget, tier_key: str) -> str:
def _pick_from_tier_value(
model: str | list[str] | TierTarget, tier_key: str
) -> str: # mutable-ok: legacy pool inputs remain lists
if isinstance(model, TierTarget):
model = model.model
if isinstance(model, str):
@ -1161,27 +1164,23 @@ class ComplexityRouter(CustomLogger):
raise ValueError(f"Empty model pool for tier {tier_key}")
return random.choice(model)
def _tier_pools(self) -> dict[str, list[str]]:
def _tier_pools(self) -> dict[str, list[str]]: # mutable-ok: adaptive router consumes mutable pools
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,)
tier: _tier_pool(target) for tier, target in self.config.tiers.items()
}
def _tier_params(self, tier: ComplexityTier | str) -> dict[str, object]:
def _tier_params(
self, tier: ComplexityTier | str
) -> dict[str, object] | None: # mutable-ok: params are merged into request kwargs
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
return target.params or None if isinstance(target, TierTarget) else None
def _params_for_model(self, model: str) -> dict[str, object]:
def _params_for_model(
self, model: str
) -> dict[str, object] | None: # mutable-ok: params are merged into request kwargs
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
return self._tier_params(tier) if tier is not None else None
async def _pick_model_for_tier(
self,
@ -1802,6 +1801,7 @@ class ComplexityRouter(CustomLogger):
# priority exactly (changing it would be a silent behavior change for
# every non-plugin user, not just a security fix).
routed_model = self.config.default_model
routed_params = None
else:
# Plugins configured: default_model must never bypass them, so it's not
# checked here at all -- _pick_model_for_tier -> get_model_for_tier still
@ -1809,12 +1809,11 @@ class ComplexityRouter(CustomLogger):
routed_model = await self._pick_model_for_tier(
ComplexityTier.MEDIUM, messages, resolved_messages, request_kwargs
)
routed_params = self._tier_params(ComplexityTier.MEDIUM)
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
params=routed_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause="default_fallback",

View file

@ -243,7 +243,7 @@ DEFAULT_TIER_MODELS: Final[dict[str, str]] = {
class TierTarget(BaseModel):
model: str | list[str]
model: str | list[str] # mutable-ok: public config accepts model pools as lists
model_config = ConfigDict(extra="allow")
@ -265,10 +265,16 @@ class TierTarget(BaseModel):
raise ValueError("model must be a string or a list of strings")
@property
def params(self) -> dict[str, object]:
def params(self) -> dict[str, object]: # mutable-ok: extras are passed through as request kwargs
return cast(dict[str, object], self.__pydantic_extra__ or {}) # cast-ok: pydantic owns extra params
def _tier_pool(value: str | list[str] | TierTarget) -> list[str]: # mutable-ok: routing pools use list semantics
if isinstance(value, TierTarget):
value = value.model
return value if isinstance(value, list) else [value]
class ClassifierLLMConfig(BaseModel):
"""Configuration for the LLM-based complexity classifier."""
@ -309,7 +315,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] | TierTarget] = Field(
tiers: dict[str, str | list[str] | TierTarget] = Field( # mutable-ok: pydantic config surface is mutable
default_factory=lambda: DEFAULT_TIER_MODELS.copy(),
description=(
"Mapping of complexity tiers to a model or model pool. "
@ -588,43 +594,21 @@ class ComplexityRouterConfig(BaseModel):
)
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: Final[
dict[str, str | list[str] | TierTarget]
dict[str, str | list[str] | TierTarget] # mutable-ok: pydantic config surface is mutable
] = { # 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,)
tier: target.model_copy(update={"model": _tier_pool(target)})
if isinstance(target, TierTarget)
else _tier_pool(target)
for tier, target in self.tiers.items()
}
if not any(target.model if isinstance(target, TierTarget) else target for target in normalized.values()):
if not any(_tier_pool(target) for target in normalized.values()):
raise ValueError("adaptive=True requires at least one non-empty tier pool")
empty: Final = [
tier
for tier, target in normalized.items()
if not (target.model if isinstance(target, TierTarget) else target)
]
empty: Final = [tier for tier, target in normalized.items() if not _tier_pool(target)]
if empty:
raise ValueError(f"adaptive=True tier pools must be non-empty; empty tiers: {empty}")
self.tiers = normalized

View file

@ -819,7 +819,7 @@ class PreRoutingHookResponse(BaseModel):
model: str
messages: list[dict[str, Any]] | None
params: dict[str, Any] | None = None
params: dict[str, Any] | None = None # mutable-ok: router merges params into request kwargs
routing_decision: StandardLoggingRoutingDecision | None = None

View file

@ -160,8 +160,8 @@ class TestComplexityRouterInit:
@pytest.mark.parametrize(
"tier_value",
[
"",
[],
{"model": ""},
{"model": []},
{"reasoning_effort": "low"},
{"model": 3},
],