mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: add groups
This commit is contained in:
parent
259e10b2fc
commit
6af8ec93c9
5 changed files with 209 additions and 27 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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: "
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue