mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
b92855d7b3
commit
2a617d48e6
3 changed files with 228 additions and 15 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue