mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
fe1e42b949
commit
e7bb85b430
5 changed files with 41 additions and 58 deletions
|
|
@ -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", {})
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -160,8 +160,8 @@ class TestComplexityRouterInit:
|
|||
@pytest.mark.parametrize(
|
||||
"tier_value",
|
||||
[
|
||||
"",
|
||||
[],
|
||||
{"model": ""},
|
||||
{"model": []},
|
||||
{"reasoning_effort": "low"},
|
||||
{"model": 3},
|
||||
],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue