fix(router): honor team and key provider weights

This commit is contained in:
Tin Chi Lo 2026-09-14 02:22:55 -07:00
parent c93708b2a5
commit 398300c4e7
25 changed files with 561 additions and 93 deletions

View file

@ -4,11 +4,12 @@ import os
from collections.abc import Callable, Mapping
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple
import httpx
from pydantic import (
BaseModel,
BeforeValidator,
ConfigDict,
Field,
Json,
@ -47,6 +48,7 @@ from litellm.types.proxy.carried_budget_state import (
)
from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry
from litellm.types.router import RouterErrors, UpdateRouterConfig
from litellm.types.router_weights import validate_router_settings_dict
from litellm.types.secret_managers.main import KeyManagementSystem
from litellm.types.utils import (
CallTypes,
@ -1981,8 +1983,14 @@ class OrgMember(MemberBase):
from litellm.models.team import TeamBase as TeamBase # noqa: E402
RouterSettingsDict = Annotated[
dict[str, object],
BeforeValidator(validate_router_settings_dict, json_schema_input_type=UpdateRouterConfig),
]
class NewTeamRequest(TeamBase):
router_settings: RouterSettingsDict | None = None
model_aliases: dict | None = None
tags: list | None = None
guardrails: list[str] | None = None
@ -2080,7 +2088,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None
enforced_batch_output_expires_after: dict | None = None
enforced_file_expires_after: dict | None = None
router_settings: dict | None = None
router_settings: RouterSettingsDict | None = None
access_group_ids: list[str] | None = None
budget_limits: list[BudgetLimitEntry] | None = None # multiple concurrent budget windows
default_team_member_models: list[str] | None = None # default allowed_models seeded onto new team members

View file

@ -249,6 +249,7 @@ async def authenticate_user(
if os.getenv("DATABASE_URL") is not None:
response = await generate_key_helper_fn(
llm_router=None,
request_type="key",
**{
"user_role": LitellmUserRoles.PROXY_ADMIN,
@ -324,6 +325,7 @@ async def authenticate_user(
await _rehash_password_if_needed(_user_row.user_id, password, _password)
if os.getenv("DATABASE_URL") is not None:
response = await generate_key_helper_fn(
llm_router=None,
request_type="key",
**{
"user_role": user_role,

View file

@ -878,6 +878,7 @@ async def _auto_register_jwt_mapping(
# the NOT NULL @id constraint. Every successful key-creation caller (e.g.
# /key/generate) passes table_name="key" explicitly.
key_data: Final = await generate_key_helper_fn(
llm_router=None,
request_type="key",
table_name="key",
team_id=team_id,

View file

@ -14,6 +14,7 @@ import httpx
import orjson
from fastapi import HTTPException, Request, status
from fastapi.responses import JSONResponse, Response, StreamingResponse
from pydantic import ValidationError
from starlette.types import Receive, Scope, Send
import litellm
@ -76,6 +77,7 @@ from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_di
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.router import RouterRateLimitError
from litellm.types.router_weights import validate_router_weights
_LateResponseT = TypeVar("_LateResponseT", bound=Response)
_LlmCallT = TypeVar("_LlmCallT")
@ -1939,6 +1941,13 @@ class ProxyBaseLLMRequestProcessing:
# This avoids expensive Router instantiation on each request
if router_settings is not None:
self.data["router_settings_override"] = router_settings
try:
self.data["_router_weights"] = validate_router_weights(router_settings.get("weights"))
except ValidationError:
self.data["_router_weights"] = None
verbose_proxy_logger.warning(
"Ignoring invalid saved router weights; update team/key router_settings"
)
alias_target: Final = await _resolve_per_request_model_group_alias(
requested_model=self.data.get("model"),
router_settings=router_settings,

View file

@ -221,6 +221,8 @@ LITELLM_TRACE_CONTROL_METADATA_FIELDS: Final = frozenset(
)
_UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
"weights",
"_router_weights",
"proxy_server_request",
"standard_logging_object",
"secret_fields",

View file

@ -567,7 +567,7 @@ async def new_user(
teams = check_if_default_team_set()
organization_ids: Final = cast(list[str] | None, data_json.pop("organizations", None))
response: Final = await generate_key_helper_fn(request_type="user", **data_json)
response: Final = await generate_key_helper_fn(request_type="user", **data_json, llm_router=None)
# Admin UI Logic
# Add User to Team and Organization
# if team_id passed add this user to the team

View file

@ -96,6 +96,7 @@ from litellm.proxy.management_endpoints.common_utils import (
from litellm.proxy.management_endpoints.model_management_endpoints import (
_add_model_to_db,
)
from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights
from litellm.proxy.management_helpers.access_group_key_sync import (
sync_key_access_group_membership,
sync_key_regeneration_access_group_membership,
@ -201,6 +202,10 @@ class _KeyUpdateResult(TypedDict):
data: ReadOnly[Mapping[str, object]]
class _StoredKeyRouterSettings(BaseModel):
router_settings: Mapping[str, object] | None = None
class _KeyRowWhere(TypedDict):
token: ReadOnly[str]
@ -1330,7 +1335,7 @@ async def _common_key_generation_helper(
prisma_client=prisma_client,
)
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key")
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router)
response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response
@ -2234,7 +2239,26 @@ async def _update_key_row_with_soft_budget(
async def prepare_key_update_data(
data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken,
*,
prisma_client: PrismaClient | None = None,
llm_router: Router | None = None,
):
if data.router_settings is not None or (
"router_settings" not in data.model_fields_set
and "team_id" in data.model_fields_set
and data.team_id != existing_key_row.team_id
):
effective_settings: Final = (
data.router_settings
if data.router_settings is not None
else _StoredKeyRouterSettings.model_validate(existing_key_row, from_attributes=True).router_settings
)
await validate_router_settings_weights(
effective_settings,
team_id=data.team_id if "team_id" in data.model_fields_set else existing_key_row.team_id,
prisma_client=prisma_client,
llm_router=llm_router,
)
data_json: Final[dict] = data.model_dump(exclude_unset=True)
data_json.pop("key", None)
data_json.pop("new_key", None)
@ -2575,7 +2599,9 @@ async def _process_single_key_update(
)
# Prepare update data
non_default_values = await prepare_key_update_data(data=update_key_request, existing_key_row=existing_key_row)
non_default_values = await prepare_key_update_data(
data=update_key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
)
# Update key in database
if prisma_client is None:
@ -3093,7 +3119,9 @@ async def update_key_fn(
# Enforce upperbound key params on update (don't fill defaults)
_enforce_upperbound_key_params(data, fill_defaults=False)
non_default_values: Final = await prepare_key_update_data(data=data, existing_key_row=existing_key_row)
non_default_values: Final = await prepare_key_update_data(
data=data, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
)
# Only validate key_alias format if it's actually being changed
new_key_alias: Final = non_default_values.get("key_alias", None)
@ -4137,15 +4165,24 @@ async def generate_key_helper_fn(
object_permission: LiteLLM_ObjectPermissionBase | None = None,
auto_rotate: bool | None = None,
rotation_interval: str | None = None,
router_settings: dict | None = None,
router_settings: dict[str, object] | None = None,
access_group_ids: list[str] | None = None,
budget_limits: list | None = None, # multiple concurrent budget windows
*,
llm_router: Router | None = None,
):
from litellm.proxy.proxy_server import premium_user, prisma_client
if prisma_client is None:
raise Exception("Connect Proxy to database to generate keys - https://docs.litellm.ai/docs/proxy/virtual_keys ")
await validate_router_settings_weights(
router_settings,
team_id=team_id,
prisma_client=prisma_client,
llm_router=llm_router,
)
if token is None:
if key is not None:
token = key
@ -5070,6 +5107,7 @@ async def _insert_deprecated_key(
async def _execute_virtual_key_regeneration(
*,
prisma_client: PrismaClient,
llm_router: Router | None = None,
key_in_db: LiteLLM_VerificationToken,
hashed_api_key: str,
key: str,
@ -5129,7 +5167,9 @@ async def _execute_virtual_key_regeneration(
if data is not None:
# Enforce upperbound key params on regenerate (don't fill defaults)
_enforce_upperbound_key_params(data, fill_defaults=False)
non_default_values = await prepare_key_update_data(data=data, existing_key_row=key_in_db)
non_default_values = await prepare_key_update_data(
data=data, existing_key_row=key_in_db, prisma_client=prisma_client, llm_router=llm_router
)
# Only validate key_alias format if it's actually being changed
new_key_alias: Final = non_default_values.get("key_alias")
if new_key_alias != key_in_db.key_alias:
@ -5268,6 +5308,7 @@ async def regenerate_key_fn(
try:
from litellm.proxy.proxy_server import (
hash_token,
llm_router,
master_key,
premium_user,
prisma_client,
@ -5456,6 +5497,7 @@ async def regenerate_key_fn(
return await _execute_virtual_key_regeneration(
prisma_client=prisma_client,
llm_router=llm_router,
key_in_db=_key_in_db,
hashed_api_key=hashed_api_key,
key=key,

View file

@ -0,0 +1,129 @@
from abc import abstractmethod
from collections.abc import Mapping
from typing import Annotated, Final, Protocol
from fastapi import HTTPException
from pydantic import BaseModel, BeforeValidator, ValidationError
from litellm.repositories.prisma_protocols import TableActions
from litellm.types.router_weights import RouterWeights
class _StoredModel(Protocol):
@property
@abstractmethod
def model_id(self) -> str:
pass
class _ModelDb(Protocol):
@property
@abstractmethod
def litellm_proxymodeltable(self) -> TableActions[_StoredModel]:
pass
class _PrismaClient(Protocol):
@property
@abstractmethod
def db(self) -> _ModelDb:
pass
class _Router(Protocol):
@abstractmethod
def get_deployment(self, model_id: str) -> object | None:
pass
class _RouterWeightSettings(BaseModel):
weights: RouterWeights | None = None
class _RouterWeightModelInfo(BaseModel):
team_id: str | None = None
db_model: bool | None = None
team_public_model_name: str | None = None
def _router_weight_model_info(value: object) -> _RouterWeightModelInfo:
if isinstance(value, str):
return _RouterWeightModelInfo.model_validate_json(value)
return _RouterWeightModelInfo.model_validate(value or {}, from_attributes=True)
class _RouterWeightDeployment(BaseModel):
model_name: str
model_info: Annotated[_RouterWeightModelInfo, BeforeValidator(_router_weight_model_info)]
def _validate_router_weight_reference(
model_group: str,
deployment_id: str,
team_id: str | None,
stored: _RouterWeightDeployment | None,
configured: object | None,
) -> None:
reference: Final = (
stored
if stored is not None
else (
_RouterWeightDeployment.model_validate(configured, from_attributes=True) if configured is not None else None
)
)
if (
reference is None
or (stored is None and reference.model_info.db_model)
or (reference.model_info.team_id is not None and reference.model_info.team_id != team_id)
):
raise HTTPException(status_code=400, detail=f"Unknown deployment ID in router weights: {deployment_id}")
canonical_group: Final = (
reference.model_info.team_public_model_name if reference.model_info.team_id is not None else None
) or reference.model_name
if model_group != canonical_group:
raise HTTPException(
status_code=400,
detail=f"Deployment {deployment_id} does not belong to model group {model_group}",
)
async def validate_router_settings_weights(
router_settings: BaseModel | Mapping[str, object] | None,
*,
team_id: str | None,
prisma_client: _PrismaClient | None,
llm_router: _Router | None,
) -> None:
try:
weights: Final = (
_RouterWeightSettings.model_validate(router_settings, from_attributes=True).weights
if router_settings is not None
else None
)
except ValidationError:
raise HTTPException(
status_code=400,
detail="Invalid router weights. Replace or clear router_settings.weights.",
) from None
if not weights:
return
deployment_ids: Final = frozenset(deployment_id for group in weights.values() for deployment_id in group)
if not deployment_ids:
return
if prisma_client is None:
raise HTTPException(status_code=503, detail="Database unavailable while validating router weights")
stored_models: Final = await prisma_client.db.litellm_proxymodeltable.find_many(
where={"model_id": {"in": list(deployment_ids)}}
)
stored_by_id: Final = {
row.model_id: _RouterWeightDeployment.model_validate(row, from_attributes=True) for row in stored_models
}
for model_group, group_weights in weights.items():
for deployment_id in group_weights:
_validate_router_weight_reference(
model_group,
deployment_id,
team_id,
stored_by_id.get(deployment_id),
llm_router.get_deployment(model_id=deployment_id) if llm_router is not None else None,
)

View file

@ -112,6 +112,7 @@ from litellm.proxy.management_endpoints.common_utils import (
from litellm.proxy.management_endpoints.organization_endpoints import (
add_member_to_organization,
)
from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights
from litellm.proxy.management_endpoints.tag_management_endpoints import (
get_daily_activity,
)
@ -1288,6 +1289,7 @@ async def new_team(
create_audit_log_for_update,
general_settings,
litellm_proxy_admin_name,
llm_router,
prisma_client,
user_api_key_cache,
)
@ -1462,6 +1464,13 @@ async def new_team(
user_api_key_dict=user_api_key_dict,
)
await validate_router_settings_weights(
data.router_settings,
team_id=data.team_id,
prisma_client=prisma_client,
llm_router=llm_router,
)
## ADD TO MODEL TABLE
_model_id = None
if data.model_aliases is not None and isinstance(data.model_aliases, dict):
@ -2075,6 +2084,13 @@ async def update_team(
user_api_key_dict=user_api_key_dict,
)
await validate_router_settings_weights(
data.router_settings,
team_id=data.team_id,
prisma_client=prisma_client,
llm_router=llm_router,
)
_existing_team_metadata: Final[object] = getattr(existing_team_row, "metadata", None)
enforce_output_token_estimates_are_admin_only(
data=data,

View file

@ -3592,6 +3592,7 @@ class SSOAuthenticationHandler:
verbose_proxy_logger.info("user_defined_values for creating ui key: %s", user_defined_values)
response: Final = await generate_key_helper_fn(
llm_router=None,
request_type="key",
duration=LITELLM_UI_SESSION_DURATION,
key_max_budget=litellm.max_ui_session_budget,

View file

@ -9546,6 +9546,7 @@ class ProxyStartupEvent:
gate the first duration window.
"""
await generate_key_helper_fn(
llm_router=llm_router,
request_type="user",
table_name="user",
user_id=LITELLM_PROXY_BUDGET_NAME,
@ -16290,6 +16291,7 @@ async def _generate_onboarding_ui_session_token(user_obj: _UserTableRow) -> str:
global master_key, general_settings
response: Final = await generate_key_helper_fn(
llm_router=llm_router,
request_type="key",
**{
"user_role": user_obj.user_role,

View file

@ -4849,6 +4849,7 @@ class Router:
model=model,
messages=messages,
specific_deployment=kwargs.pop("specific_deployment", None),
request_kwargs=kwargs,
)
data: Final = deployment["litellm_params"].copy()
@ -5163,13 +5164,11 @@ class Router:
return healthy_deployments[0]
# Use simple_shuffle for weighted selection
return cast(
GuardrailTypedDict,
simple_shuffle(
llm_router_instance=self,
healthy_deployments=healthy_deployments,
model=guardrail_name,
),
return simple_shuffle(
resolve_model_alias=self._get_model_from_alias,
healthy_deployments=healthy_deployments,
model=guardrail_name,
request_kwargs=None,
)
async def _ageneric_api_call_with_fallbacks(self, model: str, original_function: Callable, **kwargs):
@ -13045,9 +13044,10 @@ class Router:
start_time: Final = time.time()
if strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
resolve_model_alias=self._get_model_from_alias,
healthy_deployments=healthy_deployments,
model=model,
request_kwargs=request_kwargs,
)
deployment: Final = await self._select_deployment_async(
strategy=strategy,
@ -13190,9 +13190,10 @@ class Router:
start_time: Final = time.perf_counter()
if strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
resolve_model_alias=self._get_model_from_alias,
healthy_deployments=pass_through_deployments,
model=model,
request_kwargs=request_kwargs,
)
deployment: Final = await self._select_deployment_async(
strategy=strategy,
@ -13888,9 +13889,10 @@ class Router:
# 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,
resolve_model_alias=self._get_model_from_alias,
healthy_deployments=healthy_deployments,
model=model,
request_kwargs=request_kwargs,
)
deployment: Final = self._select_deployment_sync(
strategy=strategy,
@ -13958,6 +13960,7 @@ class Router:
messages=messages,
input=input,
specific_deployment=specific_deployment,
request_kwargs=request_kwargs,
)
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
@ -14040,9 +14043,10 @@ class Router:
# 6. Apply load balancing strategy
if strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
resolve_model_alias=self._get_model_from_alias,
healthy_deployments=pass_through_deployments,
model=model,
request_kwargs=request_kwargs,
)
deployment: Final = self._select_deployment_sync(
strategy=strategy,

View file

@ -1,71 +1,67 @@
"""
Returns a random deployment from the list of healthy deployments.
"""Choose among eligible deployments using request weights, then global metrics."""
If weights are provided, it will return a deployment based on the weights.
"""
from __future__ import annotations
import logging
import random
from typing import TYPE_CHECKING, Any, Final
from collections.abc import Callable, Mapping, Sequence
from itertools import chain
from typing import Final, TypeVar
from litellm._logging import verbose_router_logger
from litellm.types.router_weights import validate_router_weights
if TYPE_CHECKING:
from litellm.router import Router as _Router
_DeploymentT = TypeVar("_DeploymentT", bound=Mapping[str, object])
_ROUTER_LOGGER: Final = logging.getLogger("LiteLLM Router")
LitellmRouter = _Router
else:
LitellmRouter = Any
def _metric_weight(deployment: Mapping[str, object], metric: str) -> float:
params: Final = deployment.get("litellm_params")
value: Final = params.get(metric) if isinstance(params, Mapping) else None
if value is None:
return 0.0
if isinstance(value, (int, float)):
return float(value)
raise TypeError(f"Deployment {metric} must be numeric")
def _scoped_weights(
deployments: Sequence[Mapping[str, object]],
model: str,
request_kwargs: Mapping[str, object] | None,
) -> tuple[float, ...]:
settings: Final = validate_router_weights((request_kwargs or {}).get("_router_weights"))
model_weights: Final = settings.get(model) if settings is not None else None
if not model_weights:
return ()
return tuple(
model_weights.get(str(info.get("id")), 0.0) if isinstance(info, Mapping) else 0.0
for deployment in deployments
for info in (deployment.get("model_info"),)
)
def simple_shuffle(
llm_router_instance: LitellmRouter,
healthy_deployments: list[Any] | dict[Any, Any],
resolve_model_alias: Callable[[str], str | None],
healthy_deployments: Sequence[_DeploymentT],
model: str,
) -> dict:
"""
Returns a random deployment from the list of healthy deployments.
If weights are provided, it will return a deployment based on the weights.
If users pass `rpm` or `tpm`, we do a random weighted pick - based on `rpm`/`tpm`.
Args:
llm_router_instance: LitellmRouter instance
healthy_deployments: List of healthy deployments
model: Model name
Returns:
Dict: A single healthy deployment
"""
############## Check if 'weight' or 'rpm' or 'tpm' param set for a weighted pick #################
for weight_by in ["weight", "rpm", "tpm"]:
if any(m["litellm_params"].get(weight_by) is not None for m in healthy_deployments):
weights = [m["litellm_params"].get(weight_by, 0) for m in healthy_deployments]
verbose_router_logger.debug("\nweight %s", weights)
total_weight = sum(weights)
if total_weight <= 0:
# All remaining candidates have weight 0 for this metric (e.g.
# after a weighted-failover exclusion left only zero-weight
# backups). Skip to the next metric (rpm/tpm) which may still
# provide a meaningful weighted pick; if none do, we fall
# through to the uniform random pick at the end.
continue
weights = [weight / total_weight for weight in weights]
verbose_router_logger.debug("\n weights %s by %s", weights, weight_by)
# Perform weighted random pick
selected_index = random.choices(range(len(weights)), weights=weights)[0]
verbose_router_logger.debug("\n selected index, %s", selected_index)
deployment = healthy_deployments[selected_index]
verbose_router_logger.info(
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
model,
llm_router_instance.print_deployment(deployment) or deployment[0],
model,
)
return deployment or deployment[0]
############## No RPM/TPM passed, we do a random pick #################
item: Final = random.choice(healthy_deployments)
return item or item[0]
request_kwargs: Mapping[str, object] | None,
) -> _DeploymentT:
resolved_model: Final = resolve_model_alias(model) or model
weight_sets: Final = chain(
(_scoped_weights(healthy_deployments, resolved_model, request_kwargs),),
(
tuple(_metric_weight(deployment, metric) for deployment in healthy_deployments)
for metric in ("weight", "rpm", "tpm")
),
)
for weights in weight_sets:
largest = max(weights, default=0.0)
if largest <= 0:
continue
normalized = tuple(weight / largest for weight in weights)
if sum(normalized) <= 0:
continue
selected = random.choices(healthy_deployments, weights=normalized)[0]
_ROUTER_LOGGER.info("Selected deployment for model %s: %s", model, selected.get("model_info"))
return selected
return random.choice(healthy_deployments)

View file

@ -15,6 +15,7 @@ from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_c
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.types.router_weights import RouterWeights
if TYPE_CHECKING:
from litellm.router import Router
@ -146,6 +147,7 @@ class UpdateRouterConfig(BaseModel):
context_window_fallbacks: list[dict] | None = None
model_group_alias: dict[str, str | dict] | None = {}
enable_tag_filtering: bool | None = None
weights: RouterWeights | None = None
tag_routing_prefix: str | None = None
optional_pre_call_checks: OptionalPreCallChecks | None = None

View file

@ -0,0 +1,30 @@
from collections.abc import Mapping
from typing import Annotated, Final
from pydantic import AfterValidator, Field, TypeAdapter
def _validate_positive_router_weights(weights: Mapping[str, Mapping[str, float]]) -> Mapping[str, Mapping[str, float]]:
if any(group and not any(weight > 0 for weight in group.values()) for group in weights.values()):
raise ValueError("Each nonempty weights group must contain at least one positive weight")
return weights
RouterWeightIdentifier = Annotated[str, Field(strict=True, min_length=1, pattern=r"\S")]
RouterWeight = Annotated[float, Field(strict=True, ge=0, allow_inf_nan=False)]
RouterWeights = Annotated[
dict[RouterWeightIdentifier, dict[RouterWeightIdentifier, RouterWeight]],
AfterValidator(_validate_positive_router_weights),
]
_ROUTER_WEIGHTS_ADAPTER: Final[TypeAdapter[RouterWeights | None]] = TypeAdapter(RouterWeights | None)
_ROUTER_SETTINGS_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
def validate_router_weights(value: object) -> RouterWeights | None:
return _ROUTER_WEIGHTS_ADAPTER.validate_python(value)
def validate_router_settings_dict(value: object) -> dict[str, object]:
settings: Final = _ROUTER_SETTINGS_DICT_ADAPTER.validate_python(value)
validate_router_weights(settings.get("weights"))
return settings

View file

@ -3782,6 +3782,7 @@ all_litellm_params = (
"id",
"fallbacks",
"routing_strategy",
"_router_weights",
"azure",
"headers",
"model_list",

View file

@ -8,6 +8,10 @@ users can intentionally clear previously-set fields.
"""
from datetime import datetime, timezone
from types import SimpleNamespace
from fastapi import HTTPException
from litellm import Router
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -1120,3 +1124,41 @@ class TestUpdateMetadataFieldsPremiumCheck:
}
_update_metadata_fields(updated_kv)
mock_check.assert_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("db_model,stored_name,owner,public_name,error", [
(False, None, None, None, None),
(True, "group", None, None, None),
(True, None, None, None, "Unknown deployment ID in router weights: id"),
(False, "renamed", None, None, "Deployment id does not belong to model group group"),
(False, None, "other-team", None, "Unknown deployment ID in router weights: id"),
(True, "internal", "team", "group", None),
(True, "group", "team", "public", "Deployment id does not belong to model group group"),
(True, "group", None, "unrelated-public-name", None),
])
async def test_router_weights_validate_current_deployment_scope(
db_model: bool, stored_name: str | None, owner: str | None,
public_name: str | None, error: str | None,
) -> None:
from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights
info = {"team_id": owner, "team_public_model_name": public_name}
router = Router(model_list=[{
"model_name": "group",
"litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "test"},
"model_info": {"id": "id", "db_model": db_model, **info},
}])
rows = [SimpleNamespace(model_id="id", model_name=stored_name, model_info=info)] if stored_name else []
table = SimpleNamespace(find_many=AsyncMock(return_value=rows))
db = SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table))
validation = validate_router_settings_weights(
{"weights": {"group": {"id": 1}}}, team_id="team", prisma_client=db, llm_router=router,
)
if error:
with pytest.raises(HTTPException, match=error) as exc:
await validation
assert exc.value.status_code == 400
assert exc.value.detail == error
else:
await validation

View file

@ -1,4 +1,5 @@
from typing import Final
from types import SimpleNamespace
import json
from datetime import datetime, timedelta, timezone
@ -27,6 +28,7 @@ from litellm.proxy._types import (
Member,
ProxyException,
ResetSpendRequest,
RegenerateKeyRequest,
UpdateKeyRequest,
)
from litellm.proxy.auth.auth_checks import _delete_cache_key_object, _project_cache_key
@ -6615,6 +6617,9 @@ async def test_generate_key_with_router_settings(monkeypatch):
return_value=[]
)
mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0)
mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[
SimpleNamespace(model_id="weighted-id", model_name="gpt-4", model_info={})
])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
@ -6630,6 +6635,7 @@ async def test_generate_key_with_router_settings(monkeypatch):
"routing_strategy": "usage-based",
"num_retries": 3,
"model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}},
"weights": {"gpt-4": {"weighted-id": 1}},
}
request_data = GenerateKeyRequest(
@ -6679,21 +6685,37 @@ async def test_generate_key_with_router_settings(monkeypatch):
# Verify router_settings matches input (regardless of serialization state)
assert actual_settings == router_settings_data
mock_prisma_client.insert_data.reset_mock()
with pytest.raises(ProxyException, match="Unknown deployment ID"):
await generate_key_fn(
data=GenerateKeyRequest(router_settings={"weights": {"gpt-4": {"unknown-id": 1}}}),
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="user-router-1"),
)
mock_prisma_client.insert_data.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_key_with_router_settings(monkeypatch):
@pytest.mark.parametrize("request_type", [UpdateKeyRequest, RegenerateKeyRequest])
@pytest.mark.parametrize("target_team", ["new-team", None])
async def test_update_key_with_router_settings(
monkeypatch: pytest.MonkeyPatch,
request_type: type[UpdateKeyRequest | RegenerateKeyRequest], target_team: str | None,
) -> None:
"""
Test that /key/update correctly handles router_settings by:
1. Accepting router_settings as a dict parameter
2. Serializing router_settings to JSON when updating database
3. Updating router_settings in the key record
"""
from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.management_endpoints.key_management_endpoints import (
prepare_key_update_data,
)
model = SimpleNamespace(model_id="weighted-id", model_name="gpt-4", model_info={})
table = SimpleNamespace(find_many=AsyncMock(return_value=[model]))
db = SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table))
# Mock existing key
existing_key = LiteLLM_VerificationToken(
token="test-token-router",
@ -6710,14 +6732,16 @@ async def test_update_key_with_router_settings(monkeypatch):
router_settings_data = {
"routing_strategy": "latency-based",
"num_retries": 2,
"weights": {"gpt-4": {"weighted-id": 1}},
}
update_request = UpdateKeyRequest(
update_request = request_type(
key="test-token-router", router_settings=router_settings_data
)
result = await prepare_key_update_data(
data=update_request, existing_key_row=existing_key
data=update_request, existing_key_row=existing_key,
prisma_client=db, llm_router=None,
)
# Verify router_settings is serialized to JSON string
@ -6728,6 +6752,28 @@ async def test_update_key_with_router_settings(monkeypatch):
deserialized_settings = json.loads(result["router_settings"])
assert deserialized_settings == router_settings_data
with pytest.raises(HTTPException, match="Unknown deployment ID"):
await prepare_key_update_data(
request_type(key=existing_key.token, router_settings={"weights": {"gpt-4": {"unknown-id": 1}}}),
existing_key,
prisma_client=db, llm_router=None,
)
existing_key.team_id = "old-team"
existing_key.router_settings = router_settings_data
move = request_type(key=existing_key.token, team_id=target_team)
retained = await prepare_key_update_data(move, existing_key, prisma_client=db, llm_router=None)
assert retained["team_id"] == target_team
assert "router_settings" not in retained
model.model_info = {"team_id": "old-team"}
with pytest.raises(HTTPException, match="Unknown deployment ID"):
await prepare_key_update_data(move, existing_key, prisma_client=db, llm_router=None)
cleared = await prepare_key_update_data(
request_type(key=existing_key.token, team_id=target_team, router_settings={}), existing_key,
prisma_client=db, llm_router=None,
)
assert cleared["team_id"] == target_team
assert json.loads(cleared["router_settings"]) == {}
@pytest.mark.asyncio
async def test_validate_max_budget():

View file

@ -9476,6 +9476,9 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth):
mock_db_client.get_data = AsyncMock(return_value=None)
mock_db_client.update_data = AsyncMock(return_value=MagicMock())
mock_db_client.db = MagicMock()
mock_db_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[
SimpleNamespace(model_id="weighted-id", model_name="group", model_info={})
])
# Mock model table creation
mock_db_client.db.litellm_modeltable = MagicMock()
@ -9511,6 +9514,7 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth):
# Test router_settings with sample data
router_settings_data = {
"weights": {"group": {"weighted-id": 1}},
"routing_strategy": "usage-based",
"num_retries": 3,
"retry_policy": {"max_retries": 5},
@ -9544,6 +9548,12 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth):
deserialized_settings = json.loads(team_data["router_settings"])
assert deserialized_settings == router_settings_data
mock_team_create.reset_mock()
team_request.router_settings = {"weights": {"group": {"unknown-id": 1}}}
with pytest.raises(ProxyException, match="Unknown deployment ID"):
await new_team(data=team_request, http_request=dummy_request, user_api_key_dict=mock_admin_auth)
mock_team_create.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_team_daily_activity_member_with_permission_sees_all_spend(
@ -9739,6 +9749,9 @@ async def test_update_team_with_router_settings(
# Configure mocked prisma client
mock_db_client.jsonify_team_object = lambda db_data: db_data
mock_db_client.db = MagicMock()
mock_db_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[
SimpleNamespace(model_id="weighted-id", model_name="group", model_info={})
])
# Mock existing team row
existing_team_mock = MagicMock()
@ -9773,6 +9786,7 @@ async def test_update_team_with_router_settings(
# Test router_settings with updated data
router_settings_data = {
"weights": {"group": {"weighted-id": 1}},
"routing_strategy": "latency-based",
"num_retries": 2,
}
@ -9805,6 +9819,12 @@ async def test_update_team_with_router_settings(
deserialized_settings = json.loads(team_data["router_settings"])
assert deserialized_settings == router_settings_data
mock_team_update.reset_mock()
team_update_request.router_settings = {"weights": {"group": {"unknown-id": 1}}}
with pytest.raises(ProxyException, match="Unknown deployment ID"):
await update_team(data=team_update_request, http_request=dummy_request, user_api_key_dict=mock_admin_auth)
mock_team_update.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys(

View file

@ -6962,6 +6962,45 @@ class TestModelDeploymentsSupportStreamOptions:
assert self._support(None, None) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("key_settings, expected", [
(None, {"group": {"team": 100}}),
({"weights": {"group": {"key": 100}}}, {"group": {"key": 100}}),
({"timeout": 30}, None),
({"weights": {"group": {"key": "legacy"}}}, None),
])
async def test_saved_weights_override_caller_input_and_preserve_key_precedence(
monkeypatch: pytest.MonkeyPatch,
key_settings: dict[str, int | dict[str, dict[str, int | str]]] | None,
expected: dict[str, dict[str, int]] | None,
) -> None:
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(proxy_server, "get_team_object", AsyncMock(
return_value=SimpleNamespace(router_settings={"weights": {"group": {"team": 100}}})
))
forged = {"group": {"caller": 100}}
processor = ProxyBaseLLMRequestProcessing(data={
"model": "group", "weights": forged, "_router_weights": forged,
"router_settings_override": {"weights": forged},
})
logging = MagicMock(spec=ProxyLogging)
logging.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"])
data, _ = await processor.common_processing_pre_call_logic(
request=Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []}),
general_settings={},
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", team_id="team-a", router_settings=key_settings),
proxy_logging_obj=logging,
proxy_config=proxy_server.ProxyConfig(),
route_type="acompletion",
llm_router=litellm.Router(model_list=[]),
)
assert "weights" not in data
assert data.get("_router_weights") == expected
assert logging.pre_call_hook.call_args.kwargs["data"].get("_router_weights") == expected
class TestPerRequestModelGroupAlias:
"""``router_settings.model_group_alias`` on a key or team has to be resolved
by the proxy: the Router resolves aliases from its own shared instance

View file

@ -957,6 +957,8 @@ async def test_add_litellm_data_to_request_strips_user_control_fields():
"litellm_gateway_injected_cache": "forged-deployment-id",
"metadata": copy.deepcopy(malicious_metadata),
"litellm_metadata": copy.deepcopy(malicious_metadata),
"weights": {"gpt-3.5-turbo": {"forged-deployment-id": 100}},
"_router_weights": {"gpt-3.5-turbo": {"forged-deployment-id": 100}},
}
updated = await add_litellm_data_to_request(
@ -974,6 +976,10 @@ async def test_add_litellm_data_to_request_strips_user_control_fields():
assert "enable_prompt_caching" not in updated
assert "routing_decision" not in updated
assert "litellm_gateway_injected_cache" not in updated
assert "weights" not in updated
assert "_router_weights" not in updated
assert "weights" not in updated["proxy_server_request"]["body"]
assert "_router_weights" not in updated["proxy_server_request"]["body"]
stripped_keys = {
"disable_global_guardrails",

View file

@ -276,3 +276,13 @@ def test_a_server_only_marker_is_not_taken_from_the_caller(field, forged, defaul
auth = UserAPIKeyAuth(api_key="sk-1234", **{field: forged})
assert getattr(auth, field) == default
@pytest.mark.parametrize("weight", [True, "1", -1, 0, float("inf")])
def test_key_and_team_weights_reject_invalid_numeric_values(weight: bool | str | int | float) -> None:
from pydantic import ValidationError
from litellm.proxy._types import GenerateKeyRequest, NewTeamRequest
for request_type in (GenerateKeyRequest, NewTeamRequest):
with pytest.raises(ValidationError):
request_type(router_settings={"weights": {"group": {"id": weight}}})

View file

@ -1,4 +1,5 @@
from collections import Counter
from inspect import isawaitable
import pytest
@ -52,3 +53,52 @@ async def test_uniform_pick_when_every_configured_weight_is_zero():
assert counts["unweighted"] > 0
assert counts["standby"] > 0
@pytest.mark.asyncio
@pytest.mark.parametrize("selector", [
"get_available_deployment", "async_get_available_deployment",
"get_available_deployment_for_pass_through", "async_get_available_deployment_for_pass_through",
])
async def test_scoped_weights_are_request_local_and_respect_eligibility(selector: str) -> None:
router = Router(model_list=[
{
**_deployment(deployment_id, {
"weight": 100 if deployment_id == "global" else 0, "use_in_pass_through": True,
}),
"model_name": f"model_name_{team_id}_{deployment_id}",
"model_info": {
"id": deployment_id, "team_id": team_id, "team_public_model_name": "test-model", "blocked": blocked,
},
}
for deployment_id, team_id, blocked in (
("global", "team-a", False), ("scoped", "team-a", False),
("blocked", "team-a", True), ("foreign", "other-team", False),
)
], num_retries=0)
for weights, expected in (
({"test-model": {"global": 0, "scoped": 100, "blocked": 100, "foreign": 100}}, "scoped"),
({"test-model": {"global": 100, "scoped": 0}}, "global"),
({"test-model": {"foreign": 100}}, "global"),
({"test-model": {"blocked": 100}}, "global"),
(None, "global"),
):
result = getattr(router, selector)(
model="test-model",
request_kwargs={"metadata": {"user_api_key_team_id": "team-a"}, "_router_weights": weights},
)
deployment = await result if isawaitable(result) else result
assert deployment["model_info"]["id"] == expected
def test_scoped_weights_approximate_the_configured_split() -> None:
router = Router(model_list=[_deployment("primary"), _deployment("secondary")], num_retries=0)
counts = Counter(
router.get_available_deployment(
model="test-model",
request_kwargs={"_router_weights": {"test-model": {"primary": 80, "secondary": 20}}},
)["model_info"]["id"]
for _ in range(1000)
)
assert 700 < counts["primary"] < 900

View file

@ -4643,6 +4643,16 @@ def test_aws_bedrock_project_id_excluded_from_bedrock_optional_params():
assert result["aws_region_name"] == "us-east-1"
@pytest.mark.parametrize("filter_name", [
"get_non_default_completion_params", "get_non_default_transcription_params", "filter_out_litellm_params",
])
def test_scoped_weights_are_excluded_from_provider_params(filter_name: str) -> None:
filtered = getattr(litellm.utils, filter_name)(
{"provider_option": "kept", "_router_weights": {"group": {"deployment": 100}}}
)
assert filtered == {"provider_option": "kept"}
class TestGetOptionalParamsTencent:
"""Tests that tencent provider uses TencentChatConfig for parameter mapping."""

View file

@ -32897,9 +32897,7 @@ export interface components {
/** Prompts */
prompts?: string[] | null;
/** Router Settings */
router_settings?: {
[key: string]: unknown;
} | null;
router_settings?: components["schemas"]["UpdateRouterConfig"] | null;
/** Rpm Limit */
rpm_limit?: number | null;
/** Rpm Limit Type */
@ -33650,9 +33648,7 @@ export interface components {
/** Prompts */
prompts?: string[] | null;
/** Router Settings */
router_settings?: {
[key: string]: unknown;
} | null;
router_settings?: components["schemas"]["UpdateRouterConfig"] | null;
/** Rpm Limit */
rpm_limit?: number | null;
/** Secret Manager Settings */
@ -38605,6 +38601,12 @@ export interface components {
tag_routing_prefix?: string | null;
/** Timeout */
timeout?: number | null;
/** Weights */
weights?: {
[key: string]: {
[key: string]: number;
};
} | null;
};
/** UpdateSearchToolRequest */
UpdateSearchToolRequest: {
@ -38702,9 +38704,7 @@ export interface components {
/** Prompts */
prompts?: string[] | null;
/** Router Settings */
router_settings?: {
[key: string]: unknown;
} | null;
router_settings?: components["schemas"]["UpdateRouterConfig"] | null;
/** Rpm Limit */
rpm_limit?: number | null;
/** Secret Manager Settings */