mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(router): add auto_router/quality_router for quality-tier routing
Adds a new auto-router type that routes a request to a model at a target quality tier. The quality tier is 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. Resolution strategy: exact tier match, else round up to the next higher tier, else fall back to default_model. Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
850fe595ac
commit
b92855d7b3
6 changed files with 669 additions and 19 deletions
|
|
@ -200,12 +200,16 @@ if TYPE_CHECKING:
|
|||
from litellm.router_strategy.complexity_router.complexity_router import (
|
||||
ComplexityRouter,
|
||||
)
|
||||
from litellm.router_strategy.quality_router.quality_router import (
|
||||
QualityRouter,
|
||||
)
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
else:
|
||||
Span = Any
|
||||
AutoRouter = Any
|
||||
ComplexityRouter = Any
|
||||
QualityRouter = Any
|
||||
PreRoutingHookResponse = Any
|
||||
|
||||
|
||||
|
|
@ -464,6 +468,7 @@ class Router:
|
|||
) # {"TEAM_ID": PatternMatchRouter}
|
||||
self.auto_routers: Dict[str, "AutoRouter"] = {}
|
||||
self.complexity_routers: Dict[str, "ComplexityRouter"] = {}
|
||||
self.quality_routers: Dict[str, "QualityRouter"] = {}
|
||||
|
||||
# Initialize model_group_alias early since it's used in set_model_list
|
||||
self.model_group_alias: Dict[str, Union[str, RouterModelGroupAliasItem]] = (
|
||||
|
|
@ -3864,7 +3869,7 @@ class Router:
|
|||
self._add_deployment_model_to_endpoint_for_llm_passthrough_route(
|
||||
kwargs=kwargs, model=model, model_name=model_name
|
||||
)
|
||||
|
||||
|
||||
# Get custom_llm_provider from deployment params
|
||||
try:
|
||||
custom_llm_provider = data.get("custom_llm_provider")
|
||||
|
|
@ -3872,10 +3877,12 @@ class Router:
|
|||
model=data["model"],
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
|
||||
custom_llm_provider = (
|
||||
custom_llm_provider or inferred_custom_llm_provider
|
||||
)
|
||||
except Exception:
|
||||
custom_llm_provider = None
|
||||
|
||||
|
||||
# Build response kwargs
|
||||
response_kwargs = {
|
||||
**data,
|
||||
|
|
@ -3885,7 +3892,7 @@ class Router:
|
|||
# Only set custom_llm_provider if it's not None
|
||||
if custom_llm_provider is not None:
|
||||
response_kwargs["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
|
||||
response = original_generic_function(**response_kwargs)
|
||||
|
||||
rpm_semaphore = self._get_client(
|
||||
|
|
@ -3981,7 +3988,9 @@ class Router:
|
|||
model=data["model"],
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
|
||||
custom_llm_provider = (
|
||||
custom_llm_provider or inferred_custom_llm_provider
|
||||
)
|
||||
except Exception:
|
||||
custom_llm_provider = None
|
||||
|
||||
|
|
@ -4246,7 +4255,9 @@ class Router:
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
# Preserve explicitly stored provider, fallback to inferred
|
||||
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
|
||||
custom_llm_provider = (
|
||||
custom_llm_provider or inferred_custom_llm_provider
|
||||
)
|
||||
|
||||
## REPLACE MODEL IN FILE WITH SELECTED DEPLOYMENT ##
|
||||
purpose = cast(Optional[OpenAIFilesPurpose], kwargs.get("purpose"))
|
||||
|
|
@ -5355,9 +5366,9 @@ class Router:
|
|||
e,
|
||||
(litellm.ContextWindowExceededError, litellm.ContentPolicyViolationError),
|
||||
)
|
||||
_request_team_id: Optional[str] = (
|
||||
kwargs.get("metadata", {}) or {}
|
||||
).get("user_api_key_team_id")
|
||||
_request_team_id: Optional[str] = (kwargs.get("metadata", {}) or {}).get(
|
||||
"user_api_key_team_id"
|
||||
)
|
||||
all_deployments = self._get_all_deployments(
|
||||
model_name=original_model_group, team_id=_request_team_id
|
||||
)
|
||||
|
|
@ -6808,6 +6819,8 @@ class Router:
|
|||
"""
|
||||
if litellm_params.model.startswith("auto_router/complexity_router"):
|
||||
return False # This is handled by complexity_router
|
||||
if litellm_params.model.startswith("auto_router/quality_router"):
|
||||
return False # This is handled by quality_router
|
||||
if litellm_params.model.startswith("auto_router/"):
|
||||
return True
|
||||
return False
|
||||
|
|
@ -6914,6 +6927,58 @@ class Router:
|
|||
)
|
||||
self.complexity_routers[deployment.model_name] = complexity_router
|
||||
|
||||
def _is_quality_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
|
||||
"""
|
||||
Check if the deployment is a quality-router deployment.
|
||||
|
||||
Returns True if the litellm_params model starts with "auto_router/quality_router".
|
||||
"""
|
||||
if litellm_params.model.startswith("auto_router/quality_router"):
|
||||
return True
|
||||
return False
|
||||
|
||||
def init_quality_router_deployment(self, deployment: Deployment):
|
||||
"""
|
||||
Initialize the quality-router deployment.
|
||||
|
||||
Resolves the default model from either `quality_router_default_model` or
|
||||
`quality_router_config["default_model"]`, then instantiates the
|
||||
QualityRouter and stores it in `self.quality_routers`.
|
||||
"""
|
||||
# Import here to mirror the AutoRouter / ComplexityRouter init pattern
|
||||
# and avoid circular imports.
|
||||
from litellm.router_strategy.quality_router.quality_router import (
|
||||
QualityRouter,
|
||||
)
|
||||
|
||||
quality_router_config: Optional[dict] = (
|
||||
deployment.litellm_params.quality_router_config
|
||||
)
|
||||
|
||||
default_model: Optional[str] = (
|
||||
deployment.litellm_params.quality_router_default_model
|
||||
)
|
||||
if default_model is None and quality_router_config:
|
||||
default_model = quality_router_config.get("default_model")
|
||||
|
||||
if default_model is None:
|
||||
raise ValueError(
|
||||
"quality_router_default_model is required for quality-router deployments, "
|
||||
"or set default_model in quality_router_config. Please configure it in the litellm_params"
|
||||
)
|
||||
|
||||
quality_router: QualityRouter = QualityRouter(
|
||||
model_name=deployment.model_name,
|
||||
default_model=default_model,
|
||||
litellm_router_instance=self,
|
||||
quality_router_config=quality_router_config,
|
||||
)
|
||||
if deployment.model_name in self.quality_routers:
|
||||
raise ValueError(
|
||||
f"Quality-router deployment {deployment.model_name} already exists. Please use a different model name."
|
||||
)
|
||||
self.quality_routers[deployment.model_name] = quality_router
|
||||
|
||||
def deployment_is_active_for_environment(self, deployment: Deployment) -> bool:
|
||||
"""
|
||||
Function to check if a llm deployment is active for a given environment. Allows using the same config.yaml across multople environments
|
||||
|
|
@ -7134,6 +7199,12 @@ class Router:
|
|||
):
|
||||
self.init_complexity_router_deployment(deployment=deployment)
|
||||
|
||||
#########################################################
|
||||
# Check if this is a quality-router deployment
|
||||
#########################################################
|
||||
if self._is_quality_router_deployment(litellm_params=deployment.litellm_params):
|
||||
self.init_quality_router_deployment(deployment=deployment)
|
||||
|
||||
return deployment
|
||||
|
||||
def _initialize_deployment_for_pass_through(
|
||||
|
|
@ -9645,6 +9716,18 @@ class Router:
|
|||
specific_deployment=specific_deployment,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Check if any quality-router should be used
|
||||
#########################################################
|
||||
if model in self.quality_routers:
|
||||
return await self.quality_routers[model].async_pre_routing_hook(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def get_available_deployment(
|
||||
|
|
|
|||
21
litellm/router_strategy/quality_router/__init__.py
Normal file
21
litellm/router_strategy/quality_router/__init__.py
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
"""
|
||||
Quality-tier auto-router.
|
||||
|
||||
Re-uses the ComplexityRouter's classification to decide a request's complexity,
|
||||
then maps that complexity to an admin-configured quality tier and resolves the
|
||||
target model from each candidate's `model_info.litellm_routing_preferences`.
|
||||
"""
|
||||
|
||||
from .config import (
|
||||
DEFAULT_COMPLEXITY_TO_QUALITY,
|
||||
QualityRouterConfig,
|
||||
RoutingPreferences,
|
||||
)
|
||||
from .quality_router import QualityRouter
|
||||
|
||||
__all__ = [
|
||||
"QualityRouter",
|
||||
"QualityRouterConfig",
|
||||
"RoutingPreferences",
|
||||
"DEFAULT_COMPLEXITY_TO_QUALITY",
|
||||
]
|
||||
51
litellm/router_strategy/quality_router/config.py
Normal file
51
litellm/router_strategy/quality_router/config.py
Normal file
|
|
@ -0,0 +1,51 @@
|
|||
"""
|
||||
Configuration models for the QualityRouter.
|
||||
"""
|
||||
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
# Default mapping from ComplexityTier name (string) to quality tier (int).
|
||||
# Higher tier = higher capability requirement.
|
||||
DEFAULT_COMPLEXITY_TO_QUALITY: Dict[str, int] = {
|
||||
"SIMPLE": 1,
|
||||
"MEDIUM": 2,
|
||||
"COMPLEX": 3,
|
||||
"REASONING": 4,
|
||||
}
|
||||
|
||||
|
||||
class QualityRouterConfig(BaseModel):
|
||||
"""Configuration for the QualityRouter."""
|
||||
|
||||
available_models: List[str] = Field(
|
||||
default_factory=list,
|
||||
description=(
|
||||
"List of candidate model names this router may route to. Each model "
|
||||
"must declare its quality_tier in model_info.litellm_routing_preferences."
|
||||
),
|
||||
)
|
||||
|
||||
default_model: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Fallback model when no quality tier resolves.",
|
||||
)
|
||||
|
||||
complexity_to_quality: Dict[str, int] = Field(
|
||||
default_factory=lambda: DEFAULT_COMPLEXITY_TO_QUALITY.copy(),
|
||||
description="Mapping from ComplexityTier name to quality tier (int).",
|
||||
)
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
|
||||
class RoutingPreferences(BaseModel):
|
||||
"""Per-deployment routing preferences declared on model_info."""
|
||||
|
||||
quality_tier: int = Field(
|
||||
...,
|
||||
description="The quality tier this deployment satisfies.",
|
||||
)
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
257
litellm/router_strategy/quality_router/quality_router.py
Normal file
257
litellm/router_strategy/quality_router/quality_router.py
Normal file
|
|
@ -0,0 +1,257 @@
|
|||
"""
|
||||
Quality-tier Auto Router.
|
||||
|
||||
Routes a request to a model at a target quality tier. The quality tier is
|
||||
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`.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.router_strategy.complexity_router.complexity_router import (
|
||||
ComplexityRouter,
|
||||
)
|
||||
|
||||
from .config import QualityRouterConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
else:
|
||||
Router = Any
|
||||
PreRoutingHookResponse = Any
|
||||
|
||||
|
||||
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`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
litellm_router_instance: "Router",
|
||||
default_model: Optional[str] = None,
|
||||
quality_router_config: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.litellm_router_instance = litellm_router_instance
|
||||
|
||||
if quality_router_config:
|
||||
self.config = QualityRouterConfig(**quality_router_config)
|
||||
else:
|
||||
self.config = QualityRouterConfig()
|
||||
|
||||
# Explicit default_model arg overrides anything in the config dict.
|
||||
if default_model:
|
||||
self.config.default_model = default_model
|
||||
|
||||
# Internal scorer — re-use the existing rule-based classifier.
|
||||
self._scorer = ComplexityRouter(
|
||||
model_name=f"{model_name}::scorer",
|
||||
litellm_router_instance=litellm_router_instance,
|
||||
)
|
||||
|
||||
# Pre-built tier → models index for O(1) resolution.
|
||||
self._tier_to_models: Dict[int, List[str]] = self._build_tier_index()
|
||||
|
||||
verbose_router_logger.debug(
|
||||
f"QualityRouter initialized for {model_name} with "
|
||||
f"available_models={self.config.available_models}, "
|
||||
f"default_model={self.config.default_model}, "
|
||||
f"tier_index={self._tier_to_models}"
|
||||
)
|
||||
|
||||
def _get_routing_preferences(self, deployment: Any) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Extract litellm_routing_preferences from a deployment, handling both
|
||||
dict-shaped and Pydantic-object-shaped deployments.
|
||||
"""
|
||||
# Dict-shaped deployment.
|
||||
if isinstance(deployment, dict):
|
||||
model_info = deployment.get("model_info") or {}
|
||||
if isinstance(model_info, dict):
|
||||
return model_info.get("litellm_routing_preferences")
|
||||
# Pydantic ModelInfo nested in a dict.
|
||||
return getattr(model_info, "litellm_routing_preferences", None)
|
||||
|
||||
# Pydantic-object deployment.
|
||||
model_info = getattr(deployment, "model_info", None)
|
||||
if model_info is None:
|
||||
return None
|
||||
if isinstance(model_info, dict):
|
||||
return model_info.get("litellm_routing_preferences")
|
||||
return getattr(model_info, "litellm_routing_preferences", 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):
|
||||
return deployment.get("model_name")
|
||||
return getattr(deployment, "model_name", None)
|
||||
|
||||
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`.
|
||||
"""
|
||||
model_list = getattr(self.litellm_router_instance, "model_list", None) or []
|
||||
available = set(self.config.available_models)
|
||||
|
||||
# Track which available models we've matched so we can error on missing.
|
||||
seen: Dict[str, bool] = {name: False for name in available}
|
||||
tier_to_models: Dict[int, List[str]] = {}
|
||||
|
||||
for deployment in model_list:
|
||||
name = self._get_deployment_model_name(deployment)
|
||||
if name is None or name not in available:
|
||||
continue
|
||||
|
||||
prefs = self._get_routing_preferences(deployment)
|
||||
if prefs is None:
|
||||
raise ValueError(
|
||||
f"QualityRouter: model '{name}' is listed in available_models "
|
||||
f"but has no model_info.litellm_routing_preferences"
|
||||
)
|
||||
|
||||
# Accept dict or Pydantic-shaped prefs.
|
||||
if isinstance(prefs, dict):
|
||||
tier = prefs.get("quality_tier")
|
||||
else:
|
||||
tier = getattr(prefs, "quality_tier", None)
|
||||
|
||||
if tier is None:
|
||||
raise ValueError(
|
||||
f"QualityRouter: model '{name}' has litellm_routing_preferences "
|
||||
f"but no quality_tier field"
|
||||
)
|
||||
|
||||
tier_int = int(tier)
|
||||
tier_to_models.setdefault(tier_int, []).append(name)
|
||||
seen[name] = True
|
||||
|
||||
missing = [name for name, found in seen.items() if not found]
|
||||
if missing:
|
||||
raise ValueError(
|
||||
f"QualityRouter: the following available_models are not present in "
|
||||
f"the router's model_list (or are missing routing preferences): {missing}"
|
||||
)
|
||||
|
||||
return tier_to_models
|
||||
|
||||
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 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`.
|
||||
"""
|
||||
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:
|
||||
if self._tier_to_models[t]:
|
||||
return self._tier_to_models[t][0]
|
||||
|
||||
if self.config.default_model:
|
||||
return self.config.default_model
|
||||
|
||||
raise ValueError(
|
||||
f"QualityRouter: no model available for quality tier {tier} and "
|
||||
f"no default_model configured"
|
||||
)
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
request_kwargs: Dict,
|
||||
messages: Optional[List[Dict[str, Any]]] = None,
|
||||
input: Optional[Union[str, List]] = None,
|
||||
specific_deployment: Optional[bool] = False,
|
||||
) -> Optional["PreRoutingHookResponse"]:
|
||||
"""Classify the request, map to a quality tier, resolve the model."""
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
if messages is None or len(messages) == 0:
|
||||
verbose_router_logger.debug(
|
||||
"QualityRouter: No messages provided, skipping routing"
|
||||
)
|
||||
return None
|
||||
|
||||
# Extract last user message and last system prompt — same rules as
|
||||
# ComplexityRouter.async_pre_routing_hook.
|
||||
user_message: Optional[str] = None
|
||||
system_prompt: Optional[str] = None
|
||||
|
||||
for msg in reversed(messages):
|
||||
role = msg.get("role", "")
|
||||
content = msg.get("content") or ""
|
||||
if isinstance(content, list):
|
||||
text_parts = [
|
||||
part.get("text", "")
|
||||
for part in content
|
||||
if isinstance(part, dict) and part.get("type") == "text"
|
||||
]
|
||||
content = " ".join(text_parts).strip()
|
||||
if isinstance(content, str) and content:
|
||||
if role == "user" and user_message is None:
|
||||
user_message = content
|
||||
elif role == "system" and system_prompt is None:
|
||||
system_prompt = content
|
||||
|
||||
if user_message is None:
|
||||
verbose_router_logger.debug(
|
||||
"QualityRouter: No user message found, routing to default model"
|
||||
)
|
||||
if not self.config.default_model:
|
||||
raise ValueError(
|
||||
"QualityRouter: no user message and no default_model configured"
|
||||
)
|
||||
return PreRoutingHookResponse(
|
||||
model=self.config.default_model,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
complexity_tier, score, signals = self._scorer.classify(
|
||||
user_message, system_prompt
|
||||
)
|
||||
complexity_name = (
|
||||
complexity_tier.value
|
||||
if hasattr(complexity_tier, "value")
|
||||
else str(complexity_tier)
|
||||
)
|
||||
|
||||
quality_tier = self.config.complexity_to_quality.get(complexity_name)
|
||||
if quality_tier is None:
|
||||
raise ValueError(
|
||||
f"QualityRouter: complexity tier '{complexity_name}' not present "
|
||||
f"in complexity_to_quality mapping {self.config.complexity_to_quality}"
|
||||
)
|
||||
|
||||
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"routed_model={routed_model}"
|
||||
)
|
||||
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
messages=messages,
|
||||
)
|
||||
|
|
@ -95,16 +95,18 @@ class ModelInfo(BaseModel):
|
|||
id: Optional[
|
||||
str
|
||||
] # Allow id to be optional on input, but it will always be present as a str in the model instance
|
||||
db_model: bool = False # used for proxy - to separate models which are stored in the db vs. config.
|
||||
db_model: bool = (
|
||||
False # used for proxy - to separate models which are stored in the db vs. config.
|
||||
)
|
||||
updated_at: Optional[datetime.datetime] = None
|
||||
updated_by: Optional[str] = None
|
||||
|
||||
created_at: Optional[datetime.datetime] = None
|
||||
created_by: Optional[str] = None
|
||||
|
||||
base_model: Optional[
|
||||
str
|
||||
] = None # specify if the base model is azure/gpt-3.5-turbo etc for accurate cost tracking
|
||||
base_model: Optional[str] = (
|
||||
None # specify if the base model is azure/gpt-3.5-turbo etc for accurate cost tracking
|
||||
)
|
||||
tier: Optional[Literal["free", "paid"]] = None
|
||||
|
||||
"""
|
||||
|
|
@ -173,12 +175,12 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
|||
custom_llm_provider: Optional[str] = None
|
||||
tpm: Optional[int] = None
|
||||
rpm: Optional[int] = None
|
||||
timeout: Optional[
|
||||
Union[float, str, httpx.Timeout]
|
||||
] = None # if str, pass in as os.environ/
|
||||
stream_timeout: Optional[
|
||||
Union[float, str]
|
||||
] = None # timeout when making stream=True calls, if str, pass in as os.environ/
|
||||
timeout: Optional[Union[float, str, httpx.Timeout]] = (
|
||||
None # if str, pass in as os.environ/
|
||||
)
|
||||
stream_timeout: Optional[Union[float, str]] = (
|
||||
None # timeout when making stream=True calls, if str, pass in as os.environ/
|
||||
)
|
||||
max_retries: Optional[int] = None
|
||||
organization: Optional[str] = None # for openai orgs
|
||||
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None
|
||||
|
|
@ -219,6 +221,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
|||
complexity_router_config: Optional[Dict] = None
|
||||
complexity_router_default_model: Optional[str] = None
|
||||
|
||||
# quality-router params
|
||||
quality_router_config: Optional[Dict] = None
|
||||
quality_router_default_model: Optional[str] = None
|
||||
|
||||
# Batch/File API Params
|
||||
s3_bucket_name: Optional[str] = None
|
||||
s3_encryption_key_id: Optional[str] = None
|
||||
|
|
|
|||
232
tests/test_litellm/router_strategy/test_quality_router.py
Normal file
232
tests/test_litellm/router_strategy/test_quality_router.py
Normal file
|
|
@ -0,0 +1,232 @@
|
|||
"""
|
||||
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).
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from typing import Any, Dict, List
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.router_strategy.quality_router.config import (
|
||||
DEFAULT_COMPLEXITY_TO_QUALITY,
|
||||
)
|
||||
from litellm.router_strategy.quality_router.quality_router import QualityRouter
|
||||
|
||||
|
||||
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]}
|
||||
If quality_tier is None, the deployment is created without
|
||||
`litellm_routing_preferences`.
|
||||
"""
|
||||
out: 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"]
|
||||
}
|
||||
out.append(
|
||||
{
|
||||
"model_name": entry["model_name"],
|
||||
"litellm_params": {"model": f"openai/{entry['model_name']}"},
|
||||
"model_info": model_info,
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def four_tier_model_list() -> List[Dict[str, Any]]:
|
||||
"""A standard haiku(1)/sonnet(2)/opus(3)/opus-next(4) model list."""
|
||||
return _make_model_list(
|
||||
[
|
||||
{"model_name": "haiku", "quality_tier": 1},
|
||||
{"model_name": "sonnet", "quality_tier": 2},
|
||||
{"model_name": "opus", "quality_tier": 3},
|
||||
{"model_name": "opus-next", "quality_tier": 4},
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_router(four_tier_model_list):
|
||||
"""A MagicMock router preloaded with the four-tier model list."""
|
||||
router = MagicMock()
|
||||
router.model_list = four_tier_model_list
|
||||
return router
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def quality_router(mock_router) -> QualityRouter:
|
||||
"""Default QualityRouter wired to all four tiers."""
|
||||
config = {
|
||||
"available_models": ["haiku", "sonnet", "opus", "opus-next"],
|
||||
"complexity_to_quality": DEFAULT_COMPLEXITY_TO_QUALITY,
|
||||
}
|
||||
return QualityRouter(
|
||||
model_name="quality-router-test",
|
||||
litellm_router_instance=mock_router,
|
||||
default_model="haiku",
|
||||
quality_router_config=config,
|
||||
)
|
||||
|
||||
|
||||
# ─── Tier index ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTierIndex:
|
||||
def test_builds_correct_tier_to_models_map(self, quality_router):
|
||||
assert quality_router._tier_to_models == {
|
||||
1: ["haiku"],
|
||||
2: ["sonnet"],
|
||||
3: ["opus"],
|
||||
4: ["opus-next"],
|
||||
}
|
||||
|
||||
def test_ignores_models_not_in_available_models(self, four_tier_model_list):
|
||||
# Add a model the config doesn't list — it should be ignored.
|
||||
extra = _make_model_list([{"model_name": "ghost", "quality_tier": 5}])
|
||||
router = MagicMock()
|
||||
router.model_list = four_tier_model_list + extra
|
||||
|
||||
qr = QualityRouter(
|
||||
model_name="qr",
|
||||
litellm_router_instance=router,
|
||||
default_model="haiku",
|
||||
quality_router_config={
|
||||
"available_models": ["haiku", "sonnet", "opus", "opus-next"]
|
||||
},
|
||||
)
|
||||
|
||||
for models in qr._tier_to_models.values():
|
||||
assert "ghost" not in models
|
||||
|
||||
def test_raises_when_routing_preferences_missing(self):
|
||||
# `sonnet` is in available_models but has no preferences.
|
||||
ml = _make_model_list(
|
||||
[
|
||||
{"model_name": "haiku", "quality_tier": 1},
|
||||
{"model_name": "sonnet", "quality_tier": None},
|
||||
]
|
||||
)
|
||||
router = MagicMock()
|
||||
router.model_list = ml
|
||||
|
||||
with pytest.raises(ValueError, match="sonnet"):
|
||||
QualityRouter(
|
||||
model_name="qr",
|
||||
litellm_router_instance=router,
|
||||
default_model="haiku",
|
||||
quality_router_config={"available_models": ["haiku", "sonnet"]},
|
||||
)
|
||||
|
||||
|
||||
# ─── Resolve model for quality tier ─────────────────────────────────────────
|
||||
|
||||
|
||||
class TestResolveModelForQualityTier:
|
||||
def test_exact_match(self, quality_router):
|
||||
assert quality_router._resolve_model_for_quality_tier(2) == "sonnet"
|
||||
assert quality_router._resolve_model_for_quality_tier(4) == "opus-next"
|
||||
|
||||
def test_rounds_up_when_tier_missing(self, mock_router):
|
||||
# Available tiers: 1, 3, 4. Asking for 2 should round up to 3.
|
||||
spec = [
|
||||
{"model_name": "haiku", "quality_tier": 1},
|
||||
{"model_name": "opus", "quality_tier": 3},
|
||||
{"model_name": "opus-next", "quality_tier": 4},
|
||||
]
|
||||
router = MagicMock()
|
||||
router.model_list = _make_model_list(spec)
|
||||
|
||||
qr = QualityRouter(
|
||||
model_name="qr",
|
||||
litellm_router_instance=router,
|
||||
default_model="haiku",
|
||||
quality_router_config={"available_models": ["haiku", "opus", "opus-next"]},
|
||||
)
|
||||
|
||||
assert qr._resolve_model_for_quality_tier(2) == "opus"
|
||||
|
||||
def test_falls_back_to_default_when_nothing_higher_exists(self):
|
||||
# Only tier 1 available. Asking for tier 4 should fall back to default.
|
||||
spec = [{"model_name": "haiku", "quality_tier": 1}]
|
||||
router = MagicMock()
|
||||
router.model_list = _make_model_list(spec)
|
||||
|
||||
qr = QualityRouter(
|
||||
model_name="qr",
|
||||
litellm_router_instance=router,
|
||||
default_model="emergency-default",
|
||||
quality_router_config={"available_models": ["haiku"]},
|
||||
)
|
||||
|
||||
assert qr._resolve_model_for_quality_tier(4) == "emergency-default"
|
||||
|
||||
|
||||
# ─── Pre-routing hook ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPreRoutingHook:
|
||||
@pytest.mark.asyncio
|
||||
async def test_simple_message_routes_to_tier_1(self, quality_router):
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
resp = await quality_router.async_pre_routing_hook(
|
||||
model="quality-router-test",
|
||||
request_kwargs={},
|
||||
messages=messages,
|
||||
)
|
||||
assert resp is not None
|
||||
assert resp.model == "haiku"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reasoning_message_routes_to_tier_4(self, quality_router):
|
||||
# Two reasoning markers triggers ComplexityTier.REASONING → quality 4.
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
"Think step by step and reason through this problem. "
|
||||
"Analyze this carefully and break down each component."
|
||||
),
|
||||
}
|
||||
]
|
||||
resp = await quality_router.async_pre_routing_hook(
|
||||
model="quality-router-test",
|
||||
request_kwargs={},
|
||||
messages=messages,
|
||||
)
|
||||
assert resp is not None
|
||||
assert resp.model == "opus-next"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_messages_returns_none(self, quality_router):
|
||||
resp = await quality_router.async_pre_routing_hook(
|
||||
model="quality-router-test",
|
||||
request_kwargs={},
|
||||
messages=[],
|
||||
)
|
||||
assert resp is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_system_message_routes_to_default(self, quality_router):
|
||||
messages = [{"role": "system", "content": "You are a helpful assistant."}]
|
||||
resp = await quality_router.async_pre_routing_hook(
|
||||
model="quality-router-test",
|
||||
request_kwargs={},
|
||||
messages=messages,
|
||||
)
|
||||
assert resp is not None
|
||||
assert resp.model == "haiku" # the configured default_model
|
||||
Loading…
Add table
Reference in a new issue