diff --git a/litellm/router_strategy/quality_router/config.py b/litellm/router_strategy/quality_router/config.py index 024ac0f5e9f..36bda238b91 100644 --- a/litellm/router_strategy/quality_router/config.py +++ b/litellm/router_strategy/quality_router/config.py @@ -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") diff --git a/litellm/router_strategy/quality_router/quality_router.py b/litellm/router_strategy/quality_router/quality_router.py index abb7e0c8b60..18b657bab61 100644 --- a/litellm/router_strategy/quality_router/quality_router.py +++ b/litellm/router_strategy/quality_router/quality_router.py @@ -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}" ) diff --git a/tests/test_litellm/router_strategy/test_quality_router.py b/tests/test_litellm/router_strategy/test_quality_router.py index 1426d2253a8..0f9bbb8ae61 100644 --- a/tests/test_litellm/router_strategy/test_quality_router.py +++ b/tests/test_litellm/router_strategy/test_quality_router.py @@ -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"