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 collections.abc import Callable, Mapping
from datetime import datetime from datetime import datetime
from types import MappingProxyType 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 import httpx
from pydantic import ( from pydantic import (
BaseModel, BaseModel,
BeforeValidator,
ConfigDict, ConfigDict,
Field, Field,
Json, 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.proxy.control_plane_endpoints import WorkerRegistryEntry
from litellm.types.router import RouterErrors, UpdateRouterConfig 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.secret_managers.main import KeyManagementSystem
from litellm.types.utils import ( from litellm.types.utils import (
CallTypes, CallTypes,
@ -1981,8 +1983,14 @@ class OrgMember(MemberBase):
from litellm.models.team import TeamBase as TeamBase # noqa: E402 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): class NewTeamRequest(TeamBase):
router_settings: RouterSettingsDict | None = None
model_aliases: dict | None = None model_aliases: dict | None = None
tags: list | None = None tags: list | None = None
guardrails: list[str] | None = None guardrails: list[str] | None = None
@ -2080,7 +2088,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None
enforced_batch_output_expires_after: dict | None = None enforced_batch_output_expires_after: dict | None = None
enforced_file_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 access_group_ids: list[str] | None = None
budget_limits: list[BudgetLimitEntry] | None = None # multiple concurrent budget windows 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 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: if os.getenv("DATABASE_URL") is not None:
response = await generate_key_helper_fn( response = await generate_key_helper_fn(
llm_router=None,
request_type="key", request_type="key",
**{ **{
"user_role": LitellmUserRoles.PROXY_ADMIN, "user_role": LitellmUserRoles.PROXY_ADMIN,
@ -324,6 +325,7 @@ async def authenticate_user(
await _rehash_password_if_needed(_user_row.user_id, password, _password) await _rehash_password_if_needed(_user_row.user_id, password, _password)
if os.getenv("DATABASE_URL") is not None: if os.getenv("DATABASE_URL") is not None:
response = await generate_key_helper_fn( response = await generate_key_helper_fn(
llm_router=None,
request_type="key", request_type="key",
**{ **{
"user_role": user_role, "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. # the NOT NULL @id constraint. Every successful key-creation caller (e.g.
# /key/generate) passes table_name="key" explicitly. # /key/generate) passes table_name="key" explicitly.
key_data: Final = await generate_key_helper_fn( key_data: Final = await generate_key_helper_fn(
llm_router=None,
request_type="key", request_type="key",
table_name="key", table_name="key",
team_id=team_id, team_id=team_id,

View file

@ -14,6 +14,7 @@ import httpx
import orjson import orjson
from fastapi import HTTPException, Request, status from fastapi import HTTPException, Request, status
from fastapi.responses import JSONResponse, Response, StreamingResponse from fastapi.responses import JSONResponse, Response, StreamingResponse
from pydantic import ValidationError
from starlette.types import Receive, Scope, Send from starlette.types import Receive, Scope, Send
import litellm 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.router_utils.common_utils import resolve_model_group_alias
from litellm.types.guardrails import GuardrailEventHooks from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.router import RouterRateLimitError from litellm.types.router import RouterRateLimitError
from litellm.types.router_weights import validate_router_weights
_LateResponseT = TypeVar("_LateResponseT", bound=Response) _LateResponseT = TypeVar("_LateResponseT", bound=Response)
_LlmCallT = TypeVar("_LlmCallT") _LlmCallT = TypeVar("_LlmCallT")
@ -1939,6 +1941,13 @@ class ProxyBaseLLMRequestProcessing:
# This avoids expensive Router instantiation on each request # This avoids expensive Router instantiation on each request
if router_settings is not None: if router_settings is not None:
self.data["router_settings_override"] = router_settings 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( alias_target: Final = await _resolve_per_request_model_group_alias(
requested_model=self.data.get("model"), requested_model=self.data.get("model"),
router_settings=router_settings, router_settings=router_settings,

View file

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

View file

@ -567,7 +567,7 @@ async def new_user(
teams = check_if_default_team_set() teams = check_if_default_team_set()
organization_ids: Final = cast(list[str] | None, data_json.pop("organizations", None)) 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 # Admin UI Logic
# Add User to Team and Organization # Add User to Team and Organization
# if team_id passed add this user to the team # 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 ( from litellm.proxy.management_endpoints.model_management_endpoints import (
_add_model_to_db, _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 ( from litellm.proxy.management_helpers.access_group_key_sync import (
sync_key_access_group_membership, sync_key_access_group_membership,
sync_key_regeneration_access_group_membership, sync_key_regeneration_access_group_membership,
@ -201,6 +202,10 @@ class _KeyUpdateResult(TypedDict):
data: ReadOnly[Mapping[str, object]] data: ReadOnly[Mapping[str, object]]
class _StoredKeyRouterSettings(BaseModel):
router_settings: Mapping[str, object] | None = None
class _KeyRowWhere(TypedDict): class _KeyRowWhere(TypedDict):
token: ReadOnly[str] token: ReadOnly[str]
@ -1330,7 +1335,7 @@ async def _common_key_generation_helper(
prisma_client=prisma_client, 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 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( async def prepare_key_update_data(
data: UpdateKeyRequest | RegenerateKeyRequest, data: UpdateKeyRequest | RegenerateKeyRequest,
existing_key_row: LiteLLM_VerificationToken, 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: Final[dict] = data.model_dump(exclude_unset=True)
data_json.pop("key", None) data_json.pop("key", None)
data_json.pop("new_key", None) data_json.pop("new_key", None)
@ -2575,7 +2599,9 @@ async def _process_single_key_update(
) )
# Prepare update data # 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 # Update key in database
if prisma_client is None: 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 on update (don't fill defaults)
_enforce_upperbound_key_params(data, fill_defaults=False) _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 # Only validate key_alias format if it's actually being changed
new_key_alias: Final = non_default_values.get("key_alias", None) 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, object_permission: LiteLLM_ObjectPermissionBase | None = None,
auto_rotate: bool | None = None, auto_rotate: bool | None = None,
rotation_interval: str | 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, access_group_ids: list[str] | None = None,
budget_limits: list | None = None, # multiple concurrent budget windows budget_limits: list | None = None, # multiple concurrent budget windows
*,
llm_router: Router | None = None,
): ):
from litellm.proxy.proxy_server import premium_user, prisma_client from litellm.proxy.proxy_server import premium_user, prisma_client
if prisma_client is None: if prisma_client is None:
raise Exception("Connect Proxy to database to generate keys - https://docs.litellm.ai/docs/proxy/virtual_keys ") 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 token is None:
if key is not None: if key is not None:
token = key token = key
@ -5070,6 +5107,7 @@ async def _insert_deprecated_key(
async def _execute_virtual_key_regeneration( async def _execute_virtual_key_regeneration(
*, *,
prisma_client: PrismaClient, prisma_client: PrismaClient,
llm_router: Router | None = None,
key_in_db: LiteLLM_VerificationToken, key_in_db: LiteLLM_VerificationToken,
hashed_api_key: str, hashed_api_key: str,
key: str, key: str,
@ -5129,7 +5167,9 @@ async def _execute_virtual_key_regeneration(
if data is not None: if data is not None:
# Enforce upperbound key params on regenerate (don't fill defaults) # Enforce upperbound key params on regenerate (don't fill defaults)
_enforce_upperbound_key_params(data, fill_defaults=False) _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 # Only validate key_alias format if it's actually being changed
new_key_alias: Final = non_default_values.get("key_alias") new_key_alias: Final = non_default_values.get("key_alias")
if new_key_alias != key_in_db.key_alias: if new_key_alias != key_in_db.key_alias:
@ -5268,6 +5308,7 @@ async def regenerate_key_fn(
try: try:
from litellm.proxy.proxy_server import ( from litellm.proxy.proxy_server import (
hash_token, hash_token,
llm_router,
master_key, master_key,
premium_user, premium_user,
prisma_client, prisma_client,
@ -5456,6 +5497,7 @@ async def regenerate_key_fn(
return await _execute_virtual_key_regeneration( return await _execute_virtual_key_regeneration(
prisma_client=prisma_client, prisma_client=prisma_client,
llm_router=llm_router,
key_in_db=_key_in_db, key_in_db=_key_in_db,
hashed_api_key=hashed_api_key, hashed_api_key=hashed_api_key,
key=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 ( from litellm.proxy.management_endpoints.organization_endpoints import (
add_member_to_organization, 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 ( from litellm.proxy.management_endpoints.tag_management_endpoints import (
get_daily_activity, get_daily_activity,
) )
@ -1288,6 +1289,7 @@ async def new_team(
create_audit_log_for_update, create_audit_log_for_update,
general_settings, general_settings,
litellm_proxy_admin_name, litellm_proxy_admin_name,
llm_router,
prisma_client, prisma_client,
user_api_key_cache, user_api_key_cache,
) )
@ -1462,6 +1464,13 @@ async def new_team(
user_api_key_dict=user_api_key_dict, 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 ## ADD TO MODEL TABLE
_model_id = None _model_id = None
if data.model_aliases is not None and isinstance(data.model_aliases, dict): 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, 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) _existing_team_metadata: Final[object] = getattr(existing_team_row, "metadata", None)
enforce_output_token_estimates_are_admin_only( enforce_output_token_estimates_are_admin_only(
data=data, 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) verbose_proxy_logger.info("user_defined_values for creating ui key: %s", user_defined_values)
response: Final = await generate_key_helper_fn( response: Final = await generate_key_helper_fn(
llm_router=None,
request_type="key", request_type="key",
duration=LITELLM_UI_SESSION_DURATION, duration=LITELLM_UI_SESSION_DURATION,
key_max_budget=litellm.max_ui_session_budget, key_max_budget=litellm.max_ui_session_budget,

View file

@ -9546,6 +9546,7 @@ class ProxyStartupEvent:
gate the first duration window. gate the first duration window.
""" """
await generate_key_helper_fn( await generate_key_helper_fn(
llm_router=llm_router,
request_type="user", request_type="user",
table_name="user", table_name="user",
user_id=LITELLM_PROXY_BUDGET_NAME, 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 global master_key, general_settings
response: Final = await generate_key_helper_fn( response: Final = await generate_key_helper_fn(
llm_router=llm_router,
request_type="key", request_type="key",
**{ **{
"user_role": user_obj.user_role, "user_role": user_obj.user_role,

View file

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

View file

@ -1,71 +1,67 @@
""" """Choose among eligible deployments using request weights, then global metrics."""
Returns a random deployment from the list of healthy deployments.
If weights are provided, it will return a deployment based on the weights. from __future__ import annotations
"""
import logging
import random 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: _DeploymentT = TypeVar("_DeploymentT", bound=Mapping[str, object])
from litellm.router import Router as _Router _ROUTER_LOGGER: Final = logging.getLogger("LiteLLM Router")
LitellmRouter = _Router
else: def _metric_weight(deployment: Mapping[str, object], metric: str) -> float:
LitellmRouter = Any 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( def simple_shuffle(
llm_router_instance: LitellmRouter, resolve_model_alias: Callable[[str], str | None],
healthy_deployments: list[Any] | dict[Any, Any], healthy_deployments: Sequence[_DeploymentT],
model: str, model: str,
) -> dict: request_kwargs: Mapping[str, object] | None,
""" ) -> _DeploymentT:
Returns a random deployment from the list of healthy deployments. resolved_model: Final = resolve_model_alias(model) or model
weight_sets: Final = chain(
If weights are provided, it will return a deployment based on the weights. (_scoped_weights(healthy_deployments, resolved_model, request_kwargs),),
(
If users pass `rpm` or `tpm`, we do a random weighted pick - based on `rpm`/`tpm`. tuple(_metric_weight(deployment, metric) for deployment in healthy_deployments)
for metric in ("weight", "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] for weights in weight_sets:
largest = max(weights, default=0.0)
############## No RPM/TPM passed, we do a random pick ################# if largest <= 0:
item: Final = random.choice(healthy_deployments) continue
return item or item[0] 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._logging import verbose_logger
from litellm._uuid import uuid from litellm._uuid import uuid
from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.types.router_weights import RouterWeights
if TYPE_CHECKING: if TYPE_CHECKING:
from litellm.router import Router from litellm.router import Router
@ -146,6 +147,7 @@ class UpdateRouterConfig(BaseModel):
context_window_fallbacks: list[dict] | None = None context_window_fallbacks: list[dict] | None = None
model_group_alias: dict[str, str | dict] | None = {} model_group_alias: dict[str, str | dict] | None = {}
enable_tag_filtering: bool | None = None enable_tag_filtering: bool | None = None
weights: RouterWeights | None = None
tag_routing_prefix: str | None = None tag_routing_prefix: str | None = None
optional_pre_call_checks: OptionalPreCallChecks | 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", "id",
"fallbacks", "fallbacks",
"routing_strategy", "routing_strategy",
"_router_weights",
"azure", "azure",
"headers", "headers",
"model_list", "model_list",

View file

@ -8,6 +8,10 @@ users can intentionally clear previously-set fields.
""" """
from datetime import datetime, timezone from datetime import datetime, timezone
from types import SimpleNamespace
from fastapi import HTTPException
from litellm import Router
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
@ -1120,3 +1124,41 @@ class TestUpdateMetadataFieldsPremiumCheck:
} }
_update_metadata_fields(updated_kv) _update_metadata_fields(updated_kv)
mock_check.assert_called() 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 typing import Final
from types import SimpleNamespace
import json import json
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
@ -27,6 +28,7 @@ from litellm.proxy._types import (
Member, Member,
ProxyException, ProxyException,
ResetSpendRequest, ResetSpendRequest,
RegenerateKeyRequest,
UpdateKeyRequest, UpdateKeyRequest,
) )
from litellm.proxy.auth.auth_checks import _delete_cache_key_object, _project_cache_key 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=[] return_value=[]
) )
mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) 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) 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", "routing_strategy": "usage-based",
"num_retries": 3, "num_retries": 3,
"model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}}, "model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}},
"weights": {"gpt-4": {"weighted-id": 1}},
} }
request_data = GenerateKeyRequest( 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) # Verify router_settings matches input (regardless of serialization state)
assert actual_settings == router_settings_data 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 @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: Test that /key/update correctly handles router_settings by:
1. Accepting router_settings as a dict parameter 1. Accepting router_settings as a dict parameter
2. Serializing router_settings to JSON when updating database 2. Serializing router_settings to JSON when updating database
3. Updating router_settings in the key record 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 ( from litellm.proxy.management_endpoints.key_management_endpoints import (
prepare_key_update_data, 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 # Mock existing key
existing_key = LiteLLM_VerificationToken( existing_key = LiteLLM_VerificationToken(
token="test-token-router", token="test-token-router",
@ -6710,14 +6732,16 @@ async def test_update_key_with_router_settings(monkeypatch):
router_settings_data = { router_settings_data = {
"routing_strategy": "latency-based", "routing_strategy": "latency-based",
"num_retries": 2, "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 key="test-token-router", router_settings=router_settings_data
) )
result = await prepare_key_update_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 # 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"]) deserialized_settings = json.loads(result["router_settings"])
assert deserialized_settings == router_settings_data 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 @pytest.mark.asyncio
async def test_validate_max_budget(): 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.get_data = AsyncMock(return_value=None)
mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.update_data = AsyncMock(return_value=MagicMock())
mock_db_client.db = 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 model table creation
mock_db_client.db.litellm_modeltable = MagicMock() 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 # Test router_settings with sample data
router_settings_data = { router_settings_data = {
"weights": {"group": {"weighted-id": 1}},
"routing_strategy": "usage-based", "routing_strategy": "usage-based",
"num_retries": 3, "num_retries": 3,
"retry_policy": {"max_retries": 5}, "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"]) deserialized_settings = json.loads(team_data["router_settings"])
assert deserialized_settings == router_settings_data 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 @pytest.mark.asyncio
async def test_get_team_daily_activity_member_with_permission_sees_all_spend( 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 # Configure mocked prisma client
mock_db_client.jsonify_team_object = lambda db_data: db_data mock_db_client.jsonify_team_object = lambda db_data: db_data
mock_db_client.db = 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 existing team row # Mock existing team row
existing_team_mock = MagicMock() existing_team_mock = MagicMock()
@ -9773,6 +9786,7 @@ async def test_update_team_with_router_settings(
# Test router_settings with updated data # Test router_settings with updated data
router_settings_data = { router_settings_data = {
"weights": {"group": {"weighted-id": 1}},
"routing_strategy": "latency-based", "routing_strategy": "latency-based",
"num_retries": 2, "num_retries": 2,
} }
@ -9805,6 +9819,12 @@ async def test_update_team_with_router_settings(
deserialized_settings = json.loads(team_data["router_settings"]) deserialized_settings = json.loads(team_data["router_settings"])
assert deserialized_settings == router_settings_data 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 @pytest.mark.asyncio
async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys( 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 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: class TestPerRequestModelGroupAlias:
"""``router_settings.model_group_alias`` on a key or team has to be resolved """``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 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", "litellm_gateway_injected_cache": "forged-deployment-id",
"metadata": copy.deepcopy(malicious_metadata), "metadata": copy.deepcopy(malicious_metadata),
"litellm_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( 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 "enable_prompt_caching" not in updated
assert "routing_decision" not in updated assert "routing_decision" not in updated
assert "litellm_gateway_injected_cache" 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 = { stripped_keys = {
"disable_global_guardrails", "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}) auth = UserAPIKeyAuth(api_key="sk-1234", **{field: forged})
assert getattr(auth, field) == default 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 collections import Counter
from inspect import isawaitable
import pytest import pytest
@ -52,3 +53,52 @@ async def test_uniform_pick_when_every_configured_weight_is_zero():
assert counts["unweighted"] > 0 assert counts["unweighted"] > 0
assert counts["standby"] > 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" 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: class TestGetOptionalParamsTencent:
"""Tests that tencent provider uses TencentChatConfig for parameter mapping.""" """Tests that tencent provider uses TencentChatConfig for parameter mapping."""

View file

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