From 5bfdd07bf4629e78785f4f0ce83ef38430b1aacf Mon Sep 17 00:00:00 2001 From: tin Date: Thu, 6 Aug 2026 04:21:16 +0000 Subject: [PATCH] fix(router): finalize tier parameter routing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../complexity_router/complexity_router.py | 13 ++++-- .../complexity_router/config.py | 13 +++--- .../router_strategy/test_complexity_router.py | 42 +++++++++++++++++++ 3 files changed, 58 insertions(+), 10 deletions(-) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 6e9eafa492b..d6f67120cef 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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, diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 048afa8bb22..15d70035893 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -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 diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index faa2286d3a2..b37c574b5b0 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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(