fix(router): finalize tier parameter routing

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

View file

@ -47,7 +47,7 @@ from .config import (
ComplexityRouterConfig,
ComplexityTier,
TierTarget,
_tier_pool,
tier_pool,
)
if TYPE_CHECKING:
@ -1157,7 +1157,12 @@ class ComplexityRouter(CustomLogger):
model: str | list[str] | TierTarget, tier_key: str
) -> str: # mutable-ok: legacy pool inputs remain lists
if isinstance(model, TierTarget):
model = model.model
target_model: Final = model.model
if isinstance(target_model, str):
return target_model
if not target_model:
raise ValueError(f"Empty model pool for tier {tier_key}")
return random.choice(target_model)
if isinstance(model, str):
return model
if not model:
@ -1166,7 +1171,7 @@ class ComplexityRouter(CustomLogger):
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: _tier_pool(target) for tier, target in self.config.tiers.items()
tier: tier_pool(target) for tier, target in self.config.tiers.items()
}
def _tier_params(
@ -1883,7 +1888,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
params=None,
routing_decision=self._build_routing_decision(
routed_model=fallback_model,
conversation_continuing=conversation_continuing,

View file

@ -269,9 +269,10 @@ class TierTarget(BaseModel):
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
def tier_pool(value: str | list[str] | TierTarget) -> list[str]: # mutable-ok: routing pools use list semantics
if isinstance(value, TierTarget):
value = value.model
target_model: Final = value.model
return target_model if isinstance(target_model, list) else [target_model]
return value if isinstance(value, list) else [value]
@ -601,14 +602,14 @@ class ComplexityRouterConfig(BaseModel):
normalized: Final[
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": _tier_pool(target)})
tier: target.model_copy(update={"model": tier_pool(target)})
if isinstance(target, TierTarget)
else _tier_pool(target)
else tier_pool(target)
for tier, target in self.tiers.items()
}
if not any(_tier_pool(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 _tier_pool(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

@ -227,6 +227,48 @@ class TestComplexityRouterInit:
)
assert request_kwargs["temperature"] == 0.2
@pytest.mark.asyncio
async def test_tier_params_reach_upstream_acompletion(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"}},
]
)
response = litellm.ModelResponse(id="test-response", choices=[])
with patch("litellm.acompletion", new_callable=AsyncMock, return_value=response) as mock_acompletion:
await router.acompletion(
model="smart-router",
messages=[{"role": "user", "content": "force this request"}],
)
assert mock_acompletion.await_args is not None
assert mock_acompletion.await_args.kwargs["temperature"] == 0.2
with patch("litellm.acompletion", new_callable=AsyncMock, return_value=response) as mock_acompletion:
await router.acompletion(
model="smart-router",
messages=[{"role": "user", "content": "force this request"}],
temperature=0.3,
)
assert mock_acompletion.await_args is not None
assert mock_acompletion.await_args.kwargs["temperature"] == 0.3
@pytest.mark.asyncio
async def test_adaptive_cross_tier_model_uses_model_pool_tier_params(self, mock_router_instance):
router = ComplexityRouter(