mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(router): honor team and key provider weights
This commit is contained in:
parent
c93708b2a5
commit
398300c4e7
25 changed files with 561 additions and 93 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
129
litellm/proxy/management_endpoints/router_weights.py
Normal file
129
litellm/proxy/management_endpoints/router_weights.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
30
litellm/types/router_weights.py
Normal file
30
litellm/types/router_weights.py
Normal 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
|
||||
|
|
@ -3782,6 +3782,7 @@ all_litellm_params = (
|
|||
"id",
|
||||
"fallbacks",
|
||||
"routing_strategy",
|
||||
"_router_weights",
|
||||
"azure",
|
||||
"headers",
|
||||
"model_list",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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}}})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
18
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
18
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue