mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(quality_router): replace capabilities with keyword override
Drops the capability-based filtering in favor of a keyword-based override for v0: - RoutingPreferences.keywords: List[str] (replaces capabilities) — each deployment can declare substring keywords. - If any declared keyword (case-insensitive) appears in the user message, the router short-circuits the complexity-classification flow and routes to the matching deployment. - Tiebreaker for overlapping keyword matches: quality_tier DESC, then cheapest model_info.input_cost_per_token ASC. Unpriced models lose ties to priced ones. Decision metadata + headers now expose the override: x-litellm-quality-router-via → "keyword" | "quality_tier" x-litellm-quality-router-keyword → matched keyword (only on keyword route) x-litellm-quality-router-complexity → complexity tier (only on tier route) Removes: - request_kwargs["litellm_capabilities"] reading - _model_capabilities, _model_supports_capabilities, _first_capable_model_at_tier, capability filter in _resolve_model_for_quality_tier Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
103512f2f4
commit
1d8828665f
4 changed files with 373 additions and 222 deletions
|
|
@ -8246,15 +8246,27 @@ class Router:
|
|||
else None
|
||||
)
|
||||
if isinstance(decision, dict):
|
||||
if "routed_model" in decision:
|
||||
# Only emit headers for fields that have a meaningful value.
|
||||
# `complexity_tier` and `matched_keyword` are mutually exclusive
|
||||
# (the keyword path short-circuits classification), so each
|
||||
# request emits one or the other but not both.
|
||||
if decision.get("routed_model") is not None:
|
||||
additional_headers["x-litellm-quality-router-model"] = str(
|
||||
decision["routed_model"]
|
||||
)
|
||||
if "quality_tier" in decision:
|
||||
if decision.get("quality_tier") is not None:
|
||||
additional_headers["x-litellm-quality-router-tier"] = str(
|
||||
decision["quality_tier"]
|
||||
)
|
||||
if "complexity_tier" in decision:
|
||||
if decision.get("routed_via") is not None:
|
||||
additional_headers["x-litellm-quality-router-via"] = str(
|
||||
decision["routed_via"]
|
||||
)
|
||||
if decision.get("matched_keyword") is not None:
|
||||
additional_headers["x-litellm-quality-router-keyword"] = str(
|
||||
decision["matched_keyword"]
|
||||
)
|
||||
if decision.get("complexity_tier") is not None:
|
||||
additional_headers["x-litellm-quality-router-complexity"] = str(
|
||||
decision["complexity_tier"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -48,13 +48,13 @@ class RoutingPreferences(BaseModel):
|
|||
description="The quality tier this deployment satisfies.",
|
||||
)
|
||||
|
||||
capabilities: List[str] = Field(
|
||||
keywords: List[str] = Field(
|
||||
default_factory=list,
|
||||
description=(
|
||||
"Capability tags this deployment supports (e.g. 'vision', "
|
||||
"'function_calling', 'json_mode'). The QualityRouter will only "
|
||||
"route to deployments whose capabilities are a superset of any "
|
||||
"capabilities required by the request."
|
||||
"Substring keywords (case-insensitive) that, when present in the "
|
||||
"user message, route the request to this deployment. When multiple "
|
||||
"deployments match, ties are broken by (highest quality_tier, "
|
||||
"then cheapest model_info.input_cost_per_token)."
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -6,9 +6,17 @@ inferred by re-using the existing ComplexityRouter's classification, then
|
|||
mapped through an admin-configured `complexity_to_quality` table. Each
|
||||
candidate model declares its own `quality_tier` in
|
||||
`model_info.litellm_routing_preferences`.
|
||||
|
||||
Optional keyword override: deployments may also declare `keywords` in
|
||||
`litellm_routing_preferences`. If any declared keyword appears in the user
|
||||
message (case-insensitive substring match), the router short-circuits the
|
||||
complexity-classification flow and routes to the matching deployment. When
|
||||
multiple deployments match, ties are broken by (highest quality_tier first,
|
||||
then cheapest `model_info.input_cost_per_token`).
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, FrozenSet, List, Optional, Set, Union
|
||||
import math
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -28,15 +36,8 @@ else:
|
|||
|
||||
class QualityRouter(CustomLogger):
|
||||
"""
|
||||
Routes requests to a model at a target quality tier.
|
||||
|
||||
Pipeline:
|
||||
1. Classify the user message via ComplexityRouter to get a ComplexityTier.
|
||||
2. Map that tier name to a target quality tier (int) via
|
||||
`config.complexity_to_quality`.
|
||||
3. Resolve the target quality tier to a concrete model using the
|
||||
per-deployment `quality_tier` declared in
|
||||
`model_info.litellm_routing_preferences`.
|
||||
Routes requests to a model at a target quality tier, with an optional
|
||||
keyword override.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -64,10 +65,15 @@ class QualityRouter(CustomLogger):
|
|||
litellm_router_instance=litellm_router_instance,
|
||||
)
|
||||
|
||||
# Pre-built tier → models index for O(1) resolution. Capabilities are
|
||||
# tracked separately so resolution can filter by required capabilities
|
||||
# without complicating the tier-walk loop.
|
||||
self._model_capabilities: Dict[str, FrozenSet[str]] = {}
|
||||
# Per-model indices populated alongside the tier index. `_model_keywords`
|
||||
# stores keywords lowercased so we can substring-match against the
|
||||
# lowercased user message in O(total-keyword-count). `_model_quality`
|
||||
# and `_model_cost` are needed for keyword-match tiebreaking.
|
||||
self._model_keywords: Dict[str, List[str]] = {}
|
||||
self._model_quality: Dict[str, int] = {}
|
||||
self._model_cost: Dict[str, Optional[float]] = {}
|
||||
|
||||
# Pre-built tier → models index for O(1) tier resolution.
|
||||
self._tier_to_models: Dict[int, List[str]] = self._build_tier_index()
|
||||
|
||||
verbose_router_logger.debug(
|
||||
|
|
@ -98,6 +104,31 @@ class QualityRouter(CustomLogger):
|
|||
return model_info.get("litellm_routing_preferences")
|
||||
return getattr(model_info, "litellm_routing_preferences", None)
|
||||
|
||||
def _get_deployment_input_cost(self, deployment: Any) -> Optional[float]:
|
||||
"""
|
||||
Extract `input_cost_per_token` from a deployment's model_info.
|
||||
|
||||
Returns None when not declared — None is treated as "infinite cost"
|
||||
for the cheapest-tiebreak ordering, so unpriced models lose ties to
|
||||
priced ones. (Admins who want a model to win on price must declare it.)
|
||||
"""
|
||||
if isinstance(deployment, dict):
|
||||
model_info = deployment.get("model_info") or {}
|
||||
else:
|
||||
model_info = getattr(deployment, "model_info", None) or {}
|
||||
|
||||
if isinstance(model_info, dict):
|
||||
cost = model_info.get("input_cost_per_token")
|
||||
else:
|
||||
cost = getattr(model_info, "input_cost_per_token", None)
|
||||
|
||||
if cost is None:
|
||||
return None
|
||||
try:
|
||||
return float(cost)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _get_deployment_model_name(self, deployment: Any) -> Optional[str]:
|
||||
"""Extract `model_name` from a dict- or object-shaped deployment."""
|
||||
if isinstance(deployment, dict):
|
||||
|
|
@ -107,8 +138,9 @@ class QualityRouter(CustomLogger):
|
|||
def _build_tier_index(self) -> Dict[int, List[str]]:
|
||||
"""
|
||||
Build {quality_tier: [model_name, ...]} for every model in
|
||||
`available_models`. Raises if any listed model is missing
|
||||
`litellm_routing_preferences`.
|
||||
`available_models`, plus side indices `_model_keywords`,
|
||||
`_model_quality`, and `_model_cost`. Raises if any listed model is
|
||||
missing `litellm_routing_preferences`.
|
||||
"""
|
||||
model_list = getattr(self.litellm_router_instance, "model_list", None) or []
|
||||
available = set(self.config.available_models)
|
||||
|
|
@ -132,10 +164,10 @@ class QualityRouter(CustomLogger):
|
|||
# Accept dict or Pydantic-shaped prefs.
|
||||
if isinstance(prefs, dict):
|
||||
tier = prefs.get("quality_tier")
|
||||
capabilities = prefs.get("capabilities") or []
|
||||
keywords = prefs.get("keywords") or []
|
||||
else:
|
||||
tier = getattr(prefs, "quality_tier", None)
|
||||
capabilities = getattr(prefs, "capabilities", None) or []
|
||||
keywords = getattr(prefs, "keywords", None) or []
|
||||
|
||||
if tier is None:
|
||||
raise ValueError(
|
||||
|
|
@ -145,7 +177,9 @@ class QualityRouter(CustomLogger):
|
|||
|
||||
tier_int = int(tier)
|
||||
tier_to_models.setdefault(tier_int, []).append(name)
|
||||
self._model_capabilities[name] = frozenset(capabilities)
|
||||
self._model_keywords[name] = [str(k).lower() for k in keywords if k]
|
||||
self._model_quality[name] = tier_int
|
||||
self._model_cost[name] = self._get_deployment_input_cost(deployment)
|
||||
seen[name] = True
|
||||
|
||||
missing = [name for name, found in seen.items() if not found]
|
||||
|
|
@ -157,63 +191,54 @@ class QualityRouter(CustomLogger):
|
|||
|
||||
return tier_to_models
|
||||
|
||||
def _model_supports_capabilities(
|
||||
self, model_name: str, required: FrozenSet[str]
|
||||
) -> bool:
|
||||
"""True if the model's declared capabilities are a superset of required."""
|
||||
if not required:
|
||||
return True
|
||||
return required.issubset(self._model_capabilities.get(model_name, frozenset()))
|
||||
def _keyword_override(self, user_message: str) -> Optional[Tuple[str, str]]:
|
||||
"""
|
||||
Find a deployment whose declared keywords appear in `user_message`.
|
||||
|
||||
def _first_capable_model_at_tier(
|
||||
self, tier: int, required: FrozenSet[str]
|
||||
) -> Optional[str]:
|
||||
"""First model at `tier` that supports all `required` capabilities, or None."""
|
||||
for name in self._tier_to_models.get(tier, []):
|
||||
if self._model_supports_capabilities(name, required):
|
||||
return name
|
||||
return None
|
||||
Returns (model_name, matched_keyword) or None when no keyword matches.
|
||||
When multiple deployments match, sorts by (quality_tier DESC,
|
||||
input_cost_per_token ASC, model_name ASC) and returns the winner.
|
||||
Unpriced models are treated as `+inf` so priced models win on price.
|
||||
"""
|
||||
text = user_message.lower()
|
||||
|
||||
def _resolve_model_for_quality_tier(
|
||||
self,
|
||||
tier: int,
|
||||
required_capabilities: Optional[Set[str]] = None,
|
||||
) -> str:
|
||||
matches: List[Tuple[str, str]] = [] # (model_name, matched_keyword)
|
||||
for model_name, keywords in self._model_keywords.items():
|
||||
for kw in keywords:
|
||||
if kw and kw in text:
|
||||
matches.append((model_name, kw))
|
||||
break # one match per model is enough
|
||||
|
||||
if not matches:
|
||||
return None
|
||||
|
||||
def sort_key(match: Tuple[str, str]) -> Tuple[int, float, str]:
|
||||
name = match[0]
|
||||
quality = self._model_quality.get(name, 0)
|
||||
cost = self._model_cost.get(name)
|
||||
cost_val = cost if cost is not None else math.inf
|
||||
# Negate quality so higher tier sorts first under ASC sort.
|
||||
return (-quality, cost_val, name)
|
||||
|
||||
matches.sort(key=sort_key)
|
||||
return matches[0]
|
||||
|
||||
def _resolve_model_for_quality_tier(self, tier: int) -> str:
|
||||
"""
|
||||
Resolve a quality tier to a concrete model name.
|
||||
|
||||
Strategy:
|
||||
1. Exact tier match → first capability-matching model at that tier.
|
||||
2. Otherwise round up to the next higher tier that has a
|
||||
capability-matching model.
|
||||
3. Otherwise fall back to `config.default_model` — but only if it
|
||||
also satisfies required capabilities. Routing to a model that
|
||||
lacks a required capability would silently produce wrong results.
|
||||
1. Exact tier match → first model registered at that tier.
|
||||
2. Otherwise round up to the next higher tier that has a model.
|
||||
3. Otherwise fall back to `config.default_model`.
|
||||
"""
|
||||
required: FrozenSet[str] = (
|
||||
frozenset(required_capabilities) if required_capabilities else frozenset()
|
||||
)
|
||||
|
||||
match = self._first_capable_model_at_tier(tier, required)
|
||||
if match is not None:
|
||||
return match
|
||||
if tier in self._tier_to_models and self._tier_to_models[tier]:
|
||||
return self._tier_to_models[tier][0]
|
||||
|
||||
higher_tiers = sorted(t for t in self._tier_to_models if t > tier)
|
||||
for t in higher_tiers:
|
||||
match = self._first_capable_model_at_tier(t, required)
|
||||
if match is not None:
|
||||
return match
|
||||
|
||||
if self.config.default_model and self._model_supports_capabilities(
|
||||
self.config.default_model, required
|
||||
):
|
||||
return self.config.default_model
|
||||
|
||||
if required:
|
||||
raise ValueError(
|
||||
f"QualityRouter: no model satisfies quality tier {tier} with "
|
||||
f"required capabilities {sorted(required)}"
|
||||
)
|
||||
if self._tier_to_models[t]:
|
||||
return self._tier_to_models[t][0]
|
||||
|
||||
if self.config.default_model:
|
||||
return self.config.default_model
|
||||
|
|
@ -223,6 +248,22 @@ class QualityRouter(CustomLogger):
|
|||
f"no default_model configured"
|
||||
)
|
||||
|
||||
def _stash_decision(
|
||||
self,
|
||||
request_kwargs: Optional[Dict[str, Any]],
|
||||
decision: Dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
Stash the routing decision in request_kwargs.metadata so the Router can
|
||||
lift it into response headers (`x-litellm-quality-router-*`). The same
|
||||
dict object flows from here through to `make_call.set_response_headers`.
|
||||
"""
|
||||
if request_kwargs is None:
|
||||
return
|
||||
metadata = request_kwargs.setdefault("metadata", {})
|
||||
if isinstance(metadata, dict):
|
||||
metadata["quality_router_decision"] = decision
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -231,7 +272,7 @@ class QualityRouter(CustomLogger):
|
|||
input: Optional[Union[str, List]] = None,
|
||||
specific_deployment: Optional[bool] = False,
|
||||
) -> Optional["PreRoutingHookResponse"]:
|
||||
"""Classify the request, map to a quality tier, resolve the model."""
|
||||
"""Try keyword override first; fall back to complexity-tier routing."""
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
if messages is None or len(messages) == 0:
|
||||
|
|
@ -261,16 +302,6 @@ class QualityRouter(CustomLogger):
|
|||
elif role == "system" and system_prompt is None:
|
||||
system_prompt = content
|
||||
|
||||
# Required capabilities are an optional client-side override.
|
||||
# Accept either an iterable of strings or None. Anything else is ignored
|
||||
# rather than raising — matches the lenient style of other router params.
|
||||
raw_caps = (request_kwargs or {}).get("litellm_capabilities")
|
||||
required_capabilities: Optional[Set[str]] = (
|
||||
{str(c) for c in raw_caps}
|
||||
if isinstance(raw_caps, (list, tuple, set, frozenset)) and raw_caps
|
||||
else None
|
||||
)
|
||||
|
||||
if user_message is None:
|
||||
verbose_router_logger.debug(
|
||||
"QualityRouter: No user message found, routing to default model"
|
||||
|
|
@ -279,19 +310,38 @@ class QualityRouter(CustomLogger):
|
|||
raise ValueError(
|
||||
"QualityRouter: no user message and no default_model configured"
|
||||
)
|
||||
if required_capabilities and not self._model_supports_capabilities(
|
||||
self.config.default_model, frozenset(required_capabilities)
|
||||
):
|
||||
raise ValueError(
|
||||
f"QualityRouter: no user message and default_model "
|
||||
f"'{self.config.default_model}' does not satisfy required "
|
||||
f"capabilities {sorted(required_capabilities)}"
|
||||
)
|
||||
return PreRoutingHookResponse(
|
||||
model=self.config.default_model,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
# Try keyword override first — it short-circuits complexity classification.
|
||||
keyword_match = self._keyword_override(user_message)
|
||||
if keyword_match is not None:
|
||||
routed_model, matched_keyword = keyword_match
|
||||
verbose_router_logger.info(
|
||||
f"QualityRouter: keyword override matched='{matched_keyword}' "
|
||||
f"routed_model={routed_model} "
|
||||
f"(quality_tier={self._model_quality.get(routed_model)}, "
|
||||
f"input_cost_per_token={self._model_cost.get(routed_model)})"
|
||||
)
|
||||
self._stash_decision(
|
||||
request_kwargs,
|
||||
{
|
||||
"router_model_name": self.model_name,
|
||||
"routed_model": routed_model,
|
||||
"routed_via": "keyword",
|
||||
"matched_keyword": matched_keyword,
|
||||
"quality_tier": self._model_quality.get(routed_model),
|
||||
"complexity_tier": None,
|
||||
},
|
||||
)
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
# No keyword match → complexity classification flow.
|
||||
complexity_tier, score, signals = self._scorer.classify(
|
||||
user_message, system_prompt
|
||||
)
|
||||
|
|
@ -308,33 +358,25 @@ class QualityRouter(CustomLogger):
|
|||
f"in complexity_to_quality mapping {self.config.complexity_to_quality}"
|
||||
)
|
||||
|
||||
routed_model = self._resolve_model_for_quality_tier(
|
||||
int(quality_tier), required_capabilities=required_capabilities
|
||||
)
|
||||
routed_model = self._resolve_model_for_quality_tier(int(quality_tier))
|
||||
|
||||
verbose_router_logger.info(
|
||||
f"QualityRouter: complexity={complexity_name}, score={score:.3f}, "
|
||||
f"signals={signals}, quality_tier={quality_tier}, "
|
||||
f"required_capabilities={sorted(required_capabilities) if required_capabilities else []}, "
|
||||
f"routed_model={routed_model}"
|
||||
)
|
||||
|
||||
# Stash the decision in request_kwargs.metadata so the Router can lift
|
||||
# it into response headers (`x-litellm-quality-router-*`) for
|
||||
# transparency. The same dict object flows from here through to
|
||||
# `make_call.set_response_headers`.
|
||||
if request_kwargs is not None:
|
||||
metadata = request_kwargs.setdefault("metadata", {})
|
||||
if isinstance(metadata, dict):
|
||||
metadata["quality_router_decision"] = {
|
||||
"router_model_name": self.model_name,
|
||||
"routed_model": routed_model,
|
||||
"quality_tier": int(quality_tier),
|
||||
"complexity_tier": complexity_name,
|
||||
"required_capabilities": (
|
||||
sorted(required_capabilities) if required_capabilities else []
|
||||
),
|
||||
}
|
||||
self._stash_decision(
|
||||
request_kwargs,
|
||||
{
|
||||
"router_model_name": self.model_name,
|
||||
"routed_model": routed_model,
|
||||
"routed_via": "quality_tier",
|
||||
"matched_keyword": None,
|
||||
"quality_tier": int(quality_tier),
|
||||
"complexity_tier": complexity_name,
|
||||
},
|
||||
)
|
||||
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
|
|
|
|||
|
|
@ -4,7 +4,9 @@ Tests for the QualityRouter.
|
|||
Covers:
|
||||
- Tier index construction from `model_info.litellm_routing_preferences`.
|
||||
- Quality-tier resolution (exact, round-up, default fallback).
|
||||
- Pre-routing hook end-to-end (classification → quality tier → model).
|
||||
- Keyword override (match, tiebreaking by quality + price).
|
||||
- Pre-routing hook end-to-end.
|
||||
- Decision metadata stash + Router.set_response_headers lift.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
|
@ -29,7 +31,8 @@ def _make_model_list(spec: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
|||
spec entry shape: {
|
||||
"model_name": str,
|
||||
"quality_tier": Optional[int],
|
||||
"capabilities": Optional[List[str]], # default: omitted
|
||||
"keywords": Optional[List[str]],
|
||||
"input_cost_per_token": Optional[float],
|
||||
}
|
||||
If quality_tier is None, the deployment is created without
|
||||
`litellm_routing_preferences`.
|
||||
|
|
@ -39,9 +42,11 @@ def _make_model_list(spec: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
|||
model_info: Dict[str, Any] = {"id": f"id-{entry['model_name']}"}
|
||||
if entry.get("quality_tier") is not None:
|
||||
prefs: Dict[str, Any] = {"quality_tier": entry["quality_tier"]}
|
||||
if "capabilities" in entry:
|
||||
prefs["capabilities"] = entry["capabilities"]
|
||||
if "keywords" in entry:
|
||||
prefs["keywords"] = entry["keywords"]
|
||||
model_info["litellm_routing_preferences"] = prefs
|
||||
if "input_cost_per_token" in entry:
|
||||
model_info["input_cost_per_token"] = entry["input_cost_per_token"]
|
||||
out.append(
|
||||
{
|
||||
"model_name": entry["model_name"],
|
||||
|
|
@ -237,30 +242,44 @@ class TestPreRoutingHook:
|
|||
assert resp.model == "haiku" # the configured default_model
|
||||
|
||||
|
||||
# ─── Capabilities ───────────────────────────────────────────────────────────
|
||||
# ─── Keyword override ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def capability_router():
|
||||
def keyword_router():
|
||||
"""
|
||||
Router with mixed capabilities at each tier:
|
||||
tier 1: haiku-text (no caps), haiku-vision (vision)
|
||||
tier 2: sonnet-text (no caps), sonnet-vision (vision, function_calling)
|
||||
tier 3: opus-vision (vision, function_calling, json_mode)
|
||||
Router where multiple deployments declare overlapping keywords so we can
|
||||
exercise the (quality DESC, price ASC) tiebreak.
|
||||
|
||||
- cheap-coder tier 2, keywords [code, python], cost 0.000001
|
||||
- smart-coder tier 3, keywords [code, python], cost 0.000010
|
||||
- law-bot tier 2, keywords [legal, contract], cost 0.000005
|
||||
- default-haiku tier 1, no keywords, cost 0.0000005
|
||||
"""
|
||||
spec = [
|
||||
{"model_name": "haiku-text", "quality_tier": 1, "capabilities": []},
|
||||
{"model_name": "haiku-vision", "quality_tier": 1, "capabilities": ["vision"]},
|
||||
{"model_name": "sonnet-text", "quality_tier": 2, "capabilities": []},
|
||||
{
|
||||
"model_name": "sonnet-vision",
|
||||
"quality_tier": 2,
|
||||
"capabilities": ["vision", "function_calling"],
|
||||
"model_name": "default-haiku",
|
||||
"quality_tier": 1,
|
||||
"keywords": [],
|
||||
"input_cost_per_token": 0.0000005,
|
||||
},
|
||||
{
|
||||
"model_name": "opus-vision",
|
||||
"model_name": "cheap-coder",
|
||||
"quality_tier": 2,
|
||||
"keywords": ["code", "python"],
|
||||
"input_cost_per_token": 0.000001,
|
||||
},
|
||||
{
|
||||
"model_name": "smart-coder",
|
||||
"quality_tier": 3,
|
||||
"capabilities": ["vision", "function_calling", "json_mode"],
|
||||
"keywords": ["code", "python"],
|
||||
"input_cost_per_token": 0.000010,
|
||||
},
|
||||
{
|
||||
"model_name": "law-bot",
|
||||
"quality_tier": 2,
|
||||
"keywords": ["legal", "contract"],
|
||||
"input_cost_per_token": 0.000005,
|
||||
},
|
||||
]
|
||||
router = MagicMock()
|
||||
|
|
@ -268,103 +287,145 @@ def capability_router():
|
|||
return QualityRouter(
|
||||
model_name="qr",
|
||||
litellm_router_instance=router,
|
||||
default_model="haiku-text",
|
||||
default_model="default-haiku",
|
||||
quality_router_config={
|
||||
"available_models": [
|
||||
"haiku-text",
|
||||
"haiku-vision",
|
||||
"sonnet-text",
|
||||
"sonnet-vision",
|
||||
"opus-vision",
|
||||
"default-haiku",
|
||||
"cheap-coder",
|
||||
"smart-coder",
|
||||
"law-bot",
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class TestCapabilities:
|
||||
def test_index_records_capabilities(self, capability_router):
|
||||
assert capability_router._model_capabilities["haiku-text"] == frozenset()
|
||||
assert capability_router._model_capabilities["haiku-vision"] == frozenset(
|
||||
{"vision"}
|
||||
)
|
||||
assert capability_router._model_capabilities["opus-vision"] == frozenset(
|
||||
{"vision", "function_calling", "json_mode"}
|
||||
class TestKeywordOverride:
|
||||
def test_no_keyword_in_message_returns_none(self, keyword_router):
|
||||
assert keyword_router._keyword_override("hello there") is None
|
||||
|
||||
def test_single_match_returns_that_model(self, keyword_router):
|
||||
# Only law-bot declares "legal".
|
||||
assert keyword_router._keyword_override("review this legal doc") == (
|
||||
"law-bot",
|
||||
"legal",
|
||||
)
|
||||
|
||||
def test_no_required_capabilities_picks_first_in_tier(self, capability_router):
|
||||
# tier 2, no required caps → first registered model at tier 2.
|
||||
assert capability_router._resolve_model_for_quality_tier(2) == "sonnet-text"
|
||||
|
||||
def test_required_capabilities_filter_within_tier(self, capability_router):
|
||||
# tier 2 with vision → must pick sonnet-vision over sonnet-text.
|
||||
assert (
|
||||
capability_router._resolve_model_for_quality_tier(
|
||||
2, required_capabilities={"vision"}
|
||||
)
|
||||
== "sonnet-vision"
|
||||
def test_case_insensitive_match(self, keyword_router):
|
||||
assert keyword_router._keyword_override("LEGAL question") == (
|
||||
"law-bot",
|
||||
"legal",
|
||||
)
|
||||
|
||||
def test_round_up_when_no_capable_model_at_tier(self, capability_router):
|
||||
# tier 1 with function_calling: nothing at tier 1 has it → round up to
|
||||
# tier 2 (sonnet-vision).
|
||||
assert (
|
||||
capability_router._resolve_model_for_quality_tier(
|
||||
1, required_capabilities={"function_calling"}
|
||||
)
|
||||
== "sonnet-vision"
|
||||
def test_overlap_picks_highest_quality_tier(self, keyword_router):
|
||||
# Both cheap-coder (tier 2) and smart-coder (tier 3) declare "code".
|
||||
# Quality wins over price → smart-coder.
|
||||
assert keyword_router._keyword_override("write some code for me") == (
|
||||
"smart-coder",
|
||||
"code",
|
||||
)
|
||||
|
||||
def test_raises_when_no_model_satisfies_capabilities(self, capability_router):
|
||||
# No model anywhere has "audio".
|
||||
with pytest.raises(ValueError, match="audio"):
|
||||
capability_router._resolve_model_for_quality_tier(
|
||||
1, required_capabilities={"audio"}
|
||||
)
|
||||
|
||||
def test_default_model_used_only_if_it_satisfies_caps(self):
|
||||
# Build a router whose default model has NO capabilities, then ask for
|
||||
# a capability that nothing satisfies. Must raise rather than silently
|
||||
# routing to the default.
|
||||
spec = [{"model_name": "only-tier-1", "quality_tier": 1, "capabilities": []}]
|
||||
def test_same_tier_picks_cheapest(self):
|
||||
# Two models at the same tier, both matching "data" — cheapest wins.
|
||||
spec = [
|
||||
{
|
||||
"model_name": "expensive",
|
||||
"quality_tier": 2,
|
||||
"keywords": ["data"],
|
||||
"input_cost_per_token": 0.000050,
|
||||
},
|
||||
{
|
||||
"model_name": "cheap",
|
||||
"quality_tier": 2,
|
||||
"keywords": ["data"],
|
||||
"input_cost_per_token": 0.000005,
|
||||
},
|
||||
]
|
||||
router = MagicMock()
|
||||
router.model_list = _make_model_list(spec)
|
||||
qr = QualityRouter(
|
||||
model_name="qr",
|
||||
litellm_router_instance=router,
|
||||
default_model="only-tier-1",
|
||||
quality_router_config={"available_models": ["only-tier-1"]},
|
||||
default_model="cheap",
|
||||
quality_router_config={"available_models": ["expensive", "cheap"]},
|
||||
)
|
||||
with pytest.raises(ValueError, match="vision"):
|
||||
qr._resolve_model_for_quality_tier(1, required_capabilities={"vision"})
|
||||
match = qr._keyword_override("show me the data")
|
||||
assert match == ("cheap", "data")
|
||||
|
||||
def test_unpriced_loses_to_priced_at_same_tier(self):
|
||||
# Same quality tier, one has cost, one doesn't → priced wins.
|
||||
spec = [
|
||||
{
|
||||
"model_name": "no-price",
|
||||
"quality_tier": 2,
|
||||
"keywords": ["data"],
|
||||
# input_cost_per_token deliberately omitted
|
||||
},
|
||||
{
|
||||
"model_name": "with-price",
|
||||
"quality_tier": 2,
|
||||
"keywords": ["data"],
|
||||
"input_cost_per_token": 0.000005,
|
||||
},
|
||||
]
|
||||
router = MagicMock()
|
||||
router.model_list = _make_model_list(spec)
|
||||
qr = QualityRouter(
|
||||
model_name="qr",
|
||||
litellm_router_instance=router,
|
||||
default_model="no-price",
|
||||
quality_router_config={"available_models": ["no-price", "with-price"]},
|
||||
)
|
||||
match = qr._keyword_override("show me the data")
|
||||
assert match == ("with-price", "data")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_reads_litellm_capabilities_from_request_kwargs(
|
||||
self, capability_router
|
||||
async def test_hook_short_circuits_complexity_on_keyword_match(
|
||||
self, keyword_router
|
||||
):
|
||||
# Simple "hi" → tier 1 by complexity; with vision required, must pick
|
||||
# haiku-vision (the tier-1 model that has vision).
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
resp = await capability_router.async_pre_routing_hook(
|
||||
# A reasoning-style prompt would normally route to a high-quality model
|
||||
# via the complexity flow — but the keyword "code" should short-circuit
|
||||
# to smart-coder (highest tier among "code" models).
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
"Think step by step and reason through this code problem. "
|
||||
"Analyze this carefully and break down each component."
|
||||
),
|
||||
}
|
||||
]
|
||||
request_kwargs: Dict[str, Any] = {}
|
||||
resp = await keyword_router.async_pre_routing_hook(
|
||||
model="qr",
|
||||
request_kwargs={"litellm_capabilities": ["vision"]},
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
)
|
||||
assert resp is not None
|
||||
assert resp.model == "haiku-vision"
|
||||
assert resp.model == "smart-coder"
|
||||
|
||||
decision = request_kwargs["metadata"]["quality_router_decision"]
|
||||
assert decision["routed_via"] == "keyword"
|
||||
assert decision["matched_keyword"] == "code"
|
||||
assert decision["complexity_tier"] is None # short-circuited
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hook_with_no_capabilities_kwarg_behaves_as_before(
|
||||
self, capability_router
|
||||
):
|
||||
async def test_hook_falls_back_to_complexity_when_no_keyword(self, keyword_router):
|
||||
# No declared keyword in the message → complexity-based routing.
|
||||
# "hi" is SIMPLE → quality 1 → default-haiku (the only tier-1 model).
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
resp = await capability_router.async_pre_routing_hook(
|
||||
request_kwargs: Dict[str, Any] = {}
|
||||
resp = await keyword_router.async_pre_routing_hook(
|
||||
model="qr",
|
||||
request_kwargs={},
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
)
|
||||
assert resp is not None
|
||||
# tier 1, no caps required → first registered model at tier 1.
|
||||
assert resp.model == "haiku-text"
|
||||
assert resp.model == "default-haiku"
|
||||
|
||||
decision = request_kwargs["metadata"]["quality_router_decision"]
|
||||
assert decision["routed_via"] == "quality_tier"
|
||||
assert decision["matched_keyword"] is None
|
||||
assert decision["complexity_tier"] == "SIMPLE"
|
||||
|
||||
|
||||
# ─── Routing-decision metadata (powers x-litellm-quality-router-* headers) ──
|
||||
|
|
@ -399,7 +460,8 @@ class TestDecisionMetadata:
|
|||
assert decision["quality_tier"] == 4
|
||||
assert decision["complexity_tier"] == "REASONING"
|
||||
assert decision["router_model_name"] == "quality-router-test"
|
||||
assert decision["required_capabilities"] == []
|
||||
assert decision["routed_via"] == "quality_tier"
|
||||
assert decision["matched_keyword"] is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decision_metadata_preserves_existing_metadata(self, quality_router):
|
||||
|
|
@ -418,24 +480,6 @@ class TestDecisionMetadata:
|
|||
assert request_kwargs["metadata"]["user_id"] == "u-1"
|
||||
assert "quality_router_decision" in request_kwargs["metadata"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decision_metadata_includes_required_capabilities(
|
||||
self, capability_router
|
||||
):
|
||||
request_kwargs: Dict[str, Any] = {
|
||||
"litellm_capabilities": ["vision"],
|
||||
}
|
||||
|
||||
await capability_router.async_pre_routing_hook(
|
||||
model="qr",
|
||||
request_kwargs=request_kwargs,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
|
||||
decision = request_kwargs["metadata"]["quality_router_decision"]
|
||||
assert decision["routed_model"] == "haiku-vision"
|
||||
assert decision["required_capabilities"] == ["vision"]
|
||||
|
||||
|
||||
# ─── Router.set_response_headers lifts decision into x-litellm-quality-* ────
|
||||
|
||||
|
|
@ -477,10 +521,11 @@ class TestSetResponseHeadersLiftsDecision:
|
|||
"metadata": {
|
||||
"quality_router_decision": {
|
||||
"router_model_name": "qr",
|
||||
"routed_model": "haiku-vision",
|
||||
"quality_tier": 1,
|
||||
"complexity_tier": "SIMPLE",
|
||||
"required_capabilities": ["vision"],
|
||||
"routed_model": "smart-coder",
|
||||
"routed_via": "keyword",
|
||||
"matched_keyword": "code",
|
||||
"quality_tier": 3,
|
||||
"complexity_tier": None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -492,12 +537,65 @@ class TestSetResponseHeadersLiftsDecision:
|
|||
)
|
||||
|
||||
headers = response._hidden_params["additional_headers"]
|
||||
assert headers["x-litellm-quality-router-model"] == "haiku-vision"
|
||||
assert headers["x-litellm-quality-router-tier"] == "1"
|
||||
assert headers["x-litellm-quality-router-complexity"] == "SIMPLE"
|
||||
assert headers["x-litellm-quality-router-model"] == "smart-coder"
|
||||
assert headers["x-litellm-quality-router-tier"] == "3"
|
||||
assert headers["x-litellm-quality-router-via"] == "keyword"
|
||||
assert headers["x-litellm-quality-router-keyword"] == "code"
|
||||
# Keyword route short-circuits classification → no complexity header.
|
||||
assert "x-litellm-quality-router-complexity" not in headers
|
||||
# Existing x-litellm-model-group behavior is unchanged.
|
||||
assert headers["x-litellm-model-group"] == "qr"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_quality_tier_route_emits_complexity_not_keyword(self):
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.router import Router
|
||||
|
||||
class FakeResponse(BaseModel):
|
||||
model_config = {"arbitrary_types_allowed": True}
|
||||
_hidden_params: Dict[str, Any] = {}
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "haiku",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "sk-test",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
response = FakeResponse()
|
||||
response._hidden_params = {}
|
||||
|
||||
request_kwargs = {
|
||||
"metadata": {
|
||||
"quality_router_decision": {
|
||||
"router_model_name": "qr",
|
||||
"routed_model": "haiku",
|
||||
"routed_via": "quality_tier",
|
||||
"matched_keyword": None,
|
||||
"quality_tier": 1,
|
||||
"complexity_tier": "SIMPLE",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
await router.set_response_headers(
|
||||
response=response,
|
||||
model_group="qr",
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
headers = response._hidden_params["additional_headers"]
|
||||
assert headers["x-litellm-quality-router-via"] == "quality_tier"
|
||||
assert headers["x-litellm-quality-router-complexity"] == "SIMPLE"
|
||||
# Quality-tier route → no keyword header.
|
||||
assert "x-litellm-quality-router-keyword" not in headers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_decision_leaves_quality_router_headers_unset(self):
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -532,4 +630,3 @@ class TestSetResponseHeadersLiftsDecision:
|
|||
headers = response._hidden_params["additional_headers"]
|
||||
assert "x-litellm-quality-router-model" not in headers
|
||||
assert "x-litellm-quality-router-tier" not in headers
|
||||
assert "x-litellm-quality-router-complexity" not in headers
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue