mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
efc4e6f28c
commit
6ba2889b1e
9 changed files with 294 additions and 56 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
||||
|
|
|
|||
77
litellm/router_utils/routing_groups.py
Normal file
77
litellm/router_utils/routing_groups.py
Normal 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
|
||||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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")');
|
||||
});
|
||||
});
|
||||
|
|
@ -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}`;
|
||||
};
|
||||
Loading…
Add table
Reference in a new issue