feat(quality_router): add capability-based filtering

Each deployment can declare a `capabilities: List[str]` field in
`model_info.litellm_routing_preferences` (e.g. ["vision",
"function_calling"]). Requests can pass `litellm_capabilities` in
`request_kwargs` to require specific capabilities — the router will only
route to deployments whose declared capabilities are a superset.

Resolution still walks tier (exact → round up), but at each tier filters
by capability before picking. Falls back to default_model only when it
also satisfies the required capabilities; otherwise raises rather than
silently routing to a model that lacks a required capability.

Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-04-17 17:52:42 -07:00
parent b92855d7b3
commit 2a617d48e6
3 changed files with 228 additions and 15 deletions

View file

@ -48,4 +48,14 @@ class RoutingPreferences(BaseModel):
description="The quality tier this deployment satisfies.",
)
capabilities: 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."
),
)
model_config = ConfigDict(extra="allow")

View file

@ -8,7 +8,7 @@ candidate model declares its own `quality_tier` in
`model_info.litellm_routing_preferences`.
"""
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, FrozenSet, List, Optional, Set, Union
from litellm._logging import verbose_router_logger
from litellm.integrations.custom_logger import CustomLogger
@ -64,7 +64,10 @@ class QualityRouter(CustomLogger):
litellm_router_instance=litellm_router_instance,
)
# Pre-built tier → models index for O(1) resolution.
# 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]] = {}
self._tier_to_models: Dict[int, List[str]] = self._build_tier_index()
verbose_router_logger.debug(
@ -129,8 +132,10 @@ class QualityRouter(CustomLogger):
# Accept dict or Pydantic-shaped prefs.
if isinstance(prefs, dict):
tier = prefs.get("quality_tier")
capabilities = prefs.get("capabilities") or []
else:
tier = getattr(prefs, "quality_tier", None)
capabilities = getattr(prefs, "capabilities", None) or []
if tier is None:
raise ValueError(
@ -140,6 +145,7 @@ class QualityRouter(CustomLogger):
tier_int = int(tier)
tier_to_models.setdefault(tier_int, []).append(name)
self._model_capabilities[name] = frozenset(capabilities)
seen[name] = True
missing = [name for name, found in seen.items() if not found]
@ -151,22 +157,63 @@ class QualityRouter(CustomLogger):
return tier_to_models
def _resolve_model_for_quality_tier(self, tier: int) -> str:
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 _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
def _resolve_model_for_quality_tier(
self,
tier: int,
required_capabilities: Optional[Set[str]] = None,
) -> str:
"""
Resolve a quality tier to a concrete model name.
Strategy:
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`.
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.
"""
if tier in self._tier_to_models and self._tier_to_models[tier]:
return self._tier_to_models[tier][0]
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
higher_tiers = sorted(t for t in self._tier_to_models if t > tier)
for t in higher_tiers:
if self._tier_to_models[t]:
return self._tier_to_models[t][0]
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.config.default_model:
return self.config.default_model
@ -214,6 +261,16 @@ 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"
@ -222,6 +279,14 @@ 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,
@ -243,11 +308,14 @@ 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))
routed_model = self._resolve_model_for_quality_tier(
int(quality_tier), required_capabilities=required_capabilities
)
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}"
)

View file

@ -26,7 +26,11 @@ def _make_model_list(spec: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""
Build a router model_list from a compact spec.
spec entry shape: {"model_name": str, "quality_tier": Optional[int]}
spec entry shape: {
"model_name": str,
"quality_tier": Optional[int],
"capabilities": Optional[List[str]], # default: omitted
}
If quality_tier is None, the deployment is created without
`litellm_routing_preferences`.
"""
@ -34,9 +38,10 @@ def _make_model_list(spec: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
for entry in spec:
model_info: Dict[str, Any] = {"id": f"id-{entry['model_name']}"}
if entry.get("quality_tier") is not None:
model_info["litellm_routing_preferences"] = {
"quality_tier": entry["quality_tier"]
}
prefs: Dict[str, Any] = {"quality_tier": entry["quality_tier"]}
if "capabilities" in entry:
prefs["capabilities"] = entry["capabilities"]
model_info["litellm_routing_preferences"] = prefs
out.append(
{
"model_name": entry["model_name"],
@ -230,3 +235,133 @@ class TestPreRoutingHook:
)
assert resp is not None
assert resp.model == "haiku" # the configured default_model
# ─── Capabilities ───────────────────────────────────────────────────────────
@pytest.fixture
def capability_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)
"""
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": "opus-vision",
"quality_tier": 3,
"capabilities": ["vision", "function_calling", "json_mode"],
},
]
router = MagicMock()
router.model_list = _make_model_list(spec)
return QualityRouter(
model_name="qr",
litellm_router_instance=router,
default_model="haiku-text",
quality_router_config={
"available_models": [
"haiku-text",
"haiku-vision",
"sonnet-text",
"sonnet-vision",
"opus-vision",
],
},
)
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"}
)
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_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_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": []}]
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"]},
)
with pytest.raises(ValueError, match="vision"):
qr._resolve_model_for_quality_tier(1, required_capabilities={"vision"})
@pytest.mark.asyncio
async def test_hook_reads_litellm_capabilities_from_request_kwargs(
self, capability_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(
model="qr",
request_kwargs={"litellm_capabilities": ["vision"]},
messages=messages,
)
assert resp is not None
assert resp.model == "haiku-vision"
@pytest.mark.asyncio
async def test_hook_with_no_capabilities_kwarg_behaves_as_before(
self, capability_router
):
messages = [{"role": "user", "content": "hi"}]
resp = await capability_router.async_pre_routing_hook(
model="qr",
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"