mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
e7bb85b430
commit
5bfdd07bf4
3 changed files with 58 additions and 10 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue