diff --git a/litellm/router.py b/litellm/router.py index 9f783b3223b..09015d7ead5 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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", {}) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 0e00a33a338..6e9eafa492b 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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", diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index c25eab3f65e..048afa8bb22 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -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 diff --git a/litellm/types/router.py b/litellm/types/router.py index fd47b1dafe5..e44b9a8d1ca 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 2908e0abcb9..faa2286d3a2 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -160,8 +160,8 @@ class TestComplexityRouterInit: @pytest.mark.parametrize( "tier_value", [ - "", - [], + {"model": ""}, + {"model": []}, {"reasoning_effort": "low"}, {"model": 3}, ],