fix: add groups

This commit is contained in:
Ishaan Jaffer 2026-03-02 21:50:15 -08:00
parent 259e10b2fc
commit 6af8ec93c9
5 changed files with 209 additions and 27 deletions

View file

@ -3088,6 +3088,9 @@ class SpendLogsMetadata(TypedDict):
cost_breakdown: Optional[
CostBreakdown
] # Detailed cost breakdown (input_cost, output_cost, margin, discount, etc.)
routing_strategy: Optional[str] # Routing strategy used by the router
model_group_size: Optional[int] # Number of deployments in the model group
model_group_candidates: Optional[List[dict]] # Candidate deployments considered
class SpendLogsPayload(TypedDict):

View file

@ -10,6 +10,8 @@ from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, status
import prisma
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
@ -65,6 +67,7 @@ VALID_ROUTING_STRATEGIES = frozenset(
"cost-based-routing",
"usage-based-routing-v2",
"weighted",
"complexity-router",
}
)
@ -88,26 +91,85 @@ async def _sync_routing_group_to_router(config: RoutingGroupConfig) -> None:
group_name = config.routing_group_name
if config.routing_strategy:
llm_router.routing_group_strategies[group_name] = config.routing_strategy
for dep in config.deployments:
try:
litellm_params: dict = {"model": dep.model_name}
# Resolve the actual litellm model string by looking up the
# existing deployment. The routing group stores the user-facing
# model_name (e.g. "gpt-4..1") but the router needs the
# provider-prefixed string (e.g. "azure/gpt-4..1") plus creds.
resolved_model = dep.model_name
litellm_params: dict = {"model": resolved_model}
existing = llm_router.get_model_list(model_name=dep.model_name)
if existing:
src = existing[0].get("litellm_params", {})
resolved_model = src.get("model", dep.model_name)
litellm_params["model"] = resolved_model
for key in ("api_key", "api_base", "api_version"):
if key in src:
litellm_params[key] = src[key]
if dep.weight is not None:
litellm_params["weight"] = dep.weight
namespaced_id = f"rg:{group_name}:{dep.model_id}"
model_info: dict = {"id": namespaced_id}
if dep.priority is not None:
model_info["priority"] = dep.priority
deployment_dict = {
"model_name": group_name,
"litellm_params": litellm_params,
"model_info": {
"id": dep.model_id,
},
"model_info": model_info,
}
llm_router.add_deployment(
print(f"[RG-SYNC] {group_name}: adding {dep.model_id} resolved={resolved_model} existing={len(existing or [])}")
result = llm_router.add_deployment(
deployment=litellm.types.router.Deployment(**deployment_dict)
)
print(f"[RG-SYNC] {group_name}: {dep.model_id} -> {'ADDED' if result else 'SKIPPED (dup)'}")
except Exception as e:
verbose_proxy_logger.debug(
f"Could not add deployment {dep.model_id} to router: {e}"
)
print(f"[RG-SYNC] {group_name}: {dep.model_id} -> FAILED: {e}")
# For complexity-router groups, initialize the ComplexityRouter so it can
# classify requests and pick the right deployment by tier.
if config.routing_strategy == "complexity-router":
_init_complexity_router_for_group(llm_router, config)
def _init_complexity_router_for_group(llm_router, config: RoutingGroupConfig) -> None:
"""Initialize a ComplexityRouter for a routing group."""
from litellm.router_strategy.complexity_router.complexity_router import (
ComplexityRouter,
)
group_name = config.routing_group_name
settings = config.settings or {}
tiers = settings.get("tiers", {})
# Filter out empty tier values
tiers = {k: v for k, v in tiers.items() if v}
if not tiers:
verbose_proxy_logger.warning(
f"complexity-router group '{group_name}' has no tier mappings in settings; "
"using default tier config"
)
complexity_config: dict = {"tiers": tiers} if tiers else {}
default_model = tiers.get("MEDIUM") or tiers.get("SIMPLE") or None
complexity_router = ComplexityRouter(
model_name=group_name,
default_model=default_model,
litellm_router_instance=llm_router,
complexity_router_config=complexity_config or None,
)
llm_router.complexity_routers[group_name] = complexity_router
verbose_proxy_logger.info(
f"Initialized complexity-router for group '{group_name}' with tiers={tiers}"
)
@ -148,6 +210,7 @@ async def create_routing_group(
caller = user_api_key_dict.user_id or "unknown"
routing_group_id = str(uuid.uuid4())
deployments_json = prisma.Json([d.model_dump() for d in data.deployments]) # type: ignore[attr-defined]
try:
created = await prisma_client.db.litellm_routinggrouptable.create(
data={
@ -155,11 +218,11 @@ async def create_routing_group(
"routing_group_name": data.routing_group_name,
"description": data.description,
"routing_strategy": data.routing_strategy,
"deployments": [d.model_dump() for d in data.deployments],
"fallback_config": data.fallback_config or {},
"retry_config": data.retry_config or {},
"cooldown_config": data.cooldown_config or {},
"settings": data.settings or {},
"deployments": deployments_json,
"fallback_config": prisma.Json(data.fallback_config or {}), # type: ignore[attr-defined]
"retry_config": prisma.Json(data.retry_config or {}), # type: ignore[attr-defined]
"cooldown_config": prisma.Json(data.cooldown_config or {}), # type: ignore[attr-defined]
"settings": prisma.Json(data.settings or {}), # type: ignore[attr-defined]
"assigned_team_ids": data.assigned_team_ids or [],
"assigned_key_ids": data.assigned_key_ids or [],
"is_active": data.is_active,

View file

@ -4456,6 +4456,7 @@ class ProxyConfig:
deployments=[
RoutingGroupDeployment(**d) for d in (g.deployments or [])
],
settings=g.settings if hasattr(g, "settings") else None,
)
await _sync_routing_group_to_router(config)
verbose_proxy_logger.debug(

View file

@ -107,6 +107,9 @@ def _get_spend_logs_metadata(
attempted_retries=None,
max_retries=None,
cost_breakdown=None,
routing_strategy=None,
model_group_size=None,
model_group_candidates=None,
)
verbose_proxy_logger.debug(
"getting payload for SpendLogs, available keys in metadata: "

View file

@ -454,6 +454,7 @@ class Router:
] = {} # {"TEAM_ID": PatternMatchRouter}
self.auto_routers: Dict[str, "AutoRouter"] = {}
self.complexity_routers: Dict[str, "ComplexityRouter"] = {}
self.routing_group_strategies: Dict[str, str] = {}
# Initialize model_group_alias early since it's used in set_model_list
self.model_group_alias: Dict[str, Union[str, RouterModelGroupAliasItem]] = (
@ -2094,8 +2095,9 @@ class Router:
model_group_alias: Optional[str] = None
if self._get_model_from_alias(model=model):
model_group_alias = model
effective_strategy = self.routing_group_strategies.get(model, self.routing_strategy)
kwargs.setdefault(metadata_variable_name, {}).update(
{"model_group": model, "model_group_alias": model_group_alias}
{"model_group": model, "model_group_alias": model_group_alias, "routing_strategy": effective_strategy}
)
def _set_deployment_num_retries_on_exception(
@ -5365,12 +5367,30 @@ class Router:
model_group: Optional[str] = kwargs.get("model")
num_retries = kwargs.pop("num_retries")
## ADD MODEL GROUP SIZE TO METADATA - used for model_group_rate_limit_error tracking
_group_strategy = self.routing_group_strategies.get(model_group or "")
if _group_strategy == "priority-failover":
_group_size = len(self.get_model_list(model_name=model_group) or [])
if _group_size > 1:
num_retries = max(num_retries, _group_size - 1)
## ADD MODEL GROUP SIZE + CANDIDATES TO METADATA
_metadata: dict = kwargs.get("litellm_metadata", kwargs.get("metadata")) or {}
if "model_group" in _metadata and isinstance(_metadata["model_group"], str):
model_list = self.get_model_list(model_name=_metadata["model_group"])
if model_list is not None:
_metadata.update({"model_group_size": len(model_list)})
candidates = []
for _dep in model_list:
dep_info = _dep.get("model_info", {})
candidates.append({
"model_id": dep_info.get("id"),
"model_name": _dep.get("model_name"),
"litellm_model": _dep.get("litellm_params", {}).get("model"),
"priority": dep_info.get("priority"),
})
_metadata.update({
"model_group_size": len(model_list),
"model_group_candidates": candidates,
})
verbose_router_logger.debug(
f"async function w/ retries: original_function - {original_function}, num_retries - {num_retries}"
@ -5399,14 +5419,25 @@ class Router:
deployment_num_retries, int
):
num_retries = deployment_num_retries
# For priority-failover: track the deployment that just failed so
# the next retry skips it and picks the next-priority deployment.
if _group_strategy == "priority-failover":
_failed_id = _metadata.get("model_info", {}).get("id")
if _failed_id:
_tried = _metadata.setdefault("_pf_tried_ids", [])
if _failed_id not in _tried:
_tried.append(_failed_id)
"""
Retry Logic
"""
_retry_model = kwargs.get("model") or ""
(
_healthy_deployments,
_all_deployments,
) = await self._async_get_healthy_deployments(
model=kwargs.get("model") or "",
model=_retry_model,
parent_otel_span=parent_otel_span,
)
@ -5431,9 +5462,12 @@ class Router:
num_retries = _retry_policy_retries
_retry_policy_applies = True
# raises an exception if this error should not be retries
# raises an exception if this error should not be retried
# Skip this check if retry policy applies (retry policy takes precedence)
if not _retry_policy_applies:
# For priority-failover groups, always retry — there are other
# deployments to try and the error may be deployment-specific
# (e.g. 404 on one provider but model exists on another).
if not _retry_policy_applies and _group_strategy != "priority-failover":
self.should_retry_this_error(
error=e,
healthy_deployments=_healthy_deployments,
@ -5467,10 +5501,8 @@ class Router:
for current_attempt in range(num_retries):
try:
# Update retry tracking metadata before each retry attempt
_metadata["attempted_retries"] = current_attempt + 1
_metadata["max_retries"] = num_retries
# if the function call is successful, no exception will be raised and we'll break out of the loop
response = await self.make_call(original_function, *args, **kwargs)
if coroutine_checker.is_async_callable(
response
@ -5487,6 +5519,15 @@ class Router:
except Exception as e:
## LOGGING
kwargs = self.log_retry(kwargs=kwargs, e=e)
# For priority-failover: track the deployment that just failed
if _group_strategy == "priority-failover":
_failed_id = _metadata.get("model_info", {}).get("id")
if _failed_id:
_tried = _metadata.setdefault("_pf_tried_ids", [])
if _failed_id not in _tried:
_tried.append(_failed_id)
remaining_retries = num_retries - current_attempt - 1
_model: Optional[str] = kwargs.get("model") # type: ignore
if _model is not None:
@ -8825,9 +8866,18 @@ class Router:
input=input,
specific_deployment=specific_deployment,
)
_complexity_target_model: Optional[str] = None
if pre_routing_hook_response is not None:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
# For complexity-router routing groups, the hook returns a tier
# model (e.g. "gpt-4o") but deployments are registered under
# the group name. Keep the group name for the deployment
# lookup and filter by the classified model afterwards.
if self.routing_group_strategies.get(model) == "complexity-router":
_complexity_target_model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
else:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
#########################################################
healthy_deployments = await self.async_get_healthy_deployments(
@ -8842,7 +8892,46 @@ class Router:
return healthy_deployments
start_time = time.time()
if (
group_strategy = self.routing_group_strategies.get(model)
if group_strategy == "priority-failover" and isinstance(healthy_deployments, list):
_tried_ids = (request_kwargs or {}).get("metadata", {}).get("_pf_tried_ids") or []
if not _tried_ids:
_tried_ids = (request_kwargs or {}).get("litellm_metadata", {}).get("_pf_tried_ids") or []
untried = [
d for d in healthy_deployments
if d.get("model_info", {}).get("id") not in _tried_ids
]
candidates = untried if untried else healthy_deployments
sorted_deps = sorted(
candidates,
key=lambda d: d.get("model_info", {}).get("priority") or 999,
)
deployment = sorted_deps[0]
verbose_router_logger.info(
"priority-failover: picked %s (priority=%s, tried=%s, remaining=%d)",
deployment.get("model_info", {}).get("id"),
deployment.get("model_info", {}).get("priority"),
_tried_ids,
len(candidates),
)
elif group_strategy == "complexity-router" and _complexity_target_model and isinstance(healthy_deployments, list):
matching = [
d for d in healthy_deployments
if d.get("litellm_params", {}).get("model") == _complexity_target_model
]
if matching:
deployment = matching[0]
else:
deployment = healthy_deployments[0]
verbose_router_logger.info(
"complexity-router: target=%s, picked=%s (matched=%d/%d)",
_complexity_target_model,
deployment.get("model_info", {}).get("id"),
len(matching),
len(healthy_deployments),
)
elif (
self.routing_strategy == "usage-based-routing-v2"
and self.lowesttpm_logger_v2 is not None
):
@ -9186,13 +9275,36 @@ class Router:
cooldown_list=_cooldown_list,
)
if self.routing_strategy == "least-busy" and self.leastbusy_logger is not None:
group_strategy = self.routing_group_strategies.get(model)
if group_strategy == "priority-failover" and isinstance(healthy_deployments, list):
_tried_ids = (request_kwargs or {}).get("metadata", {}).get("_pf_tried_ids") or []
if not _tried_ids:
_tried_ids = (request_kwargs or {}).get("litellm_metadata", {}).get("_pf_tried_ids") or []
untried = [
d for d in healthy_deployments
if d.get("model_info", {}).get("id") not in _tried_ids
]
candidates = untried if untried else healthy_deployments
sorted_deps = sorted(
candidates,
key=lambda d: d.get("model_info", {}).get("priority") or 999,
)
deployment = sorted_deps[0]
elif group_strategy == "complexity-router" and model in self.complexity_routers and isinstance(healthy_deployments, list):
_cr = self.complexity_routers[model]
_target = _cr.get_model_for_tier(_cr.classify(
(messages or [{}])[-1].get("content", "") if messages else "",
)[0])
matching = [
d for d in healthy_deployments
if d.get("litellm_params", {}).get("model") == _target
]
deployment = matching[0] if matching else healthy_deployments[0]
elif self.routing_strategy == "least-busy" and self.leastbusy_logger is not None:
deployment = self.leastbusy_logger.get_available_deployments(
model_group=model, healthy_deployments=healthy_deployments # type: ignore
)
elif self.routing_strategy == "simple-shuffle":
# if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm
############## Check 'weight' param set for weighted pick #################
return simple_shuffle(
llm_router_instance=self,
healthy_deployments=healthy_deployments,