diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 90daaaeae6b..16180a696bc 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -35,7 +35,7 @@ from typing import ( import anyio import websockets import websockets.exceptions -from pydantic import BaseModel, Json, JsonValue +from pydantic import BaseModel, Json, JsonValue, TypeAdapter from typing_extensions import NotRequired, assert_never from litellm._uuid import uuid @@ -112,6 +112,7 @@ from litellm.router_utils.add_retry_fallback_headers import ( get_fallback_errors_from_headers, get_hidden_params_dict, ) +from litellm.router_utils.routing_groups import parse_routing_groups from litellm.types.utils import ( ModelResponse, ModelResponseStream, @@ -631,6 +632,7 @@ from litellm.types.realtime import RealtimeQueryParams from litellm.types.router import ( DeploymentTypedDict, RouterGeneralSettings, + RoutingGroup, RoutingPlugin, SearchToolTypedDict, updateDeployment, @@ -5926,7 +5928,32 @@ class ProxyConfig: combined_router_settings = db_router_settings.param_value if combined_router_settings: - llm_router.update_settings(**combined_router_settings) + llm_router.update_settings(**self._drop_invalid_routing_groups(combined_router_settings)) + + @staticmethod + def _drop_invalid_routing_groups(router_settings: Mapping[str, object]) -> Mapping[str, object]: + """ + A `routing_groups` value already persisted in the DB (saved before + save-time validation existed) must not take the rest of the reconcile + down with it: log it and apply every other setting, leaving whatever + groups the router already holds in place. + """ + raw_groups: Final = router_settings.get("routing_groups") + if raw_groups is None: + return router_settings + try: + parse_routing_groups( + TypeAdapter(list[RoutingGroup]).validate_python(raw_groups), + validate_strategy=Router._validate_routing_strategy, + ) + except ValueError as validation_error: + verbose_proxy_logger.error( + "Ignoring invalid router_settings.routing_groups from config/DB, all other router settings still " + "apply. Fix the routing groups in the Admin UI to load them: %s", + validation_error, + ) + return {k: v for k, v in router_settings.items() if k != "routing_groups"} + return router_settings def _add_general_settings_from_db_config( self, config_data: dict, general_settings: dict, proxy_logging_obj: ProxyLogging @@ -14881,6 +14908,15 @@ async def update_config( if prisma_client is None: raise Exception("No DB Connected") + if config_info.router_settings is not None: + try: + parse_routing_groups( + config_info.router_settings.routing_groups, + validate_strategy=Router._validate_routing_strategy, + ) + except ValueError as validation_error: + raise HTTPException(status_code=400, detail={"error": str(validation_error)}) + async def _read_section(param_name: str) -> dict: row: Final = await ConfigRepository(prisma_client).table.find_first(where={"param_name": param_name}) if row is None or row.param_value is None: diff --git a/litellm/router.py b/litellm/router.py index feaf69a44ae..a00b919cf19 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -154,6 +154,7 @@ from litellm.router_utils.router_callbacks.track_deployment_metrics import ( increment_deployment_failures_for_current_minute, increment_deployment_successes_for_current_minute, ) +from litellm.router_utils.routing_groups import parse_routing_groups from litellm.scheduler import FlowItem, Scheduler from litellm.types.llms.openai import ( AllMessageValues, @@ -898,7 +899,8 @@ class Router: return strategy.value return strategy - def _validate_routing_strategy(self, routing_strategy: RoutingStrategy | str | None) -> None: + @staticmethod + def _validate_routing_strategy(routing_strategy: RoutingStrategy | str | None) -> None: # See: https://github.com/BerriAI/litellm/issues/11330 valid_strategy_strings: Final = ["simple-shuffle", "lar1"] + [s.value for s in RoutingStrategy] if routing_strategy is None: @@ -1014,66 +1016,42 @@ class Router: at most one explicit group. Constructs per-group strategy selectors so groups with different `routing_strategy_args` track independent state. + Validation runs to completion before any router state changes, so an + invalid input raises with the previously loaded groups left intact. + Models not claimed by any explicit group are served by the implicit `"default"` group, whose selectors are the `self._logger` attributes set up in `routing_strategy_init`. """ + known_model_names: Final = frozenset(m["model_name"] for m in (self.model_list or []) if m.get("model_name")) + groups: Final = parse_routing_groups( + groups_input, + validate_strategy=self._validate_routing_strategy, + known_model_names=known_model_names, + ) + self._unregister_router_selectors( [sel for selectors in getattr(self, "_group_selectors", {}).values() for sel in selectors.values()] ) - self._routing_groups: dict[str, RoutingGroup] = {} - self._model_to_group: dict[str, str] = {} - self._group_selectors: dict[str, dict[str, Any]] = {} - - if not groups_input: - return - - known_model_names: Final = {m.get("model_name") for m in (self.model_list or []) if m.get("model_name")} - - seen_group_names: Final[set] = set() - for raw in groups_input: - group = raw if isinstance(raw, RoutingGroup) else RoutingGroup(**raw) - - if not group.group_name: - raise ValueError("routing_groups: group_name must be non-empty.") - if group.group_name == "default": - raise ValueError("routing_groups: 'default' is reserved for the implicit fallback group.") - if group.group_name in seen_group_names: - raise ValueError( - f"routing_groups: group names must be unique, duplicate group_name '{group.group_name}'." + self._routing_groups: dict[str, RoutingGroup] = {group.group_name: group for group in groups} + self._model_to_group: dict[str, str] = { + model_name: group.group_name for group in groups for model_name in group.models + } + self._group_selectors: dict[str, dict[str, Any]] = { + group.group_name: ( + {} + if ( + selector := self._build_strategy_selector( + strategy=group.routing_strategy, + routing_strategy_args=group.routing_strategy_args or {}, + ) ) - seen_group_names.add(group.group_name) - - self._validate_routing_strategy(group.routing_strategy) - - for model_name in group.models: - if model_name in self._model_to_group: - raise ValueError( - f"routing_groups: model_name '{model_name}' appears in " - f"both '{self._model_to_group[model_name]}' and " - f"'{group.group_name}'. Each model may belong to at most one group." - ) - if known_model_names and model_name not in known_model_names: - verbose_router_logger.warning( - "routing_groups: model_name '%s' (group '%s') is not in model_list; " - "the group entry will only take effect once a deployment with that " - "model_name is added.", - model_name, - group.group_name, - ) - self._model_to_group[model_name] = group.group_name - - self._routing_groups[group.group_name] = group - - strategy_value = self._normalize_strategy(group.routing_strategy) or "" - group_selector = self._build_strategy_selector( - strategy=group.routing_strategy, - routing_strategy_args=group.routing_strategy_args or {}, - ) - self._group_selectors[group.group_name] = ( - {strategy_value: group_selector} if group_selector is not None else {} + is None + else {self._normalize_strategy(group.routing_strategy) or "": selector} ) + for group in groups + } _OVERRIDABLE_ROUTING_STRATEGIES: frozenset[str] = frozenset({"simple-shuffle", *_DEFAULT_SELECTOR_ATTR_BY_STRATEGY}) diff --git a/litellm/router_utils/routing_groups.py b/litellm/router_utils/routing_groups.py new file mode 100644 index 00000000000..4b3f20c5b72 --- /dev/null +++ b/litellm/router_utils/routing_groups.py @@ -0,0 +1,77 @@ +""" +Validation for `router_settings.routing_groups`, shared by the Router and the +proxy's config-update endpoint so a config the UI saves cannot be one the +runtime refuses to load. +""" + +from collections.abc import Callable, Sequence +from typing import Final + +from litellm._logging import verbose_router_logger +from litellm.types.router import RoutingGroup, RoutingStrategy + +ValidateStrategy = Callable[[RoutingStrategy | str | None], None] + + +def parse_routing_groups( + groups_input: Sequence[RoutingGroup | dict] | None, + validate_strategy: ValidateStrategy, + known_model_names: frozenset[str] = frozenset(), +) -> tuple[RoutingGroup, ...]: + """ + Parses and validates `routing_groups`, raising `ValueError` on the first + problem found. Every check runs before the caller mutates any state, so an + invalid update can never leave a router holding a half-applied set of + groups. + """ + if not groups_input: + return () + + groups: Final = tuple(raw if isinstance(raw, RoutingGroup) else RoutingGroup(**raw) for raw in groups_input) + + if any(not group.group_name for group in groups): + raise ValueError("routing_groups: group_name must be non-empty.") + + if any(group.group_name == "default" for group in groups): + raise ValueError("routing_groups: 'default' is reserved for the implicit fallback group.") + + names: Final = [group.group_name for group in groups] + duplicate_names: Final = sorted({name for name in names if names.count(name) > 1}) + if duplicate_names: + raise ValueError(f"routing_groups: group names must be unique, duplicate group_name '{duplicate_names[0]}'.") + + for group in groups: + validate_strategy(group.routing_strategy) + + owners_by_model: Final = { + model_name: tuple(group.group_name for group in groups if model_name in group.models) + for model_name in dict.fromkeys(model_name for group in groups for model_name in group.models) + } + conflicts: Final = tuple( + f"model_name '{model_name}' appears in {' and '.join(repr(owner) for owner in owners)}" + for model_name, owners in owners_by_model.items() + if len(owners) > 1 + ) + if conflicts: + raise ValueError(f"routing_groups: {'; '.join(conflicts)}. Each model may belong to at most one group.") + + unknown_models: Final = ( + tuple( + (model_name, group.group_name) + for group in groups + for model_name in group.models + if model_name not in known_model_names + ) + if known_model_names + else () + ) + for model_name, group_name in unknown_models: + verbose_router_logger.warning( + "routing_groups: model_name '%s' (group '%s') is not in model_list; " + "the group entry will only take effect once a deployment with that " + "model_name is added.", + model_name, + group_name, + ) + + return groups diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 580a58885d9..98e29b3e48c 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11278,3 +11278,24 @@ async def test_setup_prisma_client_returns_none_when_connect_itself_fails(monkey assert result is None assert mock_client.start_db_health_watchdog_task.await_count == 0 assert mock_client.health_check.await_count == 0 + + +def test_drop_invalid_routing_groups_keeps_other_router_settings(): + from litellm.proxy.proxy_server import ProxyConfig + + valid = { + "num_retries": 3, + "routing_groups": [ + {"group_name": "g1", "models": ["m1"], "routing_strategy": "least-busy"}, + ], + } + assert ProxyConfig._drop_invalid_routing_groups(valid) == valid + + overlapping = { + "num_retries": 3, + "routing_groups": [ + {"group_name": "g1", "models": ["m1"], "routing_strategy": "least-busy"}, + {"group_name": "g2", "models": ["m1"], "routing_strategy": "least-busy"}, + ], + } + assert ProxyConfig._drop_invalid_routing_groups(overlapping) == {"num_retries": 3} diff --git a/tests/test_litellm/router_strategy/test_router_routing_groups.py b/tests/test_litellm/router_strategy/test_router_routing_groups.py index b8dcdacd8a3..8865bdf0172 100644 --- a/tests/test_litellm/router_strategy/test_router_routing_groups.py +++ b/tests/test_litellm/router_strategy/test_router_routing_groups.py @@ -726,3 +726,65 @@ def test_strategy_reinit_unregisters_override_selectors(): assert router._override_selectors == {} assert not any(id(cb) == id(override_selector) for cb in litellm.callbacks) assert router._get_override_strategy_selector("latency-based-routing") is router.lowestlatency_logger + + +def test_failed_routing_groups_update_keeps_previous_groups(): + router = _build_router( + routing_groups=[ + {"group_name": "g1", "models": ["filtered-model"], "routing_strategy": "latency-based-routing"}, + ], + ) + selector = router._group_selectors["g1"]["latency-based-routing"] + + with pytest.raises(ValueError, match="appears in"): + router.update_settings( + routing_groups=[ + {"group_name": "g1", "models": ["filtered-model"], "routing_strategy": "latency-based-routing"}, + {"group_name": "g2", "models": ["filtered-model"], "routing_strategy": "least-busy"}, + ], + ) + + assert list(router._routing_groups) == ["g1"] + assert router._model_to_group == {"filtered-model": "g1"} + assert router._group_selectors["g1"]["latency-based-routing"] is selector + assert router._get_routing_context("filtered-model", None) == ("latency-based-routing", selector) + + +def test_overlap_error_names_every_conflicting_model(): + with pytest.raises(ValueError) as exc_info: + _build_router( + routing_groups=[ + { + "group_name": "g1", + "models": ["filtered-model", "other-model"], + "routing_strategy": "latency-based-routing", + }, + { + "group_name": "g2", + "models": ["filtered-model", "other-model"], + "routing_strategy": "least-busy", + }, + ], + ) + message = str(exc_info.value) + assert "'filtered-model' appears in 'g1' and 'g2'" in message + assert "'other-model' appears in 'g1' and 'g2'" in message + + +def test_invalid_group_strategy_does_not_leak_a_selector(): + router = _build_router( + routing_groups=[ + {"group_name": "g1", "models": ["filtered-model"], "routing_strategy": "latency-based-routing"}, + ], + ) + selector = router._group_selectors["g1"]["latency-based-routing"] + + with pytest.raises(ValueError, match="Invalid routing_strategy"): + router.update_settings( + routing_groups=[ + {"group_name": "g2", "models": ["other-model"], "routing_strategy": "not-a-real-strategy"}, + ], + ) + + assert list(router._routing_groups) == ["g1"] + assert any(id(cb) == id(selector) for cb in litellm.callbacks) diff --git a/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.tsx b/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.tsx index 5b7c0fce6ef..b124217b2fb 100644 --- a/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.tsx +++ b/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.tsx @@ -3,6 +3,7 @@ import React, { useMemo } from "react"; import { Form, Input, Modal, Select, Space, Typography } from "antd"; import type { RoutingGroup, RoutingStrategy } from "./types"; +import { modelConflictError } from "./modelOwnership"; const { Text, Paragraph } = Typography; @@ -14,6 +15,7 @@ interface RoutingGroupModalProps { strategyDescriptions: Record; modelOptions: string[]; existingGroupNames: string[]; + groupNameByModel: Record; onClose: () => void; onSubmit: (group: RoutingGroup) => Promise | void; saving?: boolean; @@ -31,6 +33,16 @@ const STRATEGIES_WITH_ARGS = new Set(["latency-based-routing", "usage-ba const GROUP_NAME_PATTERN = /^[A-Za-z0-9._-]+$/; const GROUP_NAME_MAX_LENGTH = 64; +const modelRules = (groupNameByModel: Record) => [ + { required: true, message: "Select at least one model" }, + { + validator: (_: unknown, value: string[] | undefined) => { + const error = modelConflictError(value, groupNameByModel); + return error ? Promise.reject(new Error(error)) : Promise.resolve(); + }, + }, +]; + const RoutingGroupModal: React.FC = ({ open, mode, @@ -39,6 +51,7 @@ const RoutingGroupModal: React.FC = ({ strategyDescriptions, modelOptions, existingGroupNames, + groupNameByModel, onClose, onSubmit, saving, @@ -125,7 +138,7 @@ const RoutingGroupModal: React.FC = ({ }, }, ]} - extra="Use this name as the model in API calls — LiteLLM routes the request to one of the group's models." + extra="Names the group's shared routing strategy. Requests still use the model names, not this name." > @@ -133,8 +146,8 @@ const RoutingGroupModal: React.FC = ({