fix(routing_groups): validate overlapping models at save time and keep group loading atomic

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-08-09 00:00:21 +00:00
parent efc4e6f28c
commit 6ba2889b1e
9 changed files with 294 additions and 56 deletions

View file

@ -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:

View file

@ -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.<strategy>_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})

View file

@ -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

View file

@ -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}

View file

@ -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)

View file

@ -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<string, string>;
modelOptions: string[];
existingGroupNames: string[];
groupNameByModel: Record<string, string>;
onClose: () => void;
onSubmit: (group: RoutingGroup) => Promise<void> | void;
saving?: boolean;
@ -31,6 +33,16 @@ const STRATEGIES_WITH_ARGS = new Set<string>(["latency-based-routing", "usage-ba
const GROUP_NAME_PATTERN = /^[A-Za-z0-9._-]+$/;
const GROUP_NAME_MAX_LENGTH = 64;
const modelRules = (groupNameByModel: Record<string, string>) => [
{ 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<RoutingGroupModalProps> = ({
open,
mode,
@ -39,6 +51,7 @@ const RoutingGroupModal: React.FC<RoutingGroupModalProps> = ({
strategyDescriptions,
modelOptions,
existingGroupNames,
groupNameByModel,
onClose,
onSubmit,
saving,
@ -125,7 +138,7 @@ const RoutingGroupModal: React.FC<RoutingGroupModalProps> = ({
},
},
]}
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."
>
<Input placeholder="fast-chat" disabled={mode === "edit"} />
</Form.Item>
@ -133,8 +146,8 @@ const RoutingGroupModal: React.FC<RoutingGroupModalProps> = ({
<Form.Item
label="Models"
name="models"
rules={[{ required: true, message: "Select at least one model" }]}
extra="Models from your model list that this group routes between."
rules={modelRules(groupNameByModel)}
extra="Models from your model list that this group routes between. A model can only be in one group."
>
<Select
mode="multiple"

View file

@ -11,6 +11,7 @@ import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySett
import RoutingGroupsTable from "./RoutingGroupsTable";
import RoutingGroupModal from "./RoutingGroupModal";
import NotificationsManager from "../molecules/notifications_manager";
import { groupNameByModel } from "./modelOwnership";
import type { RoutingGroup } from "./types";
const { Text } = Typography;
@ -56,6 +57,8 @@ const RoutingGroups: React.FC = () => {
return Array.from(new Set(names));
}, [modelHub]);
const ownerByModel = groupNameByModel(groups, drawerMode === "edit" ? editingGroup?.group_name : undefined);
const openCreate = () => {
setDrawerMode("create");
setEditingGroup(null);
@ -141,6 +144,7 @@ const RoutingGroups: React.FC = () => {
strategyDescriptions={strategyDescriptions}
modelOptions={modelOptions}
existingGroupNames={groups.map((g) => g.group_name)}
groupNameByModel={ownerByModel}
onClose={() => setDrawerOpen(false)}
onSubmit={handleSubmit}
saving={saveMutation.isPending}

View file

@ -0,0 +1,29 @@
import { groupNameByModel, modelConflictError } from "./modelOwnership";
import type { RoutingGroup } from "./types";
const groups: RoutingGroup[] = [
{ group_name: "cheap", models: ["m1", "m2"], routing_strategy: "latency-based-routing" },
{ group_name: "security", models: ["m3"], routing_strategy: "least-busy" },
];
describe("groupNameByModel", () => {
it("maps every claimed model to its owning group", () => {
expect(groupNameByModel(groups)).toEqual({ m1: "cheap", m2: "cheap", m3: "security" });
});
it("excludes the group being edited so its own models stay selectable", () => {
expect(groupNameByModel(groups, "cheap")).toEqual({ m3: "security" });
});
});
describe("modelConflictError", () => {
it("passes models that no other group claims", () => {
expect(modelConflictError(["m4"], groupNameByModel(groups, "cheap"))).toBeNull();
expect(modelConflictError(undefined, groupNameByModel(groups))).toBeNull();
});
it("names every model already claimed by another group", () => {
const error = modelConflictError(["m1", "m3", "m4"], groupNameByModel(groups));
expect(error).toBe('Each model may belong to at most one group. Already claimed: m1 (in "cheap"), m3 (in "security")');
});
});

View file

@ -0,0 +1,18 @@
import type { RoutingGroup } from "./types";
export const groupNameByModel = (groups: RoutingGroup[], excludeGroupName?: string): Record<string, string> =>
Object.fromEntries(
groups
.filter((group) => group.group_name !== excludeGroupName)
.flatMap((group) => group.models.map((model) => [model, group.group_name] as const)),
);
export const modelConflictError = (
models: string[] | undefined,
ownerByModel: Record<string, string>,
): string | null => {
const conflicts = (models ?? []).filter((model) => ownerByModel[model] !== undefined);
if (conflicts.length === 0) return null;
const detail = conflicts.map((model) => `${model} (in "${ownerByModel[model]}")`).join(", ");
return `Each model may belong to at most one group. Already claimed: ${detail}`;
};