From 6af8ec93c9893545347b36721091698a814bf1fe Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 2 Mar 2026 21:50:15 -0800 Subject: [PATCH] fix: add groups --- litellm/proxy/_types.py | 3 + .../routing_group_endpoints.py | 89 +++++++++-- litellm/proxy/proxy_server.py | 1 + .../spend_tracking/spend_tracking_utils.py | 3 + litellm/router.py | 140 ++++++++++++++++-- 5 files changed, 209 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index e408abb3c1b..405ee297ebe 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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): diff --git a/litellm/proxy/management_endpoints/routing_group_endpoints.py b/litellm/proxy/management_endpoints/routing_group_endpoints.py index ec468fd92a9..2d18024564a 100644 --- a/litellm/proxy/management_endpoints/routing_group_endpoints.py +++ b/litellm/proxy/management_endpoints/routing_group_endpoints.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index dd65bc39b46..e426e4639bc 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 31615a768d7..57d4cab1aff 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -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: " diff --git a/litellm/router.py b/litellm/router.py index 34408e38e40..6dfca536a67 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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,