diff --git a/litellm/models/team.py b/litellm/models/team.py index da526515e6e..8edf10703b1 100644 --- a/litellm/models/team.py +++ b/litellm/models/team.py @@ -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 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d7e65f93121..827e850c8ff 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 093865e8f71..88d6dfca368 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 39e6ca9a369..ca25f07eb87 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -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 diff --git a/litellm/proxy/auth/litellm_license.py b/litellm/proxy/auth/litellm_license.py index 677f1a0fdda..40eda903da2 100644 --- a/litellm/proxy/auth/litellm_license.py +++ b/litellm/proxy/auth/litellm_license.py @@ -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 diff --git a/litellm/proxy/auth/team_grants.py b/litellm/proxy/auth/team_grants.py index 8db7e642728..efb67bb1202 100644 --- a/litellm/proxy/auth/team_grants.py +++ b/litellm/proxy/auth/team_grants.py @@ -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, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e92d090a2fb..dfea802497b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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. diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index cf9146d7a65..90ea597eee9 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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() diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index c6d7975b75e..2db0d66df6d 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 613508da22b..aa11a9c2e12 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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): diff --git a/litellm/router.py b/litellm/router.py index 63937ce5e76..1a2114cdab7 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index c85d201fb65..96a05c61ffc 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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): diff --git a/litellm/types/router.py b/litellm/types/router.py index 97bd93f3f47..8295444fa9f 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 3dea89ed67b..41b2f8d6c13 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -126,14 +126,10 @@ def invalid_sso_user_defined_values(): def test_get_experimental_ui_login_jwt_auth_token_valid(valid_sso_user_defined_values): """Test generating JWT token with valid user role""" - token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token( - valid_sso_user_defined_values - ) + token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(valid_sso_user_defined_values) # Decrypt and verify token contents - decrypted_token = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") # Check that decrypted_token is not None before using json.loads assert decrypted_token is not None token_data = json.loads(decrypted_token) @@ -159,9 +155,7 @@ def test_get_cli_jwt_auth_token_includes_team_alias(valid_sso_user_defined_value team_alias="test-team", ) - decrypted_token = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted_token is not None token_data = json.loads(decrypted_token) @@ -188,9 +182,7 @@ def test_get_cli_jwt_auth_token_carries_team_grants_not_user_allowlist( team_model_aliases={"team-fast": "gpt-4.1-mini"}, ) - decrypted_token = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted_token is not None token_data = json.loads(decrypted_token) @@ -207,9 +199,7 @@ def test_get_cli_jwt_auth_token_keeps_user_allowlist_when_no_team( """A session token with no team bound still carries the user's own allowlist.""" token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) - decrypted_token = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted_token is not None token_data = json.loads(decrypted_token) @@ -222,12 +212,8 @@ def test_get_experimental_ui_login_jwt_auth_token_uses_10_min_expiry( valid_sso_user_defined_values, ): """Test that Experimental UI token uses fixed 10-minute expiry (does not use LITELLM_UI_SESSION_DURATION).""" - token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token( - valid_sso_user_defined_values - ) - decrypted_token = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(valid_sso_user_defined_values) + decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted_token is not None token_data = json.loads(decrypted_token) expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00")) @@ -244,43 +230,33 @@ def test_experimental_ui_token_ignores_litellm_ui_session_duration( Experimental UI intentionally uses fixed 10-min expiry. If this test fails, the constant was incorrectly wired to the experimental flow.""" # Default LITELLM_UI_SESSION_DURATION is "24h" - token must still expire in ~10 min - token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token( - valid_sso_user_defined_values - ) - decrypted_token = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(valid_sso_user_defined_values) + decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted_token is not None token_data = json.loads(decrypted_token) expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00")) now = get_utc_datetime() # Must be ~10 min, NOT 24h. If LITELLM_UI_SESSION_DURATION were incorrectly used, this would fail. - assert expires <= now + timedelta( - minutes=11 - ), "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION" + assert expires <= now + timedelta(minutes=11), ( + "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION" + ) def test_get_experimental_ui_login_jwt_auth_token_invalid( invalid_sso_user_defined_values, ): """Test generating JWT token with missing user role""" - with pytest.raises(Exception, match='User role is required for experimental UI login') as exc_info: - ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token( - invalid_sso_user_defined_values - ) + with pytest.raises(Exception, match="User role is required for experimental UI login") as exc_info: + ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(invalid_sso_user_defined_values) assert str(exc_info.value) == "User role is required for experimental UI login" -def test_get_key_object_from_ui_hash_key_valid( - valid_sso_user_defined_values, monkeypatch -): +def test_get_key_object_from_ui_hash_key_valid(valid_sso_user_defined_values, monkeypatch): """Test getting key object from valid UI hash key""" monkeypatch.setenv("EXPERIMENTAL_UI_LOGIN", "True") # Generate a valid token - token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token( - valid_sso_user_defined_values - ) + token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(valid_sso_user_defined_values) # Get key object key_object = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) @@ -309,9 +285,7 @@ def test_get_key_object_from_ui_hash_key_invalid(): ("project", ProxyErrorTypes.project_model_access_denied), ], ) -def test_can_object_call_model_denials_return_forbidden( - object_type, expected_error_type -): +def test_can_object_call_model_denials_return_forbidden(object_type, expected_error_type): with pytest.raises(ProxyException) as exc_info: _can_object_call_model( model="restricted-model", @@ -568,9 +542,7 @@ async def test_get_key_object_should_reconnect_once_on_db_connection_error(): @pytest.mark.asyncio async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_error(): mock_prisma_client = MagicMock() - mock_prisma_client.get_data = AsyncMock( - side_effect=httpx.ConnectError("db not reachable after outage") - ) + mock_prisma_client.get_data = AsyncMock(side_effect=httpx.ConnectError("db not reachable after outage")) mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=False) mock_cache = MagicMock() @@ -613,9 +585,7 @@ class TestAuthCacheRedisWritePolicy: @pytest.mark.asyncio async def test_get_key_object_db_load_publishes_to_redis(self): mock_prisma_client = MagicMock() - mock_prisma_client.get_data = AsyncMock( - return_value=UserAPIKeyAuth(token="hashed-token-db") - ) + mock_prisma_client.get_data = AsyncMock(return_value=UserAPIKeyAuth(token="hashed-token-db")) fake_redis = _fake_redis_cache() cache = UserApiKeyCache() @@ -630,8 +600,7 @@ class TestAuthCacheRedisWritePolicy: assert key_obj.token == "hashed-token-db" fake_redis.async_set_cache.assert_awaited_once() assert ( - fake_redis.async_set_cache.await_args.kwargs.get("key") - or fake_redis.async_set_cache.await_args.args[0] + fake_redis.async_set_cache.await_args.kwargs.get("key") or fake_redis.async_set_cache.await_args.args[0] ) == "hashed-token-db" @@ -640,9 +609,7 @@ def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) # Decrypt and verify token contents - decrypted_token = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted_token is not None token_data = json.loads(decrypted_token) @@ -664,9 +631,7 @@ def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values assert expires >= get_utc_datetime() + timedelta(hours=23, minutes=59) -def test_get_cli_jwt_auth_token_custom_expiration( - valid_sso_user_defined_values, monkeypatch -): +def test_get_cli_jwt_auth_token_custom_expiration(valid_sso_user_defined_values, monkeypatch): """Test generating CLI JWT token with custom expiration via environment variable""" import importlib @@ -681,14 +646,10 @@ def test_get_cli_jwt_auth_token_custom_expiration( # Also reload auth_checks to pick up the new constant value importlib.reload(auth_checks) - token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token( - valid_sso_user_defined_values - ) + token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) # Decrypt and verify token contents - decrypted_token = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + decrypted_token = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted_token is not None token_data = json.loads(decrypted_token) @@ -706,18 +667,12 @@ def test_get_cli_jwt_auth_token_unique_per_session(valid_sso_user_defined_values from litellm.constants import CLI_SESSION_KEY_PREFIX def _decode(token: str) -> dict: - decrypted = decrypt_value_helper( - token, key="ui_hash_key", exception_type="debug" - ) + decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted is not None return json.loads(decrypted) - first = _decode( - ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) - ) - second = _decode( - ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) - ) + first = _decode(ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)) + second = _decode(ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)) assert first["token"].startswith(f"{CLI_SESSION_KEY_PREFIX}-") assert second["token"].startswith(f"{CLI_SESSION_KEY_PREFIX}-") @@ -740,9 +695,7 @@ def test_get_cli_jwt_auth_token_applies_fallback_budget(valid_sso_user_defined_v def test_get_cli_jwt_auth_token_no_fallback_when_budget_provided( valid_sso_user_defined_values, ): - token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( - valid_sso_user_defined_values, max_budget=None - ) + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values, max_budget=None) decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug") assert decrypted is not None assert json.loads(decrypted).get("max_budget") is None @@ -945,9 +898,7 @@ async def test_get_user_object_upsert_includes_user_email(): mock_prisma_client.db.litellm_usertable.create.assert_called_once() creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"] - assert ( - "user_email" in creation_args - ), "user_email should be included when upserting a new user" + assert "user_email" in creation_args, "user_email should be included when upserting a new user" assert creation_args["user_email"] == "test@example.com" assert creation_args["user_id"] == "new_test_user" @@ -962,12 +913,8 @@ async def test_get_user_object_backfills_null_email_from_cache_hit(): was returned unchanged and the DB was never updated. """ cache = UserApiKeyCache() - existing = LiteLLM_UserTable( - user_id="jwt-user-1", user_email=None, user_role="internal_user" - ) - await cache.async_set_cache( - key="jwt-user-1", value=existing, model_type=LiteLLM_UserTable - ) + existing = LiteLLM_UserTable(user_id="jwt-user-1", user_email=None, user_role="internal_user") + await cache.async_set_cache(key="jwt-user-1", value=existing, model_type=LiteLLM_UserTable) mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1) @@ -996,9 +943,7 @@ async def test_get_user_object_backfills_null_email_from_cache_hit(): assert update_kwargs["where"] == {"user_id": "jwt-user-1", "user_email": None} assert update_kwargs["data"]["user_email"] == "jwt-user-1@example.com" - refreshed = await cache.async_get_cache( - key="jwt-user-1", model_type=LiteLLM_UserTable - ) + refreshed = await cache.async_get_cache(key="jwt-user-1", model_type=LiteLLM_UserTable) assert refreshed is not None assert refreshed.user_email == "jwt-user-1@example.com" @@ -1010,9 +955,7 @@ async def test_get_user_object_backfills_null_email_from_db_read(): backfilled from the JWT-provided email before it is cached and returned. """ cache = UserApiKeyCache() - db_row = LiteLLM_UserTable( - user_id="jwt-user-3", user_email=None, user_role="internal_user" - ) + db_row = LiteLLM_UserTable(user_id="jwt-user-3", user_email=None, user_role="internal_user") backfilled_row = LiteLLM_UserTable( user_id="jwt-user-3", user_email="jwt-user-3@example.com", @@ -1020,15 +963,13 @@ async def test_get_user_object_backfills_null_email_from_db_read(): ) mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - side_effect=[db_row, backfilled_row] - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=[db_row, backfilled_row]) mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1) with patch( "litellm.proxy.auth.auth_checks._should_check_db", return_value=True - ): + ): # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise result = await get_user_object( user_id="jwt-user-3", prisma_client=mock_prisma_client, @@ -1042,9 +983,7 @@ async def test_get_user_object_backfills_null_email_from_db_read(): assert result.user_email == "jwt-user-3@example.com" mock_prisma_client.db.litellm_usertable.update_many.assert_called_once() - refreshed = await cache.async_get_cache( - key="jwt-user-3", model_type=LiteLLM_UserTable - ) + refreshed = await cache.async_get_cache(key="jwt-user-3", model_type=LiteLLM_UserTable) assert refreshed is not None assert refreshed.user_email == "jwt-user-3@example.com" @@ -1062,9 +1001,7 @@ async def test_get_user_object_does_not_overwrite_existing_email(): user_email="operator-set@example.com", user_role="internal_user", ) - await cache.async_set_cache( - key="jwt-user-2", value=existing, model_type=LiteLLM_UserTable - ) + await cache.async_set_cache(key="jwt-user-2", value=existing, model_type=LiteLLM_UserTable) mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=0) @@ -1091,12 +1028,8 @@ async def test_get_user_object_backfill_race_prefers_db_email(): with the value the DB accepted, not this request's proposed email. """ cache = UserApiKeyCache() - existing = LiteLLM_UserTable( - user_id="jwt-user-4", user_email=None, user_role="internal_user" - ) - await cache.async_set_cache( - key="jwt-user-4", value=existing, model_type=LiteLLM_UserTable - ) + existing = LiteLLM_UserTable(user_id="jwt-user-4", user_email=None, user_role="internal_user") + await cache.async_set_cache(key="jwt-user-4", value=existing, model_type=LiteLLM_UserTable) winner_row = LiteLLM_UserTable( user_id="jwt-user-4", @@ -1105,9 +1038,7 @@ async def test_get_user_object_backfill_race_prefers_db_email(): ) mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=0) - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=winner_row - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=winner_row) result = await get_user_object( user_id="jwt-user-4", @@ -1121,9 +1052,7 @@ async def test_get_user_object_backfill_race_prefers_db_email(): assert result is not None assert result.user_email == "winner@example.com" - refreshed = await cache.async_get_cache( - key="jwt-user-4", model_type=LiteLLM_UserTable - ) + refreshed = await cache.async_get_cache(key="jwt-user-4", model_type=LiteLLM_UserTable) assert refreshed is not None assert refreshed.user_email == "winner@example.com" @@ -1138,12 +1067,8 @@ async def test_get_user_object_backfill_caches_persisted_email_not_proposed(): optimistically caching the proposed email would serve a stale value. """ cache = UserApiKeyCache() - existing = LiteLLM_UserTable( - user_id="jwt-user-5", user_email=None, user_role="internal_user" - ) - await cache.async_set_cache( - key="jwt-user-5", value=existing, model_type=LiteLLM_UserTable - ) + existing = LiteLLM_UserTable(user_id="jwt-user-5", user_email=None, user_role="internal_user") + await cache.async_set_cache(key="jwt-user-5", value=existing, model_type=LiteLLM_UserTable) persisted_row = LiteLLM_UserTable( user_id="jwt-user-5", @@ -1152,9 +1077,7 @@ async def test_get_user_object_backfill_caches_persisted_email_not_proposed(): ) mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1) - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=persisted_row - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=persisted_row) result = await get_user_object( user_id="jwt-user-5", @@ -1168,9 +1091,7 @@ async def test_get_user_object_backfill_caches_persisted_email_not_proposed(): assert result is not None assert result.user_email == "admin-edited@example.com" - refreshed = await cache.async_get_cache( - key="jwt-user-5", model_type=LiteLLM_UserTable - ) + refreshed = await cache.async_get_cache(key="jwt-user-5", model_type=LiteLLM_UserTable) assert refreshed is not None assert refreshed.user_email == "admin-edited@example.com" @@ -1224,10 +1145,7 @@ async def test_get_user_object_upsert_routes_default_team_to_membership(monkeypa mock_add_to_team.assert_awaited_once() passed_teams = mock_add_to_team.await_args[1]["teams"] assert [team.team_id for team in passed_teams] == ["default-team"] - assert ( - mock_add_to_team.await_args[1]["user_api_key_dict"].user_role - == LitellmUserRoles.PROXY_ADMIN - ) + assert mock_add_to_team.await_args[1]["user_api_key_dict"].user_role == LitellmUserRoles.PROXY_ADMIN def test_log_budget_lookup_failure_dry_run(): @@ -1254,7 +1172,7 @@ def test_log_budget_lookup_failure_skips_user_not_found(): @pytest.mark.asyncio @patch( "litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock -) +) # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeypatch): """ Test that _get_team_db_check correctly calls the `new_team` function @@ -1290,10 +1208,10 @@ async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeyp @pytest.mark.asyncio @patch( "litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock -) +) # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise async def test_get_team_db_check_does_not_call_new_team_if_exists( mock_new_team, monkeypatch -): +): # test-quality-ok: [TQ002] collaborator injected via its import site; there is no seam to patch otherwise """ Test that _get_team_db_check does NOT call the `new_team` function if the team already exists in the database. @@ -1327,9 +1245,7 @@ async def test_get_team_db_check_does_not_call_new_team_if_exists( (MagicMock(), MagicMock(), True), # No vector stores to run ], ) -async def test_vector_store_access_check_early_returns( - prisma_client, vector_store_registry, expected_result -): +async def test_vector_store_access_check_early_returns(prisma_client, vector_store_registry, expected_result): """Test vector_store_access_check returns True for early exit conditions""" request_body = {"messages": [{"role": "user", "content": "test"}]} @@ -1379,9 +1295,7 @@ async def test_vector_store_access_check_early_returns( ), # Partial access ], ) -def test_can_object_call_vector_stores_scenarios( - object_permissions, vector_store_ids, should_raise, error_type -): +def test_can_object_call_vector_stores_scenarios(object_permissions, vector_store_ids, should_raise, error_type): """Test _can_object_call_vector_stores with various permission scenarios""" # Convert dict to object if not None if object_permissions is not None: @@ -1389,11 +1303,7 @@ def test_can_object_call_vector_stores_scenarios( mock_permissions.vector_stores = object_permissions["vector_stores"] object_permissions = mock_permissions - object_type = ( - "key" - if error_type == ProxyErrorTypes.key_vector_store_access_denied - else "team" - ) + object_type = "key" if error_type == ProxyErrorTypes.key_vector_store_access_denied else "team" if should_raise: with pytest.raises(ProxyException) as exc_info: @@ -1428,9 +1338,7 @@ async def test_vector_store_access_check_with_permissions(): mock_prisma_client = MagicMock() mock_permissions = MagicMock() mock_permissions.vector_stores = ["store-1", "store-2"] - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=mock_permissions - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=mock_permissions) mock_vector_store_registry = MagicMock() mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["store-1"] @@ -1476,14 +1384,10 @@ async def test_vector_store_access_check_with_team_permissions(): mock_prisma_client = MagicMock() team_permissions = MagicMock() team_permissions.vector_stores = ["team-store-allowed"] - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=team_permissions - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=team_permissions) mock_vector_store_registry = MagicMock() - mock_vector_store_registry.get_vector_store_ids_to_run.return_value = [ - "team-store-allowed" - ] + mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["team-store-allowed"] with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -1497,9 +1401,7 @@ async def test_vector_store_access_check_with_team_permissions(): assert result is True - mock_vector_store_registry.get_vector_store_ids_to_run.return_value = [ - "team-store-denied" - ] + mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["team-store-denied"] with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -2060,9 +1962,7 @@ async def test_get_tag_objects_batch(): mock_cache.async_set_cache = AsyncMock() # Mock DB to return all uncached tags in ONE query - mock_prisma.db.litellm_tagtable.find_many = AsyncMock( - return_value=[uncached_tag_1, uncached_tag_2, uncached_tag_3] - ) + mock_prisma.db.litellm_tagtable.find_many = AsyncMock(return_value=[uncached_tag_1, uncached_tag_2, uncached_tag_3]) # Call batch fetch tag_objects = await get_tag_objects_batch( @@ -2164,9 +2064,7 @@ async def test_get_tag_objects_batch_never_queries_db_for_unregistered_tags(): from litellm.proxy.auth.auth_checks import get_tag_objects_batch mock_prisma = MagicMock() - mock_prisma.db.litellm_tagtable.find_many = AsyncMock( - return_value=[_tag_registry_row("some-other-tag")] - ) + mock_prisma.db.litellm_tagtable.find_many = AsyncMock(return_value=[_tag_registry_row("some-other-tag")]) cache = UserApiKeyCache() first = await get_tag_objects_batch( @@ -2177,9 +2075,7 @@ async def test_get_tag_objects_batch_never_queries_db_for_unregistered_tags(): assert first == {} # The only query is the names-only registry fetch; the tag itself is never looked up. - mock_prisma.db.litellm_tagtable.find_many.assert_called_once_with( - take=TAG_REGISTRY_MAX_SIZE + 1 - ) + mock_prisma.db.litellm_tagtable.find_many.assert_called_once_with(take=TAG_REGISTRY_MAX_SIZE + 1) second = await get_tag_objects_batch( tag_names=["unregistered-tag"], @@ -2348,9 +2244,7 @@ async def test_get_tag_objects_batch_oversized_registry_falls_back_and_stops_ref """Past the cap the registry is unusable: keep the old per-tag path, but stop rebuilding it.""" from litellm.proxy.auth.auth_checks import get_tag_objects_batch - oversized = [ - _tag_registry_row(f"tag-{index}") for index in range(TAG_REGISTRY_MAX_SIZE + 1) - ] + oversized = [_tag_registry_row(f"tag-{index}") for index in range(TAG_REGISTRY_MAX_SIZE + 1)] async def fake_find_many(**kwargs): if "where" not in kwargs: @@ -2367,10 +2261,7 @@ async def test_get_tag_objects_batch_oversized_registry_falls_back_and_stops_ref user_api_key_cache=cache, ) assert list(first) == ["tag-a"] - assert ( - await cache.async_get_cache(key=tag_registry_cache_key()) - == TAG_REGISTRY_OVERFLOW_SENTINEL - ) + assert await cache.async_get_cache(key=tag_registry_cache_key()) == TAG_REGISTRY_OVERFLOW_SENTINEL second = await get_tag_objects_batch( tag_names=["tag-b"], @@ -2395,17 +2286,12 @@ async def test_tag_max_budget_check_still_enforces_registered_tag_over_budget(): async def fake_find_many(**kwargs): if "where" not in kwargs: return [_tag_registry_row("paid-tag")] - return [ - _tag_db_row(name, max_budget=1.0) - for name in kwargs["where"]["tag_name"]["in"] - ] + return [_tag_db_row(name, max_budget=1.0) for name in kwargs["where"]["tag_name"]["in"]] mock_prisma = MagicMock() mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:tag:paid-tag": return 1.5 return fallback_spend @@ -2464,6 +2350,44 @@ def _mock_prisma_for_team_lookup(find_unique): return mock_prisma_client +_TEAM_ALIAS_TABLE_ROW = {"id": 1, "model_aliases": '{"fast": "gpt-4o"}', "created_by": "admin", "updated_by": "admin"} + + +def _prisma_team_row(include): + """Mimics Prisma: the `litellm_model_table` relation rides on the row only when the query `include`s it.""" + columns = {"team_id": "team-aliases", "team_alias": "aliases", "models": ["gpt-4o"]} + row = ( + {**columns, "litellm_model_table": _TEAM_ALIAS_TABLE_ROW} + if (include or {}).get("litellm_model_table") + else columns + ) + return SimpleNamespace(dict=lambda: row, model_dump=lambda: row) + + +@pytest.mark.asyncio +async def test_get_team_object_loads_model_aliases_relation(): + """LIT-5858: the auth path read teams without `include`ing `litellm_model_table`, so every JWT + team came back with `model_aliases=None` and alias requests 403'd.""" + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.auth.team_grants import team_model_aliases + + async def find_unique(where, include=None): + return _prisma_team_row(include) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + team = await get_team_object( + team_id="team-aliases", + prisma_client=_mock_prisma_for_team_lookup(AsyncMock(side_effect=find_unique)), + user_api_key_cache=mock_cache, + check_db_only=True, + ) + + assert team_model_aliases(team) == {"fast": "gpt-4o"} + + @pytest.mark.asyncio async def test_get_team_object_distinguishes_absent_team_from_unreadable_row(): """A deleted team and a database that would not answer both surface as a 404, @@ -2814,8 +2738,7 @@ def _pass_through_request() -> Request: LITELLM_PASS_THROUGH_ENDPOINT_MARKER, ) - def pass_through_endpoint(): - ... + def pass_through_endpoint(): ... setattr(pass_through_endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True) return Request(scope={"type": "http", "headers": [], "endpoint": pass_through_endpoint}) @@ -2825,8 +2748,7 @@ def _builtin_request() -> Request: """A Request dispatched to a built-in (non-pass-through) handler, e.g. what a custom path colliding with a core route actually resolves to.""" - def chat_completions(): - ... + def chat_completions(): ... return Request(scope={"type": "http", "headers": [], "endpoint": chat_completions}) @@ -2991,9 +2913,7 @@ async def test_virtual_key_soft_budget_check_without_user_obj(): ], ) @pytest.mark.asyncio -async def test_virtual_key_soft_budget_check_scenarios( - spend, soft_budget, expect_alert -): +async def test_virtual_key_soft_budget_check_scenarios(spend, soft_budget, expect_alert): """Test _virtual_key_soft_budget_check with various spend and soft_budget scenarios""" alert_triggered = False @@ -3022,9 +2942,9 @@ async def test_virtual_key_soft_budget_check_scenarios( await asyncio.sleep(0.1) - assert ( - alert_triggered == expect_alert - ), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}" + assert alert_triggered == expect_alert, ( + f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}" + ) @pytest.mark.asyncio @@ -3135,9 +3055,7 @@ async def test_virtual_key_max_budget_alert_check_without_user_obj(): ], ) @pytest.mark.asyncio -async def test_virtual_key_max_budget_alert_check_scenarios( - spend, max_budget, expect_alert -): +async def test_virtual_key_max_budget_alert_check_scenarios(spend, max_budget, expect_alert): """Test _virtual_key_max_budget_alert_check with various spend and max_budget scenarios""" alert_triggered = False @@ -3166,9 +3084,9 @@ async def test_virtual_key_max_budget_alert_check_scenarios( await asyncio.sleep(0.1) - assert ( - alert_triggered == expect_alert - ), f"Expected alert_triggered to be {expect_alert} for spend={spend}, max_budget={max_budget}" + assert alert_triggered == expect_alert, ( + f"Expected alert_triggered to be {expect_alert} for spend={spend}, max_budget={max_budget}" + ) @pytest.mark.asyncio @@ -3427,9 +3345,7 @@ async def test_custom_auth_common_checks_opt_in(): "prisma_client": None, "user_api_key_cache": MagicMock(), "proxy_logging_obj": MagicMock(), - "general_settings": ( - {"custom_auth_run_common_checks": True} if flag else {} - ), + "general_settings": ({"custom_auth_run_common_checks": True} if flag else {}), "llm_router": None, "user_custom_auth": user_custom_auth, "litellm_proxy_admin_name": "admin", @@ -3501,9 +3417,7 @@ async def test_virtual_key_budget_check_reads_from_spend_counter(): proxy_logging_obj = ProxyLogging(user_api_key_cache=None) proxy_logging_obj.budget_alerts = AsyncMock() - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:key:test-hashed-token": return 1.5 return fallback_spend @@ -3537,9 +3451,7 @@ async def test_virtual_key_budget_check_fallback_no_counter(): proxy_logging_obj.budget_alerts = AsyncMock() # get_current_spend returns fallback_spend when no counter exists - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): return fallback_spend with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): @@ -3569,9 +3481,7 @@ def _over_budget_token(**overrides) -> UserAPIKeyAuth: def _patched_spend(value: float): - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): return value return patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend) @@ -3632,9 +3542,7 @@ async def test_budget_throttle_decision_cleared_before_caching(): otherwise it would re-apply (and compound) on every subsequent request.""" from litellm.proxy.auth.auth_checks import _copy_user_api_key_auth_for_cache - valid_token = _over_budget_token( - tpm_limit=1000, rpm_limit=100, metadata={"throttle_on_budget_exceeded": True} - ) + valid_token = _over_budget_token(tpm_limit=1000, rpm_limit=100, metadata={"throttle_on_budget_exceeded": True}) valid_token.budget_throttle_pct = 0.1 cached = _copy_user_api_key_auth_for_cache(user_api_key_obj=valid_token) @@ -3730,9 +3638,7 @@ async def test_team_budget_check_reads_from_spend_counter(): proxy_logging_obj = ProxyLogging(user_api_key_cache=None) proxy_logging_obj.budget_alerts = AsyncMock() - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:team:test-team": return 1.5 return fallback_spend @@ -3759,9 +3665,7 @@ async def test_end_user_budget_check_reads_from_spend_counter(): litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), ) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:end_user:customer-1": return 1.5 return fallback_spend @@ -3789,9 +3693,7 @@ async def test_tag_budget_check_reads_from_spend_counter(): litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), ) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:tag:paid-tag": return 1.5 return fallback_spend @@ -3841,9 +3743,7 @@ async def test_team_member_budget_check_reads_from_spend_counter(): proxy_logging_obj = ProxyLogging(user_api_key_cache=None) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:team_member:test-user:test-team": return 1.5 return fallback_spend @@ -3883,9 +3783,7 @@ class TestGuardrailModificationCheck: team_object = MagicMock() team_object.metadata = {} # no permission - return _guardrail_modification_check( - request_body=request_body, team_object=team_object - ) + return _guardrail_modification_check(request_body=request_body, team_object=team_object) def test_noop_when_no_guardrail_keys_present(self): # no-op — should return silently @@ -3933,9 +3831,7 @@ class TestGuardrailModificationCheck: return_value=False, ): with pytest.raises(HTTPException) as exc: - self._call( - {"metadata": {"opted_out_global_guardrails": ["some_guardrail"]}} - ) + self._call({"metadata": {"opted_out_global_guardrails": ["some_guardrail"]}}) assert exc.value.status_code == 403 @pytest.mark.parametrize( @@ -4069,18 +3965,12 @@ async def test_team_member_budget_check_falls_back_to_team_default_budget_id(): fake_budget_row = MagicMock() fake_budget_row.max_budget = 50.0 - fake_budget_row.dict = MagicMock( - return_value={"budget_id": "budget-default", "max_budget": 50.0} - ) + fake_budget_row.dict = MagicMock(return_value={"budget_id": "budget-default", "max_budget": 50.0}) prisma_client = MagicMock() - prisma_client.db.litellm_budgettable.find_unique = AsyncMock( - return_value=fake_budget_row - ) + prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_budget_row) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:team_member:test-user:test-team": return 70.0 return fallback_spend @@ -4171,15 +4061,11 @@ async def test_team_member_budget_check_per_member_override_wins_over_team_defau fake_budget_row.max_budget = 50.0 prisma_client = MagicMock() - prisma_client.db.litellm_budgettable.find_unique = AsyncMock( - return_value=fake_budget_row - ) + prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_budget_row) mocked_spend = 70.0 - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:team_member:test-user:test-team": return mocked_spend return fallback_spend @@ -4260,18 +4146,12 @@ async def test_team_member_budget_check_null_clone_falls_back_to_team_default(): fake_default_row = MagicMock() fake_default_row.max_budget = 65.0 - fake_default_row.dict = MagicMock( - return_value={"budget_id": "budget-default", "max_budget": 65.0} - ) + fake_default_row.dict = MagicMock(return_value={"budget_id": "budget-default", "max_budget": 65.0}) prisma_client = MagicMock() - prisma_client.db.litellm_budgettable.find_unique = AsyncMock( - return_value=fake_default_row - ) + prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_default_row) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:team_member:test-user:test-team": return 500.0 return fallback_spend @@ -4329,18 +4209,12 @@ async def test_team_member_budget_check_null_clone_with_null_default_skips_enfor fake_default_row = MagicMock() fake_default_row.max_budget = None - fake_default_row.dict = MagicMock( - return_value={"budget_id": "budget-default", "max_budget": None} - ) + fake_default_row.dict = MagicMock(return_value={"budget_id": "budget-default", "max_budget": None}) prisma_client = MagicMock() - prisma_client.db.litellm_budgettable.find_unique = AsyncMock( - return_value=fake_default_row - ) + prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_default_row) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:team_member:test-user:test-team": return 1000.0 return fallback_spend @@ -4398,18 +4272,12 @@ async def test_team_member_budget_check_zero_team_default_treated_as_no_cap(): # Team default budget row with max_budget=0.0 (the regression trigger). fake_default_row = MagicMock() fake_default_row.max_budget = 0.0 - fake_default_row.dict = MagicMock( - return_value={"budget_id": "budget-default", "max_budget": 0.0} - ) + fake_default_row.dict = MagicMock(return_value={"budget_id": "budget-default", "max_budget": 0.0}) prisma_client = MagicMock() - prisma_client.db.litellm_budgettable.find_unique = AsyncMock( - return_value=fake_default_row - ) + prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_default_row) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:team_member:test-user:test-team": return 0.0 return fallback_spend @@ -4467,9 +4335,7 @@ async def test_team_member_budget_check_zero_per_member_row_still_blocks(): prisma_client = MagicMock() prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) - async def mock_get_current_spend( - counter_key, fallback_spend, max_budget=None, **kwargs - ): + async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs): if counter_key == "spend:team_member:test-user:test-team": return 0.0 return fallback_spend @@ -4517,19 +4383,13 @@ def _patch_validation_helpers(monkeypatch, *, end_user=None, user=None, fuzzy=No """Stub out the DB helpers resolve_and_validate_end_user_id delegates to.""" from litellm.proxy.auth import auth_checks - monkeypatch.setattr( - auth_checks, "get_end_user_object", AsyncMock(return_value=end_user) - ) + monkeypatch.setattr(auth_checks, "get_end_user_object", AsyncMock(return_value=end_user)) monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user)) - monkeypatch.setattr( - auth_checks, "_get_fuzzy_user_object", AsyncMock(return_value=fuzzy) - ) + monkeypatch.setattr(auth_checks, "_get_fuzzy_user_object", AsyncMock(return_value=fuzzy)) @pytest.mark.asyncio -async def test_resolve_end_user_returns_none_for_none_input( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_returns_none_for_none_input(_validate_flag_on, monkeypatch): from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id _patch_validation_helpers(monkeypatch) @@ -4564,9 +4424,7 @@ async def test_resolve_end_user_passes_through_when_flag_disabled(monkeypatch): @pytest.mark.asyncio -async def test_resolve_end_user_passes_through_when_no_prisma_client( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_passes_through_when_no_prisma_client(_validate_flag_on, monkeypatch): from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id _patch_validation_helpers(monkeypatch) @@ -4600,9 +4458,7 @@ async def test_resolve_end_user_matches_end_user_table(_validate_flag_on, monkey @pytest.mark.asyncio -async def test_resolve_end_user_matches_user_table_by_user_id( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_matches_user_table_by_user_id(_validate_flag_on, monkeypatch): from litellm.proxy.auth import auth_checks from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id @@ -4620,9 +4476,7 @@ async def test_resolve_end_user_matches_user_table_by_user_id( @pytest.mark.asyncio -async def test_resolve_end_user_matches_user_table_by_email( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_matches_user_table_by_email(_validate_flag_on, monkeypatch): """Email-shaped ids route through get_user_object with user_email set. The fuzzy lookup must happen inside get_user_object so it shares the @@ -4650,9 +4504,7 @@ async def test_resolve_end_user_matches_user_table_by_email( @pytest.mark.asyncio -async def test_resolve_end_user_non_email_id_does_not_pass_user_email( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_non_email_id_does_not_pass_user_email(_validate_flag_on, monkeypatch): """Non-email ids skip the email fuzzy path to avoid a pointless DB hit.""" from litellm.proxy.auth import auth_checks from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id @@ -4671,9 +4523,7 @@ async def test_resolve_end_user_non_email_id_does_not_pass_user_email( @pytest.mark.asyncio -async def test_resolve_end_user_drops_codex_opaque_identifier( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_drops_codex_opaque_identifier(_validate_flag_on, monkeypatch): from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id _patch_validation_helpers(monkeypatch) # all helpers return None @@ -4695,9 +4545,7 @@ async def test_resolve_end_user_drops_codex_opaque_identifier( @pytest.mark.asyncio -async def test_resolve_end_user_preserves_id_when_default_budget_configured( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_preserves_id_when_default_budget_configured(_validate_flag_on, monkeypatch): """Don't drop unregistered ids when litellm.max_end_user_budget_id is set. The default end-user budget is applied downstream when the id is present @@ -4734,9 +4582,7 @@ async def test_resolve_end_user_drops_unknown_email(_validate_flag_on, monkeypat @pytest.mark.asyncio -async def test_resolve_end_user_uses_cached_valid_result( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_uses_cached_valid_result(_validate_flag_on, monkeypatch): from litellm.proxy.auth import auth_checks from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id @@ -4756,9 +4602,7 @@ async def test_resolve_end_user_uses_cached_valid_result( @pytest.mark.asyncio -async def test_resolve_end_user_uses_cached_invalid_result( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_uses_cached_invalid_result(_validate_flag_on, monkeypatch): from litellm.proxy.auth import auth_checks from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id @@ -4777,9 +4621,7 @@ async def test_resolve_end_user_uses_cached_invalid_result( @pytest.mark.asyncio -async def test_resolve_end_user_swallows_db_errors_and_returns_none( - _validate_flag_on, monkeypatch -): +async def test_resolve_end_user_swallows_db_errors_and_returns_none(_validate_flag_on, monkeypatch): from litellm.proxy.auth import auth_checks from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id @@ -4900,19 +4742,13 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): ) # (1) team_id-keyed write fires with the refreshed object - written_keys = [ - (c.kwargs.get("key") or c.args[0]) - for c in cache.async_set_cache.await_args_list - ] + written_keys = [(c.kwargs.get("key") or c.args[0]) for c in cache.async_set_cache.await_args_list] assert written_keys == ["team_id:team-1234"], ( "Only the team_id-keyed write should fire; the alias key must be " "deleted, NOT written. " f"Got writes: {written_keys}" ) - written_value = ( - cache.async_set_cache.await_args.kwargs.get("value") - or cache.async_set_cache.await_args.args[1] - ) + written_value = cache.async_set_cache.await_args.kwargs.get("value") or cache.async_set_cache.await_args.args[1] assert written_value is team_table # (2) team_alias-keyed entry is deleted in BOTH the in-memory cache @@ -4946,10 +4782,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with( key="team_id:team-no-alias" ) - written_keys_aliasless = [ - (c.kwargs.get("key") or c.args[0]) - for c in cache2.async_set_cache.await_args_list - ] + written_keys_aliasless = [(c.kwargs.get("key") or c.args[0]) for c in cache2.async_set_cache.await_args_list] assert written_keys_aliasless == ["team_id:team-no-alias"] @@ -5029,9 +4862,7 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): await _cache_team_object( team_id=team_id, - team_table=LiteLLM_TeamTableCachedObj( - team_id=team_id, models=["model-a", "model-b"] - ), + team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a", "model-b"]), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -5141,9 +4972,7 @@ async def test_cache_team_object_tolerates_cache_invalidation_failures(): cache.async_set_cache = AsyncMock() cache.delete_cache = MagicMock(side_effect=Exception("redis down")) logging_obj = MagicMock() - logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock( - side_effect=Exception("redis down") - ) + logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock(side_effect=Exception("redis down")) await _cache_team_object( team_id="team-cache-outage", @@ -5156,10 +4985,7 @@ async def test_cache_team_object_tolerates_cache_invalidation_failures(): proxy_logging_obj=logging_obj, ) - written_keys = [ - (c.kwargs.get("key") or c.args[0]) - for c in cache.async_set_cache.await_args_list - ] + written_keys = [(c.kwargs.get("key") or c.args[0]) for c in cache.async_set_cache.await_args_list] assert written_keys == ["team_id:team-cache-outage"] @@ -5435,8 +5261,11 @@ async def test_common_checks_budget_reads_run_concurrently(): probe = _BudgetSpendConcurrencyProbe(expected=4) - with patch("litellm.proxy.proxy_server.prisma_client", None), patch( - "litellm.proxy.proxy_server.get_current_spend", probe + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: common_checks has no database seam + patch( + "litellm.proxy.proxy_server.get_current_spend", probe + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): task = asyncio.create_task( common_checks( @@ -5504,8 +5333,11 @@ async def test_common_checks_budget_gather_raises_highest_priority_scope(): request=MagicMock(spec=Request), ) - with patch("litellm.proxy.proxy_server.prisma_client", None), patch( - "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: common_checks has no database seam + patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): # Both team and end-user over budget: team wins on priority. _spend_by_counter.team = 999.0 @@ -5540,8 +5372,11 @@ async def test_common_checks_personal_user_budget_blocks_in_gather(): async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs): return 999.0 if counter_key == "spend:user:u1" else 0.0 - with patch("litellm.proxy.proxy_server.prisma_client", None), patch( - "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: common_checks has no database seam + patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): with pytest.raises(litellm.BudgetExceededError) as over: await common_checks( @@ -5583,9 +5418,15 @@ async def test_common_checks_personal_user_budget_skipped_for_team_key(): async def _no_membership(*args, **kwargs): return None - with patch("litellm.proxy.proxy_server.prisma_client", None), patch( - "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter - ), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership): + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: common_checks has no database seam + patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", _no_membership + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ): result = await common_checks( request_body={"messages": [{"role": "user", "content": "hi"}]}, team_object=team, @@ -5624,9 +5465,15 @@ async def test_common_checks_personal_user_budget_enforced_on_team_key_when_flag async def _no_membership(*args, **kwargs): return None - with patch("litellm.proxy.proxy_server.prisma_client", None), patch( - "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter - ), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership): + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: common_checks has no database seam + patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", _no_membership + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ): with pytest.raises(litellm.BudgetExceededError) as exc_info: await common_checks( request_body={"messages": [{"role": "user", "content": "hi"}]}, @@ -5657,8 +5504,11 @@ async def test_common_checks_personal_user_budget_still_enforced_on_personal_key async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs): return 999.0 if counter_key == "spend:user:u1" else 0.0 - with patch("litellm.proxy.proxy_server.prisma_client", None), patch( - "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: common_checks has no database seam + patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): with pytest.raises(litellm.BudgetExceededError): await common_checks( @@ -5741,10 +5591,19 @@ async def test_budget_checks_only_run_on_llm_api_routes(scope, route, expect_blo request=MagicMock(spec=Request), ) - with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( - "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter - ), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership), patch( - "litellm.proxy.auth.auth_checks.get_org_object", _get_org + with ( + patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", _no_membership + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + patch( + "litellm.proxy.auth.auth_checks.get_org_object", _get_org + ), # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): if expect_blocked: with pytest.raises(litellm.BudgetExceededError): @@ -6243,9 +6102,7 @@ async def test_get_end_user_object_token_budget_gate_keeps_fetching_unrestricted mock_prisma = MagicMock() mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[]) - mock_prisma.db.litellm_endusertable.find_unique = AsyncMock( - return_value=_end_user_db_row("eu-anon-1", spend=100.0) - ) + mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1", spend=100.0)) cache = UserApiKeyCache() result = await get_end_user_object( @@ -6376,6 +6233,32 @@ async def test_get_team_object_by_alias_db_fetch_returns_cached_obj(): assert result.models == ["gpt-4"] +@pytest.mark.asyncio +async def test_get_team_object_by_alias_loads_model_aliases_relation(): + """LIT-5858: same regression as `test_get_team_object_loads_model_aliases_relation`, for the + `team_alias_jwt_field` lookup.""" + from litellm.proxy.auth.auth_checks import get_team_object_by_alias + from litellm.proxy.auth.team_grants import team_model_aliases + + async def find_many(where, include=None): + return [_prisma_team_row(include)] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + team = await get_team_object_by_alias( + team_alias="aliases", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert team_model_aliases(team) == {"fast": "gpt-4o"} + + @pytest.mark.asyncio async def test_get_org_object_by_alias_db_fetch_returns_validated_org(): from litellm.proxy._types import LiteLLM_OrganizationTable @@ -6927,9 +6810,7 @@ async def test_common_checks_ignores_non_llm_route_when_enabled(monkeypatch): monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", True) router = _router_with_priced_and_unpriced_models() - result = await _run_common_checks( - model="unpriced-group", llm_router=router, route="/model/new" - ) + result = await _run_common_checks(model="unpriced-group", llm_router=router, route="/model/new") assert result is True @@ -7045,11 +6926,15 @@ def test_team_allowed_routes_exact_route_does_not_become_a_prefix_grant(): roles = LiteLLM_JWTAuth(team_allowed_routes=["/internal-models/model-a"]) assert ( - allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/internal-models/model-a", litellm_proxy_roles=roles) + allowed_routes_check( + user_role=LitellmUserRoles.TEAM, user_route="/internal-models/model-a", litellm_proxy_roles=roles + ) is True ) assert ( - allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/internal-models/model-b", litellm_proxy_roles=roles) + allowed_routes_check( + user_role=LitellmUserRoles.TEAM, user_route="/internal-models/model-b", litellm_proxy_roles=roles + ) is False ) @@ -7127,8 +7012,7 @@ async def test_invalidate_team_member_spend_state_sets_the_spend_counter_and_cle assert await real_cache.async_get_cache(key="team_membership:user-1:team-1") is None assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 0.0 assert ( - real_spend_counter_cache.in_memory_cache.get_cache(key="spend_db_floor:spend:team_member:user-1:team-1") - == 0.0 + real_spend_counter_cache.in_memory_cache.get_cache(key="spend_db_floor:spend:team_member:user-1:team-1") == 0.0 ), "the DB-floor marker kept the pre-reset value; a stale-floor read can raise the counter right back up" @@ -7324,9 +7208,9 @@ async def test_invalidate_team_member_spend_state_broadcasts_the_spend_counter_t ) assert remote_spend_counter_in_memory_cache.get_cache("spend:team_member:user-1:team-1") == 0.0 - assert ( - remote_spend_counter_in_memory_cache.get_cache("spend_db_floor:spend:team_member:user-1:team-1") == 0.0 - ), "the DB-floor marker was not broadcast; a remote worker can re-raise the counter off its stale floor" + assert remote_spend_counter_in_memory_cache.get_cache("spend_db_floor:spend:team_member:user-1:team-1") == 0.0, ( + "the DB-floor marker was not broadcast; a remote worker can re-raise the counter off its stale floor" + ) @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 99a0a4c0a8b..f0f9adfb580 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -13,6 +13,7 @@ from litellm.proxy._types import ( DEFAULT_JWKS_STALE_TTL, JWTLiteLLMRoleMap, LiteLLM_JWTAuth, + LiteLLM_ModelTable, LiteLLM_TeamMembership, LiteLLM_TeamTable, LiteLLM_UserTable, @@ -96,9 +97,7 @@ async def test_map_user_to_teams_handles_already_in_team_exception(): ) as mock_add: with patch("litellm.proxy.auth.handle_jwt.verbose_proxy_logger") as mock_logger: # This should not raise an exception - result = await JWTAuthManager.map_user_to_teams( - user_object=user, team_object=team - ) + result = await JWTAuthManager.map_user_to_teams(user_object=user, team_object=team) # Verify the method completed successfully assert result is None @@ -135,14 +134,10 @@ async def test_map_user_to_teams_reraises_other_proxy_exceptions(): async def test_map_user_to_teams_null_inputs(): """Test that method handles null inputs gracefully""" # Test with null user - await JWTAuthManager.map_user_to_teams( - user_object=None, team_object=LiteLLM_TeamTable(team_id="test_team_1") - ) + await JWTAuthManager.map_user_to_teams(user_object=None, team_object=LiteLLM_TeamTable(team_id="test_team_1")) # Test with null team - await JWTAuthManager.map_user_to_teams( - user_object=LiteLLM_UserTable(user_id="test_user_1"), team_object=None - ) + await JWTAuthManager.map_user_to_teams(user_object=LiteLLM_UserTable(user_id="test_user_1"), team_object=None) # Test with both null await JWTAuthManager.map_user_to_teams(user_object=None, team_object=None) @@ -198,9 +193,7 @@ async def test_find_team_with_model_access_reports_passthrough_allowlist_denial( assert exc_info.value.status_code == 403 assert "allowed_passthrough_routes" in exc_info.value.detail assert "requested model" not in exc_info.value.detail - mock_is_auth_enforced_pass_through_route.assert_called_once_with( - route="/my-pass-through", method="POST" - ) + mock_is_auth_enforced_pass_through_route.assert_called_once_with(route="/my-pass-through", method="POST") user_api_key_dict = mock_passthrough_check.call_args.kwargs["user_api_key_dict"] assert user_api_key_dict.metadata == {} @@ -293,9 +286,7 @@ async def test_auth_builder_proxy_admin_user_role(): route = "/chat/completions" # Create user object with PROXY_ADMIN role - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.PROXY_ADMIN - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.PROXY_ADMIN) # Create mock JWT handler jwt_handler = JWTHandler() @@ -306,12 +297,10 @@ async def test_auth_builder_proxy_admin_user_role(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object( JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + ) as mock_check_rbac, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -319,9 +308,7 @@ async def test_auth_builder_proxy_admin_user_role(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -336,7 +323,7 @@ async def test_auth_builder_proxy_admin_user_role(): ) as mock_find_team, patch.object( JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + ) as mock_get_all_team_ids, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "find_team_with_model_access", @@ -351,10 +338,10 @@ async def test_auth_builder_proxy_admin_user_role(): ) as mock_get_objects, patch.object( JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, + ) as mock_map_user, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, + ) as mock_validate_object, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -388,9 +375,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): route = "/chat/completions" # Create user object with regular USER role - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create mock JWT handler jwt_handler = JWTHandler() @@ -401,12 +386,10 @@ async def test_auth_builder_non_proxy_admin_user_role(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object( JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + ) as mock_check_rbac, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -414,9 +397,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -431,7 +412,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): ) as mock_find_team, patch.object( JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + ) as mock_get_all_team_ids, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "find_team_with_model_access", @@ -446,10 +427,10 @@ async def test_auth_builder_non_proxy_admin_user_role(): ) as mock_get_objects, patch.object( JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, + ) as mock_map_user, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, + ) as mock_validate_object, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -598,11 +579,7 @@ async def test_sync_user_role_and_teams(): prisma_client=None, user_api_key_cache=mock_user_api_key_cache, litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -611,9 +588,7 @@ async def test_sync_user_role_and_teams(): token = {"roles": ["ADMIN"], "my_id_teams": ["team1", "team2"]} - user = LiteLLM_UserTable( - user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER.value, teams=["team2"] - ) + user = LiteLLM_UserTable(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER.value, teams=["team2"]) prisma = AsyncMock() prisma.db.litellm_usertable.update = AsyncMock() @@ -640,11 +615,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -661,9 +632,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): prisma = AsyncMock() prisma.db.litellm_usertable.update = AsyncMock() - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args @@ -683,11 +652,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -708,9 +673,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ): - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args @@ -730,11 +693,7 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -750,9 +709,7 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): prisma = AsyncMock() - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_not_called() @@ -777,9 +734,7 @@ def test_get_all_jwt_team_ids_unions_singular_and_plural(): assert jwt_handler.get_all_jwt_team_ids({"teams": ["a", "b"]}) == ["a", "b"] # both populated, no overlap - assert jwt_handler.get_all_jwt_team_ids( - {"team_id": "primary", "teams": ["a", "b"]} - ) == ["a", "b", "primary"] + assert jwt_handler.get_all_jwt_team_ids({"team_id": "primary", "teams": ["a", "b"]}) == ["a", "b", "primary"] # both populated with overlap — singular dedup'd assert jwt_handler.get_all_jwt_team_ids({"team_id": "a", "teams": ["a", "b"]}) == [ @@ -788,9 +743,7 @@ def test_get_all_jwt_team_ids_unions_singular_and_plural(): ] # singular field as multi-element list (some IdPs) — merge all, preserve plural-first order - assert jwt_handler.get_all_jwt_team_ids( - {"team_id": ["primary", "secondary"], "teams": ["a"]} - ) == [ + assert jwt_handler.get_all_jwt_team_ids({"team_id": ["primary", "secondary"], "teams": ["a"]}) == [ "a", "primary", "secondary", @@ -850,19 +803,11 @@ async def test_map_jwt_role_to_litellm_role(): litellm_jwtauth=LiteLLM_JWTAuth( jwt_litellm_role_map=[ # Exact match - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ), + JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN), # Wildcard patterns - JWTLiteLLMRoleMap( - jwt_role="user_*", litellm_role=LitellmUserRoles.INTERNAL_USER - ), - JWTLiteLLMRoleMap( - jwt_role="team_?", litellm_role=LitellmUserRoles.TEAM - ), - JWTLiteLLMRoleMap( - jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER - ), + JWTLiteLLMRoleMap(jwt_role="user_*", litellm_role=LitellmUserRoles.INTERNAL_USER), + JWTLiteLLMRoleMap(jwt_role="team_?", litellm_role=LitellmUserRoles.TEAM), + JWTLiteLLMRoleMap(jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER), ], roles_jwt_field="roles", ), @@ -934,9 +879,7 @@ async def test_map_jwt_role_to_litellm_role(): # Test patterns that don't match character classes jwt_handler.litellm_jwtauth.jwt_litellm_role_map = [ - JWTLiteLLMRoleMap( - jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER - ), + JWTLiteLLMRoleMap(jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER), ] token = {"roles": ["dev_4"]} # 4 is not in [123] result = jwt_handler.map_jwt_role_to_litellm_role(token) @@ -1031,25 +974,19 @@ async def test_nested_jwt_field_access(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="profile.object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(nested_token, None) == "obj789" # Test 5b: object_id_jwt_field with flat access (backward compatibility) jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(flat_token, None) == "obj789" # Test 6: end_user_id_jwt_field with nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - end_user_id_jwt_field="customer.end_user_id" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") assert jwt_handler.get_end_user_id(nested_token, None) == "customer123" # Test 6b: end_user_id_jwt_field with flat access (backward compatibility) @@ -1065,9 +1002,7 @@ async def test_nested_jwt_field_access(): assert jwt_handler.get_team_id(flat_token, None) == "team456" # Test 8: roles_jwt_field with deeply nested access (already supported, but testing) - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - roles_jwt_field="resource_access.my-client.roles" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") assert jwt_handler.get_jwt_role(nested_token, []) == ["admin", "user"] # Test 9: user_roles_jwt_field with nested access (already supported, but testing) @@ -1113,10 +1048,7 @@ async def test_nested_jwt_field_missing_paths(): # Test 2: Missing user.email should return default jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="user.email") - assert ( - jwt_handler.get_user_email(incomplete_token, "default@example.com") - == "default@example.com" - ) + assert jwt_handler.get_user_email(incomplete_token, "default@example.com") == "default@example.com" # Test 3: Missing groups should return empty list jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups") @@ -1131,43 +1063,28 @@ async def test_nested_jwt_field_missing_paths(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="profile.object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(incomplete_token, "default_obj") == "default_obj" # Test 6: Missing customer.end_user_id should return default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - end_user_id_jwt_field="customer.end_user_id" - ) - assert ( - jwt_handler.get_end_user_id(incomplete_token, "default_customer") - == "default_customer" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") + assert jwt_handler.get_end_user_id(incomplete_token, "default_customer") == "default_customer" # Test 7: Missing tenant.team_id should use team_id_default fallback - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_id_jwt_field="tenant.team_id", team_id_default="fallback_team" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="tenant.team_id", team_id_default="fallback_team") assert jwt_handler.get_team_id(incomplete_token, "default_team") == "fallback_team" # Test 8: Missing resource_access.my-client.roles should return default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - roles_jwt_field="resource_access.my-client.roles" - ) - assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == [ - "default_role" - ] + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") + assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == ["default_role"] # Test 9: Missing nested user roles should return default jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( user_roles_jwt_field="resource_access.my-client.roles", user_allowed_roles=["admin", "user"], ) - assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == [ - "default_user_role" - ] + assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == ["default_user_role"] @pytest.mark.asyncio @@ -1192,9 +1109,7 @@ async def test_metadata_prefix_handling_in_nested_fields(): } # Test 1: metadata.user.email should access user.email after prefix removal - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - user_email_jwt_field="metadata.user.email" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="metadata.user.email") # The get_nested_value function removes "metadata." prefix, so "metadata.user.email" becomes "user.email" assert jwt_handler.get_user_email(token, None) == "user@example.com" @@ -1230,9 +1145,7 @@ async def test_find_team_with_model_access_model_group(monkeypatch): async def mock_get_team_object(*args, **kwargs): # type: ignore return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -1255,6 +1168,57 @@ async def test_find_team_with_model_access_model_group(monkeypatch): assert team_obj.team_id == "team-1" +@pytest.mark.asyncio +@pytest.mark.parametrize( + "model_aliases", + ['{"fast": "gpt-4o"}', {"fast": "gpt-4o"}], + ids=["json-string", "dict"], +) +async def test_find_team_with_model_access_resolves_team_model_alias(monkeypatch, model_aliases): + """LIT-5858: a JWT team that grants `gpt-4o` under the alias `fast` must resolve a request + for `fast`. The JWT path used to pass `team_model_aliases=None`, so every alias request 403'd.""" + import sys + import types + + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + from litellm.router import Router + + router = Router(model_list=[{"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}]) + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + team = LiteLLM_TeamTable( + team_id="team-aliases", + models=["gpt-4o"], + litellm_model_table=LiteLLM_ModelTable(model_aliases=model_aliases, created_by="admin", updated_by="admin"), + ) + + async def mock_get_team_object(*args, **kwargs): + return team + + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + user_api_key_cache = DualCache() + + team_id, team_obj = await JWTAuthManager.find_team_with_model_access( + team_ids={"team-aliases"}, + requested_model="fast", + route="/chat/completions", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=ProxyLogging(user_api_key_cache=user_api_key_cache), + ) + + assert team_id == "team-aliases" + assert team_obj is team + + @pytest.mark.asyncio async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatch): """Regression for #31189: a single-team JWT that grants the requested model @@ -1289,9 +1253,7 @@ async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatc async def mock_get_team_object(*args, **kwargs): # type: ignore return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -1357,14 +1319,10 @@ async def test_auth_builder_returns_team_membership_object(): team_id=_team_id, budget_id="budget_123", spend=10.5, - litellm_budget_table=LiteLLM_BudgetTable( - budget_id="budget_123", rpm_limit=100, tpm_limit=5000 - ), + litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget_123", rpm_limit=100, tpm_limit=5000), ) - user_object = LiteLLM_UserTable( - user_id=_user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id=_user_id, user_role=LitellmUserRoles.INTERNAL_USER) team_object = LiteLLM_TeamTable(team_id=_team_id) @@ -1377,12 +1335,10 @@ async def test_auth_builder_returns_team_membership_object(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object( JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + ) as mock_check_rbac, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1390,9 +1346,7 @@ async def test_auth_builder_returns_team_membership_object(): return_value=(_user_id, "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1407,7 +1361,7 @@ async def test_auth_builder_returns_team_membership_object(): ) as mock_find_team, patch.object( JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + ) as mock_get_all_team_ids, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1428,13 +1382,13 @@ async def test_auth_builder_returns_team_membership_object(): ) as mock_get_objects, patch.object( JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, + ) as mock_map_user, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, + ) as mock_validate_object, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + ) as mock_sync_user, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": _user_id, "scope": ""} @@ -1453,24 +1407,12 @@ async def test_auth_builder_returns_team_membership_object(): ) # Verify that team_membership_object is returned - assert result["team_membership"] is not None, ( - "team_membership should be present" - ) - assert result["team_membership"] == mock_team_membership, ( - "team_membership should match the mock object" - ) - assert result["team_membership"].user_id == _user_id, ( - "team_membership user_id should match" - ) - assert result["team_membership"].team_id == _team_id, ( - "team_membership team_id should match" - ) - assert result["team_membership"].budget_id == "budget_123", ( - "team_membership budget_id should match" - ) - assert result["team_membership"].spend == 10.5, ( - "team_membership spend should match" - ) + assert result["team_membership"] is not None, "team_membership should be present" + assert result["team_membership"] == mock_team_membership, "team_membership should match the mock object" + assert result["team_membership"].user_id == _user_id, "team_membership user_id should match" + assert result["team_membership"].team_id == _team_id, "team_membership team_id should match" + assert result["team_membership"].budget_id == "budget_123", "team_membership budget_id should match" + assert result["team_membership"].spend == 10.5, "team_membership spend should match" @pytest.mark.asyncio @@ -1487,9 +1429,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create JWT handler with OIDC UserInfo enabled jwt_handler = JWTHandler() @@ -1516,18 +1456,14 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): # Mock all the dependencies with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object( JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + ) as mock_check_rbac, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1535,9 +1471,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1552,7 +1486,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): ) as mock_find_team, patch.object( JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + ) as mock_get_all_team_ids, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1567,13 +1501,13 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): ) as mock_get_objects, patch.object( JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, + ) as mock_map_user, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, + ) as mock_validate_object, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object( JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + ) as mock_sync_user, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise ): # Set up mock return values mock_get_userinfo.return_value = userinfo_response @@ -1614,9 +1548,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create JWT handler with OIDC UserInfo disabled jwt_handler = JWTHandler() @@ -1640,18 +1572,14 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): # Mock all the dependencies with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object( JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + ) as mock_check_rbac, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1659,9 +1587,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): return_value=("test_user_1", None, None), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1674,9 +1600,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1689,15 +1613,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, ): # Set up mock return values mock_auth_jwt.return_value = jwt_response @@ -1743,9 +1661,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) jwt_handler = JWTHandler() user_api_key_cache = DualCache() @@ -1764,9 +1680,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() jwt_response = {"sub": "test_user_1", "scope": ""} with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), @@ -1807,9 +1721,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), ): mock_auth_jwt.return_value = jwt_response @@ -1879,9 +1791,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): ) team_object = LiteLLM_TeamTable(team_id="team-2") - user_object = LiteLLM_UserTable( - user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER) with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, @@ -1892,9 +1802,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): new_callable=AsyncMock, return_value=None, ), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, patch.object( JWTAuthManager, "get_objects", @@ -1902,9 +1810,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): return_value=(user_object, None, None, None, user_object.user_id), ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), ): mock_auth_jwt.return_value = { "sub": "user-1", @@ -2121,9 +2027,7 @@ async def test_auth_builder_rbac_team_loads_team_for_passthrough_allowlist(): return_value=(None, None, None, None, None), ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.RouteChecks.is_auth_enforced_pass_through_route", return_value=True, @@ -2152,9 +2056,7 @@ async def test_auth_builder_rbac_team_loads_team_for_passthrough_allowlist(): mock_get_team.assert_awaited_once() assert mock_get_team.await_args.kwargs["team_id"] == "team-rbac" user_api_key_dict = mock_passthrough_check.call_args.kwargs["user_api_key_dict"] - assert user_api_key_dict.team_metadata == { - "allowed_passthrough_routes": ["/my-pass-through"] - } + assert user_api_key_dict.team_metadata == {"allowed_passthrough_routes": ["/my-pass-through"]} @pytest.mark.asyncio @@ -2248,9 +2150,7 @@ async def test_auth_builder_admin_on_llm_route_honors_team_header(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2302,9 +2202,7 @@ async def test_auth_builder_admin_on_mgmt_route_ignores_team_header(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2358,9 +2256,7 @@ async def test_auth_builder_admin_on_llm_route_without_header_unchanged(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2404,9 +2300,7 @@ async def test_get_team_alias_with_nested_fields(): } # Test nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_alias_jwt_field="organization.team.name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_alias_jwt_field="organization.team.name") assert jwt_handler.get_team_alias(nested_token, None) == "engineering-team" # Test flat access (backward compatibility) @@ -2414,9 +2308,7 @@ async def test_get_team_alias_with_nested_fields(): assert jwt_handler.get_team_alias(nested_token, None) == "flat-team" # Test missing field returns default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_alias_jwt_field="nonexistent.field" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_alias_jwt_field="nonexistent.field") assert jwt_handler.get_team_alias(nested_token, "default-team") == "default-team" # Test with team_alias_jwt_field not configured @@ -2447,9 +2339,7 @@ async def test_is_required_team_id_with_team_alias_field(): assert jwt_handler.is_required_team_id() is True # Both fields set - should return True - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_id_jwt_field="team_id", team_alias_jwt_field="team_name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_name") assert jwt_handler.is_required_team_id() is True @@ -2481,9 +2371,7 @@ async def test_find_and_validate_specific_team_id_with_team_alias(): # Mock team object returned by get_team_object_by_alias team_object = LiteLLM_TeamTable(team_id="resolved-team-id", team_alias="my-team") - with patch( - "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock - ) as mock_get_by_alias: + with patch("litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias: mock_get_by_alias.return_value = team_object team_id, result_team = await JWTAuthManager.find_and_validate_specific_team_id( @@ -2526,9 +2414,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): jwt_handler.update_environment( prisma_client=None, user_api_key_cache=user_api_key_cache, - litellm_jwtauth=LiteLLM_JWTAuth( - team_id_jwt_field="team_id", team_alias_jwt_field="team_alias" - ), + litellm_jwtauth=LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"), ) # Token with both team_id and team name @@ -2538,9 +2424,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): team_object = LiteLLM_TeamTable(team_id="direct-team-id") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -2620,9 +2504,7 @@ async def test_get_org_alias_with_nested_fields(): } # Test nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - org_alias_jwt_field="company.organization.name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_alias_jwt_field="company.organization.name") assert jwt_handler.get_org_alias(nested_token, None) == "acme-corp" # Test flat access @@ -2630,9 +2512,7 @@ async def test_get_org_alias_with_nested_fields(): assert jwt_handler.get_org_alias(nested_token, None) == "flat-org" # Test missing field returns default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - org_alias_jwt_field="nonexistent.field" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_alias_jwt_field="nonexistent.field") assert jwt_handler.get_org_alias(nested_token, "default-org") == "default-org" # Test with org_alias_jwt_field not configured @@ -2670,9 +2550,7 @@ async def test_get_objects_resolves_org_by_name(): models=[], ) - with patch( - "litellm.proxy.auth.handle_jwt.get_org_object_by_alias", new_callable=AsyncMock - ) as mock_get_by_alias: + with patch("litellm.proxy.auth.handle_jwt.get_org_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias: mock_get_by_alias.return_value = org_object ( @@ -2750,9 +2628,7 @@ async def test_resolve_jwks_url_resolves_oidc_discovery_document(): litellm_jwtauth=LiteLLM_JWTAuth(), ) - discovery_url = ( - "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" - ) + discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" mock_response = MagicMock() @@ -2783,9 +2659,7 @@ async def test_resolve_jwks_url_caches_resolved_jwks_uri(): litellm_jwtauth=LiteLLM_JWTAuth(), ) - discovery_url = ( - "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" - ) + discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" mock_response = MagicMock() @@ -2929,9 +2803,7 @@ async def test_find_and_validate_specific_team_id_hints_bracket_notation(): error_msg = str(exc_info.value) # Should mention the bad field name and suggest the fix assert "roles.0" in error_msg, f"Expected field name in: {error_msg}" - assert "roles" in error_msg and "list" in error_msg, ( - f"Expected hint about using 'roles' instead: {error_msg}" - ) + assert "roles" in error_msg and "list" in error_msg, f"Expected hint about using 'roles' instead: {error_msg}" @pytest.mark.asyncio @@ -2959,9 +2831,7 @@ async def test_find_and_validate_specific_team_id_hints_bracket_index_notation() error_msg = str(exc_info.value) assert "roles[0]" in error_msg, f"Expected field name in: {error_msg}" - assert "roles" in error_msg and "list" in error_msg, ( - f"Expected hint about using 'roles' instead: {error_msg}" - ) + assert "roles" in error_msg and "list" in error_msg, f"Expected hint about using 'roles' instead: {error_msg}" @pytest.mark.asyncio @@ -3071,9 +2941,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( if len(user_teams) == 1 and get_team_object_return == "resolved_row": only = user_teams[0] team_table = LiteLLM_TeamTable(team_id=only) - membership = LiteLLM_TeamMembership( - user_id=user_id, team_id=only, litellm_budget_table=None - ) + membership = LiteLLM_TeamMembership(user_id=user_id, team_id=only, litellm_budget_table=None) get_team_return_value = team_table membership_return_value = membership else: @@ -3132,9 +3000,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -3151,9 +3017,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( code = 404 if get_team_object_return == "http_404" else 500 mock_get_team.side_effect = HTTPException( status_code=code, - detail={ - "error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call." - }, + detail={"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."}, ) else: mock_get_team.return_value = get_team_return_value @@ -3250,9 +3114,7 @@ async def test_auth_builder_single_team_fallback_membership_error_skips_no_raise ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -3264,9 +3126,7 @@ async def test_auth_builder_single_team_fallback_membership_error_skips_no_raise ): mock_auth_jwt.return_value = {"sub": user_id, "scope": ""} mock_get_team.return_value = team_table - mock_get_membership.side_effect = HTTPException( - status_code=500, detail="membership lookup failed" - ) + mock_get_membership.side_effect = HTTPException(status_code=500, detail="membership lookup failed") result = await JWTAuthManager.auth_builder( api_key="test_jwt_token", @@ -3301,9 +3161,7 @@ def _reset_unscoped_warning_flag(): JWTHandler._unscoped_jwt_warning_emitted = False -def test_build_decode_kwargs_no_env_disables_both_verifications( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_no_env_disables_both_verifications(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.delenv("JWT_AUDIENCE", raising=False) monkeypatch.delenv("JWT_ISSUER", raising=False) @@ -3314,9 +3172,7 @@ def test_build_decode_kwargs_no_env_disables_both_verifications( assert kwargs["options"] == {"verify_aud": False, "verify_iss": False} -def test_build_decode_kwargs_audience_only_enables_aud_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_audience_only_enables_aud_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") monkeypatch.delenv("JWT_ISSUER", raising=False) @@ -3328,9 +3184,7 @@ def test_build_decode_kwargs_audience_only_enables_aud_verification( assert kwargs["options"] == {"verify_iss": False} -def test_build_decode_kwargs_issuer_only_enables_iss_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_issuer_only_enables_iss_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.delenv("JWT_AUDIENCE", raising=False) monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") @@ -3341,9 +3195,7 @@ def test_build_decode_kwargs_issuer_only_enables_iss_verification( assert kwargs["options"] == {"verify_aud": False} -def test_build_decode_kwargs_both_set_enables_full_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_both_set_enables_full_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") @@ -3355,9 +3207,7 @@ def test_build_decode_kwargs_both_set_enables_full_verification( assert kwargs["options"] is None -def test_build_decode_kwargs_warns_once_when_unscoped( - monkeypatch, _reset_unscoped_warning_flag, caplog -): +def test_build_decode_kwargs_warns_once_when_unscoped(monkeypatch, _reset_unscoped_warning_flag, caplog): """The warning about unscoped JWT auth should fire on the first call but not on every subsequent decode.""" import logging @@ -3373,17 +3223,12 @@ def test_build_decode_kwargs_warns_once_when_unscoped( matching = [ r for r in caplog.records - if "JWT auth is enabled" in r.getMessage() - and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + if "JWT auth is enabled" in r.getMessage() and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() ] - assert len(matching) == 1, ( - f"Expected exactly one warning across 3 calls, got {len(matching)}" - ) + assert len(matching) == 1, f"Expected exactly one warning across 3 calls, got {len(matching)}" -def test_build_decode_kwargs_no_warning_when_scoped( - monkeypatch, _reset_unscoped_warning_flag, caplog -): +def test_build_decode_kwargs_no_warning_when_scoped(monkeypatch, _reset_unscoped_warning_flag, caplog): import logging monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") @@ -3392,11 +3237,7 @@ def test_build_decode_kwargs_no_warning_when_scoped( JWTHandler._build_decode_kwargs() - matching = [ - r - for r in caplog.records - if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() - ] + matching = [r for r in caplog.records if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()] assert matching == [] @@ -3454,11 +3295,7 @@ async def test_find_team_with_model_access_unresolved_group_claim_returns_none( from litellm.router import Router - router = Router( - model_list=[ - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} - ] - ) + router = Router(model_list=[{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}]) proxy_server_module = types.ModuleType("proxy_server") proxy_server_module.llm_router = router monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) @@ -3530,9 +3367,7 @@ async def test_find_and_validate_specific_team_id_non_404_http_exception_propaga "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, ) as mock_get_team: - mock_get_team.side_effect = HTTPException( - status_code=status_code, detail="non-404 failure" - ) + mock_get_team.side_effect = HTTPException(status_code=status_code, detail="non-404 failure") with pytest.raises(HTTPException) as exc_info: await JWTAuthManager.find_and_validate_specific_team_id( @@ -3605,9 +3440,7 @@ async def test_find_team_with_model_access_resolved_team_without_model_still_rai async def mock_get_team_object(*_args, **_kwargs): return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -3672,11 +3505,7 @@ async def test_find_team_with_model_access_unresolved_group_claim_default_raises from litellm.router import Router - router = Router( - model_list=[ - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} - ] - ) + router = Router(model_list=[{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}]) proxy_server_module = types.ModuleType("proxy_server") proxy_server_module.llm_router = router monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) @@ -3713,12 +3542,7 @@ def test_canonical_user_id_rebinds_to_legacy_uuid(): jwt_email = "matt@example.com" user_object = LiteLLM_UserTable(user_id=legacy_uuid, user_email=jwt_email) - assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id=jwt_email, user_object=user_object - ) - == legacy_uuid - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=jwt_email, user_object=user_object) == legacy_uuid def test_canonical_user_id_no_change_when_ids_match(): @@ -3726,28 +3550,20 @@ def test_canonical_user_id_no_change_when_ids_match(): same = "alice@example.com" user_object = LiteLLM_UserTable(user_id=same, user_email=same) - assert ( - JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object) - == same - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object) == same def test_canonical_user_id_returns_claim_when_no_user_object(): """No resolved row (e.g. upsert disabled / brand new) -> keep the claim.""" assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id="newcomer@example.com", user_object=None - ) + JWTAuthManager._canonical_user_id_from_db(user_id="newcomer@example.com", user_object=None) == "newcomer@example.com" ) def test_canonical_user_id_returns_none_when_claim_none_and_no_object(): """Defensive: no claim and no row -> stays None, never invents an id.""" - assert ( - JWTAuthManager._canonical_user_id_from_db(user_id=None, user_object=None) - is None - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=None, user_object=None) is None def test_canonical_user_id_no_change_when_db_user_id_falsy(): @@ -3757,10 +3573,7 @@ def test_canonical_user_id_no_change_when_db_user_id_falsy(): user_id = "" assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id="jwt@example.com", user_object=_Stub() - ) - == "jwt@example.com" + JWTAuthManager._canonical_user_id_from_db(user_id="jwt@example.com", user_object=_Stub()) == "jwt@example.com" ) @@ -3776,9 +3589,7 @@ async def test_auth_jwt_expired_token_raises_401_jwk_path(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() with ( - patch.object( - jwt_handler, "get_public_key", new_callable=AsyncMock - ) as mock_get_public_key, + patch.object(jwt_handler, "get_public_key", new_callable=AsyncMock) as mock_get_public_key, patch( "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", return_value={"kid": "test-kid"}, @@ -3814,9 +3625,7 @@ async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): mock_cert.public_key.return_value.public_bytes.return_value = b"fake-key" with ( - patch.object( - jwt_handler, "get_public_key", new_callable=AsyncMock - ) as mock_get_public_key, + patch.object(jwt_handler, "get_public_key", new_callable=AsyncMock) as mock_get_public_key, patch( "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", return_value={"kid": "test-kid"}, @@ -3830,9 +3639,7 @@ async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): side_effect=jwt_lib.ExpiredSignatureError("Signature has expired"), ), ): - mock_get_public_key.return_value = ( - "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" - ) + mock_get_public_key.return_value = "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" with pytest.raises(ProxyException) as exc_info: await jwt_handler.auth_jwt(token="expired.jwt.token") @@ -3946,9 +3753,7 @@ async def test_get_public_key_fetches_and_caches_jwks_response(): ) assert public_key == jwk - cached_keys = await cache.async_get_cache( - key="litellm_jwt_auth_keys_https://issuer.example.com/keys" - ) + cached_keys = await cache.async_get_cache(key="litellm_jwt_auth_keys_https://issuer.example.com/keys") assert cached_keys == [jwk] @@ -4193,9 +3998,7 @@ async def test_lowering_public_key_stale_ttl_stops_serving_a_copy_cached_under_t # The operator tightens the window and restarts; the cache, and its long-lived copy, survive. endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) - tightened = _get_jwt_handler_with_scripted_endpoint( - cache, endpoint, public_key_stale_ttl=lowered_stale_ttl - ) + tightened = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_stale_ttl=lowered_stale_ttl) await cache.async_set_cache( key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}", value=time.time() - 7200, @@ -4543,9 +4346,7 @@ def test_get_jwks_url_for_issuer_falls_back_to_discovery_document(): jwks_url = jwt_handler._get_jwks_url_for_issuer(issuer_config=issuer_config) - assert ( - jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" - ) + assert jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" @pytest.mark.asyncio @@ -4569,9 +4370,7 @@ async def test_get_objects_team_membership_uses_rebound_user_id(): return None jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - user_id_jwt_field="email", user_id_upsert=True - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="email", user_id_upsert=True) with ( patch( @@ -4667,9 +4466,7 @@ async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims( assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-org" - assert jwt_handler.get_team_id(token=claims, default_value=None) == ( - "example-org/litellm-fork" - ) + assert jwt_handler.get_team_id(token=claims, default_value=None) == ("example-org/litellm-fork") @pytest.mark.asyncio @@ -4772,9 +4569,7 @@ async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): claims = await jwt_handler.auth_jwt(token=token) - assert ( - jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" - ) + assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" @pytest.mark.asyncio @@ -4807,7 +4602,7 @@ async def test_multi_issuer_jwt_unknown_issuer_falls_back_to_global_jwks(monkeyp kid="issuer-key", ) - with pytest.raises(Exception, match='Missing JWT Public Key URL from environment\\.') as exc: + with pytest.raises(Exception, match="Missing JWT Public Key URL from environment\\.") as exc: await jwt_handler.auth_jwt(token=token) assert "Missing JWT Public Key URL from environment." in str(exc.value) @@ -4881,7 +4676,7 @@ async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch) kid=shared_kid, ) - with pytest.raises(Exception, match='Validation fails: Signature verification failed') as exc: + with pytest.raises(Exception, match="Validation fails: Signature verification failed") as exc: await jwt_handler.auth_jwt(token=token) assert "Validation fails" in str(exc.value) @@ -4936,7 +4731,7 @@ def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( issuer = "https://issuer.example.com" jwks_url = f"{issuer}/keys" - with pytest.raises(Exception, match='must configure audience or set') as exc: + with pytest.raises(Exception, match="must configure audience or set") as exc: LiteLLM_JWTAuth( issuers=[ { @@ -4953,7 +4748,7 @@ def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): issuer = "https://issuer.example.com" jwks_url = f"{issuer}/keys" - with pytest.raises(Exception, match='cannot set audience and disable_audience_validation=True') as exc: + with pytest.raises(Exception, match="cannot set audience and disable_audience_validation=True") as exc: LiteLLM_JWTAuth( issuers=[ { @@ -4965,9 +4760,7 @@ def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): ] ) - assert "cannot set audience and disable_audience_validation=True together" in str( - exc.value - ) + assert "cannot set audience and disable_audience_validation=True together" in str(exc.value) @pytest.mark.asyncio @@ -5020,21 +4813,15 @@ async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): claims = await jwt_handler.auth_jwt(token=token) - assert jwt_handler.get_user_id(token=claims, default_value=None) == ( - "real-user@example.com" - ) - assert jwt_handler.get_user_email(token=claims, default_value=None) == ( - "real-user@example.com" - ) + assert jwt_handler.get_user_id(token=claims, default_value=None) == ("real-user@example.com") + assert jwt_handler.get_user_email(token=claims, default_value=None) == ("real-user@example.com") assert jwt_handler.get_team_id(token=claims, default_value=None) == "real-team" assert jwt_handler.get_team_ids_from_jwt(token=claims) == [ "real-team", "secondary-team", ] assert jwt_handler.get_org_id(token=claims, default_value=None) == "real-org" - assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ( - "real-end-user" - ) + assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ("real-end-user") @pytest.mark.asyncio @@ -5074,15 +4861,11 @@ async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): assert JWTHandler.LITELLM_TEAM_ID_CLAIM not in claims assert jwt_handler.get_user_id(token=claims, default_value=None) is None assert jwt_handler.get_team_id(token=claims, default_value=None) is None - assert jwt_handler.get_user_email(token=claims, default_value=None) == ( - "real-user@example.com" - ) + assert jwt_handler.get_user_email(token=claims, default_value=None) == ("real-user@example.com") @pytest.mark.asyncio -async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( - monkeypatch, caplog -): +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning(monkeypatch, caplog): import logging monkeypatch.delenv("JWT_AUDIENCE", raising=False) @@ -5132,11 +4915,7 @@ def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deploym JWTHandler._build_decode_kwargs() - matching = [ - r - for r in caplog.records - if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() - ] + matching = [r for r in caplog.records if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()] assert len(matching) == 1 @@ -5273,9 +5052,7 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership(): "expect_403", ), [ - pytest.param( - True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team" - ), + pytest.param(True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"), pytest.param( True, ["team_a", "team_b"], @@ -5353,9 +5130,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( async def call_auth_builder(): with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock - ) as mock_auth_jwt, + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -5395,9 +5170,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -5501,9 +5274,7 @@ async def test_resolve_db_team_fallback_skips_team_without_model_access(): teams=["restricted_team", "allowed_team"], ) teams = { - "restricted_team": LiteLLM_TeamTable( - team_id="restricted_team", models=["claude-3"] - ), + "restricted_team": LiteLLM_TeamTable(team_id="restricted_team", models=["claude-3"]), "allowed_team": LiteLLM_TeamTable(team_id="allowed_team", models=["gpt-4"]), } @@ -5645,9 +5416,7 @@ async def _run_auth_builder_with_header_team( jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = jwt_auth_config with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value=token - ), + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value=token), patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -5666,9 +5435,7 @@ async def _run_auth_builder_with_header_team( new_callable=AsyncMock, return_value=None, ), - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=allowed_team_ids - ), + patch.object(JWTAuthManager, "get_all_team_ids", return_value=allowed_team_ids), patch.object( JWTAuthManager, "get_objects", @@ -5677,9 +5444,7 @@ async def _run_auth_builder_with_header_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -5708,9 +5473,7 @@ async def _team_lookup_404(team_id, **kwargs): @pytest.mark.asyncio -async def test_auth_builder_header_team_not_found_matches_non_membership_denial() -> ( - None -): +async def test_auth_builder_header_team_not_found_matches_non_membership_denial() -> None: """A provisional x-litellm-team-id naming a nonexistent team must produce the exact same 403 shape as one naming an existing team outside the caller's memberships. Letting get_team_object's 404 surface would give any @@ -5730,19 +5493,15 @@ async def test_auth_builder_header_team_not_found_matches_non_membership_denial( return LiteLLM_TeamTable(team_id=team_id) with pytest.raises(HTTPException) as missing_exc: - await _run_auth_builder_with_header_team( - config, token, "team_ghost", user_object, _team_lookup_404, set() - ) + await _run_auth_builder_with_header_team(config, token, "team_ghost", user_object, _team_lookup_404, set()) with pytest.raises(HTTPException) as outsider_exc: - await _run_auth_builder_with_header_team( - config, token, "team_other", user_object, team_exists, set() - ) + await _run_auth_builder_with_header_team(config, token, "team_other", user_object, team_exists, set()) assert missing_exc.value.status_code == 403 assert outsider_exc.value.status_code == 403 - assert missing_exc.value.detail.replace( - "team_ghost", "" - ) == outsider_exc.value.detail.replace("team_other", "") + assert missing_exc.value.detail.replace("team_ghost", "") == outsider_exc.value.detail.replace( + "team_other", "" + ) assert "exist" not in missing_exc.value.detail @@ -5939,9 +5698,7 @@ async def test_auth_builder_db_fallback_does_not_validate_rbac_team_against_db_m ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6093,9 +5850,7 @@ async def test_auth_builder_db_fallback_runs_when_only_team_id_default_set(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6177,9 +5932,7 @@ async def test_auth_builder_alias_only_token_resolves_alias_not_db_fallback(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6236,14 +5989,10 @@ async def test_find_and_validate_specific_team_id_alias_wins_over_team_id_defaul ) jwt_token = {"sub": "user-1", "team_alias": "my-team"} - alias_team = LiteLLM_TeamTable( - team_id="alias_resolved_team", team_alias="my-team" - ) + alias_team = LiteLLM_TeamTable(team_id="alias_resolved_team", team_alias="my-team") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -6290,9 +6039,7 @@ async def test_find_and_validate_specific_team_id_team_id_default_used_without_a default_team = LiteLLM_TeamTable(team_id="config_default_team") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -6361,9 +6108,7 @@ async def test_auth_builder_db_fallback_enforces_passthrough_route_access(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6433,18 +6178,14 @@ async def test_sync_user_role_and_teams_singular_claim_reconciles_memberships(): "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ) as mock_patch: - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, AsyncMock() - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, AsyncMock()) mock_patch.assert_awaited_once() assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == { "team_stale_a", "team_stale_b", } - assert set(mock_patch.call_args.kwargs["teams_ids_to_add_user_to"]) == { - "team_primary" - } + assert set(mock_patch.call_args.kwargs["teams_ids_to_add_user_to"]) == {"team_primary"} assert user.teams == ["team_primary"] @@ -6503,9 +6244,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6580,9 +6319,7 @@ async def test_auth_builder_header_cannot_override_rbac_team_under_db_fallback() ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6632,9 +6369,7 @@ async def test_auth_builder_header_team_enforces_team_allowed_routes_under_db_fa async def call(route: str): with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock - ) as mock_auth_jwt, + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -6662,9 +6397,7 @@ async def test_auth_builder_header_team_enforces_team_allowed_routes_under_db_fa ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6724,13 +6457,9 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ) as mock_patch: - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, AsyncMock() - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, AsyncMock()) mock_patch.assert_awaited_once() - assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == { - "team_existing" - } + assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == {"team_existing"} assert mock_patch.call_args.kwargs["teams_ids_to_add_user_to"] == [] assert user.teams == [] diff --git a/tests/test_litellm/proxy/auth/test_team_grants.py b/tests/test_litellm/proxy/auth/test_team_grants.py new file mode 100644 index 00000000000..447fc1c93a1 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_team_grants.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index d44f96d95bf..f7bfea6c7ed 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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"} diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index b9a884e910c..b9affeccebb 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -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" diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index ffa6bc601e9..c2f29eb9bad 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -173,9 +173,7 @@ mock_prisma_client.db.litellm_auditlog.create = AsyncMock() # Fixture to provide the mock prisma client @pytest.fixture(autouse=True) def mock_db_client(): - with patch( - "litellm.proxy.proxy_server.prisma_client", mock_prisma_client - ): # Mock in both places if necessary + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): # Mock in both places if necessary yield mock_prisma_client mock_prisma_client.reset_mock() @@ -218,27 +216,17 @@ async def test_validate_team_org_change_same_org_id(): organization.organization_id = org_id organization.models = [] organization.litellm_budget_table = MagicMock() - organization.litellm_budget_table.max_budget = ( - 50.0 # This would normally fail validation - ) - organization.litellm_budget_table.tpm_limit = ( - 500 # This would normally fail validation - ) - organization.litellm_budget_table.rpm_limit = ( - 50 # This would normally fail validation - ) + organization.litellm_budget_table.max_budget = 50.0 # This would normally fail validation + organization.litellm_budget_table.tpm_limit = 500 # This would normally fail validation + organization.litellm_budget_table.rpm_limit = 50 # This would normally fail validation organization.members = [] # Mock Router mock_router = MagicMock(spec=Router) # Use patch to ensure the model access check is never called - with patch( - "litellm.proxy.management_endpoints.team_endpoints.can_org_access_model" - ) as mock_access_check: - result = validate_team_org_change( - team=team, organization=organization, llm_router=mock_router - ) + with patch("litellm.proxy.management_endpoints.team_endpoints.can_org_access_model") as mock_access_check: + result = validate_team_org_change(team=team, organization=organization, llm_router=mock_router) # Assert the function returns True without checking anything assert result is True @@ -290,9 +278,7 @@ async def test_validate_team_org_change_members_in_org(): mock_router = MagicMock(spec=Router) # Test should pass - all team members are in org members - result = validate_team_org_change( - team=team, organization=organization, llm_router=mock_router - ) + result = validate_team_org_change(team=team, organization=organization, llm_router=mock_router) assert result is True @@ -344,9 +330,7 @@ async def test_validate_team_org_change_member_not_in_org(): # Test should fail - user_id_not_in_org is not in org members with pytest.raises(HTTPException) as exc_info: - validate_team_org_change( - team=team, organization=organization, llm_router=mock_router - ) + validate_team_org_change(team=team, organization=organization, llm_router=mock_router) assert exc_info.value.status_code == 403 assert "not a member of the organization" in str(exc_info.value.detail) @@ -390,10 +374,7 @@ async def test_get_team_permissions_list_success(mock_db_client, mock_admin_auth assert response.status_code == 200 response_data = response.json() assert response_data["team_id"] == test_team_id - assert ( - response_data["team_member_permissions"] - == mock_team_data["team_member_permissions"] - ) + assert response_data["team_member_permissions"] == mock_team_data["team_member_permissions"] assert ( response_data["all_available_permissions"] == TeamMemberPermissionChecks.get_all_available_team_member_permissions() @@ -456,9 +437,7 @@ async def test_update_team_permissions_success(mock_db_client, mock_admin_auth): return_value=mock_existing_team_row, ): # Mock the database update function - mock_db_client.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team_row - ) + mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team_row) # Override the dependency for this test app.dependency_overrides[user_api_key_auth] = lambda: mock_admin_auth @@ -483,9 +462,7 @@ async def test_update_team_permissions_success(mock_db_client, mock_admin_auth): @pytest.mark.asyncio @pytest.mark.parametrize("field", ["budget_duration", "team_member_budget_duration"]) @pytest.mark.parametrize("bad_duration", ["0s", "-5m"]) -async def test_new_team_rejects_a_duration_that_never_advances( - mock_db_client, mock_admin_auth, field, bad_duration -): +async def test_new_team_rejects_a_duration_that_never_advances(mock_db_client, mock_admin_auth, field, bad_duration): """A zero-length window resets to "now", so the team row is due again the moment it is written. The reset job re-reads such rows on every tick, and a tenant with enough of them fills each batch and starves other tenants. @@ -515,9 +492,7 @@ async def test_new_team_rejects_a_duration_that_never_advances( @pytest.mark.asyncio @pytest.mark.parametrize("field", ["budget_duration", "team_member_budget_duration"]) -async def test_update_team_rejects_a_duration_that_never_advances( - mock_db_client, mock_admin_auth, field -): +async def test_update_team_rejects_a_duration_that_never_advances(mock_db_client, mock_admin_auth, field): """/team/update must reject the same never-advancing durations /team/new does.""" from fastapi import Request @@ -556,17 +531,13 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth): mock_db_client.db = MagicMock() # Mock object permission table creation - mock_object_perm_create = AsyncMock( - return_value=MagicMock(object_permission_id="objperm123") - ) + mock_object_perm_create = AsyncMock(return_value=MagicMock(object_permission_id="objperm123")) mock_db_client.db.litellm_objectpermissiontable = MagicMock() mock_db_client.db.litellm_objectpermissiontable.create = mock_object_perm_create # Mock model table creation mock_db_client.db.litellm_modeltable = MagicMock() - mock_db_client.db.litellm_modeltable.create = AsyncMock( - return_value=MagicMock(id="model123") - ) + mock_db_client.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model123")) # Capture team table creation team_create_result = MagicMock( @@ -583,9 +554,7 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth): mock_db_client.db.litellm_teamtable.create = mock_team_create _wire_team_create_tx(mock_db_client) mock_db_client.db.litellm_teamtable.count = mock_team_count - mock_db_client.db.litellm_teamtable.update = AsyncMock( - return_value=team_create_result - ) + mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=team_create_result) # Mock user table mock_db_client.db.litellm_usertable = MagicMock() @@ -654,9 +623,7 @@ async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_aut # Mock model table mock_db_client.db.litellm_modeltable = MagicMock() - mock_db_client.db.litellm_modeltable.create = AsyncMock( - return_value=MagicMock(id="model456") - ) + mock_db_client.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model456")) # Mock team table team_create_result = MagicMock( @@ -668,14 +635,10 @@ async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_aut "object_permission_id": "objperm_team_mcp_456", } mock_db_client.db.litellm_teamtable = MagicMock() - mock_db_client.db.litellm_teamtable.create = AsyncMock( - return_value=team_create_result - ) + mock_db_client.db.litellm_teamtable.create = AsyncMock(return_value=team_create_result) _wire_team_create_tx(mock_db_client) mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0) - mock_db_client.db.litellm_teamtable.update = AsyncMock( - return_value=team_create_result - ) + mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=team_create_result) # Mock user table mock_db_client.db.litellm_usertable = MagicMock() @@ -734,20 +697,14 @@ def test_should_auto_add_team_creator(user_role, user_id, flag_value, expected): _should_auto_add_team_creator, ) - general_settings = ( - {} if flag_value is None else {"disable_auto_add_proxy_admin_to_teams": flag_value} - ) + general_settings = {} if flag_value is None else {"disable_auto_add_proxy_admin_to_teams": flag_value} auth = UserAPIKeyAuth(user_role=user_role, user_id=user_id) assert _should_auto_add_team_creator(auth, general_settings) is expected @pytest.mark.asyncio -@pytest.mark.parametrize( - "disable_flag,expect_creator_added", [(True, False), (False, True)] -) -async def test_new_team_disable_auto_add_proxy_admin_flag( - mock_db_client, disable_flag, expect_creator_added -): +@pytest.mark.parametrize("disable_flag,expect_creator_added", [(True, False), (False, True)]) +async def test_new_team_disable_auto_add_proxy_admin_flag(mock_db_client, disable_flag, expect_creator_added): """ When general_settings.disable_auto_add_proxy_admin_to_teams is True, a proxy admin calling /team/new must NOT be auto-added to the team's members. When @@ -762,9 +719,7 @@ async def test_new_team_disable_auto_add_proxy_admin_flag( team_create_result = MagicMock(team_id="team-789") team_create_result.model_dump.return_value = {"team_id": "team-789"} mock_db_client.db.litellm_teamtable = MagicMock() - mock_db_client.db.litellm_teamtable.create = AsyncMock( - return_value=team_create_result - ) + mock_db_client.db.litellm_teamtable.create = AsyncMock(return_value=team_create_result) _wire_team_create_tx(mock_db_client) mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_db_client.db.litellm_usertable = MagicMock() @@ -775,17 +730,18 @@ async def test_new_team_disable_auto_add_proxy_admin_flag( from litellm.proxy._types import NewTeamRequest from litellm.proxy.management_endpoints.team_endpoints import new_team - admin_auth = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user-1" - ) + admin_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user-1") - with patch( - "litellm.proxy.proxy_server.general_settings", - {"disable_auto_add_proxy_admin_to_teams": disable_flag}, - ), patch( - "litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team", - new_callable=AsyncMock, - ) as mock_add_members: + with ( + patch( + "litellm.proxy.proxy_server.general_settings", + {"disable_auto_add_proxy_admin_to_teams": disable_flag}, + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team", + new_callable=AsyncMock, + ) as mock_add_members, + ): await new_team( data=NewTeamRequest(team_alias="flag-test-team"), http_request=MagicMock(spec=Request), @@ -834,16 +790,12 @@ async def test_team_update_object_permissions_existing_permission(monkeypatch): "vector_stores": ["old_store_1", "old_store_2"], } - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=existing_object_permission - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_object_permission) # Mock upsert operation updated_permission = MagicMock() updated_permission.object_permission_id = "existing_perm_id_123" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=updated_permission - ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(return_value=updated_permission) # Test data with new object permission data_json = { @@ -899,21 +851,17 @@ async def test_team_update_object_permissions_no_existing_permission(monkeypatch ) # Mock find_unique to return None (no existing permission) - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) # Mock upsert to create new record new_permission = MagicMock() new_permission.object_permission_id = "new_perm_id_456" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=new_permission - ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(return_value=new_permission) data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["brand_new_store"] - ).model_dump(exclude_unset=True, exclude_none=True), + "object_permission": LiteLLM_ObjectPermissionBase(vector_stores=["brand_new_store"]).model_dump( + exclude_unset=True, exclude_none=True + ), "team_alias": "updated_team_2", } @@ -959,21 +907,17 @@ async def test_team_update_object_permissions_missing_permission_record(monkeypa ) # Mock find_unique to return None (permission record not found) - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) # Mock upsert to create new record new_permission = MagicMock() new_permission.object_permission_id = "recreated_perm_id_789" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=new_permission - ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(return_value=new_permission) data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["recreated_store"] - ).model_dump(exclude_unset=True, exclude_none=True), + "object_permission": LiteLLM_ObjectPermissionBase(vector_stores=["recreated_store"]).model_dump( + exclude_unset=True, exclude_none=True + ), "team_alias": "updated_team_3", } @@ -1079,14 +1023,10 @@ async def test_add_team_member_budget_table_success(): mock_budget_record.budget_id = "budget-123" mock_budget_record.max_budget = 1000.0 - mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock( - return_value=mock_budget_record - ) + mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=mock_budget_record) # Create team info response object - team_info_response = TeamInfoResponseObjectTeamTable( - team_id="test-team-123", team_alias="Test Team" - ) + team_info_response = TeamInfoResponseObjectTeamTable(team_id="test-team-123", team_alias="Test Team") # Call the function result = await _add_team_member_budget_table( @@ -1097,15 +1037,11 @@ async def test_add_team_member_budget_table_success(): # Verify the result assert result.team_member_budget_table == mock_budget_record - assert result == team_info_response.model_copy( - update={"team_member_budget_table": mock_budget_record} - ) + assert result == team_info_response.model_copy(update={"team_member_budget_table": mock_budget_record}) assert team_info_response.team_member_budget_table is None # Verify database call was made correctly - mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with( - where={"budget_id": "budget-123"} - ) + mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with(where={"budget_id": "budget-123"}) @pytest.mark.asyncio @@ -1125,14 +1061,10 @@ async def test_add_team_member_budget_table_exception_handling(): ) # Create team info response object - team_info_response = TeamInfoResponseObjectTeamTable( - team_id="test-team-456", team_alias="Test Team 2" - ) + team_info_response = TeamInfoResponseObjectTeamTable(team_id="test-team-456", team_alias="Test Team 2") # Mock the verbose_proxy_logger to capture log calls - with patch( - "litellm.proxy.management_endpoints.team_endpoints.verbose_proxy_logger" - ) as mock_logger: + with patch("litellm.proxy.management_endpoints.team_endpoints.verbose_proxy_logger") as mock_logger: # Call the function result = await _add_team_member_budget_table( team_member_budget_id="nonexistent-budget-456", @@ -1144,10 +1076,7 @@ async def test_add_team_member_budget_table_exception_handling(): assert result == team_info_response # Verify team_member_budget_table is not set when exception occurs - assert ( - not hasattr(result, "team_member_budget_table") - or result.team_member_budget_table is None - ) + assert not hasattr(result, "team_member_budget_table") or result.team_member_budget_table is None # Verify the error was logged mock_logger.info.assert_called_once_with( @@ -1176,9 +1105,7 @@ async def test_add_team_member_budget_table_budget_not_found(): mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) # Create team info response object - team_info_response = TeamInfoResponseObjectTeamTable( - team_id="test-team-789", team_alias="Test Team 3" - ) + team_info_response = TeamInfoResponseObjectTeamTable(team_id="test-team-789", team_alias="Test Team 3") # Call the function result = await _add_team_member_budget_table( @@ -1395,9 +1322,7 @@ async def test_available_team_self_join_blocks_other_user_id(): await _validate_team_member_add_permissions( user_api_key_dict=user, complete_team_data=team, - data=_make_team_member_add_request( - member_user_id="bob-victim", role="user" - ), + data=_make_team_member_add_request(member_user_id="bob-victim", role="user"), ) assert exc_info.value.status_code == 403 @@ -1836,9 +1761,7 @@ async def test_update_team_members_list_duplicate_prevention(): # Create mock team with existing members mock_team = MagicMock(spec=LiteLLM_TeamTable) - mock_team.members_with_roles = [ - Member(user_id="existing-user", user_email="existing@example.com", role="admin") - ] + mock_team.members_with_roles = [Member(user_id="existing-user", user_email="existing@example.com", role="admin")] # Try to add the same member again duplicate_member = Member(user_id="existing-user", role="user") @@ -1968,9 +1891,7 @@ async def test_add_team_members_runs_member_writes_on_the_lock_holding_transacti tx = MagicMock() tx.query_raw = AsyncMock(return_value=[{"members_with_roles": []}]) - tx.litellm_teamtable.update = AsyncMock( - return_value=LiteLLM_TeamTable(team_id="team-pool", members_with_roles=[]) - ) + tx.litellm_teamtable.update = AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-pool", members_with_roles=[])) tx.litellm_usertable.upsert = AsyncMock(return_value=added_user) tx.litellm_usertable.update_many = AsyncMock() tx.litellm_budgettable.create = AsyncMock(return_value=created_budget) @@ -2106,9 +2027,7 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name): ) mock_request = Mock(spec=Request) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") existing_team = MagicMock() existing_team.model_dump.return_value = { @@ -2145,12 +2064,8 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name): new_callable=AsyncMock, ) as mock_cache_team, ): - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=updated_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=updated_team) mock_prisma_client.db.execute_raw = AsyncMock(return_value=None) if endpoint_name == "team_model_add": @@ -2192,18 +2107,58 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name): # "no team-level restriction" and stop enforcing the team's # search-tool allowlist on key issuance. assert call_kwargs["team_table"].object_permission is not None - assert call_kwargs["team_table"].object_permission.search_tools == [ - "allowed-tool-A" - ] + assert call_kwargs["team_table"].object_permission.search_tools == ["allowed-tool-A"] # Pin the Prisma call shape too — the regression is in *what the # update returns*, so the contract that the update asks for # `object_permission` belongs in this test. - update_call_kwargs = ( - mock_prisma_client.db.litellm_teamtable.update.call_args.kwargs - ) + update_call_kwargs = mock_prisma_client.db.litellm_teamtable.update.call_args.kwargs assert update_call_kwargs.get("include", {}).get("object_permission") is True +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint_name", ["team_model_add", "team_model_delete"]) +async def test_team_model_add_delete_keep_model_aliases_in_team_cache(endpoint_name, monkeypatch): + """LIT-5858: Prisma only returns `litellm_model_table` when the `update` asks for it, so the refreshed + cache entry lost the team's model aliases and JWT alias requests 403'd until the next DB read.""" + from litellm.proxy._types import TeamModelAddRequest, TeamModelDeleteRequest + from litellm.proxy.auth.team_grants import team_model_aliases + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.management_endpoints.team_endpoints import team_model_add, team_model_delete + + columns = {"team_id": "team-1234", "models": ["gpt-4o", "openai/*"]} + alias_table = {"id": 1, "model_aliases": '{"fast": "gpt-4o"}', "created_by": "admin", "updated_by": "admin"} + + async def update(where, data, include=None): + row = {**columns, "litellm_model_table": alias_table} if (include or {}).get("litellm_model_table") else columns + return SimpleNamespace(team_id="team-1234", model_dump=lambda: row) + + prisma_client = MagicMock() + prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(model_dump=lambda: columns)) + prisma_client.db.litellm_teamtable.update = AsyncMock(side_effect=update) + prisma_client.db.execute_raw = AsyncMock(return_value=None) + cache = UserApiKeyCache() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", None) + + admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin") + if endpoint_name == "team_model_add": + await team_model_add( + data=TeamModelAddRequest(team_id="team-1234", models=["team-byok-1"]), + http_request=MagicMock(), + user_api_key_dict=admin, + ) + else: + await team_model_delete( + data=TeamModelDeleteRequest(team_id="team-1234", models=["openai/*"]), + http_request=MagicMock(), + user_api_key_dict=admin, + ) + + cached_team = await cache.async_get_cache(key="team_id:team-1234", model_type=LiteLLM_TeamTableCachedObj) + assert team_model_aliases(cached_team) == {"fast": "gpt-4o"} + + @pytest.mark.asyncio @pytest.mark.parametrize( "endpoint_name", @@ -2237,9 +2192,7 @@ async def test_team_write_404s_when_row_vanishes_before_update(endpoint_name): ) mock_request = Mock(spec=Request) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") existing_team = MagicMock() existing_team.team_id = "team-1234" @@ -2272,9 +2225,15 @@ async def test_team_write_404s_when_row_vanishes_before_update(endpoint_name): }[endpoint_name] with ( - patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, # test-quality-ok: proxy_server module global is the endpoint's only injection point - patch("litellm.proxy.proxy_server.user_api_key_cache"), # test-quality-ok: proxy_server module global is the endpoint's only injection point - patch("litellm.proxy.proxy_server.proxy_logging_obj"), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch( + "litellm.proxy.proxy_server.prisma_client" + ) as mock_prisma_client, # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch( + "litellm.proxy.proxy_server.user_api_key_cache" + ), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch( + "litellm.proxy.proxy_server.proxy_logging_obj" + ), # test-quality-ok: proxy_server module global is the endpoint's only injection point patch( # test-quality-ok: stubs the cache write so the test observes only the DB result handling "litellm.proxy.management_endpoints.team_endpoints._cache_team_object", new_callable=AsyncMock, @@ -2285,9 +2244,7 @@ async def test_team_write_404s_when_row_vanishes_before_update(endpoint_name): return_value=existing_team, ), ): - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=None) mock_prisma_client.db.execute_raw = AsyncMock(return_value=None) @@ -2318,9 +2275,7 @@ async def test_update_team_team_member_budget_not_passed_to_db( # Mock dependencies mock_request = Mock(spec=Request) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") with ( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, @@ -2328,9 +2283,7 @@ async def test_update_team_team_member_budget_not_passed_to_db( patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object" - ) as mock_cache_team, + patch("litellm.proxy.management_endpoints.team_endpoints._cache_team_object") as mock_cache_team, patch( "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" ) as mock_upsert_budget, @@ -2342,20 +2295,14 @@ async def test_update_team_team_member_budget_not_passed_to_db( "team_alias": "test_team", "metadata": {"team_member_budget_id": "budget_123"}, } - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) # Mock the update return value mock_updated_team = MagicMock() mock_updated_team.team_id = "test_team_id" mock_updated_team.model_dump.return_value = {"team_id": "test_team_id"} - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) - mock_prisma_client.jsonify_team_object = MagicMock( - side_effect=lambda db_data: db_data - ) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) + mock_prisma_client.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) # Mock budget upsert to return updated_kv without team_member_budget def mock_upsert_side_effect( @@ -2394,14 +2341,14 @@ async def test_update_team_team_member_budget_not_passed_to_db( update_data = call_args[1]["data"] # data parameter from the update call # Verify team_member_budget is NOT in the update data - assert ( - "team_member_budget" not in update_data - ), f"team_member_budget should not be in update data, but found: {update_data}" + assert "team_member_budget" not in update_data, ( + f"team_member_budget should not be in update data, but found: {update_data}" + ) # Verify other fields are present (team_alias should be there) - assert "team_alias" in update_data or "team_id" in str( - call_args - ), "Expected team update fields should be present" + assert "team_alias" in update_data or "team_id" in str(call_args), ( + "Expected team update fields should be present" + ) # Reset mock for second test mock_prisma_client.db.litellm_teamtable.update.reset_mock() @@ -2427,9 +2374,9 @@ async def test_update_team_team_member_budget_not_passed_to_db( update_data = call_args[1]["data"] # data parameter from the update call # Verify team_member_budget is NOT in the update data - assert ( - "team_member_budget" not in update_data - ), f"team_member_budget should not be in update data, but found: {update_data}" + assert "team_member_budget" not in update_data, ( + f"team_member_budget should not be in update data, but found: {update_data}" + ) # Test Case 3: No team_member_budget field at all (excluded from request) mock_prisma_client.db.litellm_teamtable.update.reset_mock() @@ -2454,14 +2401,12 @@ async def test_update_team_team_member_budget_not_passed_to_db( update_data = call_args[1]["data"] # data parameter from the update call # Verify team_member_budget is NOT in the update data - assert ( - "team_member_budget" not in update_data - ), f"team_member_budget should not be in update data, but found: {update_data}" - - print( - "✅ All test cases passed: team_member_budget is properly excluded from database update operations" + assert "team_member_budget" not in update_data, ( + f"team_member_budget should not be in update data, but found: {update_data}" ) + print("✅ All test cases passed: team_member_budget is properly excluded from database update operations") + def test_clean_team_member_fields(): """ @@ -2523,9 +2468,7 @@ async def test_create_team_member_budget_table(): TeamMemberBudgetHandler, ) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") data = NewTeamRequest( team_id="test_team_id", @@ -2592,9 +2535,7 @@ async def test_create_team_member_budget_table_without_team_alias(): TeamMemberBudgetHandler, ) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") data = NewTeamRequest(team_id="test_team_id") new_team_data_json = { @@ -2638,9 +2579,7 @@ async def test_upsert_team_member_budget_table_existing_budget(): TeamMemberBudgetHandler, ) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") team_table = MagicMock(spec=LiteLLM_TeamTable) team_table.metadata = {"team_member_budget_id": "existing_budget_123"} @@ -2699,9 +2638,7 @@ async def test_upsert_team_member_budget_table_no_existing_budget(): TeamMemberBudgetHandler, ) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") team_table = MagicMock(spec=LiteLLM_TeamTable) team_table.metadata = {} @@ -2749,16 +2686,12 @@ async def test_upsert_team_member_budget_table_clears_duration_kept_budget(mock_ TeamMemberBudgetHandler, ) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") team_table = MagicMock(spec=LiteLLM_TeamTable) team_table.metadata = {"team_member_budget_id": "existing_budget_123"} - mock_db_client.db.litellm_budgettable.update = AsyncMock( - side_effect=lambda where, data: SimpleNamespace(**data) - ) + mock_db_client.db.litellm_budgettable.update = AsyncMock(side_effect=lambda where, data: SimpleNamespace(**data)) result = await TeamMemberBudgetHandler.upsert_team_member_budget_table( team_table=team_table, @@ -2799,18 +2732,14 @@ async def test_create_team_member_budget_table_explicit_null_duration_does_not_i TeamMemberBudgetHandler, ) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") team_table = MagicMock(spec=LiteLLM_TeamTable) team_table.metadata = {} team_table.team_alias = "Test Team" team_table.budget_duration = "30d" - mock_db_client.db.litellm_budgettable.create = AsyncMock( - side_effect=lambda data: SimpleNamespace(**data) - ) + mock_db_client.db.litellm_budgettable.create = AsyncMock(side_effect=lambda data: SimpleNamespace(**data)) result = await TeamMemberBudgetHandler.create_team_member_budget_table( data=team_table, @@ -2844,18 +2773,14 @@ async def test_create_team_member_budget_table_inherits_team_duration_when_durat TeamMemberBudgetHandler, ) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") team_table = MagicMock(spec=LiteLLM_TeamTable) team_table.metadata = {} team_table.team_alias = "Test Team" team_table.budget_duration = "30d" - mock_db_client.db.litellm_budgettable.create = AsyncMock( - side_effect=lambda data: SimpleNamespace(**data) - ) + mock_db_client.db.litellm_budgettable.create = AsyncMock(side_effect=lambda data: SimpleNamespace(**data)) result = await TeamMemberBudgetHandler.create_team_member_budget_table( data=team_table, @@ -2886,9 +2811,7 @@ async def test_update_team_with_team_member_budget_duration( from litellm.proxy.management_endpoints.team_endpoints import update_team mock_request = Mock(spec=Request) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id") with ( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, @@ -2896,9 +2819,7 @@ async def test_update_team_with_team_member_budget_duration( patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object" - ) as mock_cache_team, + patch("litellm.proxy.management_endpoints.team_endpoints._cache_team_object") as mock_cache_team, patch( "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" ) as mock_upsert_budget, @@ -2910,19 +2831,13 @@ async def test_update_team_with_team_member_budget_duration( "metadata": {"team_member_budget_id": "budget_123"}, } mock_existing_team.metadata = {"team_member_budget_id": "budget_123"} - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_updated_team = MagicMock() mock_updated_team.team_id = "test_team_id" mock_updated_team.model_dump.return_value = {"team_id": "test_team_id"} - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) - mock_prisma_client.jsonify_team_object = MagicMock( - side_effect=lambda db_data: db_data - ) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) + mock_prisma_client.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) def mock_upsert_side_effect( team_table, @@ -2990,9 +2905,7 @@ async def test_backfill_team_member_budget_entries_creates_missing_memberships() existing_membership.user_id = "user-A" mock_prisma = MagicMock() - mock_prisma.db.litellm_teammembership.find_many = AsyncMock( - return_value=[existing_membership] - ) + mock_prisma.db.litellm_teammembership.find_many = AsyncMock(return_value=[existing_membership]) mock_prisma.db.litellm_teammembership.create_many = AsyncMock(return_value=None) mock_prisma.db.litellm_teammembership.update_many = AsyncMock(return_value=0) @@ -3010,9 +2923,7 @@ async def test_backfill_team_member_budget_entries_creates_missing_memberships() ) # find_many should have been called to fetch existing memberships - mock_prisma.db.litellm_teammembership.find_many.assert_awaited_once_with( - where={"team_id": team_id} - ) + mock_prisma.db.litellm_teammembership.find_many.assert_awaited_once_with(where={"team_id": team_id}) # create_many should only create an entry for user-B (user-A already has one) mock_prisma.db.litellm_teammembership.create_many.assert_awaited_once_with( @@ -3065,9 +2976,7 @@ async def test_backfill_team_member_budget_entries_no_op_when_all_exist(): existing_b.user_id = "user-B" mock_prisma = MagicMock() - mock_prisma.db.litellm_teammembership.find_many = AsyncMock( - return_value=[existing_a, existing_b] - ) + mock_prisma.db.litellm_teammembership.find_many = AsyncMock(return_value=[existing_a, existing_b]) mock_prisma.db.litellm_teammembership.create_many = AsyncMock(return_value=None) mock_prisma.db.litellm_teammembership.update_many = AsyncMock(return_value=0) @@ -3112,9 +3021,7 @@ async def test_backfill_team_member_budget_entries_populates_null_budget_id_on_e existing_b.user_id = "user-B" mock_prisma = MagicMock() - mock_prisma.db.litellm_teammembership.find_many = AsyncMock( - return_value=[existing_a, existing_b] - ) + mock_prisma.db.litellm_teammembership.find_many = AsyncMock(return_value=[existing_a, existing_b]) mock_prisma.db.litellm_teammembership.create_many = AsyncMock(return_value=None) mock_prisma.db.litellm_teammembership.update_many = AsyncMock(return_value=2) @@ -3303,9 +3210,7 @@ async def test_bulk_team_member_add_batch_size_limit(): from litellm.proxy.management_endpoints.team_endpoints import bulk_team_member_add # Create more than 500 members (the max batch size) - large_member_list = [ - Member(user_email=f"user{i}@example.com", role="user") for i in range(501) - ] + large_member_list = [Member(user_email=f"user{i}@example.com", role="user") for i in range(501)] bulk_request = BulkTeamMemberAddRequest( team_id="test-team-123", @@ -3360,9 +3265,7 @@ async def test_bulk_team_member_add_all_users_flag(): ) as mock_team_member_add, ): # Mock the database find_many call - mock_prisma.db.litellm_usertable.find_many = AsyncMock( - return_value=mock_db_users - ) + mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=mock_db_users) mock_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -3372,9 +3275,7 @@ async def test_bulk_team_member_add_all_users_flag(): ) # Verify that find_many was called to get all users - mock_prisma.db.litellm_usertable.find_many.assert_called_once_with( - order={"created_at": "desc"} - ) + mock_prisma.db.litellm_usertable.find_many.assert_called_once_with(order={"created_at": "desc"}) # Verify team_member_add was called with users from database mock_team_member_add.assert_called_once() @@ -3497,9 +3398,7 @@ async def test_list_team_v2_security_check_non_admin_user(): ) assert exc_info.value.status_code == 401 - assert "Only admin users can query all teams/other teams" in str( - exc_info.value.detail - ) + assert "Only admin users can query all teams/other teams" in str(exc_info.value.detail) assert LitellmUserRoles.INTERNAL_USER.value in str(exc_info.value.detail) @@ -3547,9 +3446,7 @@ async def test_list_team_v2_security_check_non_admin_user_other_user(): ) assert exc_info.value.status_code == 401 - assert "Only admin users can query all teams/other teams" in str( - exc_info.value.detail - ) + assert "Only admin users can query all teams/other teams" in str(exc_info.value.detail) @pytest.mark.asyncio @@ -3699,17 +3596,11 @@ async def test_list_team_v2_with_status_deleted(): mock_prisma_client.db = mock_db # Mock deleted teams - mock_deleted_team1 = Mock( - model_dump=lambda: {"team_id": "team_1", "team_alias": "Deleted Team 1"} - ) - mock_deleted_team2 = Mock( - model_dump=lambda: {"team_id": "team_2", "team_alias": "Deleted Team 2"} - ) + mock_deleted_team1 = Mock(model_dump=lambda: {"team_id": "team_1", "team_alias": "Deleted Team 1"}) + mock_deleted_team2 = Mock(model_dump=lambda: {"team_id": "team_2", "team_alias": "Deleted Team 2"}) # Mock deleted teams table (should be called) - mock_db.litellm_deletedteamtable.find_many = AsyncMock( - return_value=[mock_deleted_team1, mock_deleted_team2] - ) + mock_db.litellm_deletedteamtable.find_many = AsyncMock(return_value=[mock_deleted_team1, mock_deleted_team2]) mock_db.litellm_deletedteamtable.count = AsyncMock(return_value=2) # Mock regular teams table (should NOT be called) @@ -3898,9 +3789,7 @@ async def test_list_team_v2_org_admin_own_user_id_sees_all_org_teams(): "organization_id": "org_A", "members_with_roles": [{"user_id": "other_user", "role": "user"}], } - mock_db.litellm_teamtable.find_many = AsyncMock( - return_value=[mock_team_1, mock_team_2] - ) + mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team_1, mock_team_2]) mock_db.litellm_teamtable.count = AsyncMock(return_value=2) mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[]) @@ -3996,10 +3885,7 @@ async def test_list_team_v2_org_admin_cannot_view_other_orgs(): ) assert exc_info.value.status_code == 403 - assert ( - "only view teams within your organizations" - in str(exc_info.value.detail).lower() - ) + assert "only view teams within your organizations" in str(exc_info.value.detail).lower() @pytest.mark.asyncio @@ -4159,9 +4045,7 @@ async def test_list_team_v2_search_builds_or_clause(): from litellm.proxy.management_endpoints.team_endpoints import list_team_v2 mock_request = Mock(spec=Request) - mock_admin = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" - ) + mock_admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user") with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: mock_db = Mock() @@ -4206,9 +4090,7 @@ async def test_list_team_v2_search_team_id_match_prefix(): from litellm.proxy.management_endpoints.team_endpoints import list_team_v2 mock_request = Mock(spec=Request) - mock_admin = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" - ) + mock_admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user") with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: mock_db = Mock() @@ -4259,9 +4141,7 @@ async def test_list_team_v2_search_composes_with_user_id_filter(): from litellm.proxy.management_endpoints.team_endpoints import list_team_v2 mock_request = Mock(spec=Request) - mock_user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id="member_user" - ) + mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="member_user") mock_user = LiteLLM_UserTable( user_id="member_user", @@ -4446,9 +4326,7 @@ async def test_list_team_v2_keys_count_skipped_for_deleted_status(): "team_alias": "Deleted Team", } - mock_db.litellm_deletedteamtable.find_many = AsyncMock( - return_value=[mock_deleted] - ) + mock_db.litellm_deletedteamtable.find_many = AsyncMock(return_value=[mock_deleted]) mock_db.litellm_deletedteamtable.count = AsyncMock(return_value=1) mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[]) @@ -4481,9 +4359,7 @@ async def test_team_member_delete_cleans_membership(mock_db_client, mock_admin_a mock_team_row = MagicMock() mock_team_row.model_dump.return_value = { "team_id": test_team_id, - "members_with_roles": [ - {"user_id": test_user_id, "user_email": None, "role": "user"} - ], + "members_with_roles": [{"user_id": test_user_id, "user_email": None, "role": "user"}], "team_member_permissions": [], "metadata": {}, "models": [], @@ -4491,32 +4367,24 @@ async def test_team_member_delete_cleans_membership(mock_db_client, mock_admin_a } # Configure DB mocks used by team_member_delete - mock_db_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_team_row - ) + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team_row) mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row) # User row to allow removal from user's teams list mock_user_row = MagicMock() mock_user_row.user_id = test_user_id mock_user_row.teams = [test_team_id] - mock_db_client.db.litellm_usertable.find_many = AsyncMock( - return_value=[mock_user_row] - ) + mock_db_client.db.litellm_usertable.find_many = AsyncMock(return_value=[mock_user_row]) mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock()) # Membership deletion should be called mock_db_client.db.litellm_teammembership = MagicMock() - mock_db_client.db.litellm_teammembership.delete_many = AsyncMock( - return_value=MagicMock() - ) + mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(return_value=MagicMock()) # Verification token deletion should be called mock_db_client.db.litellm_verificationtoken = MagicMock() mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) - mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock( - return_value=MagicMock() - ) + mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(return_value=MagicMock()) _wire_member_delete_tx(mock_db_client) @@ -4533,9 +4401,7 @@ async def test_team_member_delete_cleans_membership(mock_db_client, mock_admin_a @pytest.mark.asyncio -async def test_team_member_delete_cleans_verification_tokens( - mock_db_client, mock_admin_auth -): +async def test_team_member_delete_cleans_verification_tokens(mock_db_client, mock_admin_auth): from litellm.proxy._types import TeamMemberDeleteRequest from litellm.proxy.management_endpoints.team_endpoints import team_member_delete @@ -4545,38 +4411,28 @@ async def test_team_member_delete_cleans_verification_tokens( mock_team_row = MagicMock() mock_team_row.model_dump.return_value = { "team_id": test_team_id, - "members_with_roles": [ - {"user_id": test_user_id, "user_email": None, "role": "user"} - ], + "members_with_roles": [{"user_id": test_user_id, "user_email": None, "role": "user"}], "team_member_permissions": [], "metadata": {}, "models": [], "spend": 0.0, } - mock_db_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_team_row - ) + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team_row) mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row) mock_user_row = MagicMock() mock_user_row.user_id = test_user_id mock_user_row.teams = [test_team_id] - mock_db_client.db.litellm_usertable.find_many = AsyncMock( - return_value=[mock_user_row] - ) + mock_db_client.db.litellm_usertable.find_many = AsyncMock(return_value=[mock_user_row]) mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock()) mock_db_client.db.litellm_teammembership = MagicMock() - mock_db_client.db.litellm_teammembership.delete_many = AsyncMock( - return_value=MagicMock() - ) + mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(return_value=MagicMock()) mock_db_client.db.litellm_verificationtoken = MagicMock() mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) - mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock( - return_value=MagicMock() - ) + mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(return_value=MagicMock()) _wire_member_delete_tx(mock_db_client) @@ -4594,9 +4450,7 @@ async def test_team_member_delete_cleans_verification_tokens( @pytest.mark.asyncio -async def test_team_member_delete_reads_on_the_lock_holding_transaction( - mock_db_client, mock_admin_auth -): +async def test_team_member_delete_reads_on_the_lock_holding_transaction(mock_db_client, mock_admin_auth): """ Regression pin against exhausting the connection pool with advisory-lock waiters. @@ -4622,9 +4476,7 @@ async def test_team_member_delete_reads_on_the_lock_holding_transaction( "models": [], "spend": 0.0, } - mock_db_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_team_row - ) + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team_row) user_row = MagicMock() user_row.user_id = test_user_id @@ -4656,18 +4508,14 @@ async def test_team_member_delete_reads_on_the_lock_holding_transaction( user_api_key_dict=mock_admin_auth, ) - tx.litellm_usertable.find_many.assert_awaited_once_with( - where={"user_id": {"in": [test_user_id]}} - ) + tx.litellm_usertable.find_many.assert_awaited_once_with(where={"user_id": {"in": [test_user_id]}}) tx.litellm_verificationtoken.find_many.assert_awaited_once_with( where={"user_id": {"in": [test_user_id]}, "team_id": test_team_id} ) pooled_user_read.assert_not_awaited() pooled_token_read.assert_not_awaited() - tx.litellm_usertable.update.assert_awaited_once_with( - where={"user_id": test_user_id}, data={"teams": {"set": []}} - ) + tx.litellm_usertable.update.assert_awaited_once_with(where={"user_id": test_user_id}, data={"teams": {"set": []}}) tx.litellm_teammembership.delete_many.assert_awaited_once_with( where={"team_id": test_team_id, "user_id": test_user_id} ) @@ -4708,18 +4556,14 @@ async def test_team_member_delete_by_email_the_user_row_does_not_carry( mock_team_row = MagicMock() mock_team_row.model_dump.return_value = { "team_id": test_team_id, - "members_with_roles": [ - {"user_id": test_user_id, "user_email": roster_email, "role": "user"} - ], + "members_with_roles": [{"user_id": test_user_id, "user_email": roster_email, "role": "user"}], "team_member_permissions": [], "metadata": {}, "models": [], "spend": 0.0, } - mock_db_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_team_row - ) + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team_row) mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row) mock_user_row = MagicMock() @@ -4731,29 +4575,21 @@ async def test_team_member_delete_by_email_the_user_row_does_not_carry( if not user_row_exists: return [] user_id_filter = where.get("user_id") - if isinstance(user_id_filter, dict) and test_user_id in user_id_filter.get( - "in", [] - ): + if isinstance(user_id_filter, dict) and test_user_id in user_id_filter.get("in", []): return [mock_user_row] if where.get("user_email") == user_row_email: return [mock_user_row] return [] - mock_db_client.db.litellm_usertable.find_many = AsyncMock( - side_effect=find_user_rows - ) + mock_db_client.db.litellm_usertable.find_many = AsyncMock(side_effect=find_user_rows) mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock()) mock_db_client.db.litellm_teammembership = MagicMock() - mock_db_client.db.litellm_teammembership.delete_many = AsyncMock( - return_value=MagicMock() - ) + mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(return_value=MagicMock()) mock_db_client.db.litellm_verificationtoken = MagicMock() mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) - mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock( - return_value=MagicMock() - ) + mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(return_value=MagicMock()) _wire_member_delete_tx(mock_db_client) @@ -4780,9 +4616,7 @@ class _InjectedMemberDeleteFailure(Exception): @pytest.mark.asyncio -async def test_team_member_delete_is_atomic_across_its_four_writes( - mock_db_client, mock_admin_auth -): +async def test_team_member_delete_is_atomic_across_its_four_writes(mock_db_client, mock_admin_auth): """ /team/member_delete's four cleanups (team roster, user.teams, team membership, verification tokens) run as one transaction, so a failure @@ -4803,25 +4637,19 @@ async def test_team_member_delete_is_atomic_across_its_four_writes( mock_team_row = MagicMock() mock_team_row.model_dump.return_value = { "team_id": test_team_id, - "members_with_roles": [ - {"user_id": test_user_id, "user_email": None, "role": "user"} - ], + "members_with_roles": [{"user_id": test_user_id, "user_email": None, "role": "user"}], "team_member_permissions": [], "metadata": {}, "models": [], "spend": 0.0, } - mock_db_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_team_row - ) + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team_row) mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row) mock_user_row = MagicMock() mock_user_row.user_id = test_user_id mock_user_row.teams = [test_team_id] - mock_db_client.db.litellm_usertable.find_many = AsyncMock( - return_value=[mock_user_row] - ) + mock_db_client.db.litellm_usertable.find_many = AsyncMock(return_value=[mock_user_row]) mock_db_client.db.litellm_usertable.update = AsyncMock( side_effect=_InjectedMemberDeleteFailure("boom between writes 1 and 2") ) @@ -4885,9 +4713,7 @@ async def test_new_team_max_budget_exceeds_user_max_budget(): patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4915,9 +4741,7 @@ async def test_new_team_max_budget_exceeds_user_max_budget(): # ProxyException stores status_code in 'code' attribute assert exc_info.value.code == "400" assert "max budget higher than user max" in str(exc_info.value.message) - assert "100.0" in str( - exc_info.value.message - ) # User's user_max_budget should be mentioned + assert "100.0" in str(exc_info.value.message) # User's user_max_budget should be mentioned assert LitellmUserRoles.INTERNAL_USER.value in str(exc_info.value.message) @@ -4954,9 +4778,7 @@ async def test_new_team_max_budget_within_user_limit(): patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4988,19 +4810,13 @@ async def test_new_team_max_budget_within_user_limit(): "max_budget": 50.0, "members_with_roles": [], } - mock_prisma.db.litellm_teamtable.create = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.create = AsyncMock(return_value=mock_created_team) _wire_team_create_tx(mock_prisma) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_created_team) # Mock model table mock_prisma.db.litellm_modeltable = MagicMock() - mock_prisma.db.litellm_modeltable.create = AsyncMock( - return_value=MagicMock(id="model123") - ) + mock_prisma.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model123")) # Mock user table operations for adding the creator as a member mock_user = MagicMock() @@ -5023,9 +4839,7 @@ async def test_new_team_max_budget_within_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( - return_value=mock_membership - ) + mock_prisma.db.litellm_teammembership.create = AsyncMock(return_value=mock_membership) # Should NOT raise an exception result = await new_team( @@ -5086,12 +4900,8 @@ async def test_new_team_org_scoped_budget_bypasses_user_limit(): patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, - patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object" - ) as mock_get_org, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, + patch("litellm.proxy.management_endpoints.team_endpoints.get_org_object") as mock_get_org, ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -5131,19 +4941,13 @@ async def test_new_team_org_scoped_budget_bypasses_user_limit(): "organization_id": "test-org-123", "members_with_roles": [], } - mock_prisma.db.litellm_teamtable.create = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.create = AsyncMock(return_value=mock_created_team) _wire_team_create_tx(mock_prisma) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_created_team) # Mock model table mock_prisma.db.litellm_modeltable = MagicMock() - mock_prisma.db.litellm_modeltable.create = AsyncMock( - return_value=MagicMock(id="model123") - ) + mock_prisma.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model123")) # Mock user table operations mock_user = MagicMock() @@ -5166,9 +4970,7 @@ async def test_new_team_org_scoped_budget_bypasses_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( - return_value=mock_membership - ) + mock_prisma.db.litellm_teammembership.create = AsyncMock(return_value=mock_membership) # Should NOT raise an exception - the fix should bypass user budget validation for org-scoped teams result = await new_team( @@ -5219,9 +5021,7 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): # Create team request with models that are within org's allowed models but not user's team_request = NewTeamRequest( team_alias="org-scoped-models-team", - models=[ - "gpt-4" - ], # Within org's allowed models, but not in user's personal models + models=["gpt-4"], # Within org's allowed models, but not in user's personal models organization_id="test-org-456", # This makes it an org-scoped team ) @@ -5232,12 +5032,8 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, - patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object" - ) as mock_get_org, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, + patch("litellm.proxy.management_endpoints.team_endpoints.get_org_object") as mock_get_org, ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -5279,19 +5075,13 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): "models": ["gpt-4"], "members_with_roles": [], } - mock_prisma.db.litellm_teamtable.create = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.create = AsyncMock(return_value=mock_created_team) _wire_team_create_tx(mock_prisma) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_created_team) # Mock model table mock_prisma.db.litellm_modeltable = MagicMock() - mock_prisma.db.litellm_modeltable.create = AsyncMock( - return_value=MagicMock(id="model123") - ) + mock_prisma.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model123")) # Mock user table operations mock_user = MagicMock() @@ -5314,9 +5104,7 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( - return_value=mock_membership - ) + mock_prisma.db.litellm_teammembership.create = AsyncMock(return_value=mock_membership) # Should NOT raise an exception - the fix should bypass user model validation for org-scoped teams result = await new_team( @@ -5376,9 +5164,7 @@ async def test_new_team_standalone_validates_against_user_models(monkeypatch): patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -5445,9 +5231,7 @@ async def test_new_team_standalone_validates_against_user_budget(): patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -5472,9 +5256,7 @@ async def test_new_team_standalone_validates_against_user_budget(): # Verify exception details assert exc_info.value.code == "400" assert "max budget higher than user max" in str(exc_info.value.message) - assert "3.0" in str( - exc_info.value.message - ) # User's max_budget should be mentioned + assert "3.0" in str(exc_info.value.message) # User's max_budget should be mentioned @pytest.mark.asyncio @@ -5518,12 +5300,8 @@ async def test_new_team_org_scoped_budget_exceeds_org_limit(): patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, - patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object" - ) as mock_get_org, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, + patch("litellm.proxy.management_endpoints.team_endpoints.get_org_object") as mock_get_org, ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -5597,12 +5375,8 @@ async def test_new_team_org_scoped_models_not_in_org_models(): patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, - patch( - "litellm.proxy.management_endpoints.team_endpoints.get_org_object" - ) as mock_get_org, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, + patch("litellm.proxy.management_endpoints.team_endpoints.get_org_object") as mock_get_org, ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -5626,10 +5400,7 @@ async def test_new_team_org_scoped_models_not_in_org_models(): # Verify exception details assert exc_info.value.code == "400" - assert ( - "claude-3-opus" in str(exc_info.value.message) - or "organization" in str(exc_info.value.message).lower() - ) + assert "claude-3-opus" in str(exc_info.value.message) or "organization" in str(exc_info.value.message).lower() @pytest.mark.asyncio @@ -5673,9 +5444,7 @@ async def test_update_team_standalone_budget_raise_blocked_for_team_admin(): patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), ): mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-123" @@ -5686,13 +5455,9 @@ async def test_update_team_standalone_budget_raise_blocked_for_team_admin(): "team_id": "standalone-team-123", "organization_id": None, "max_budget": 30.0, - "members_with_roles": [ - {"user_id": "non-admin-update-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "non-admin-update-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_cache.async_get_cache = AsyncMock(return_value=None) with pytest.raises(ProxyException) as exc_info: @@ -5741,9 +5506,7 @@ async def test_update_team_standalone_budget_raise_allowed_for_proxy_admin( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), ): mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-123" @@ -5754,13 +5517,9 @@ async def test_update_team_standalone_budget_raise_allowed_for_proxy_admin( "team_id": "standalone-team-123", "organization_id": None, "max_budget": 30.0, - "members_with_roles": [ - {"user_id": "proxy-admin-update-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "proxy-admin-update-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() @@ -5775,9 +5534,7 @@ async def test_update_team_standalone_budget_raise_allowed_for_proxy_admin( "organization_id": None, "max_budget": 100.0, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) result = await update_team( data=update_request, @@ -5829,9 +5586,7 @@ async def test_update_team_standalone_budget_removal_blocked_for_team_admin(): patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), ): mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-123" @@ -5842,13 +5597,9 @@ async def test_update_team_standalone_budget_removal_blocked_for_team_admin(): "team_id": "standalone-team-123", "organization_id": None, "max_budget": 500.0, - "members_with_roles": [ - {"user_id": "budget-removal-admin", "role": "admin"} - ], + "members_with_roles": [{"user_id": "budget-removal-admin", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_cache.async_get_cache = AsyncMock(return_value=None) with pytest.raises(ProxyException) as exc_info: @@ -5898,9 +5649,7 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), ): mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-uncapped-123" @@ -5911,13 +5660,9 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed( "team_id": "standalone-uncapped-123", "organization_id": None, "max_budget": None, - "members_with_roles": [ - {"user_id": "uncapped-team-admin", "role": "admin"} - ], + "members_with_roles": [{"user_id": "uncapped-team-admin", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() @@ -5932,9 +5677,7 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed( "organization_id": None, "max_budget": 1000.0, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) result = await update_team( data=update_request, @@ -5992,9 +5735,7 @@ async def test_update_team_standalone_unchanged_budget_allowed( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): # Mock existing standalone team (no organization_id) with budget=$500 mock_existing_team = MagicMock() @@ -6006,13 +5747,9 @@ async def test_update_team_standalone_unchanged_budget_allowed( "team_id": "standalone-unchanged-budget-123", "organization_id": None, "max_budget": 500.0, - "members_with_roles": [ - {"user_id": "standalone-unchanged-budget-admin", "role": "admin"} - ], + "members_with_roles": [{"user_id": "standalone-unchanged-budget-admin", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data # User has a restrictive personal budget that is lower than the team's. @@ -6034,9 +5771,7 @@ async def test_update_team_standalone_unchanged_budget_allowed( "max_budget": 500.0, "tpm_limit": 50000, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) # Should NOT raise - unchanged budget skips the personal-budget check. result = await update_team( @@ -6090,9 +5825,7 @@ async def test_update_team_standalone_lower_budget_allowed( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-lower-budget-123" @@ -6103,13 +5836,9 @@ async def test_update_team_standalone_lower_budget_allowed( "team_id": "standalone-lower-budget-123", "organization_id": None, "max_budget": 500.0, - "members_with_roles": [ - {"user_id": "standalone-lower-budget-admin", "role": "admin"} - ], + "members_with_roles": [{"user_id": "standalone-lower-budget-admin", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data mock_user_obj = LiteLLM_UserTable( @@ -6129,9 +5858,7 @@ async def test_update_team_standalone_lower_budget_allowed( "organization_id": None, "max_budget": 300.0, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) result = await update_team( data=update_request, @@ -6192,9 +5919,7 @@ async def test_update_team_org_scoped_budget_exceeds_org_limit(): patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, patch( "litellm.proxy.management_endpoints.team_endpoints.get_org_object", new=AsyncMock(return_value=mock_org), @@ -6209,13 +5934,9 @@ async def test_update_team_org_scoped_budget_exceeds_org_limit(): "team_id": "org-team-456", "organization_id": "test-org-update", "max_budget": 80.0, - "members_with_roles": [ - {"user_id": "org-admin-update-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-update-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) # Should raise ProxyException because new budget exceeds org's max_budget with pytest.raises(ProxyException) as exc_info: @@ -6227,10 +5948,7 @@ async def test_update_team_org_scoped_budget_exceeds_org_limit(): # Verify exception details assert exc_info.value.code == "400" - assert ( - "organization" in str(exc_info.value.message).lower() - or "budget" in str(exc_info.value.message).lower() - ) + assert "organization" in str(exc_info.value.message).lower() or "budget" in str(exc_info.value.message).lower() @pytest.mark.asyncio @@ -6272,9 +5990,7 @@ async def test_update_team_standalone_models_not_gated_by_user_limit( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): # Mock existing standalone team (no organization_id) mock_existing_team = MagicMock() @@ -6286,13 +6002,9 @@ async def test_update_team_standalone_models_not_gated_by_user_limit( "team_id": "standalone-team-models-123", "organization_id": None, "models": ["gpt-3.5-turbo"], - "members_with_roles": [ - {"user_id": "non-admin-update-models-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "non-admin-update-models-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() @@ -6306,9 +6018,7 @@ async def test_update_team_standalone_models_not_gated_by_user_limit( "organization_id": None, "models": ["gpt-4"], } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) result = await update_team( data=update_request, @@ -6371,9 +6081,7 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, patch( "litellm.proxy.management_endpoints.team_endpoints.get_org_object", new=AsyncMock(return_value=mock_org), @@ -6389,13 +6097,9 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit( "team_id": "org-team-update-budget-123", "organization_id": "test-org-update-budget", "max_budget": 30.0, - "members_with_roles": [ - {"user_id": "org-admin-update-budget-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-update-budget-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data # Mock user cache to return user with restrictive budget @@ -6404,9 +6108,7 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit( max_budget=3.0, # Restrictive personal budget ) mock_cache.async_get_cache = AsyncMock(return_value=mock_user_obj) - mock_cache.async_set_cache = ( - AsyncMock() - ) # Mock cache set for _cache_team_object + mock_cache.async_set_cache = AsyncMock() # Mock cache set for _cache_team_object # Mock team update mock_updated_team = MagicMock() @@ -6419,9 +6121,7 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit( "organization_id": "test-org-update-budget", "max_budget": 50.0, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) # Should NOT raise an exception - bypass user budget validation for org-scoped teams result = await update_team( @@ -6482,9 +6182,7 @@ async def test_update_team_org_scoped_models_bypasses_user_limit( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, patch( "litellm.proxy.management_endpoints.team_endpoints.get_org_object", new=AsyncMock(return_value=mock_org), @@ -6500,17 +6198,11 @@ async def test_update_team_org_scoped_models_bypasses_user_limit( "team_id": "org-team-update-models-123", "organization_id": "test-org-update-models", "models": ["gpt-3.5-turbo"], - "members_with_roles": [ - {"user_id": "org-admin-update-models-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-update-models-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data - mock_cache.async_set_cache = ( - AsyncMock() - ) # Mock cache set for _cache_team_object + mock_cache.async_set_cache = AsyncMock() # Mock cache set for _cache_team_object # Mock team update mock_updated_team = MagicMock() @@ -6523,9 +6215,7 @@ async def test_update_team_org_scoped_models_bypasses_user_limit( "organization_id": "test-org-update-models", "models": ["gpt-4"], } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) # Should NOT raise an exception - bypass user models validation for org-scoped teams result = await update_team( @@ -6584,9 +6274,7 @@ async def test_update_team_org_scoped_models_not_in_org_models(): patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, patch( "litellm.proxy.management_endpoints.team_endpoints.get_org_object", new=AsyncMock(return_value=mock_org), @@ -6601,13 +6289,9 @@ async def test_update_team_org_scoped_models_not_in_org_models(): "team_id": "org-team-update-models-fail-123", "organization_id": "test-org-update-models-fail", "models": ["gpt-4"], - "members_with_roles": [ - {"user_id": "org-admin-update-models-fail-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-update-models-fail-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) # Should raise ProxyException because claude-3-opus is not in org's allowed models with pytest.raises(ProxyException) as exc_info: @@ -6619,10 +6303,7 @@ async def test_update_team_org_scoped_models_not_in_org_models(): # Verify exception details assert exc_info.value.code == "400" - assert ( - "claude-3-opus" in str(exc_info.value.message) - or "organization" in str(exc_info.value.message).lower() - ) + assert "claude-3-opus" in str(exc_info.value.message) or "organization" in str(exc_info.value.message).lower() @pytest.mark.asyncio @@ -6673,9 +6354,7 @@ async def test_update_team_org_scoped_models_with_all_proxy_models( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, patch( "litellm.proxy.management_endpoints.team_endpoints.get_org_object", new=AsyncMock(return_value=mock_org), @@ -6691,17 +6370,11 @@ async def test_update_team_org_scoped_models_with_all_proxy_models( "team_id": "org-team-all-proxy-models-123", "organization_id": "test-org-all-proxy-models", "models": ["gpt-4"], - "members_with_roles": [ - {"user_id": "org-admin-all-proxy-models-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-all-proxy-models-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data - mock_cache.async_set_cache = ( - AsyncMock() - ) # Mock cache set for _cache_team_object + mock_cache.async_set_cache = AsyncMock() # Mock cache set for _cache_team_object # Mock team update mock_updated_team = MagicMock() @@ -6722,9 +6395,7 @@ async def test_update_team_org_scoped_models_with_all_proxy_models( "gpt-4o-mini-test", ], } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) # Should NOT raise an exception - 'all-proxy-models' allows all models result = await update_team( @@ -6782,9 +6453,7 @@ async def test_update_team_tpm_limit_not_gated_by_user_limit( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), ): # Mock existing standalone team mock_existing_team = MagicMock() @@ -6798,9 +6467,7 @@ async def test_update_team_tpm_limit_not_gated_by_user_limit( "tpm_limit": 500, "members_with_roles": [{"user_id": "tpm-limit-user", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() @@ -6814,9 +6481,7 @@ async def test_update_team_tpm_limit_not_gated_by_user_limit( "organization_id": None, "tpm_limit": 5000, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) result = await update_team( data=update_request, @@ -6864,9 +6529,7 @@ async def test_update_team_rpm_limit_not_gated_by_user_limit( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), ): # Mock existing standalone team mock_existing_team = MagicMock() @@ -6880,9 +6543,7 @@ async def test_update_team_rpm_limit_not_gated_by_user_limit( "rpm_limit": 50, "members_with_roles": [{"user_id": "rpm-limit-user", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_prisma.jsonify_team_object = lambda db_data: db_data mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() @@ -6896,9 +6557,7 @@ async def test_update_team_rpm_limit_not_gated_by_user_limit( "organization_id": None, "rpm_limit": 500, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) result = await update_team( data=update_request, @@ -7115,9 +6774,7 @@ async def test_new_team_org_scoped_tpm_rpm_bypasses_user_limit(): patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), patch("litellm.proxy.proxy_server._license_check") as mock_license, - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), patch( "litellm.proxy.management_endpoints.team_endpoints.get_org_object", new=AsyncMock(return_value=mock_org), @@ -7148,13 +6805,9 @@ async def test_new_team_org_scoped_tpm_rpm_bypasses_user_limit(): "metadata": None, "members_with_roles": [], } - mock_prisma.db.litellm_teamtable.create = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.create = AsyncMock(return_value=mock_created_team) _wire_team_create_tx(mock_prisma) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_created_team) mock_prisma.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) # Should succeed - bypasses user limits since org-scoped @@ -7233,13 +6886,9 @@ async def test_update_team_org_scoped_tpm_exceeds_org_limit(): "team_id": "org-team-update-tpm-123", "organization_id": "test-org-update-tpm", "tpm_limit": 5000, - "members_with_roles": [ - {"user_id": "org-admin-update-tpm-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-update-tpm-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) # Should raise ProxyException because TPM exceeds org limit with pytest.raises(ProxyException) as exc_info: @@ -7319,13 +6968,9 @@ async def test_update_team_org_scoped_rpm_exceeds_org_limit(): "team_id": "org-team-update-rpm-123", "organization_id": "test-org-update-rpm", "rpm_limit": 500, - "members_with_roles": [ - {"user_id": "org-admin-update-rpm-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-update-rpm-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) # Should raise ProxyException because RPM exceeds org limit with pytest.raises(ProxyException) as exc_info: @@ -7414,13 +7059,9 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit( "organization_id": "test-org-update-bypass", "tpm_limit": 5000, "rpm_limit": 500, - "members_with_roles": [ - {"user_id": "org-admin-update-bypass-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-update-bypass-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_cache.async_set_cache = AsyncMock() mock_logging.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() @@ -7435,9 +7076,7 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit( "tpm_limit": 10000, "rpm_limit": 1000, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) mock_prisma.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) # Should succeed - bypasses user limits since org-scoped @@ -7524,9 +7163,7 @@ async def test_update_team_guardrails_with_org_id( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ), + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), patch( "litellm.proxy.proxy_server.premium_user", True, # Required for guardrails feature @@ -7551,20 +7188,14 @@ async def test_update_team_guardrails_with_org_id( "max_budget": None, "tpm_limit": None, "rpm_limit": None, - "members_with_roles": [ - {"user_id": "org-admin-guardrails-test", "role": "admin"} - ], + "members_with_roles": [{"user_id": "org-admin-guardrails-test", "role": "admin"}], } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) mock_cache.async_set_cache = AsyncMock() # Mock organization fetch - this is where the bug occurred # The fix ensures 'teams: True' is in the include clause - mock_prisma.db.litellm_organizationtable.find_unique = AsyncMock( - return_value=mock_org - ) + mock_prisma.db.litellm_organizationtable.find_unique = AsyncMock(return_value=mock_org) # Destination-org guard in update_team queries for the caller's # ORG_ADMIN membership on the destination org. Return a match so @@ -7572,17 +7203,13 @@ async def test_update_team_guardrails_with_org_id( mock_org_admin_membership = MagicMock() mock_org_admin_membership.user_id = "org-admin-guardrails-test" mock_org_admin_membership.organization_id = "test-org-guardrails" - mock_prisma.db.litellm_organizationmembership.find_many = AsyncMock( - return_value=[mock_org_admin_membership] - ) + mock_prisma.db.litellm_organizationmembership.find_many = AsyncMock(return_value=[mock_org_admin_membership]) # Mock team update mock_updated_team = MagicMock(spec=LiteLLM_TeamTable) mock_updated_team.team_id = "team-guardrails-123" mock_updated_team.organization_id = "test-org-guardrails" - mock_updated_team.metadata = { - "guardrails": ["aporia-pre-call", "aporia-post-call"] - } + mock_updated_team.metadata = {"guardrails": ["aporia-pre-call", "aporia-post-call"]} mock_updated_team.litellm_model_table = None mock_updated_team.access_group_ids = None mock_updated_team.model_dump.return_value = { @@ -7590,9 +7217,7 @@ async def test_update_team_guardrails_with_org_id( "organization_id": "test-org-guardrails", "metadata": {"guardrails": ["aporia-pre-call", "aporia-post-call"]}, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) mock_prisma.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) # async_get_cache must be an AsyncMock so `await` in get_org_object works mock_cache.async_get_cache = AsyncMock(return_value=None) @@ -7622,11 +7247,7 @@ async def test_update_team_guardrails_with_org_id( assert mock_prisma.db.litellm_organizationtable.find_unique.call_count >= 1 # Get the first call (from fetch_and_validate_organization) - first_call_kwargs = ( - mock_prisma.db.litellm_organizationtable.find_unique.call_args_list[ - 0 - ].kwargs - ) + first_call_kwargs = mock_prisma.db.litellm_organizationtable.find_unique.call_args_list[0].kwargs # Verify that 'teams' is included in the fetch assert "include" in first_call_kwargs @@ -7677,9 +7298,7 @@ def test_transform_teams_to_deleted_records(): assert all("litellm_changed_by" in record for record in records) assert all(record["deleted_by"] == "user-123" for record in records) # UserAPIKeyAuth hashes the api_key, so we check against the hashed value - assert all( - record["deleted_by_api_key"] == user_api_key_dict.api_key for record in records - ) + assert all(record["deleted_by_api_key"] == user_api_key_dict.api_key for record in records) assert all(record["litellm_changed_by"] == "admin-user" for record in records) record1 = records[0] @@ -7824,9 +7443,7 @@ async def test_delete_team_persists_deleted_teams( mock_prisma_client.db.litellm_deletedteamtable.create_many = mock_create_many_teams mock_create_many_keys = AsyncMock() - mock_prisma_client.db.litellm_deletedverificationtoken.create_many = ( - mock_create_many_keys - ) + mock_prisma_client.db.litellm_deletedverificationtoken.create_many = mock_create_many_keys mock_find_many_keys = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many_keys @@ -7944,9 +7561,7 @@ async def test_delete_team_sweeps_references_outside_members_with_roles( ("team-doomed", "doomed-team"), ("team-kept", "kept-team"), ): - cached_obj = LiteLLM_TeamTableCachedObj( - team_id=cached_team_id, team_alias=cached_alias - ) + cached_obj = LiteLLM_TeamTableCachedObj(team_id=cached_team_id, team_alias=cached_alias) fresh_cache.set_cache(key=f"team_id:{cached_team_id}", value=cached_obj) fresh_cache.set_cache(key=f"team_alias:{cached_alias}", value=cached_obj) @@ -8040,7 +7655,9 @@ async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes( _wire_team_delete_tx(mock_prisma_client) fresh_cache = UserApiKeyCache() - fresh_cache.set_cache(key="hashed-doomed-key", value=UserAPIKeyAuth(token="hashed-doomed-key", team_id="team-doomed")) + fresh_cache.set_cache( + key="hashed-doomed-key", value=UserAPIKeyAuth(token="hashed-doomed-key", team_id="team-doomed") + ) fresh_cache.set_cache(key="hashed-unrelated-key", value=UserAPIKeyAuth(token="hashed-unrelated-key")) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -8358,9 +7975,7 @@ async def test_team_member_delete_persists_deleted_keys(monkeypatch): mock_prisma_client.db.litellm_verificationtoken.delete_many = mock_delete_keys mock_create_many_keys = AsyncMock() - mock_prisma_client.db.litellm_deletedverificationtoken.create_many = ( - mock_create_many_keys - ) + mock_prisma_client.db.litellm_deletedverificationtoken.create_many = mock_create_many_keys _wire_member_delete_tx(mock_prisma_client) @@ -8516,9 +8131,7 @@ async def test_new_team_soft_budget_validation( patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server._license_check") as mock_license, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -8552,19 +8165,13 @@ async def test_new_team_soft_budget_validation( "max_budget": expected_max_budget, "members_with_roles": [], } - mock_prisma.db.litellm_teamtable.create = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.create = AsyncMock(return_value=mock_created_team) _wire_team_create_tx(mock_prisma) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_created_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_created_team) # Mock model table mock_prisma.db.litellm_modeltable = MagicMock() - mock_prisma.db.litellm_modeltable.create = AsyncMock( - return_value=MagicMock(id="model123") - ) + mock_prisma.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model123")) # Mock user table operations mock_user = MagicMock() @@ -8587,9 +8194,7 @@ async def test_new_team_soft_budget_validation( "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( - return_value=mock_membership - ) + mock_prisma.db.litellm_teammembership.create = AsyncMock(return_value=mock_membership) if should_succeed: # Should NOT raise an exception @@ -8718,9 +8323,7 @@ async def test_update_team_soft_budget_validation( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch( - "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() - ) as mock_audit, + patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()) as mock_audit, ): # Mock existing team with existing budgets mock_existing_team = MagicMock() @@ -8734,9 +8337,7 @@ async def test_update_team_soft_budget_validation( "soft_budget": existing_soft_budget, "max_budget": existing_max_budget, } - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) # Mock user cache mock_user_obj = LiteLLM_UserTable( @@ -8746,14 +8347,8 @@ async def test_update_team_soft_budget_validation( mock_cache.async_get_cache = AsyncMock(return_value=mock_user_obj) # Mock updated team - preserve existing values if not being updated - final_soft_budget = ( - update_soft_budget - if update_soft_budget is not None - else existing_soft_budget - ) - final_max_budget = ( - update_max_budget if update_max_budget is not None else existing_max_budget - ) + final_soft_budget = update_soft_budget if update_soft_budget is not None else existing_soft_budget + final_max_budget = update_max_budget if update_max_budget is not None else existing_max_budget mock_updated_team = MagicMock() mock_updated_team.team_id = "test-team-123" @@ -8766,13 +8361,9 @@ async def test_update_team_soft_budget_validation( "soft_budget": final_soft_budget, "max_budget": final_max_budget, } - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) mock_prisma.jsonify_team_object = lambda db_data: db_data - mock_cache.async_set_cache = ( - AsyncMock() - ) # Mock cache set for _cache_team_object + mock_cache.async_set_cache = AsyncMock() # Mock cache set for _cache_team_object if should_succeed: # Should NOT raise an exception @@ -8818,9 +8409,7 @@ async def test_new_team_positive_budgets_accepted(): from litellm.proxy._types import NewTeamRequest # Should not raise any errors - request = NewTeamRequest( - team_alias="test-team", max_budget=100.0, team_member_budget=50.0 - ) + request = NewTeamRequest(team_alias="test-team", max_budget=100.0, team_member_budget=50.0) assert request.max_budget == 100.0 assert request.team_member_budget == 50.0 @@ -8841,9 +8430,7 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth): # Mock model table creation mock_db_client.db.litellm_modeltable = MagicMock() - mock_db_client.db.litellm_modeltable.create = AsyncMock( - return_value=MagicMock(id="model123") - ) + mock_db_client.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model123")) # Capture team table creation team_create_result = MagicMock( @@ -8858,9 +8445,7 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth): mock_db_client.db.litellm_teamtable.create = mock_team_create _wire_team_create_tx(mock_db_client) mock_db_client.db.litellm_teamtable.count = mock_team_count - mock_db_client.db.litellm_teamtable.update = AsyncMock( - return_value=team_create_result - ) + mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=team_create_result) # Mock user table mock_db_client.db.litellm_usertable = MagicMock() @@ -8922,9 +8507,7 @@ async def test_get_team_daily_activity_member_with_permission_sees_all_spend( # Create a non-admin user user_id = "test_user_with_perm_123" team_id = "test_team_789" - user_api_key_dict = UserAPIKeyAuth( - user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) + user_api_key_dict = UserAPIKeyAuth(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER) # Mock user info mock_user_info = LiteLLM_UserTable( @@ -9009,9 +8592,7 @@ async def test_get_team_daily_activity_member_without_permission_filters_by_keys # Create a non-admin user user_id = "test_user_no_perm_123" team_id = "test_team_789" - user_api_key_dict = UserAPIKeyAuth( - user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) + user_api_key_dict = UserAPIKeyAuth(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER) # Mock user info mock_user_info = LiteLLM_UserTable( @@ -9045,9 +8626,7 @@ async def test_get_team_daily_activity_member_without_permission_filters_by_keys # Setup mocks mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team]) - mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[user_api_key_1, user_api_key_2] - ) + mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[user_api_key_1, user_api_key_2]) # Mock get_user_object with patch( @@ -9183,9 +8762,7 @@ async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys( # Create a non-admin user user_id = "test_user_123" team_id = "test_team_456" - user_api_key_dict = UserAPIKeyAuth( - user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) + user_api_key_dict = UserAPIKeyAuth(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER) # Mock user info mock_user_info = LiteLLM_UserTable( @@ -9217,9 +8794,7 @@ async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys( # Setup mocks mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team]) - mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[user_api_key_1, user_api_key_2] - ) + mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[user_api_key_1, user_api_key_2]) # Mock get_user_object with patch( @@ -9256,9 +8831,7 @@ async def test_get_team_daily_activity_non_admin_filters_by_user_api_keys( # Verify user's API keys were fetched mock_db_client.db.litellm_verificationtoken.find_many.assert_called_once() - api_key_call_kwargs = ( - mock_db_client.db.litellm_verificationtoken.find_many.call_args[1] - ) + api_key_call_kwargs = mock_db_client.db.litellm_verificationtoken.find_many.call_args[1] assert api_key_call_kwargs["where"] == {"user_id": user_id} @@ -9275,9 +8848,7 @@ async def test_get_team_daily_activity_team_admin_sees_all_spend(mock_db_client) # Create a team admin user user_id = "test_admin_123" team_id = "test_team_456" - user_api_key_dict = UserAPIKeyAuth( - user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) + user_api_key_dict = UserAPIKeyAuth(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER) # Mock user info mock_user_info = LiteLLM_UserTable( @@ -9402,9 +8973,7 @@ async def test_validate_and_populate_member_user_info_only_email_provided(): mock_user_find_first.user_email = "test@example.com" # Mock find_first to return the user - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - return_value=mock_user_find_first - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=mock_user_find_first) # Mock get_data to return single user (no duplicates) mock_prisma_client.get_data = AsyncMock(return_value=[mock_user_find_first]) @@ -9460,9 +9029,7 @@ async def test_validate_and_populate_member_user_info_only_user_id_not_found(): assert result.role == "user" # Verify find_unique was called with correct parameters - mock_prisma_client.db.litellm_usertable.find_unique.assert_called_once_with( - where={"user_id": "nonexistent-user"} - ) + mock_prisma_client.db.litellm_usertable.find_unique.assert_called_once_with(where={"user_id": "nonexistent-user"}) @pytest.mark.asyncio @@ -9558,9 +9125,7 @@ async def test_list_team_v1_batches_key_queries(): return [key3] return [key1, key2, key3] - mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( - side_effect=filtered_find_many - ) + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=filtered_find_many) result = await list_team( http_request=mock_request, @@ -9657,9 +9222,7 @@ class TestBatchResolveAccessGroupResources: fake_row.access_agent_ids = ["agent-1", "agent-2"] fake_prisma = MagicMock() - fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock( - return_value=[fake_row] - ) + fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[fake_row]) with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma): result = await _batch_resolve_access_group_resources(["ag-1"]) @@ -9688,9 +9251,7 @@ class TestBatchResolveAccessGroupResources: row2.access_agent_ids = ["agent-2"] fake_prisma = MagicMock() - fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock( - return_value=[row1, row2] - ) + fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[row1, row2]) with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma): result = await _batch_resolve_access_group_resources(["ag-1", "ag-2"]) @@ -9712,9 +9273,7 @@ class TestBatchResolveAccessGroupResources: row1.access_agent_ids = [] fake_prisma = MagicMock() - fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock( - return_value=[row1] - ) + fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[row1]) with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma): result = await _batch_resolve_access_group_resources(["ag-1", "ag-missing"]) @@ -9752,9 +9311,7 @@ class TestBatchResolveAccessGroupResources: fake_prisma.db.litellm_accessgrouptable.find_many = fake_find_many with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma): - result = await _batch_resolve_access_group_resources( - ["ag-1", "ag-1", "ag-1"] - ) + result = await _batch_resolve_access_group_resources(["ag-1", "ag-1", "ag-1"]) # Should have been called with deduplicated list call_args = fake_find_many.call_args @@ -9791,9 +9348,7 @@ class TestResolveTeamAccessGroupResources: row2.access_agent_ids = ["agent-1"] fake_prisma = MagicMock() - fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock( - return_value=[row1, row2] - ) + fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[row1, row2]) team_info = TeamInfoResponseObjectTeamTable( team_id="team-1", access_group_ids=["ag-1", "ag-2", "ag-1", "ag-missing"] @@ -9809,10 +9364,7 @@ class TestResolveTeamAccessGroupResources: ] assert resolved.access_group_mcp_server_ids == ["mcp-1"] assert resolved.access_group_agent_ids == ["agent-1"] - assert [ - (d.access_group_id, d.access_group_name, d.models) - for d in (resolved.access_group_details or []) - ] == [ + assert [(d.access_group_id, d.access_group_name, d.models) for d in (resolved.access_group_details or [])] == [ ("ag-1", "shared-models", ("gpt-4", "claude-3")), ("ag-2", "extra-models", ("claude-3", "gemini")), ] @@ -9905,9 +9457,7 @@ async def test_update_team_rejects_unauthorized_caller(): ], "organization_id": "org-456", } - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) update_request = UpdateTeamRequest( team_id="team-123", @@ -9989,9 +9539,7 @@ async def test_team_member_me_returns_caller_membership(mock_db_client): team_id = "team-me-1" caller_id = "alice@example.com" other_id = "bob@example.com" - caller_auth = UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id=caller_id - ) + caller_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id=caller_id) team = _build_team_for_me( team_id, @@ -10003,9 +9551,7 @@ async def test_team_member_me_returns_caller_membership(mock_db_client): membership = _build_membership_for_me(caller_id, team_id, spend=42.0) user = LiteLLM_UserTable(user_id=caller_id, user_email=caller_id, max_budget=None) - p_team, p_membership, p_user = _patch_member_me_helpers( - team=team, membership=membership, user=user - ) + p_team, p_membership, p_user = _patch_member_me_helpers(team=team, membership=membership, user=user) with p_team, p_membership as mock_get_membership, p_user: response = await team_member_me( http_request=MagicMock(spec=Request), @@ -10023,9 +9569,7 @@ async def test_team_member_me_returns_caller_membership(mock_db_client): # budget_reset_at must survive end-to-end — proves the BudgetTableFull # variant of the Union is selected (created_at is present), not the base # LiteLLM_BudgetTable which would silently strip this field. - assert response.litellm_budget_table.budget_reset_at == datetime( - 2026, 5, 1, tzinfo=timezone.utc - ) + assert response.litellm_budget_table.budget_reset_at == datetime(2026, 5, 1, tzinfo=timezone.utc) # Membership lookup must scope to caller_id, not just team_id — proves the # endpoint cannot return another member's row. @@ -10060,13 +9604,9 @@ async def test_team_member_me_matches_email_only_member(mock_db_client): [{"user_id": None, "user_email": caller_email, "role": "user"}], ) membership = _build_membership_for_me(caller_id, team_id, spend=7.0) - user = LiteLLM_UserTable( - user_id=caller_id, user_email=caller_email, max_budget=None - ) + user = LiteLLM_UserTable(user_id=caller_id, user_email=caller_email, max_budget=None) - p_team, p_membership, p_user = _patch_member_me_helpers( - team=team, membership=membership, user=user - ) + p_team, p_membership, p_user = _patch_member_me_helpers(team=team, membership=membership, user=user) with p_team, p_membership, p_user: response = await team_member_me( http_request=MagicMock(spec=Request), @@ -10088,9 +9628,7 @@ async def test_team_member_me_returns_404_for_non_member(mock_db_client): team_id = "team-me-2" caller_id = "outsider@example.com" - caller_auth = UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id=caller_id - ) + caller_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id=caller_id) team = _build_team_for_me( team_id, @@ -10109,9 +9647,7 @@ async def test_team_member_me_returns_404_for_non_member(mock_db_client): @pytest.mark.asyncio -async def test_team_member_me_returns_404_for_proxy_admin_not_in_team( - mock_db_client, mock_admin_auth -): +async def test_team_member_me_returns_404_for_proxy_admin_not_in_team(mock_db_client, mock_admin_auth): """ Proxy admins get 404 if they are not actually a member of the team. `me` only resolves for actual team members; admins use /team/info instead. @@ -10151,9 +9687,7 @@ async def test_team_member_me_returns_defaults_when_no_membership_row(mock_db_cl team_id = "team-me-4" caller_id = "newmember@example.com" - caller_auth = UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id=caller_id - ) + caller_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id=caller_id) team = _build_team_for_me( team_id, @@ -10199,18 +9733,12 @@ async def test_team_member_me_returns_404_for_unknown_team(mock_db_client): from litellm.proxy.management_endpoints.team_endpoints import team_member_me - caller_auth = UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id="alice@example.com" - ) + caller_auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="alice@example.com") # get_team_object raises 404 directly when the team is missing. with patch( "litellm.proxy.management_endpoints.team_endpoints.get_team_object", - AsyncMock( - side_effect=HTTPException( - status_code=404, detail={"error": "Team doesn't exist in db."} - ) - ), + AsyncMock(side_effect=HTTPException(status_code=404, detail={"error": "Team doesn't exist in db."})), ): with pytest.raises(HTTPException) as exc_info: await team_member_me( @@ -10222,9 +9750,7 @@ async def test_team_member_me_returns_404_for_unknown_team(mock_db_client): @pytest.mark.asyncio -async def test_new_team_encrypts_callback_vars( - mock_db_client, mock_admin_auth, monkeypatch -): +async def test_new_team_encrypts_callback_vars(mock_db_client, mock_admin_auth, monkeypatch): """/team/new must encrypt callback_vars values before they reach the DB.""" from fastapi import Request @@ -10239,9 +9765,7 @@ async def test_new_team_encrypts_callback_vars( # actual JSON serialization production uses (catches non-serializable # ciphertext, missing fields, etc.). mock_db_client.jsonify_object = PrismaClient.jsonify_object.__get__(mock_db_client) - mock_db_client.jsonify_team_object = PrismaClient.jsonify_team_object.__get__( - mock_db_client - ) + mock_db_client.jsonify_team_object = PrismaClient.jsonify_team_object.__get__(mock_db_client) mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.db = MagicMock() mock_db_client.db.litellm_teamtable = MagicMock() @@ -10251,9 +9775,7 @@ async def test_new_team_encrypts_callback_vars( mock_db_client.db.litellm_teamtable.create = mock_team_create _wire_team_create_tx(mock_db_client) mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0) - mock_db_client.db.litellm_teamtable.update = AsyncMock( - return_value=team_create_result - ) + mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=team_create_result) mock_db_client.db.litellm_usertable = MagicMock() mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock()) @@ -10290,9 +9812,7 @@ async def test_new_team_encrypts_callback_vars( def _non_admin_auth(): - return UserAPIKeyAuth( - user_id="u-team-admin", user_role=LitellmUserRoles.INTERNAL_USER - ) + return UserAPIKeyAuth(user_id="u-team-admin", user_role=LitellmUserRoles.INTERNAL_USER) def test_check_passthrough_routes_caller_permission_team(): @@ -10308,12 +9828,8 @@ def test_check_passthrough_routes_caller_permission_team(): NewTeamRequest(allowed_passthrough_routes=["/foo/*"]), admin, entity="team" ) - _check_passthrough_routes_caller_permission( - NewTeamRequest(), non_admin, entity="team" - ) - _check_passthrough_routes_caller_permission( - NewTeamRequest(allowed_passthrough_routes=[]), non_admin, entity="team" - ) + _check_passthrough_routes_caller_permission(NewTeamRequest(), non_admin, entity="team") + _check_passthrough_routes_caller_permission(NewTeamRequest(allowed_passthrough_routes=[]), non_admin, entity="team") with pytest.raises(HTTPException) as exc: _check_passthrough_routes_caller_permission( @@ -10350,9 +9866,7 @@ async def test_new_team_blocks_non_admin_passthrough_routes(mock_db_client): ): with pytest.raises(ProxyException) as exc: await new_team( - data=NewTeamRequest( - team_alias="t", allowed_passthrough_routes=["/admin/*"] - ), + data=NewTeamRequest(team_alias="t", allowed_passthrough_routes=["/admin/*"]), http_request=MagicMock(spec=Request), user_api_key_dict=_non_admin_auth(), ) @@ -10379,9 +9893,7 @@ async def test_update_team_blocks_non_admin_passthrough_routes(mock_db_client): ): with pytest.raises(ProxyException) as exc: await update_team( - data=UpdateTeamRequest( - team_id="t1", allowed_passthrough_routes=["/admin/*"] - ), + data=UpdateTeamRequest(team_id="t1", allowed_passthrough_routes=["/admin/*"]), http_request=MagicMock(spec=Request), user_api_key_dict=_non_admin_auth(), ) @@ -10725,16 +10237,12 @@ async def test_team_info_forwards_key_limit_to_get_data(): from litellm.proxy.management_endpoints import team_endpoints mock_prisma = MagicMock() - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=LiteLLM_TeamTable(team_id="team-1") - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=LiteLLM_TeamTable(team_id="team-1")) mock_prisma.get_data = AsyncMock(return_value=[]) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), - patch.object( - team_endpoints, "get_all_team_memberships", AsyncMock(return_value=[]) - ), + patch.object(team_endpoints, "get_all_team_memberships", AsyncMock(return_value=[])), ): await team_endpoints.team_info( http_request=MagicMock(spec=Request), @@ -10772,9 +10280,7 @@ async def test_team_info_returns_model_aliases(): with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), - patch.object( - team_endpoints, "get_all_team_memberships", AsyncMock(return_value=[]) - ), + patch.object(team_endpoints, "get_all_team_memberships", AsyncMock(return_value=[])), ): response = await team_endpoints.team_info( http_request=MagicMock(spec=Request), @@ -10857,12 +10363,8 @@ async def test_update_model_table_clears_aliases_with_empty_map(): """ mock_prisma = MagicMock() mock_prisma.db.litellm_modeltable.create = AsyncMock() - mock_prisma.db.litellm_modeltable.upsert = AsyncMock( - return_value=MagicMock(id="model-123") - ) - user_api_key_dict = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin" - ) + mock_prisma.db.litellm_modeltable.upsert = AsyncMock(return_value=MagicMock(id="model-123")) + user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin") returned_model_id = await _update_model_table( data=UpdateTeamRequest(team_id="team-1", model_aliases={}), @@ -10910,9 +10412,7 @@ class TestEmitTeamMembersMetric: return LiteLLM_TeamTable( team_id="team-x", team_alias="X", - members_with_roles=[ - Member(user_id=f"u{i}", role="user") for i in range(member_count) - ], + members_with_roles=[Member(user_id=f"u{i}", role="user") for i in range(member_count)], ) def test_emits_with_team_when_logger_registered(self, restore_callbacks): @@ -10986,9 +10486,7 @@ async def test_new_team_rejects_reserved_ui_session_team_id(): await new_team( data=team_request, http_request=dummy_request, - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN - ), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), ) assert exc_info.value.code == "400" @@ -11066,9 +10564,7 @@ async def _drive_team_write( new=AsyncMock(), ), ): - pc.db.litellm_teamtable.find_unique = AsyncMock( - return_value=None if find_returns_none else existing - ) + pc.db.litellm_teamtable.find_unique = AsyncMock(return_value=None if find_returns_none else existing) pc.db.litellm_teamtable.update = AsyncMock( return_value=LiteLLM_TeamTable(team_id=_PATCH_TEAM_ID, team_alias="t") ) @@ -11169,9 +10665,7 @@ _METADATA_MAPPING = [ _METADATA_MAPPING, ids=[row[0] for row in _METADATA_MAPPING], ) -async def test_post_vs_patch_metadata_write_mapping( - label, existing_metadata, body, expected_post, expected_patch -): +async def test_post_vs_patch_metadata_write_mapping(label, existing_metadata, body, expected_post, expected_patch): """Exhaustive map: POST replaces metadata wholesale, PATCH merges per RFC 7386.""" post_meta = await _written_metadata("post", existing_metadata, body) patch_meta = await _written_metadata("patch", existing_metadata, body) @@ -11299,9 +10793,7 @@ async def test_patch_team_not_found_returns_404(): # metadata present -> patch_team does its own existence check with pytest.raises(ProxyException) as exc: - await _drive_team_write( - "patch", raw_body={"metadata": {"cost_center": "1"}}, find_returns_none=True - ) + await _drive_team_write("patch", raw_body={"metadata": {"cost_center": "1"}}, find_returns_none=True) assert exc.value.code == "404" or exc.value.code == 404 # metadata absent -> existence check happens in the delegated update_team @@ -11318,9 +10810,7 @@ async def test_patch_enforces_team_access_via_delegation(): outsider = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="outsider") with pytest.raises(ProxyException) as exc: - await _drive_team_write( - "patch", raw_body={"tpm_limit": 5}, user=outsider - ) + await _drive_team_write("patch", raw_body={"tpm_limit": 5}, user=outsider) assert exc.value.code == "403" or exc.value.code == 403 @@ -11330,9 +10820,7 @@ async def test_patch_returns_full_team_object_not_wrapper(): {"team_id", "data"} envelope.""" from litellm.proxy._types import LiteLLM_TeamTable - result, _ = await _drive_team_write( - "patch", existing_metadata={"a": 1}, raw_body={"metadata": {"b": 2}} - ) + result, _ = await _drive_team_write("patch", existing_metadata={"a": 1}, raw_body={"metadata": {"b": 2}}) assert isinstance(result, LiteLLM_TeamTable) assert result.team_id == _PATCH_TEAM_ID @@ -11649,9 +11137,7 @@ def test_patch_body_reshaping_adds_no_keys_the_caller_did_not_send(body): @pytest.mark.asyncio async def test_patch_ignores_unknown_body_keys(): """Unknown keys were silently dropped by the previous construction; keep that.""" - _, update_mock = await _drive_team_write( - "patch", raw_body={"tpm_limit": 5, "not_a_team_field": "x"} - ) + _, update_mock = await _drive_team_write("patch", raw_body={"tpm_limit": 5, "not_a_team_field": "x"}) written = update_mock.call_args.kwargs["data"] assert written["tpm_limit"] == 5 @@ -11966,9 +11452,7 @@ async def test_resolve_existing_member_user_ids_skips_the_query_when_no_user_ids def _user_row(user_id: str, user_email: str | None) -> LiteLLM_UserTable: - return LiteLLM_UserTable( - user_id=user_id, user_email=user_email, max_budget=None, spend=0.0, models=[] - ) + return LiteLLM_UserTable(user_id=user_id, user_email=user_email, max_budget=None, spend=0.0, models=[]) @pytest.mark.asyncio @@ -12506,9 +11990,7 @@ async def test_get_team_daily_activity_aggregated_scopes_and_flags(mock_db_clien user_id = "test_user_123" team_id = "test_team_456" - user_api_key_dict = UserAPIKeyAuth( - user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) + user_api_key_dict = UserAPIKeyAuth(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER) mock_user_info = LiteLLM_UserTable( user_id=user_id, @@ -12534,9 +12016,7 @@ async def test_get_team_daily_activity_aggregated_scopes_and_flags(mock_db_clien user_api_key_1.token = "user_key_1" mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team]) - mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[user_api_key_1] - ) + mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[user_api_key_1]) with patch( "litellm.proxy.management_endpoints.team_endpoints.get_user_object", @@ -12565,9 +12045,7 @@ async def test_get_team_daily_activity_aggregated_scopes_and_flags(mock_db_clien call_kwargs = mock_aggregated.call_args[1] assert call_kwargs["api_key"] == ["user_key_1"] assert call_kwargs["entity_id"] == [team_id] - assert call_kwargs["entity_metadata_field"] == { - team_id: {"team_alias": "Test Team"} - } + assert call_kwargs["entity_metadata_field"] == {team_id: {"team_alias": "Test Team"}} assert call_kwargs["include_entity_breakdown"] is True assert call_kwargs["timezone_offset_minutes"] == 480 assert call_kwargs["table_name"] == "litellm_dailyteamspend" @@ -12606,9 +12084,7 @@ async def test_get_team_daily_activity_aggregated_rejects_bad_ranges( api_key=None, exclude_team_ids=None, timezone=None, - user_api_key_dict=UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ), + user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), ) assert exc_info.value.status_code == 400 @@ -12757,16 +12233,12 @@ class _FakeMirrorDb: if "array_append" in sql: self.transactions[-1].append("attach") - changed = [ - g for g in desired if g in self._access_groups and team_id not in self._team_ids(g) - ] + changed = [g for g in desired if g in self._access_groups and team_id not in self._team_ids(g)] for group_id in changed: self._team_ids(group_id).append(team_id) else: self.transactions[-1].append("detach") - changed = [ - g for g in self._access_groups if team_id in self._team_ids(g) and g not in desired - ] + changed = [g for g in self._access_groups if team_id in self._team_ids(g) and g not in desired] for group_id in changed: self._team_ids(group_id).remove(team_id) return [{"access_group_id": group_id} for group_id in changed] @@ -13013,7 +12485,9 @@ async def test_new_team_and_delete_team_both_drive_the_mirror( patch("litellm.proxy.proxy_server.prisma_client") as prisma, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), patch("litellm.proxy.proxy_server.llm_router", None), - patch("litellm.proxy.management_endpoints.team_endpoints._persist_deleted_team_records", new_callable=AsyncMock), + patch( + "litellm.proxy.management_endpoints.team_endpoints._persist_deleted_team_records", new_callable=AsyncMock + ), patch("litellm.proxy.management_endpoints.team_endpoints._verify_team_access", new_callable=AsyncMock), patch( "litellm.proxy.management_endpoints.team_endpoints.sync_team_access_group_membership", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 788a28b0772..a8c2c84b5e7 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */