fix(proxy): port the member auto-router write path and grant plumbing to stable/1.100.x

Bugbot on #42668 flagged that the backport's picks left the member
auto-router management path unwired on this line. This ports the pieces
that make it work, mirroring main: the member write slot in
model_management_endpoints (FOR UPDATE lock, team reload with the model
table include, identity and name-collision checks, post-commit config
publish), StoredAutoRouterIdentity wiring, the license feature helpers,
team tpd_limit, Router.config_deployments, the member_auto_router
ModelInfo flag, the _TEAM_GRANT_RELATIONS include on team lookups, and
the UserAPIKeyAuth fields the team_grants unpack needs. Test files were
rebuilt as line content plus the picks' own additions, and
ui_sso/test_team_grants carry the pick's grant assertions.

Gate-clearing edits stay local to what the picks added: prisma TypedDict
arguments replace mutable dict literals, remaining dict/mapping
arguments carry reasoned mutable-ok comments, test-quality-ok comments
mark the picks' internal-seam patches, and the regenerated dashboard api
types are staged. The only remaining make check failure is pre-existing
staging drift in untouched tests/test_litellm/test_router.py:2969.
This commit is contained in:
devin-ai-integration[bot] 2026-09-23 06:04:07 +00:00
parent a392f7bc66
commit 54b42e3f05
20 changed files with 1701 additions and 2056 deletions

View file

@ -71,6 +71,7 @@ class TeamBase(LiteLLMPydanticObjectBase):
metadata: dict | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
tpd_limit: int | None = None
max_budget: float | None = None
soft_budget: float | None = None
budget_duration: str | None = None

View file

@ -2779,8 +2779,10 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken):
team_alias: str | None = None
team_tpm_limit: int | None = None
team_rpm_limit: int | None = None
team_tpd_limit: int | None = None
team_max_budget: float | None = None
team_soft_budget: float | None = None
team_model_max_budget: dict[str, object] | None = None
team_models: list = []
team_blocked: bool = False
soft_budget: float | None = None

View file

@ -473,6 +473,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None
_NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({})
_TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True})
def _has_ptu_flat_cost(model: str, llm_router: "Router") -> bool:
@ -549,11 +550,12 @@ def _model_group_has_pricing(model: str, llm_router: "Router") -> bool:
return True
model_id = (deployment.get("model_info") or _EMPTY_COST_ENTRY).get("id")
if model_id is None:
if not isinstance(model_id, str):
continue
model_name = litellm_params.get("model")
model_info = llm_router.get_deployment_model_info(
model_id=model_id, model_name=litellm_params.get("model") or ""
model_id=model_id, model_name=model_name if isinstance(model_name, str) else ""
)
if model_info is not None and _entry_has_priced_metric(model_info):
return True
@ -2847,7 +2849,10 @@ class TeamNotFoundError(HTTPException):
async def _get_team_db_check(
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
) -> "_PrismaTeamRow | None":
response = await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id})
response = await _team_table(TeamRepository(prisma_client)).find_unique(
where={"team_id": team_id}, # mutable-ok: prisma where clause
include=_TEAM_GRANT_RELATIONS,
)
if response is None and team_id_upsert:
from litellm.proxy.management_endpoints.team_endpoints import new_team
@ -3147,7 +3152,10 @@ async def get_team_object_by_alias(
# Query database by team_alias
try:
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(where={"team_alias": team_alias})
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(
where={"team_alias": team_alias}, # mutable-ok: prisma where clause
include=_TEAM_GRANT_RELATIONS,
)
if not teams:
raise HTTPException(

View file

@ -52,6 +52,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.team_grants import team_model_aliases
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
get_management_object_ttl,
@ -1553,7 +1554,9 @@ class JWTAuthManager:
model=requested_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=None,
team_model_aliases=dict(aliases)
if (aliases := team_model_aliases(team_object)) is not None
else None,
)
):
is_allowed = allowed_routes_check(
@ -2090,7 +2093,9 @@ class JWTAuthManager:
model=requested_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=None,
team_model_aliases=dict(aliases)
if (aliases := team_model_aliases(team_object)) is not None
else None,
)
except ProxyException:
continue

View file

@ -15,6 +15,10 @@ from litellm.llms.custom_httpx.http_handler import HTTPHandler
if TYPE_CHECKING:
from litellm.proxy._types import EnterpriseLicenseData
AUTO_ROUTER_LICENSE_FEATURE: Final = "auto_router"
LICENSE_ALL_FEATURES: Final = "*"
AUTO_ROUTER_LICENSE_REMEDY: Final = "A LiteLLM license with the 'auto_router' feature lifts the limit."
class LicenseCheck:
"""
@ -149,6 +153,24 @@ class LicenseCheck:
return False
return team_count > _max_teams_in_license
def grants_feature(self, feature: str) -> bool:
if self.airgapped_license_data is None:
return False
allowed_features: Final = self.airgapped_license_data.get("allowed_features")
granted: Final = allowed_features if isinstance(allowed_features, list) else (allowed_features,)
return feature in granted or LICENSE_ALL_FEATURES in granted
def auto_router_capability_limit(self) -> int | None:
"""
How many auto-routers may claim each gated classifier or customization capability:
unlimited (None) only when the signed license lists the auto_router feature or the
"*" wildcard that grants every feature, otherwise one per capability. A license verified
through the API carries no feature list, so it does not lift the limit either.
"""
if self.grants_feature(AUTO_ROUTER_LICENSE_FEATURE):
return None
return 1
def verify_license_without_api_request(self, public_key, license_key):
try:
from cryptography.hazmat.primitives import hashes

View file

@ -6,7 +6,7 @@ callers kept losing grants (aliases, permissions, limits) one field at a time. B
``team_grants`` and the two paths cannot drift.
"""
from collections.abc import Mapping, Sequence
from collections.abc import Mapping
from types import MappingProxyType
from typing import Annotated, Final
@ -61,10 +61,10 @@ class TeamGrants(TypedDict, total=False):
team_soft_budget: ReadOnly[float | None]
team_model_max_budget: ReadOnly[dict[str, object] | None] # mutable-ok: prisma table field typed loosely
team_spend: ReadOnly[float | None]
team_models: ReadOnly[Sequence[str]]
team_models: ReadOnly[list[str]] # mutable-ok: UserAPIKeyAuth declares a list field
team_blocked: ReadOnly[bool]
team_metadata: ReadOnly[Mapping[str, object] | None]
team_model_aliases: ReadOnly[Mapping[str, str] | None]
team_metadata: ReadOnly[dict[str, object] | None] # mutable-ok: UserAPIKeyAuth declares a dict field
team_model_aliases: ReadOnly[dict[str, str] | None] # mutable-ok: UserAPIKeyAuth declares a dict field
team_object_permission_id: ReadOnly[str | None]
team_object_permission: ReadOnly[LiteLLM_ObjectPermissionTable | None]
team_member: ReadOnly[Member | None]
@ -104,11 +104,14 @@ def team_grants(
team_soft_budget=team_object.soft_budget,
team_model_max_budget=team_object.model_max_budget,
team_spend=team_object.spend,
team_models=tuple(team_object.models),
team_models=list(team_object.models),
team_blocked=team_object.blocked,
team_metadata=json_columns.metadata,
team_metadata=dict(json_columns.metadata) if json_columns.metadata is not None else None,
team_model_aliases=(
json_columns.litellm_model_table.model_aliases if json_columns.litellm_model_table is not None else None
dict(json_columns.litellm_model_table.model_aliases)
if json_columns.litellm_model_table is not None
and json_columns.litellm_model_table.model_aliases is not None
else None
),
team_object_permission_id=team_object.object_permission_id,
team_object_permission=team_object.object_permission,

View file

@ -77,6 +77,7 @@ from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
from litellm.proxy.auth.resolvers import CredentialRef, Principal
from litellm.proxy.auth.resolvers.store import IdentityStore
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.team_grants import team_grants
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
from litellm.proxy.common_utils.http_parsing_utils import (
@ -1469,24 +1470,16 @@ async def _user_api_key_auth_builder(
user_id=user_id,
user_email=user_email,
team_id=team_id,
team_alias=(team_object.team_alias if team_object is not None else None),
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
team_models=(team_object.models if team_object is not None else []),
team_metadata=(team_object.metadata if team_object is not None else None),
org_id=org_id,
end_user_id=end_user_id,
parent_otel_span=parent_otel_span,
jwt_claims=jwt_claims,
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
valid_token = UserAPIKeyAuth(
api_key=None,
team_id=team_id,
team_alias=(team_object.team_alias if team_object is not None else None),
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
team_models=(team_object.models if team_object is not None else []),
user_role=(
LitellmUserRoles(user_object.user_role)
if user_object is not None and user_object.user_role is not None
@ -1500,17 +1493,8 @@ async def _user_api_key_auth_builder(
user_tpm_limit=(user_object.tpm_limit if user_object is not None else None),
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
team_member_rpm_limit=(
team_membership.safe_get_team_member_rpm_limit() if team_membership is not None else None
),
team_member_tpm_limit=(
team_membership.safe_get_team_member_tpm_limit() if team_membership is not None else None
),
team_metadata=(team_object.metadata if team_object is not None else None),
jwt_claims=jwt_claims,
)
valid_token.team_object_permission = (
team_object.object_permission if team_object is not None else None
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.

View file

@ -13,10 +13,13 @@ model/{model_id}/update - PATCH endpoint for model update.
import asyncio
import datetime
import json
from collections.abc import Awaitable, Mapping, Sequence
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from dataclasses import dataclass
from fnmatch import fnmatchcase
from json import JSONDecodeError
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeVar, cast, runtime_checkable
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError, field_validator
@ -54,7 +57,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import (
coordination_redis_cache,
publish_config_change,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin
from litellm.proxy.management_endpoints.team_endpoints import (
@ -66,6 +69,13 @@ from litellm.proxy.management_endpoints.team_endpoints import (
update_team as _legacy_update_team,
)
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
from litellm.proxy.management_helpers.auto_router_permissions import (
MemberAutoRouterWrite,
StoredAutoRouterIdentity,
authorize_member_auto_router_dependencies,
authorize_member_auto_router_team,
authorize_member_auto_router_write,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import (
PTU_COST_ATTRIBUTION_ENV_VAR,
is_ptu_cost_attribution_enabled,
@ -103,11 +113,13 @@ from litellm.types.router import (
GenericLiteLLMParams,
ModelInfo,
updateDeployment,
updateLiteLLMParams,
)
from litellm.utils import get_utc_datetime
if TYPE_CHECKING:
from prisma import models as prisma_models
from prisma import types as prisma_types
router: Final = APIRouter()
@ -155,8 +167,39 @@ class _ProxyModelTable(Protocol):
def delete_many(self, *, where: Mapping[str, object]) -> Awaitable[int]: ...
class _TxTable(Protocol):
def find_unique(
self, *, where: Mapping[str, object], include: Mapping[str, bool] | None = None
) -> Awaitable[BaseModel | None]: ...
class _TxModelTables(Protocol):
litellm_proxymodeltable: _ProxyModelTable
litellm_teamtable: _TxTable
litellm_teammembership: _TxTable
litellm_organizationtable: _TxTable
litellm_projecttable: _TxTable
async def query_raw(self, query: str, *args: object) -> Sequence[Mapping[str, object]]: ...
@runtime_checkable
class _TransactionFactory(Protocol):
def __call__(self, *, timeout: datetime.timedelta = ...) -> AbstractAsyncContextManager[_TxModelTables]: ...
class _ModelTransactionClient(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True, from_attributes=True)
tx: _TransactionFactory
@dataclass(frozen=True, slots=True)
class _TransactionClient:
db: _TxModelTables
_RowT = TypeVar("_RowT")
class _ExistingModelRow(Protocol):
@ -241,6 +284,142 @@ def _effective_complexity_router_config(
}
def _effective_model(
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
) -> str | None:
"""The model a write leaves on the row, decrypting an existing value only when the patch omits it."""
incoming: Final = None if incoming_params is None else incoming_params.model
if incoming is not None:
return incoming
existing: Final = None if existing_params is None else existing_params.model
if existing is None:
return None
decrypted: Final = decrypt_value_helper(
value=existing,
key="model",
exception_type="debug",
return_original_value=True,
)
return decrypted if isinstance(decrypted, str) else None
def _member_auto_router_marker_for_update(
*,
incoming_params: updateLiteLLMParams | None,
existing: Deployment,
member_write: MemberAutoRouterWrite | None,
) -> bool | None:
if member_write is not None:
return True
if not existing.model_info.member_auto_router:
return None
if incoming_params is None:
return True
if any(getattr(incoming_params, field, None) is not None for field in STRATEGY_ROUTER_PARAM_FIELDS):
return False
return incoming_params.model is None or incoming_params.model == _effective_model(None, existing.litellm_params)
@asynccontextmanager
async def _member_auto_router_write_slot(
prisma_client: PrismaClient,
*,
member_write: MemberAutoRouterWrite | None,
) -> AsyncGenerator[_ProxyModelTable, None]:
"""Hand out the model table a member write goes through.
Member writes to a team auto router recheck their authorization inside one
transaction that locks the row first, so two concurrent member writes
cannot both pass the ownership and name checks. Non-member writes keep the
direct table. The transaction write bypasses the repository's
publish-on-write, so the config change is published once after commit.
"""
if member_write is None:
yield _proxy_model_table(prisma_client)
return
import litellm
from litellm.proxy.auth.team_grants import team_model_aliases
from litellm.proxy.proxy_server import llm_router, premium_user
transaction_client: Final = _ModelTransactionClient.model_validate(prisma_client.db)
async with transaction_client.tx(timeout=datetime.timedelta(seconds=30)) as tx_ctx:
tables: Final[_TxModelTables] = tx_ctx
config_rows: Final = () if llm_router is None else tuple(llm_router.config_deployments())
if member_write.model_id is not None:
await tx_ctx.query_raw(
'SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_id = $1 FOR UPDATE',
member_write.model_id,
)
pinned_client: Final = _TransactionClient(tx_ctx)
team_where: Final[prisma_types.LiteLLM_TeamTableWhereUniqueInput] = {"team_id": member_write.team_id}
team_row: Final = await tx_ctx.litellm_teamtable.find_unique(
where=team_where,
include={"litellm_model_table": True}, # mutable-ok: prisma include clause
)
if team_row is None or llm_router is None:
raise HTTPException(status_code=403, detail="The auto router's team or model catalog is unavailable.")
team: Final = LiteLLM_TeamTable.model_validate(team_row.model_dump())
authorize_member_auto_router_team(user_api_key_dict=member_write.actor, team=team, premium_user=premium_user)
if member_write.model_id is not None:
model_where: Final[prisma_types.LiteLLM_ProxyModelTableWhereInput] = {"model_id": member_write.model_id}
current_row: Final = await tables.litellm_proxymodeltable.find_unique(where=model_where)
current_identity: Final = (
StoredAutoRouterIdentity.model_validate(current_row.model_dump()) if current_row is not None else None
)
current_model: Final = (
Deployment.model_validate(current_row.model_dump()) if current_row is not None else None
)
if (
current_identity is None
or current_identity.created_by != member_write.actor.user_id
or current_model is None
or current_model.model_info.team_id != member_write.team_id
):
raise HTTPException(status_code=403, detail="Team members can update only their own auto routers.")
if current_identity.updated_at != member_write.updated_at:
raise HTTPException(status_code=409, detail="This auto router changed. Reload it before updating.")
else:
all_models: Final[prisma_types.LiteLLM_ProxyModelTableWhereInput] = {}
rows_for_names: Final = await tables.litellm_proxymodeltable.find_many(where=all_models)
stored_names: Final = tuple(
(
row.model_name,
model_info_as_mapping(row.model_info),
)
for row in rows_for_names
)
config_names: Final = tuple(
(str(row.get("model_name", "")), model_info_as_mapping(row.get("model_info"))) for row in config_rows
)
team_aliases: Final = team_model_aliases(team)
aliases: Final = (
*(llm_router.model_group_alias or ()),
*(litellm.model_alias_map or ()),
*(team_aliases or ()),
)
if member_write.public_name in aliases or any(
fnmatchcase(
member_write.public_name,
str(info.get("team_public_model_name") or name)
if info is not None and info.get("team_id") == member_write.team_id
else name,
)
for name, info in (*stored_names, *config_names)
if info is None or info.get("team_id") in (None, member_write.team_id)
):
raise HTTPException(status_code=409, detail="This auto-router name is already used by a model.")
await authorize_member_auto_router_dependencies(
config=member_write.config,
default_model=member_write.default_model,
user_api_key_dict=member_write.actor,
team=team,
prisma_client=pinned_client,
llm_router=llm_router,
)
yield tables.litellm_proxymodeltable
await publish_config_change(redis_cache=coordination_redis_cache(), object_type="litellm_proxymodeltable")
def _strategy_router_write_violation(
incoming_params: GenericLiteLLMParams | None,
existing_params: GenericLiteLLMParams | None,
@ -724,11 +903,39 @@ async def patch_model(
param=None,
)
await ModelManagementAuthChecks.can_user_make_model_call(
write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call(
model_params=db_model,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
member_operation="update",
incoming_model_params=patch_data,
)
member_write: Final = write_authorization if isinstance(write_authorization, MemberAutoRouterWrite) else None
member_marker: Final = _member_auto_router_marker_for_update(
incoming_params=patch_data.litellm_params, existing=db_model, member_write=member_write
)
marker_info: Final = (
ModelInfo(id=db_model.model_info.id)
if member_write is not None
else patch_data.model_info or ModelInfo(id=db_model.model_info.id)
)
effective_info: Final = (
marker_info.model_copy(update=MappingProxyType({"member_auto_router": member_marker}))
if member_marker is not None
else patch_data.model_info
)
effective_patch: Final = (
patch_data.model_copy(
update=MappingProxyType(
{
"model_name": None if member_write is not None else patch_data.model_name,
"model_info": effective_info,
}
)
)
if member_marker is not None
else patch_data
)
# Pause/resume (`blocked`) is a proxy-admin-only privilege. Team admins
@ -747,22 +954,23 @@ async def patch_model(
existing_params=db_model.litellm_params,
)
async def write_row(update_data: PrismaCompatibleUpdateDBModel) -> _ProxyModelRow | None:
update_data["updated_by"] = (
user_api_key_dict.user_id or litellm_proxy_admin_name
) # mutable-ok: prisma update payload is dict-shaped
update_data["updated_at"] = cast(
str, get_utc_datetime()
) # mutable-ok: prisma update payload is dict-shaped
async with _member_auto_router_write_slot(prisma_client, member_write=member_write) as table:
return await table.update(where={"model_id": model_id}, data=update_data)
# Handle team model updates with proper alias management
update_data: Final = await _update_team_model_in_db(
updated_model: Final = await _update_team_model_in_db(
db_model=db_model,
patch_data=patch_data,
patch_data=effective_patch,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
# Add metadata about update
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
update_data["updated_at"] = cast(str, get_utc_datetime())
# Perform partial update
updated_model: Final = await _proxy_model_table(prisma_client).update(
where={"model_id": model_id},
data=update_data,
write_row=write_row,
)
if updated_model is None:
@ -993,6 +1201,7 @@ async def _add_model_to_db(
prisma_client: PrismaClient,
new_encryption_key: str | None = None,
should_create_model_in_db: bool = True,
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
# encrypt litellm params #
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
@ -1011,17 +1220,19 @@ async def _add_model_to_db(
if model_params.model_info.id is not None:
_data["model_id"] = model_params.model_info.id
_create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above
if should_create_model_in_db:
model_response = await ModelRepository(prisma_client).table.create(data=_create_data)
else:
model_response = LiteLLM_ProxyModelTable(**_data)
return model_response
if not should_create_model_in_db:
return LiteLLM_ProxyModelTable(**_data)
if slot is None:
return await ModelRepository(prisma_client).table.create(data=_create_data)
async with slot as table:
return await table.create(data=_create_data)
async def _add_team_model_to_db(
model_params: Deployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
"""
If 'team_id' is provided,
@ -1053,6 +1264,7 @@ async def _add_team_model_to_db(
model_params=model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
slot=slot,
)
if original_model_name:
@ -1073,7 +1285,8 @@ async def _update_team_model_in_db(
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
) -> PrismaCompatibleUpdateDBModel:
write_row: Callable[[PrismaCompatibleUpdateDBModel], Awaitable[_RowT]],
) -> _RowT:
"""
Handle team model updates with proper alias management.
@ -1081,6 +1294,9 @@ async def _update_team_model_in_db(
- Creates unique internal model_name and team alias
- Adds model to team object
- Preserves team_public_model_name for external reference
The row is written through ``write_row`` before the team's model list is touched, so a
refused or failed write leaves the team as it was (the create path orders itself the same way).
"""
# Validate team_id if present in patch_data
from litellm.proxy.proxy_server import premium_user
@ -1114,7 +1330,7 @@ async def _update_team_model_in_db(
# No team_id in patch, proceed with standard update
if patch_team_id is None:
return update_db_model(db_model=db_model, updated_patch=patch_data)
return await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
# Determine public model name
public_model_name: Final = _get_public_model_name(
@ -1133,6 +1349,10 @@ async def _update_team_model_in_db(
db_team_id: Final = db_model.model_info.team_id if db_model.model_info else None
is_new_team_assignment: Final = db_team_id != patch_team_id
# Team rows keep their internal UUID-based model_name; the public name lives in model_info
patch_data.model_name = f"model_name_{patch_team_id}_{uuid.uuid4()}" if is_new_team_assignment else None
row: Final = await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
if is_new_team_assignment:
await _setup_new_team_model_assignment(
team_id=patch_team_id,
@ -1150,7 +1370,7 @@ async def _update_team_model_in_db(
prisma_client=prisma_client,
)
return update_db_model(db_model=db_model, updated_patch=patch_data)
return row
def _get_public_model_name(
@ -1313,7 +1533,7 @@ async def _get_team_public_model_names(
model_info = model_info_as_mapping(row.model_info)
if model_info is not None:
public_name = model_info.get("team_public_model_name")
if public_name:
if isinstance(public_name, str) and public_name:
public_names.add(public_name)
return public_names
@ -1546,7 +1766,14 @@ class ModelManagementAuthChecks:
prisma_client: PrismaClient,
premium_user: bool,
allow_missing_team: bool = False,
) -> Literal[True]:
member_operation: Literal["create", "update"] | None = None,
incoming_model_params: updateDeployment | None = None,
) -> Literal[True] | MemberAutoRouterWrite:
if user_api_key_dict.user_role in (
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
):
raise HTTPException(status_code=403, detail="View-only users cannot manage models.")
## Check team model auth
if model_params.model_info is not None and model_params.model_info.team_id is not None:
team_obj_row: Final = await _repo_team_table(prisma_client).find_unique(
@ -1569,6 +1796,27 @@ class ModelManagementAuthChecks:
)
team_obj: Final = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump())
if (
member_operation is not None
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
and not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj)
):
from litellm.proxy.proxy_server import llm_router
if llm_router is None or (member_operation == "update" and incoming_model_params is None):
raise HTTPException(
status_code=400, detail="An auto-router configuration and model catalog are required."
)
return await authorize_member_auto_router_write(
incoming=incoming_model_params if incoming_model_params is not None else model_params,
existing=model_params if member_operation == "update" else None,
user_api_key_dict=user_api_key_dict,
team=team_obj,
premium_user=premium_user,
prisma_client=prisma_client,
llm_router=llm_router,
)
return ModelManagementAuthChecks.can_user_make_team_model_call(
team_id=model_params.model_info.team_id,
user_api_key_dict=user_api_key_dict,
@ -1820,12 +2068,18 @@ async def add_new_model(
)
## Auth check
await ModelManagementAuthChecks.can_user_make_model_call(
write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call(
model_params=model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
member_operation="create",
)
member_write: Final = write_authorization if isinstance(write_authorization, MemberAutoRouterWrite) else None
if member_write is not None and model_params.model_info is not None:
model_params.model_info = model_params.model_info.model_copy( # rebind-ok: downstream team-model handling mutates this same object
update=MappingProxyType({"member_auto_router": True})
)
_raise_on_strategy_router_write_violation(
incoming_params=model_params.litellm_params,
@ -1854,17 +2108,20 @@ async def add_new_model(
reload_outcome: ReconcileOutcome = ReconcileOutcome(still_desired=None, live_after=None)
try:
_original_litellm_model_name: Final = model_params.model_name
add_slot: Final = _member_auto_router_write_slot(prisma_client, member_write=member_write)
if model_params.model_info.team_id is None:
model_response = await _add_model_to_db(
model_params=priced_model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
slot=add_slot,
)
else:
model_response = await _add_team_model_to_db(
model_params=priced_model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
slot=add_slot,
)
reload_outcome = await proxy_config.add_deployment(
prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
@ -1992,11 +2249,17 @@ async def update_model(
raise Exception("model not found")
deployment: Final = Deployment(**_existing_litellm_params.model_dump())
await ModelManagementAuthChecks.can_user_make_model_call(
write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call(
model_params=deployment,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
member_operation="update",
incoming_model_params=model_params,
)
member_write: Final = write_authorization if isinstance(write_authorization, MemberAutoRouterWrite) else None
member_marker: Final = _member_auto_router_marker_for_update(
incoming_params=model_params.litellm_params, existing=deployment, member_write=member_write
)
_raise_on_strategy_router_write_violation(
@ -2038,14 +2301,24 @@ async def update_model(
if value is not None or _existing_litellm_params_dict.get(key) is not None
}
_data: Final[dict[str, str]] = {
_data: Final[dict[str, str]] = { # mutable-ok: prisma update payload is dict-shaped
"litellm_params": json.dumps(merged_dictionary),
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
**(
{
"model_info": deployment.model_info.model_copy(
update=MappingProxyType({"member_auto_router": member_marker})
).model_dump_json(exclude_none=True)
}
if member_marker is not None
else {}
),
}
model_response: Final = await _proxy_model_table(prisma_client).update(
where={"model_id": _model_id},
data=_data,
)
async with _member_auto_router_write_slot(prisma_client, member_write=member_write) as update_table:
model_response: Final = await update_table.update(
where={"model_id": _model_id},
data=_data,
)
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
live_before_reload: Final = live_model_ids_snapshot()

View file

@ -4364,7 +4364,7 @@ async def team_info(
try:
team_info: BaseModel | None = await _team_db(prisma_client).find_unique(
where={"team_id": team_id},
include={"litellm_model_table": True, "object_permission": True},
include={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause
)
if team_info is None:
raise Exception
@ -5567,7 +5567,7 @@ async def team_model_add(
updated_team: Final = await _team_db(prisma_client).update(
where={"team_id": data.team_id},
data={"updated_at": datetime.now(timezone.utc)},
include={"object_permission": True},
include={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause
)
if updated_team is None:
raise HTTPException(
@ -5654,7 +5654,7 @@ async def team_model_delete(
updated_team: Final = await _team_db(prisma_client).update(
where={"team_id": data.team_id},
data={"models": updated_models},
include={"object_permission": True},
include={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause
)
if updated_team is None:
raise HTTPException(

View file

@ -22,7 +22,6 @@ from html import escape
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
Annotated,
Any,
Final,
Literal,
@ -41,7 +40,7 @@ if TYPE_CHECKING:
import jwt
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
from fastapi.responses import RedirectResponse
from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError
from pydantic import BaseModel, TypeAdapter, ValidationError
import litellm
from litellm._logging import verbose_proxy_logger
@ -92,6 +91,7 @@ from litellm.proxy.auth.auth_utils import (
_has_user_setup_sso,
)
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.team_grants import TeamModelAliasTable
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.admin_ui_utils import (
admin_ui_disabled,
@ -202,31 +202,14 @@ def _team_detail_db(repo: TeamRepository) -> "TableActions[_TeamDetailRow]":
return repo.table
_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def _decode_model_aliases(value: object) -> object:
"""``/team/new`` stores team model aliases as a JSON-encoded string in the Json column."""
if not isinstance(value, str):
return value
try:
return _MODEL_ALIASES_ADAPTER.validate_json(value)
except ValidationError:
return None
class _TeamModelAliasTable(BaseModel):
model_config = ConfigDict(protected_namespaces=())
model_aliases: Annotated[Mapping[str, str] | None, BeforeValidator(_decode_model_aliases)] = None
class _TeamRowGrants(BaseModel):
team_id: str
team_alias: str | None = None
models: tuple[str, ...] = ()
litellm_model_table: _TeamModelAliasTable | None = None
litellm_model_table: TeamModelAliasTable | None = None
class CliSsoTeamDetail(BaseModel):

View file

@ -21,7 +21,7 @@ import time
import traceback
import weakref
from collections import defaultdict
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Mapping, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Iterator, Mapping, Sequence
from functools import lru_cache
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
@ -9446,6 +9446,15 @@ class Router:
if _budget_limiter is not None:
_budget_limiter.register_deployment_budget(deployment=deployment.to_json(exclude_none=True))
def config_deployments(self) -> Iterator[Mapping[str, object]]:
"""The model_list rows that came from config.yaml rather than the DB (``model_info.db_model`` unset)."""
for deployment in self.model_list:
if not isinstance(deployment, Mapping):
continue
model_info = deployment.get("model_info")
if not (isinstance(model_info, Mapping) and model_info.get("db_model")):
yield deployment
def get_deployment(self, model_id: str) -> Deployment | None:
"""
Returns -> Deployment or None

View file

@ -922,7 +922,9 @@ def _is_classifier_timeout(exc: BaseException) -> bool:
# LiteLLM still supports 3.10, where they are distinct exception classes.
if isinstance(exc, (TimeoutError, asyncio.TimeoutError)):
return True
return type(exc).__name__.endswith("TimeoutError")
from litellm.exceptions import Timeout as LiteLLMTimeout
return isinstance(exc, LiteLLMTimeout)
class _SessionAffinityPin(NamedTuple):

View file

@ -160,6 +160,7 @@ class ModelInfo(MirroredPricingParams):
# the model_name that can be used by the team when making LLM calls
team_public_model_name: str | None = None
member_auto_router: bool = False
# admin-toggled pause flag; mirrors LiteLLM_ProxyModelTable.blocked
blocked: bool | None = None

File diff suppressed because it is too large Load diff

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,129 @@
import pytest
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
LiteLLM_VerificationTokenView,
Member,
UserAPIKeyAuth,
)
from litellm.models.team import LiteLLM_ModelTable
from litellm.proxy.auth.team_grants import team_grants, team_model_aliases
TEAM_ID = "team-grants"
USER_ID = "user-in-team"
ALIASES = {"fast": "gpt-4o-mini", "smart": "gpt-4o"}
def _alias_table(model_aliases) -> LiteLLM_ModelTable:
return LiteLLM_ModelTable(model_aliases=model_aliases, created_by="admin", updated_by="admin")
def _full_team(model_aliases=ALIASES) -> LiteLLM_TeamTable:
return LiteLLM_TeamTable(
team_id=TEAM_ID,
team_alias="grants-team",
tpm_limit=1000,
rpm_limit=10,
max_budget=50.0,
soft_budget=25.0,
spend=12.5,
models=["gpt-4o", "gpt-4o-mini"],
blocked=True,
metadata={"tier": "gold"},
litellm_model_table=_alias_table(model_aliases),
object_permission_id="op-1",
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-1", mcp_servers=["mcp-a"]),
members_with_roles=[
Member(user_id="someone-else", role="user"),
Member(user_id=USER_ID, role="admin"),
],
)
def _membership() -> LiteLLM_TeamMembership:
return LiteLLM_TeamMembership(
user_id=USER_ID,
team_id=TEAM_ID,
spend=3.25,
litellm_budget_table=LiteLLM_BudgetTable(tpm_limit=500, rpm_limit=5),
)
def test_team_grants_cover_every_team_field_the_key_path_gets():
"""Class guard for LIT-5858 and its siblings: every ``team_*`` column the combined-view SQL hands the
virtual-key path must come out of the projection too, with the team's actual value, so adding a column
to ``LiteLLM_VerificationTokenView`` without teaching ``team_grants`` fails here instead of in prod."""
team = _full_team()
grants = team_grants(team_object=team, team_membership=_membership(), user_id=USER_ID)
token = UserAPIKeyAuth(team_id=TEAM_ID, **grants)
view_team_fields = {name for name in LiteLLM_VerificationTokenView.model_fields if name.startswith("team_")}
assert view_team_fields - {"team_id"} <= set(grants)
assert all(grants[name] is not None for name in view_team_fields - {"team_id"})
assert token.team_alias == "grants-team"
assert token.team_tpm_limit == 1000
assert token.team_rpm_limit == 10
assert token.team_max_budget == 50.0
assert token.team_soft_budget == 25.0
assert token.team_spend == 12.5
assert token.team_models == ["gpt-4o", "gpt-4o-mini"]
assert token.team_blocked is True
assert token.team_metadata == {"tier": "gold"}
assert token.team_model_aliases == ALIASES
assert token.team_object_permission_id == "op-1"
assert token.team_object_permission is not None
assert token.team_object_permission.mcp_servers == ["mcp-a"]
assert token.team_member == Member(user_id=USER_ID, role="admin")
assert token.team_member_spend == 3.25
assert token.team_member_tpm_limit == 500
assert token.team_member_rpm_limit == 5
def test_team_grants_without_team_leave_token_defaults():
token = UserAPIKeyAuth(**team_grants(team_object=None, team_membership=None, user_id=USER_ID))
assert token == UserAPIKeyAuth()
@pytest.mark.parametrize(
"stored_aliases",
[ALIASES, '{"fast": "gpt-4o-mini", "smart": "gpt-4o"}'],
ids=["json-object", "json-string-as-written-by-team-new"],
)
def test_team_model_aliases_decode_both_storage_shapes(stored_aliases):
team = _full_team(model_aliases=stored_aliases)
assert team_model_aliases(team) == ALIASES
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_model_aliases"] == ALIASES
@pytest.mark.parametrize("stored_aliases", [None, "not json", '["a", "b"]', {"fast": 3}], ids=str)
def test_team_model_aliases_treat_unusable_column_as_no_aliases(stored_aliases):
team = _full_team(model_aliases=stored_aliases)
assert team_model_aliases(team) is None
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_model_aliases"] is None
def test_team_model_aliases_none_without_relation_loaded():
team = _full_team()
team.litellm_model_table = None
assert team_model_aliases(team) is None
assert team_model_aliases(None) is None
def test_team_member_is_the_callers_row_only():
team = _full_team()
assert team_grants(team_object=team, team_membership=None, user_id="someone-else")["team_member"] == Member(
user_id="someone-else", role="user"
)
assert team_grants(team_object=team, team_membership=None, user_id="stranger")["team_member"] is None
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_member"] is None
def test_membership_limits_absent_without_membership_row():
grants = team_grants(team_object=_full_team(), team_membership=None, user_id=USER_ID)
assert grants["team_member_spend"] is None
assert grants["team_member_tpm_limit"] is None
assert grants["team_member_rpm_limit"] is None

View file

@ -187,7 +187,7 @@ async def test_budget_reservation_runs_when_not_disabled():
)
async def test_fail_closed_budget_enforcement_reaches_reservation(
general_settings, expected_flag
):
): # test-quality-ok: [TQ002] collaborator injected via its import site; there is no seam to patch otherwise
"""#33923: the strict flag must be threaded into reserve_budget_for_request so a
failed reservation write can reject instead of failing open."""
user_api_key_auth_obj = UserAPIKeyAuth(token="test_token")
@ -210,10 +210,7 @@ async def test_fail_closed_budget_enforcement_reaches_reservation(
general_settings=general_settings,
)
assert (
mock_reserve.await_args.kwargs["fail_closed_budget_enforcement"]
is expected_flag
)
assert mock_reserve.await_args.kwargs["fail_closed_budget_enforcement"] is expected_flag
@pytest.mark.asyncio
@ -227,7 +224,7 @@ async def test_fail_closed_budget_enforcement_reaches_reservation(
)
async def test_apply_user_budget_to_team_keys_reaches_reservation(
general_settings, expected_flag
):
): # test-quality-ok: [TQ002] collaborator injected via its import site; there is no seam to patch otherwise
"""The opt-in lives in general_settings but is consumed inside
_get_budget_counters, so it has to be threaded through reserve_budget_for_request
or the reservation path keeps exempting team keys while the read path enforces."""
@ -251,9 +248,7 @@ async def test_apply_user_budget_to_team_keys_reaches_reservation(
general_settings=general_settings,
)
assert (
mock_reserve.await_args.kwargs["apply_user_budget_to_team_keys"] is expected_flag
)
assert mock_reserve.await_args.kwargs["apply_user_budget_to_team_keys"] is expected_flag
@pytest.mark.asyncio
@ -353,9 +348,7 @@ async def test_custom_auth_honors_key_level_model_access_restriction_allowed_wit
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
new_callable=AsyncMock,
) as mock_can_key,
patch(
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
),
patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock),
patch(
"litellm.proxy.proxy_server.general_settings",
{"custom_auth_run_common_checks": True},
@ -386,9 +379,7 @@ async def test_custom_auth_enforces_key_model_access_from_file_route_header_with
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
new_callable=AsyncMock,
) as mock_can_key,
patch(
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
),
patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock),
patch(
"litellm.proxy.proxy_server.general_settings",
{"custom_auth_run_common_checks": True},
@ -419,9 +410,7 @@ async def test_custom_auth_honors_key_level_model_access_restriction_denied_with
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
new_callable=AsyncMock,
) as mock_can_key,
patch(
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
),
patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock),
patch(
"litellm.proxy.proxy_server.general_settings",
{"custom_auth_run_common_checks": True},
@ -457,9 +446,7 @@ def _proxy_server_attrs_for_custom_auth(*, user_custom_auth):
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
return {
@ -721,9 +708,7 @@ async def test_enterprise_custom_auth_runs_post_custom_auth_checks_when_opt_in()
litellm.enable_post_custom_auth_checks = original_flag
def _assert_get_api_key_with_custom_litellm_key_header(
custom_litellm_key_header, api_key, passed_in_key
):
def _assert_get_api_key_with_custom_litellm_key_header(custom_litellm_key_header, api_key, passed_in_key):
assert get_api_key(
custom_litellm_key_header=custom_litellm_key_header,
api_key=None,
@ -780,9 +765,7 @@ def _assert_get_api_key_with_custom_litellm_key_header(
("App:LiteLLM", None, False, False),
],
)
def test_routing_selector_matches_claim_parametrized(
selector_value, claim_value, expected, split_space_delimited
):
def test_routing_selector_matches_claim_parametrized(selector_value, claim_value, expected, split_space_delimited):
assert (
_routing_selector_matches_claim(
selector_value=selector_value,
@ -876,10 +859,7 @@ def test_routing_selector_matches_claim_parametrized(
],
)
def test_matches_routing_override_parametrized(override, token_claims, expected):
assert (
_matches_routing_override(token_claims=token_claims, override=override)
is expected
)
assert _matches_routing_override(token_claims=token_claims, override=override) is expected
def test_get_api_key_with_custom_litellm_key_header_bearer_prefix():
@ -958,12 +938,9 @@ def test_team_metadata_with_tags_flows_through_jwt_auth():
)
# Verify team_metadata is set
assert (
user_api_key_auth.team_metadata is not None
), "team_metadata should be populated"
assert user_api_key_auth.team_metadata is not None, "team_metadata should be populated"
assert user_api_key_auth.team_metadata == team_object.metadata, (
f"team_metadata not correctly mapped. "
f"Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}"
f"team_metadata not correctly mapped. Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}"
)
# Specifically verify tags are present
@ -1002,9 +979,7 @@ def test_route_checks_is_llm_api_route():
]
for route in openai_routes:
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
# Test Anthropic routes
anthropic_routes = [
@ -1013,9 +988,7 @@ def test_route_checks_is_llm_api_route():
]
for route in anthropic_routes:
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
# Test passthrough routes (this is the key improvement over the old route checking)
passthrough_routes = [
@ -1035,9 +1008,7 @@ def test_route_checks_is_llm_api_route():
]
for route in passthrough_routes:
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
# Test MCP routes
mcp_routes = [
@ -1047,9 +1018,7 @@ def test_route_checks_is_llm_api_route():
]
for route in mcp_routes:
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
# Test LiteLLM native RAG routes
rag_routes = [
@ -1059,9 +1028,7 @@ def test_route_checks_is_llm_api_route():
"/v1/rag/query",
]
for route in rag_routes:
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
# Test routes with placeholders
placeholder_routes = [
@ -1076,9 +1043,7 @@ def test_route_checks_is_llm_api_route():
]
for route in placeholder_routes:
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
# Test Azure OpenAI routes
azure_routes = [
@ -1089,9 +1054,7 @@ def test_route_checks_is_llm_api_route():
]
for route in azure_routes:
assert RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should be identified as LLM API route"
assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route"
# Test non-LLM routes (should return False)
non_llm_routes = [
@ -1110,9 +1073,7 @@ def test_route_checks_is_llm_api_route():
]
for route in non_llm_routes:
assert not RouteChecks.is_llm_api_route(
route=route
), f"Route {route} should NOT be identified as LLM API route"
assert not RouteChecks.is_llm_api_route(route=route), f"Route {route} should NOT be identified as LLM API route"
# Test invalid inputs
invalid_inputs = [
@ -1124,9 +1085,9 @@ def test_route_checks_is_llm_api_route():
]
for invalid_input in invalid_inputs:
assert not RouteChecks.is_llm_api_route(
route=invalid_input
), f"Invalid input {invalid_input} should return False"
assert not RouteChecks.is_llm_api_route(route=invalid_input), (
f"Invalid input {invalid_input} should return False"
)
@pytest.mark.asyncio
@ -1173,9 +1134,7 @@ async def test_proxy_admin_expired_key_from_cache():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
# Mock post_call_failure_hook as async function returning None (no transformation)
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
@ -1212,9 +1171,7 @@ async def test_proxy_admin_expired_key_from_cache():
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
@ -1238,36 +1195,30 @@ async def test_proxy_admin_expired_key_from_cache():
)
# Verify that ProxyException was raised with expired_key type
assert hasattr(
exc_info.value, "type"
), "Exception should have 'type' attribute"
assert (
exc_info.value.type == ProxyErrorTypes.expired_key
), f"Expected expired_key error type, got {exc_info.value.type}"
assert hasattr(exc_info.value, "type"), "Exception should have 'type' attribute"
assert exc_info.value.type == ProxyErrorTypes.expired_key, (
f"Expected expired_key error type, got {exc_info.value.type}"
)
assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED
assert "Expired Key" in str(
exc_info.value.message
), f"Exception message should mention 'Expired Key', got: {exc_info.value.message}"
assert "Expired Key" in str(exc_info.value.message), (
f"Exception message should mention 'Expired Key', got: {exc_info.value.message}"
)
# Verify that the param field does NOT leak the full API key (Issue #18731)
# The param should be abbreviated like "sk-...XXXX" not the full plaintext key
assert (
exc_info.value.param is not None
), "Exception should have 'param' attribute"
assert exc_info.value.param is not None, "Exception should have 'param' attribute"
assert exc_info.value.param != api_key, (
f"SECURITY: Full API key should NOT be in param field! "
f"Got: {exc_info.value.param}, Expected abbreviated format like 'sk-...XXXX'"
)
assert exc_info.value.param.startswith(
"sk-..."
), f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}"
assert exc_info.value.param.startswith("sk-..."), (
f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}"
)
# Verify that cache deletion was called
mock_delete_cache.assert_called_once()
call_args = mock_delete_cache.call_args
assert (
call_args[1]["hashed_token"] == hashed_key
), "Cache deletion should be called with the hashed key"
assert call_args[1]["hashed_token"] == hashed_key, "Cache deletion should be called with the hashed key"
finally:
# Restore all module-level attributes so subsequent tests are not affected
for attr, val in _original_values.items():
@ -1305,9 +1256,7 @@ async def test_scim_deactivated_user_key_is_rejected():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
mock_prisma_client = MagicMock()
@ -1328,9 +1277,7 @@ async def test_scim_deactivated_user_key_is_rejected():
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
@ -1397,9 +1344,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
@ -1418,9 +1363,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker():
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
@ -1472,9 +1415,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
@ -1493,9 +1434,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker():
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
@ -1548,9 +1487,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
mock_prisma_client = MagicMock()
@ -1571,9 +1508,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker():
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
_original_values = {attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
@ -1993,10 +1928,7 @@ class TestJWTOAuth2Coexistence:
def test_is_jwt_detects_jwt_tokens(self):
"""JWT tokens have 3 dot-separated parts."""
assert JWTHandler.is_jwt("header.payload.signature") is True
assert (
JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123")
is True
)
assert JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123") is True
def test_is_jwt_rejects_opaque_tokens(self):
"""Opaque OAuth2 tokens do not have 3 dot-separated parts."""
@ -2105,10 +2037,7 @@ class TestJWTOAuth2Coexistence:
assert exc_info.value.type == ProxyErrorTypes.auth_error
assert exc_info.value.code == "403"
assert (
"Oauth2 token validation is only available for premium users"
in exc_info.value.message
)
assert "Oauth2 token validation is only available for premium users" in exc_info.value.message
mock_oauth2.assert_not_called()
@pytest.mark.asyncio
@ -2300,9 +2229,7 @@ class TestJWTOAuth2Coexistence:
assert mock_auto_register.call_args.kwargs["team_id"] == "validated-team"
assert mock_auto_register.call_args.kwargs["user_id"] == "validated-user"
assert mock_auto_register.call_args.kwargs["org_id"] == "validated-org"
assert (
mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user"
)
assert mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user"
assert result.org_id == "validated-org"
assert result.user_email == "validated@example.com"
@ -2380,10 +2307,7 @@ class TestJWTOAuth2Coexistence:
assert result.user_id == "mapped-user"
assert result.user_email == "mapped@example.com"
assert (
mock_get_user_object.call_args_list[0].kwargs["user_email"]
== "mapped@example.com"
)
assert mock_get_user_object.call_args_list[0].kwargs["user_email"] == "mapped@example.com"
@pytest.mark.asyncio
async def test_mapped_virtual_key_does_not_backfill_mismatched_owner(self):
@ -2459,8 +2383,7 @@ class TestJWTOAuth2Coexistence:
assert result.user_id == "other-owner"
assert result.user_email is None
assert all(
call.kwargs.get("user_email") != "principal@example.com"
for call in mock_get_user_object.call_args_list
call.kwargs.get("user_email") != "principal@example.com" for call in mock_get_user_object.call_args_list
)
@pytest.mark.asyncio
@ -3265,9 +3188,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
@ -3399,9 +3320,7 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
@ -3451,9 +3370,9 @@ async def test_team_metadata_refreshed_from_team_object_during_auth():
request_data={},
)
assert result.team_metadata == {
"guardrails": ["test-guardrail-333"]
}, f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}"
assert result.team_metadata == {"guardrails": ["test-guardrail-333"]}, (
f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}"
)
finally:
for k, v in _originals.items():
@ -3778,9 +3697,7 @@ async def test_auth_flow_fallback_team_object_permission_none_when_unreadable():
# ---------------------------------------------------------------------------
def _proxy_attrs_for_centralized_checks(
user_custom_auth=None, flag=False, master_key="sk-test-master"
):
def _proxy_attrs_for_centralized_checks(user_custom_auth=None, flag=False, master_key="sk-test-master"):
"""Build the minimal proxy_server module attributes that
_run_centralized_common_checks reads.
@ -3909,9 +3826,7 @@ async def test_centralized_common_checks_skipped_for_custom_auth_without_flag():
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
attrs = _proxy_attrs_for_centralized_checks(
user_custom_auth=AsyncMock(), flag=False
)
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=AsyncMock(), flag=False)
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
try:
for k, v in attrs.items():
@ -4316,9 +4231,7 @@ async def test_centralized_common_checks_reserves_request_end_user_budget():
"applied_adjustment": 0.0,
}
]
assert counter_cache.in_memory_cache.get_cache(
key="spend:end_user:alice"
) == pytest.approx(0.6)
assert counter_cache.in_memory_cache.get_cache(key="spend:end_user:alice") == pytest.approx(0.6)
@pytest.mark.asyncio
@ -4333,9 +4246,7 @@ async def test_centralized_common_checks_short_circuits_when_master_key_unset():
from litellm.proxy._types import LitellmUserRoles
token = UserAPIKeyAuth(
api_key="sk-test", user_id="u", user_role=LitellmUserRoles.INTERNAL_USER
)
token = UserAPIKeyAuth(api_key="sk-test", user_id="u", user_role=LitellmUserRoles.INTERNAL_USER)
request = Request(scope={"type": "http"})
request._url = URL(url="/get/config/callbacks")
@ -5136,9 +5047,7 @@ async def test_centralized_common_checks_user_http_exception_isolates_to_user_on
request._url = URL(url="/chat/completions")
request._body = json.dumps({"user": "alice", "model": "gpt-4o"}).encode()
fetched_team = LiteLLM_TeamTableCachedObj(
team_id="t1", max_budget=20.0, models=["gpt-4o"]
)
fetched_team = LiteLLM_TeamTableCachedObj(team_id="t1", max_budget=20.0, models=["gpt-4o"])
fetched_end_user = LiteLLM_EndUserTable(user_id="alice", blocked=False, spend=1.0)
fetched_project = LiteLLM_ProjectTableCachedObj(
project_id="proj-1",
@ -5433,9 +5342,7 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it():
}
)
request._url = URL(url="/chat/completions")
request._body = json.dumps(
{"model": "gpt-4o", "user": "alice@example.com"}
).encode()
request._body = json.dumps({"model": "gpt-4o", "user": "alice@example.com"}).encode()
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
@ -5479,9 +5386,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder()
import litellm.proxy.proxy_server as _proxy_server_mod
builder_token = UserAPIKeyAuth(
api_key="sk-test", user_id="u1", end_user_id="builder-resolved-id"
)
builder_token = UserAPIKeyAuth(api_key="sk-test", user_id="u1", end_user_id="builder-resolved-id")
request = Request(
scope={
@ -5491,9 +5396,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder()
}
)
request._url = URL(url="/chat/completions")
request._body = json.dumps(
{"model": "gpt-4o", "user": "different-id-from-body"}
).encode()
request._body = json.dumps({"model": "gpt-4o", "user": "different-id-from-body"}).encode()
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
@ -5891,9 +5794,7 @@ def _mint_cli_session_token(monkeypatch, *, user_id="cli-admin"):
models=["gpt-3.5-turbo"],
max_budget=100.0,
)
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(
user_info, team_id="cli-team", team_alias="cli-team-alias"
)
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info, team_id="cli-team", team_alias="cli-team-alias")
@pytest.mark.asyncio
@ -5943,7 +5844,7 @@ async def test_random_non_sk_token_is_rejected(monkeypatch):
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
):
with pytest.raises(Exception, match='LiteLLM Virtual Key expected\\.') as exc_info:
with pytest.raises(Exception, match="LiteLLM Virtual Key expected\\.") as exc_info:
await user_api_key_auth(
request=mock_request,
api_key="Bearer not-a-real-token",
@ -6022,9 +5923,7 @@ async def test_non_admin_cli_session_token_reaches_production_auth_path(monkeypa
user_role=LitellmUserRoles.INTERNAL_USER.value,
models=[],
)
cli_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
user_info, team_id="team-abc", team_alias="my-team"
)
cli_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info, team_id="team-abc", team_alias="my-team")
import litellm.proxy.proxy_server as _proxy_server_mod
from fastapi import Request
@ -6145,7 +6044,7 @@ async def test_real_jwt_still_requires_license_when_jwt_auth_enabled(monkeypatch
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.prisma_client", None),
):
with pytest.raises(Exception, match='JWT Auth is an enterprise only feature\\. You must be a') as exc_info:
with pytest.raises(Exception, match="JWT Auth is an enterprise only feature\\. You must be a") as exc_info:
await user_api_key_auth(
request=mock_request,
api_key=f"Bearer {jwt_token}",
@ -6184,13 +6083,9 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache():
metadata={"model_rpm_limit": {"gpt-5.4-mini": 3}},
last_refreshed_at=1000.0,
)
await key_cache.async_set_cache(
key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth
)
await key_cache.async_set_cache(key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth)
fetch_from_db = AsyncMock(
side_effect=AssertionError("cache-hit auth must not touch the DB")
)
fetch_from_db = AsyncMock(side_effect=AssertionError("cache-hit auth must not touch the DB"))
proxy_logging_obj = MagicMock()
proxy_logging_obj.internal_usage_cache = MagicMock()
@ -6237,9 +6132,7 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache():
assert result.token == hashed_key
fetch_from_db.assert_not_called()
cached_after = await key_cache.async_get_cache(
key=hashed_key, model_type=UserAPIKeyAuth
)
cached_after = await key_cache.async_get_cache(key=hashed_key, model_type=UserAPIKeyAuth)
assert cached_after is not None
assert cached_after.last_refreshed_at == 1000.0
assert cached_after.metadata == {"model_rpm_limit": {"gpt-5.4-mini": 3}}
@ -6352,9 +6245,7 @@ class TestCheckKeyModelBudgetWithFallback:
@pytest.mark.asyncio
async def test_within_budget_does_not_reroute(self):
valid_token = UserAPIKeyAuth(
token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]}
)
valid_token = UserAPIKeyAuth(token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]})
limiter = AsyncMock()
limiter.is_key_within_model_budget.return_value = True
request_data = {"model": "gpt-4o"}
@ -6379,9 +6270,7 @@ class TestCheckKeyModelBudgetWithFallback:
budget_fallbacks={"gpt-4o": ["gpt-4o-mini", "claude-haiku"]},
)
limiter = AsyncMock()
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(
current_cost=10, max_budget=5
)
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5)
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
request_data = {"model": "gpt-4o"}
request = self._make_request()
@ -6395,9 +6284,7 @@ class TestCheckKeyModelBudgetWithFallback:
)
assert request_data["model"] == "gpt-4o-mini"
limiter.get_fallback_model_within_budget.assert_awaited_once_with(
user_api_key_dict=valid_token, model="gpt-4o"
)
limiter.get_fallback_model_within_budget.assert_awaited_once_with(user_api_key_dict=valid_token, model="gpt-4o")
# the rerouted model must be visible to a later, separate
# `_read_request_body` call on the same `request` (route handlers
# re-parse the body from this cache instead of reusing the dict).
@ -6406,9 +6293,7 @@ class TestCheckKeyModelBudgetWithFallback:
@pytest.mark.asyncio
async def test_raises_when_every_fallback_also_exceeded(self):
valid_token = UserAPIKeyAuth(
token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]}
)
valid_token = UserAPIKeyAuth(token="test-key", budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]})
limiter = AsyncMock()
original_error = litellm.BudgetExceededError(current_cost=10, max_budget=5)
limiter.is_key_within_model_budget.side_effect = original_error
@ -6478,9 +6363,7 @@ class TestCheckKeyModelBudgetWithFallback:
budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]},
)
limiter = AsyncMock()
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(
current_cost=10, max_budget=5
)
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5)
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
request_data = {"model": "gpt-4o"}
request = self._make_request()
@ -6548,9 +6431,7 @@ class TestCheckKeyModelBudgetWithFallback:
budget_fallbacks={"gpt-4o": ["gpt-4o-mini"]},
)
limiter = AsyncMock()
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(
current_cost=10, max_budget=5
)
limiter.is_key_within_model_budget.side_effect = litellm.BudgetExceededError(current_cost=10, max_budget=5)
limiter.get_fallback_model_within_budget.return_value = "gpt-4o-mini"
request_data = {"model": "gpt-4o"}
request = self._make_request()
@ -6630,9 +6511,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row():
)
assert result == 42.5
prisma_client.db.litellm_usertable.find_unique.assert_awaited_once_with(
where={"user_id": "litellm-proxy-budget"}
)
prisma_client.db.litellm_usertable.find_unique.assert_awaited_once_with(where={"user_id": "litellm-proxy-budget"})
@pytest.mark.asyncio
@ -6798,9 +6677,7 @@ async def test_jwt_shaped_key_error_names_enable_jwt_auth_when_disabled():
Prometheus invalid-key filter and the admin UI both substring-match it.
Keys that are not JWT-shaped must not pick up the hint.
"""
jwt_error = await _proxy_exception_for_key(
"eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9.c2lnbmF0dXJl", {}, True
)
jwt_error = await _proxy_exception_for_key("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9.c2lnbmF0dXJl", {}, True)
assert jwt_error.code == "401"
assert "enable_jwt_auth" in jwt_error.message
@ -6810,9 +6687,7 @@ async def test_jwt_shaped_key_error_names_enable_jwt_auth_when_disabled():
assert "is a JWT" not in jwt_error.message
opaque_error = await _proxy_exception_for_key("not-a-jwt-at-all", {}, True)
two_segment_error = await _proxy_exception_for_key(
"eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9", {}, True
)
two_segment_error = await _proxy_exception_for_key("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJzdmMtMSJ9", {}, True)
assert "enable_jwt_auth" not in opaque_error.message
assert "enable_jwt_auth" not in two_segment_error.message
@ -6841,9 +6716,7 @@ class TestLitellmReceivedAtStamping:
on OTEL being configured to see a true request-arrival timestamp."""
def test_stamped_even_when_otel_is_not_configured(self, monkeypatch):
monkeypatch.setattr(
"litellm.proxy.proxy_server.open_telemetry_logger", None
)
monkeypatch.setattr("litellm.proxy.proxy_server.open_telemetry_logger", None)
request = MagicMock()
request.state = SimpleNamespace()
@ -6872,3 +6745,119 @@ class TestLitellmReceivedAtStamping:
assert result == earlier
assert request.state.litellm_received_at == earlier
@pytest.mark.asyncio
@pytest.mark.parametrize("is_proxy_admin", [False, True], ids=["standard-return", "proxy-admin-return"])
async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_admin):
"""LIT-5858: the team-based JWT path hand-built ``UserAPIKeyAuth`` from a short list of team fields, so the
team's model aliases (and on the admin return, its object permission) never reached the token and alias
requests 403'd. Both returns now go through ``team_grants``; pin the fields that used to be dropped."""
import litellm.proxy.proxy_server as _proxy_server_mod
from fastapi import Request
from starlette.datastructures import URL
from litellm.models.team import LiteLLM_ModelTable
from litellm.proxy._types import (
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
Member,
)
class _AcceptEveryJwt(JWTHandler):
def is_jwt(self, token: str) -> bool:
return True
jwt_handler = _AcceptEveryJwt()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
team = LiteLLM_TeamTable(
team_id="team-jwt-aliases",
team_alias="jwt-aliases",
models=["gpt-4o"],
max_budget=40.0,
spend=4.0,
blocked=False,
metadata={"tier": "gold"},
litellm_model_table=LiteLLM_ModelTable(
model_aliases='{"fast": "gpt-4o"}', created_by="admin", updated_by="admin"
),
object_permission_id="op-jwt",
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-jwt", mcp_servers=["mcp-a"]),
members_with_roles=[Member(user_id="jwt-user", role="admin")],
)
membership = LiteLLM_TeamMembership(user_id="jwt-user", team_id="team-jwt-aliases", spend=1.5)
builder_result = {
"is_proxy_admin": is_proxy_admin,
"team_object": team,
"user_object": None,
"end_user_object": None,
"org_object": None,
"token": "jwt",
"team_id": "team-jwt-aliases",
"user_id": "jwt-user",
"user_email": "jwt-user@example.com",
"end_user_id": None,
"org_id": None,
"team_membership": membership,
"jwt_claims": {"sub": "jwt-user"},
}
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
attrs = {
"prisma_client": MagicMock(),
"user_api_key_cache": DualCache(),
"proxy_logging_obj": mock_proxy_logging_obj,
"master_key": "sk-master-key",
"general_settings": {"enable_jwt_auth": True},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": jwt_handler,
"premium_user": True,
"litellm_proxy_admin_name": "admin",
}
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
try:
for k, v in attrs.items():
setattr(_proxy_server_mod, k, v)
request = Request(scope={"type": "http", "headers": [], "method": "POST"})
request._url = URL(url="/chat/completions")
with patch( # test-quality-ok: auth_builder is the claim-resolution seam; the regression is how its result is projected onto the token
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
new_callable=AsyncMock,
return_value=builder_result,
):
token = await _user_api_key_auth_builder(
request=request,
api_key="Bearer header.payload.signature",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
finally:
for k, v in originals.items():
setattr(_proxy_server_mod, k, v)
assert token.team_id == "team-jwt-aliases"
assert token.user_role == (LitellmUserRoles.PROXY_ADMIN if is_proxy_admin else LitellmUserRoles.INTERNAL_USER)
assert token.team_model_aliases == {"fast": "gpt-4o"}
assert token.team_object_permission is not None
assert token.team_object_permission.mcp_servers == ["mcp-a"]
assert token.team_object_permission_id == "op-jwt"
assert token.team_alias == "jwt-aliases"
assert token.team_models == ["gpt-4o"]
assert token.team_max_budget == 40.0
assert token.team_spend == 4.0
assert token.team_metadata == {"tier": "gold"}
assert token.team_member == Member(user_id="jwt-user", role="admin")
assert token.team_member_spend == 1.5
assert token.jwt_claims == {"sub": "jwt-user"}

View file

@ -31,6 +31,10 @@ from litellm.router import Router
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo, updateDeployment, updateLiteLLMParams
async def _passthrough_row(update_data):
return update_data
class MockPrismaClient:
def __init__(
self,
@ -1027,7 +1031,7 @@ class TestTeamModelSiblingRouting:
team_id = "team_no_alias"
public_name = "gpt-4.1-mini"
async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client):
async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client, slot=None):
return MagicMock(model_id=str(uuid.uuid4()))
mock_team_model_add = AsyncMock()
@ -1207,6 +1211,7 @@ class TestTeamModelUpdate:
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client, # type: ignore
write_row=_passthrough_row,
)
assert result.get("model_name", "").startswith("model_name_test_team_123_")
@ -1437,6 +1442,7 @@ class TestTeamModelUpdate:
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client, # type: ignore
write_row=_passthrough_row,
)
assert "403" in str(exc_info.value)
@ -1697,6 +1703,7 @@ class TestTeamModelUpdate:
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client, # type: ignore
write_row=_passthrough_row,
)
# team ACL must not be touched on a no-op edit
@ -4480,3 +4487,122 @@ class TestTeamMemberAutoRouterWrites:
assert saved == expected
assert row.litellm_params["complexity_router_config"] == stored_config
assert request.litellm_params.complexity_router_config == config
async def test_admin_router_changes_release_member_scope(self, endpoint: str, change: str) -> None:
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model, update_model
original: Final = self._row()
row: Final = original.model_copy(update={"model_info": {**original.model_info, "member_auto_router": True}})
database: Final = self._database(self._team(), row)
params: Final = {
"config": {"complexity_router_config": {"tiers": {"SIMPLE": "allowed"}, "session_affinity": True}},
"strategy": {"model": "auto_router/quality_router", "quality_router_default_model": "allowed"},
"unrelated": {"model": "auto_router/complexity_router", "max_tokens": 100},
}
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams.model_validate(params[change]),
model_info=ModelInfo(id=row.model_id) if endpoint == "legacy" or change == "unrelated" else None,
)
with self._environment(database, row):
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
if endpoint == "patch":
await patch_model(row.model_id, request, actor)
else:
await update_model(request, actor)
written: Final = database.db.litellm_proxymodeltable.update.await_args.kwargs["data"]
saved_info: Final = json.loads(written["model_info"]) if "model_info" in written else row.model_info
assert saved_info["member_auto_router"] is (change == "unrelated")
assert saved_info["team_id"] == "member-team"
assert saved_info["access_groups"] == ["retained-admin-group"]
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize("access", ["owner", "peer", "limited-key"])
async def test_both_update_entries_enforce_creator_and_stamp_member_scope(self, endpoint: str, access: str) -> None:
from fastapi import HTTPException
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model, update_model
row: Final = self._row()
database: Final = self._database(self._team(), row)
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams(
complexity_router_config={"tiers": {"SIMPLE": "allowed"}, "session_affinity": True}
),
model_info=ModelInfo(id=row.model_id, team_id="member-team"),
)
actor: Final = UserAPIKeyAuth(
user_id="peer" if access == "peer" else "owner",
user_role=LitellmUserRoles.INTERNAL_USER,
models=["personal-router"] if access == "limited-key" else ["allowed"],
config={"timeout": 60},
)
with self._environment(database, row):
operation: Final = (
patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)
)
if access != "owner":
with pytest.raises((HTTPException, ProxyException)):
await operation
database.transaction.litellm_proxymodeltable.update.assert_not_awaited()
return
await operation
written: Final = database.transaction.litellm_proxymodeltable.update.await_args.kwargs["data"]
saved_info: Final = json.loads(written["model_info"])
assert saved_info["member_auto_router"] is True
assert saved_info["team_id"] == "member-team"
assert saved_info["access_groups"] == ["retained-admin-group"]
assert "created_by" not in written
assert json.loads(written["litellm_params"])["complexity_router_config"]["session_affinity"] is True
assert written.get("model_name", row.model_name) == row.model_name
@pytest.mark.asyncio
@pytest.mark.parametrize("access", ["allowed", "opt-out", "limited-key"])
async def test_create_entry_requires_opt_in_and_appends_only_its_router(self, access: str) -> None:
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model
row: Final = self._row()
database: Final = self._database(self._team(enabled=access != "opt-out"), row)
actor: Final = UserAPIKeyAuth(
user_id="owner",
user_role=LitellmUserRoles.INTERNAL_USER,
models=["personal-router"] if access == "limited-key" else ["allowed"],
config={"timeout": 60},
)
deployment: Final = Deployment(
model_name="new-personal-router",
litellm_params=LiteLLM_Params(
model="auto_router/complexity_router", complexity_router_config={"tiers": {"SIMPLE": "allowed"}}
),
model_info=ModelInfo(id=row.model_id, team_id="member-team"),
)
with (
self._environment(database, row),
patch(
"litellm.proxy.proxy_server.proxy_config.add_deployment",
new=AsyncMock(
return_value=ReconcileOutcome( # test-quality-ok: [TQ008] model reload I/O boundary
still_desired=frozenset((row.model_id, "allowed-id")),
live_after=frozenset((row.model_id, "allowed-id")),
)
),
),
patch(
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_add", new=AsyncMock()
) as appended, # test-quality-ok: [TQ008] persistence boundary; the appended scope is asserted
):
if access != "allowed":
with pytest.raises(ProxyException) as denied:
await add_new_model(deployment, actor)
assert denied.value.code == "403"
database.transaction.litellm_proxymodeltable.create.assert_not_awaited()
appended.assert_not_awaited()
return
await add_new_model(deployment, actor)
written: Final = database.transaction.litellm_proxymodeltable.create.await_args.kwargs["data"]
assert written["created_by"] == "owner"
assert json.loads(written["model_info"])["member_auto_router"] is True
assert appended.await_args.kwargs["data"].models == ["new-personal-router"]
assert appended.await_args.kwargs["data"].team_id == "member-team"

View file

@ -28154,6 +28154,8 @@ export interface components {
team_id: string;
/** Team Member Permissions */
team_member_permissions?: string[] | null;
/** Tpd Limit */
tpd_limit?: number | null;
/** Tpm Limit */
tpm_limit?: number | null;
/** Updated At */
@ -29477,6 +29479,8 @@ export interface components {
team_id: string;
/** Team Member Permissions */
team_member_permissions?: string[] | null;
/** Tpd Limit */
tpd_limit?: number | null;
/** Tpm Limit */
tpm_limit?: number | null;
/** Updated At */
@ -31959,6 +31963,8 @@ export interface components {
team_member_rpm_limit?: number | null;
/** Team Member Tpm Limit */
team_member_tpm_limit?: number | null;
/** Tpd Limit */
tpd_limit?: number | null;
/** Tpm Limit */
tpm_limit?: number | null;
/** Tpm Limit Type */
@ -35923,6 +35929,8 @@ export interface components {
team_id: string;
/** Team Member Permissions */
team_member_permissions?: string[] | null;
/** Tpd Limit */
tpd_limit?: number | null;
/** Tpm Limit */
tpm_limit?: number | null;
/** Updated At */
@ -36063,6 +36071,8 @@ export interface components {
team_id: string;
/** Team Member Permissions */
team_member_permissions?: string[] | null;
/** Tpd Limit */
tpd_limit?: number | null;
/** Tpm Limit */
tpm_limit?: number | null;
/** Updated At */
@ -38034,6 +38044,10 @@ export interface components {
team_model_aliases?: {
[key: string]: unknown;
} | null;
/** Team Model Max Budget */
team_model_max_budget?: {
[key: string]: unknown;
} | null;
/**
* Team Models
* @default []
@ -38048,6 +38062,8 @@ export interface components {
team_soft_budget?: number | null;
/** Team Spend */
team_spend?: number | null;
/** Team Tpd Limit */
team_tpd_limit?: number | null;
/** Team Tpm Limit */
team_tpm_limit?: number | null;
/** Token */
@ -38537,6 +38553,11 @@ export interface components {
input_cost_per_character?: number | null;
/** Input Cost Per Token */
input_cost_per_token?: number | null;
/**
* Member Auto Router
* @default false
*/
member_auto_router: boolean;
/** Output Cost Per Character */
output_cost_per_character?: number | null;
/** Output Cost Per Token */