From 55d159764581f33217c43544bd5f96e26330146f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 14:25:39 +0000 Subject: [PATCH] revert(proxy): drop the member auto-router write-path port from stable/1.100.x This reverts commits 54b42e3f05, 4864227716, 77915d43b8 and e39e1c8dea. Jev does not use the member auto-router write path and no other stable line ships it, so dropping the port leaves 1.100.x matching stable/1.101.x and stable/1.102.x The classifier circuit breaker's litellm Timeout detection that 54b42e3f05 carried is kept, since main and the other lines have it This reverts commit e39e1c8deaefe71d9ddea6025370939afaf98598. This reverts commit 77915d43b82d16025c91f6924354e8f5b40cf27d. This reverts commit 4864227716989aaa4787fa1ee68b7c0f1df58ad1. This reverts commit 54b42e3f05e1230db8197fcfce6533d1e6fd153b. --- litellm/models/team.py | 1 - litellm/proxy/_types.py | 4 - litellm/proxy/auth/auth_checks.py | 16 +- litellm/proxy/auth/handle_jwt.py | 9 +- litellm/proxy/auth/litellm_license.py | 22 - litellm/proxy/auth/team_grants.py | 21 +- litellm/proxy/auth/user_api_key_auth.py | 22 +- .../model_management_endpoints.py | 357 +---- .../management_endpoints/team_endpoints.py | 58 +- litellm/proxy/management_endpoints/ui_sso.py | 23 +- litellm/router.py | 11 +- litellm/types/router.py | 1 - .../proxy/auth/test_auth_checks.py | 604 ++++---- .../proxy/auth/test_handle_jwt.py | 727 +++++++--- .../proxy/auth/test_team_grants.py | 130 -- .../proxy/auth/test_user_api_key_auth.py | 405 +++--- .../test_model_management_endpoints.py | 132 +- .../test_team_endpoints.py | 1270 ++++++++++++----- ui/litellm-dashboard/src/lib/http/schema.d.ts | 21 - 19 files changed, 2059 insertions(+), 1775 deletions(-) delete mode 100644 tests/test_litellm/proxy/auth/test_team_grants.py diff --git a/litellm/models/team.py b/litellm/models/team.py index 8edf10703b1..da526515e6e 100644 --- a/litellm/models/team.py +++ b/litellm/models/team.py @@ -71,7 +71,6 @@ 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 69897660d93..d7e65f93121 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2779,12 +2779,8 @@ 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 # mutable-ok: mirrors LiteLLM_TeamTable.model_max_budget, a JSON dict column - ) 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 88d6dfca368..093865e8f71 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -473,7 +473,6 @@ 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: @@ -550,12 +549,11 @@ 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 not isinstance(model_id, str): + if model_id is None: continue - model_name = litellm_params.get("model") model_info = llm_router.get_deployment_model_info( - model_id=model_id, model_name=model_name if isinstance(model_name, str) else "" + model_id=model_id, model_name=litellm_params.get("model") or "" ) if model_info is not None and _entry_has_priced_metric(model_info): return True @@ -2849,10 +2847,7 @@ 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}, # mutable-ok: prisma where clause - include=_TEAM_GRANT_RELATIONS, - ) + response = await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id}) if response is None and team_id_upsert: from litellm.proxy.management_endpoints.team_endpoints import new_team @@ -3152,10 +3147,7 @@ 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}, # mutable-ok: prisma where clause - include=_TEAM_GRANT_RELATIONS, - ) + teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(where={"team_alias": team_alias}) if not teams: raise HTTPException( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 1c1303aca82..39e6ca9a369 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -52,7 +52,6 @@ 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, @@ -1554,9 +1553,7 @@ class JWTAuthManager: model=requested_model, team_object=team_object, llm_router=llm_router, - team_model_aliases=dict(aliases) # mutable-ok: can_team_access_model takes a dict - if (aliases := team_model_aliases(team_object)) is not None - else None, + team_model_aliases=None, ) ): is_allowed = allowed_routes_check( @@ -2093,9 +2090,7 @@ class JWTAuthManager: model=requested_model, team_object=team_object, llm_router=llm_router, - team_model_aliases=dict(aliases) # mutable-ok: can_team_access_model takes a dict - if (aliases := team_model_aliases(team_object)) is not None - else None, + team_model_aliases=None, ) except ProxyException: continue diff --git a/litellm/proxy/auth/litellm_license.py b/litellm/proxy/auth/litellm_license.py index 40eda903da2..677f1a0fdda 100644 --- a/litellm/proxy/auth/litellm_license.py +++ b/litellm/proxy/auth/litellm_license.py @@ -15,10 +15,6 @@ 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: """ @@ -153,24 +149,6 @@ 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 606fbcc4b15..8db7e642728 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 +from collections.abc import Mapping, Sequence 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[list[str]] # mutable-ok: UserAPIKeyAuth declares a list field + team_models: ReadOnly[Sequence[str]] team_blocked: ReadOnly[bool] - 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_metadata: ReadOnly[Mapping[str, object] | None] + team_model_aliases: ReadOnly[Mapping[str, str] | None] team_object_permission_id: ReadOnly[str | None] team_object_permission: ReadOnly[LiteLLM_ObjectPermissionTable | None] team_member: ReadOnly[Member | None] @@ -104,18 +104,11 @@ 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=list(team_object.models), # mutable-ok: UserAPIKeyAuth declares a list field + team_models=tuple(team_object.models), team_blocked=team_object.blocked, - team_metadata=( - dict(json_columns.metadata) - if json_columns.metadata is not None - else None # mutable-ok: UserAPIKeyAuth declares a dict field - ), + team_metadata=json_columns.metadata, team_model_aliases=( - dict(json_columns.litellm_model_table.model_aliases) # mutable-ok: UserAPIKeyAuth declares a dict field - if json_columns.litellm_model_table is not None - and json_columns.litellm_model_table.model_aliases is not None - else None + json_columns.litellm_model_table.model_aliases if json_columns.litellm_model_table 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 dfea802497b..e92d090a2fb 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -77,7 +77,6 @@ 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 ( @@ -1470,16 +1469,24 @@ 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 @@ -1493,8 +1500,17 @@ 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, - **team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id), + ) + valid_token.team_object_permission = ( + team_object.object_permission if team_object is not None else None ) # 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 4e27efad609..cf9146d7a65 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -13,13 +13,10 @@ model/{model_id}/update - PATCH endpoint for model update. import asyncio import datetime import json -from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence -from contextlib import AbstractAsyncContextManager, asynccontextmanager -from dataclasses import dataclass -from fnmatch import fnmatchcase +from collections.abc import Awaitable, Mapping, Sequence from json import JSONDecodeError from types import MappingProxyType -from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeVar, cast, runtime_checkable +from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError, field_validator @@ -57,12 +54,11 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( coordination_redis_cache, publish_config_change, ) -from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.common_utils.encrypt_decrypt_utils import 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 ( _refresh_cached_team, - append_team_models, team_model_add, team_model_delete, ) @@ -70,13 +66,6 @@ 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, @@ -114,13 +103,11 @@ 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() @@ -168,39 +155,8 @@ 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): @@ -285,147 +241,6 @@ 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) - - -AUTO_ROUTER_WRITE_SLOT_LOCK_KEY: Final = 5_872_301 -_WRITE_SLOT_LOCK_SQL: Final = "SELECT 1 AS locked FROM pg_advisory_xact_lock($1)" - - -@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 - await tx_ctx.query_raw(_WRITE_SLOT_LOCK_SQL, AUTO_ROUTER_WRITE_SLOT_LOCK_KEY) - 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, @@ -909,39 +724,11 @@ async def patch_model( param=None, ) - write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call( + 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 @@ -960,26 +747,22 @@ 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}, # mutable-ok: prisma where clause - data=update_data, - ) - # Handle team model updates with proper alias management - updated_model: Final = await _update_team_model_in_db( + update_data: Final = await _update_team_model_in_db( db_model=db_model, - patch_data=effective_patch, + patch_data=patch_data, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, - write_row=write_row, + ) + + # 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, ) if updated_model is None: @@ -1210,7 +993,6 @@ 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) @@ -1229,19 +1011,17 @@ 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 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) + 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 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, @@ -1250,8 +1030,6 @@ async def _add_team_model_to_db( - store the model in the db with the unique 'model_name' - add the public model name to the team's allowed models list """ - from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache - _team_id: Final = model_params.model_info.team_id if _team_id is None: return None @@ -1275,18 +1053,16 @@ 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: - await append_team_models( + await team_model_add( data=TeamModelAddRequest( team_id=_team_id, models=[original_model_name], ), - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, + http_request=Request(scope={"type": "http"}), + user_api_key_dict=user_api_key_dict, ) return model_response @@ -1297,8 +1073,7 @@ async def _update_team_model_in_db( patch_data: updateDeployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, - write_row: Callable[[PrismaCompatibleUpdateDBModel], Awaitable[_RowT]], -) -> _RowT: +) -> PrismaCompatibleUpdateDBModel: """ Handle team model updates with proper alias management. @@ -1306,9 +1081,6 @@ 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 @@ -1342,7 +1114,7 @@ async def _update_team_model_in_db( # No team_id in patch, proceed with standard update if patch_team_id is None: - return await write_row(update_db_model(db_model=db_model, updated_patch=patch_data)) + return update_db_model(db_model=db_model, updated_patch=patch_data) # Determine public model name public_model_name: Final = _get_public_model_name( @@ -1361,10 +1133,6 @@ 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, @@ -1382,7 +1150,7 @@ async def _update_team_model_in_db( prisma_client=prisma_client, ) - return row + return update_db_model(db_model=db_model, updated_patch=patch_data) def _get_public_model_name( @@ -1545,7 +1313,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 isinstance(public_name, str) and public_name: + if public_name: public_names.add(public_name) return public_names @@ -1778,14 +1546,7 @@ class ModelManagementAuthChecks: prisma_client: PrismaClient, premium_user: bool, allow_missing_team: bool = False, - 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.") + ) -> Literal[True]: ## 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( @@ -1808,27 +1569,6 @@ 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, @@ -2080,18 +1820,12 @@ async def add_new_model( ) ## Auth check - write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call( + 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, @@ -2120,20 +1854,17 @@ 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 @@ -2261,17 +1992,11 @@ async def update_model( raise Exception("model not found") deployment: Final = Deployment(**_existing_litellm_params.model_dump()) - write_authorization: Final = await ModelManagementAuthChecks.can_user_make_model_call( + 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( @@ -2313,24 +2038,14 @@ async def update_model( if value is not None or _existing_litellm_params_dict.get(key) is not None } - _data: Final[dict[str, str]] = { # mutable-ok: prisma update payload is dict-shaped + _data: Final[dict[str, str]] = { "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 {} - ), } - 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, - ) + model_response: Final = await _proxy_model_table(prisma_client).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 e7f4217634c..c6d7975b75e 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -425,58 +425,6 @@ async def _refresh_cached_team( ) -async def append_team_models( - *, - data: TeamModelAddRequest, - prisma_client: PrismaClient, - user_api_key_cache: UserApiKeyCache, - proxy_logging_obj: ProxyLogging, -) -> "prisma_models.LiteLLM_TeamTable": - # Atomic array append with dedup at the database level so concurrent - # BYOK model creates don't overwrite each other's team.models entries. - # When the team currently has models=[] (unrestricted access), the - # CASE expression inserts the 'all-proxy-models' sentinel first. - models_to_add: Final = list(data.models) - await prisma_client.db.execute_raw( - 'UPDATE "LiteLLM_TeamTable" ' - "SET models = (" - " SELECT ARRAY(SELECT DISTINCT unnest(" - " CASE WHEN cardinality(COALESCE(models, ARRAY[]::text[])) = 0 " - " THEN ARRAY['all-proxy-models']::text[] " - " ELSE models " - " END || $1::text[]" - " ))" - ") " - "WHERE team_id = $2", - models_to_add, - data.team_id, - ) - # Re-fetch via update (write-routed) instead of find_unique (read-routed) - # to avoid returning stale data from a read replica. The models column - # was already set by execute_raw above; this just retrieves the row from - # the writer and lets Prisma bump updated_at. - # `include` mirrors the relations the auth path consumes off the cached - # team object so that `_refresh_cached_team` doesn't null them out. - updated_team: Final = await _team_db(prisma_client).update( - where={"team_id": data.team_id}, - data={"updated_at": datetime.now(timezone.utc)}, - include={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause - ) - if updated_team is None: - raise HTTPException( - status_code=404, - detail={"error": f"Team not found, passed team_id={data.team_id}"}, - ) - - await _refresh_cached_team( - team_row=updated_team, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - - return updated_team - - async def _verify_team_access( team_obj: LiteLLM_TeamTable, user_api_key_dict: UserAPIKeyAuth, @@ -4416,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}, # mutable-ok: prisma include clause + include={"litellm_model_table": True, "object_permission": True}, ) if team_info is None: raise Exception @@ -5619,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={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause + include={"object_permission": True}, ) if updated_team is None: raise HTTPException( @@ -5706,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={"litellm_model_table": True, "object_permission": True}, # mutable-ok: prisma include clause + include={"object_permission": True}, ) 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 aa11a9c2e12..613508da22b 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -22,6 +22,7 @@ from html import escape from types import MappingProxyType from typing import ( TYPE_CHECKING, + Annotated, Any, Final, Literal, @@ -40,7 +41,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, TypeAdapter, ValidationError +from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError import litellm from litellm._logging import verbose_proxy_logger @@ -91,7 +92,6 @@ 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,14 +202,31 @@ 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 1a2114cdab7..63937ce5e76 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, Iterator, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, 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,15 +9446,6 @@ 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/types/router.py b/litellm/types/router.py index 8295444fa9f..97bd93f3f47 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -160,7 +160,6 @@ 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 41b2f8d6c13..3dea89ed67b 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -126,10 +126,14 @@ 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) @@ -155,7 +159,9 @@ 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) @@ -182,7 +188,9 @@ 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) @@ -199,7 +207,9 @@ 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) @@ -212,8 +222,12 @@ 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")) @@ -230,33 +244,43 @@ 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) @@ -285,7 +309,9 @@ 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", @@ -542,7 +568,9 @@ 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() @@ -585,7 +613,9 @@ 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() @@ -600,7 +630,8 @@ 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" @@ -609,7 +640,9 @@ 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) @@ -631,7 +664,9 @@ 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 @@ -646,10 +681,14 @@ def test_get_cli_jwt_auth_token_custom_expiration(valid_sso_user_defined_values, # 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) @@ -667,12 +706,18 @@ 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}-") @@ -695,7 +740,9 @@ 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 @@ -898,7 +945,9 @@ 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" @@ -913,8 +962,12 @@ 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) @@ -943,7 +996,9 @@ 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" @@ -955,7 +1010,9 @@ 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", @@ -963,13 +1020,15 @@ 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, @@ -983,7 +1042,9 @@ 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" @@ -1001,7 +1062,9 @@ 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) @@ -1028,8 +1091,12 @@ 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", @@ -1038,7 +1105,9 @@ 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", @@ -1052,7 +1121,9 @@ 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" @@ -1067,8 +1138,12 @@ 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", @@ -1077,7 +1152,9 @@ 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", @@ -1091,7 +1168,9 @@ 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" @@ -1145,7 +1224,10 @@ 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(): @@ -1172,7 +1254,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 @@ -1208,10 +1290,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. @@ -1245,7 +1327,9 @@ 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"}]} @@ -1295,7 +1379,9 @@ async def test_vector_store_access_check_early_returns(prisma_client, vector_sto ), # 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: @@ -1303,7 +1389,11 @@ def test_can_object_call_vector_stores_scenarios(object_permissions, vector_stor 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: @@ -1338,7 +1428,9 @@ 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"] @@ -1384,10 +1476,14 @@ 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), @@ -1401,7 +1497,9 @@ 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), @@ -1962,7 +2060,9 @@ 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( @@ -2064,7 +2164,9 @@ 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( @@ -2075,7 +2177,9 @@ 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"], @@ -2244,7 +2348,9 @@ 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: @@ -2261,7 +2367,10 @@ 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"], @@ -2286,12 +2395,17 @@ 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 @@ -2350,44 +2464,6 @@ 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, @@ -2738,7 +2814,8 @@ 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}) @@ -2748,7 +2825,8 @@ 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}) @@ -2913,7 +2991,9 @@ 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 @@ -2942,9 +3022,9 @@ async def test_virtual_key_soft_budget_check_scenarios(spend, soft_budget, expec 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 @@ -3055,7 +3135,9 @@ 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 @@ -3084,9 +3166,9 @@ async def test_virtual_key_max_budget_alert_check_scenarios(spend, max_budget, e 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 @@ -3345,7 +3427,9 @@ 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", @@ -3417,7 +3501,9 @@ 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 @@ -3451,7 +3537,9 @@ 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): @@ -3481,7 +3569,9 @@ 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) @@ -3542,7 +3632,9 @@ 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) @@ -3638,7 +3730,9 @@ 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 @@ -3665,7 +3759,9 @@ 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 @@ -3693,7 +3789,9 @@ 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 @@ -3743,7 +3841,9 @@ 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 @@ -3783,7 +3883,9 @@ 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 @@ -3831,7 +3933,9 @@ 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( @@ -3965,12 +4069,18 @@ 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 @@ -4061,11 +4171,15 @@ 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 @@ -4146,12 +4260,18 @@ 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 @@ -4209,12 +4329,18 @@ 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 @@ -4272,12 +4398,18 @@ 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 @@ -4335,7 +4467,9 @@ 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 @@ -4383,13 +4517,19 @@ 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) @@ -4424,7 +4564,9 @@ 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) @@ -4458,7 +4600,9 @@ 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 @@ -4476,7 +4620,9 @@ async def test_resolve_end_user_matches_user_table_by_user_id(_validate_flag_on, @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 @@ -4504,7 +4650,9 @@ async def test_resolve_end_user_matches_user_table_by_email(_validate_flag_on, m @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 @@ -4523,7 +4671,9 @@ async def test_resolve_end_user_non_email_id_does_not_pass_user_email(_validate_ @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 @@ -4545,7 +4695,9 @@ async def test_resolve_end_user_drops_codex_opaque_identifier(_validate_flag_on, @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 @@ -4582,7 +4734,9 @@ 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 @@ -4602,7 +4756,9 @@ async def test_resolve_end_user_uses_cached_valid_result(_validate_flag_on, monk @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 @@ -4621,7 +4777,9 @@ async def test_resolve_end_user_uses_cached_invalid_result(_validate_flag_on, mo @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 @@ -4742,13 +4900,19 @@ 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 @@ -4782,7 +4946,10 @@ 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"] @@ -4862,7 +5029,9 @@ 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, ) @@ -4972,7 +5141,9 @@ 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", @@ -4985,7 +5156,10 @@ 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"] @@ -5261,11 +5435,8 @@ async def test_common_checks_budget_reads_run_concurrently(): probe = _BudgetSpendConcurrencyProbe(expected=4) - 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 + with patch("litellm.proxy.proxy_server.prisma_client", None), patch( + "litellm.proxy.proxy_server.get_current_spend", probe ): task = asyncio.create_task( common_checks( @@ -5333,11 +5504,8 @@ async def test_common_checks_budget_gather_raises_highest_priority_scope(): request=MagicMock(spec=Request), ) - 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 patch("litellm.proxy.proxy_server.prisma_client", None), patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter ): # Both team and end-user over budget: team wins on priority. _spend_by_counter.team = 999.0 @@ -5372,11 +5540,8 @@ 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), # 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 patch("litellm.proxy.proxy_server.prisma_client", None), patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter ): with pytest.raises(litellm.BudgetExceededError) as over: await common_checks( @@ -5418,15 +5583,9 @@ 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), # 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 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): result = await common_checks( request_body={"messages": [{"role": "user", "content": "hi"}]}, team_object=team, @@ -5465,15 +5624,9 @@ 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), # 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 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 pytest.raises(litellm.BudgetExceededError) as exc_info: await common_checks( request_body={"messages": [{"role": "user", "content": "hi"}]}, @@ -5504,11 +5657,8 @@ 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), # 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 patch("litellm.proxy.proxy_server.prisma_client", None), patch( + "litellm.proxy.proxy_server.get_current_spend", _spend_by_counter ): with pytest.raises(litellm.BudgetExceededError): await common_checks( @@ -5591,19 +5741,10 @@ 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() - ), # 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 + 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 ): if expect_blocked: with pytest.raises(litellm.BudgetExceededError): @@ -6102,7 +6243,9 @@ 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( @@ -6233,32 +6376,6 @@ 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 @@ -6810,7 +6927,9 @@ 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 @@ -6926,15 +7045,11 @@ 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 ) @@ -7012,7 +7127,8 @@ 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" @@ -7208,9 +7324,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 f0f9adfb580..99a0a4c0a8b 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -13,7 +13,6 @@ from litellm.proxy._types import ( DEFAULT_JWKS_STALE_TTL, JWTLiteLLMRoleMap, LiteLLM_JWTAuth, - LiteLLM_ModelTable, LiteLLM_TeamMembership, LiteLLM_TeamTable, LiteLLM_UserTable, @@ -97,7 +96,9 @@ 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 @@ -134,10 +135,14 @@ 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) @@ -193,7 +198,9 @@ 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 == {} @@ -286,7 +293,9 @@ 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() @@ -297,10 +306,12 @@ 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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_check_rbac, 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", @@ -308,7 +319,9 @@ 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", @@ -323,7 +336,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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -338,10 +351,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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_map_user, patch.object( JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_validate_object, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -375,7 +388,9 @@ 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() @@ -386,10 +401,12 @@ 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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_check_rbac, 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", @@ -397,7 +414,9 @@ 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", @@ -412,7 +431,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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -427,10 +446,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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_map_user, patch.object( JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_validate_object, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -579,7 +598,11 @@ 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, @@ -588,7 +611,9 @@ 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() @@ -615,7 +640,11 @@ 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, @@ -632,7 +661,9 @@ 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 @@ -652,7 +683,11 @@ 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, @@ -673,7 +708,9 @@ 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 @@ -693,7 +730,11 @@ 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, @@ -709,7 +750,9 @@ 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() @@ -734,7 +777,9 @@ 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"]}) == [ @@ -743,7 +788,9 @@ 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", @@ -803,11 +850,19 @@ 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", ), @@ -879,7 +934,9 @@ 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) @@ -974,19 +1031,25 @@ 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) @@ -1002,7 +1065,9 @@ 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) @@ -1048,7 +1113,10 @@ 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") @@ -1063,28 +1131,43 @@ 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 @@ -1109,7 +1192,9 @@ 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" @@ -1145,7 +1230,9 @@ 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() @@ -1168,57 +1255,6 @@ 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 @@ -1253,7 +1289,9 @@ 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() @@ -1319,10 +1357,14 @@ 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) @@ -1335,10 +1377,12 @@ 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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_check_rbac, 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", @@ -1346,7 +1390,9 @@ 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", @@ -1361,7 +1407,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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1382,13 +1428,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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_map_user, patch.object( JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_validate_object, patch.object( JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_sync_user, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": _user_id, "scope": ""} @@ -1407,12 +1453,24 @@ 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 @@ -1429,7 +1487,9 @@ 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() @@ -1456,14 +1516,18 @@ 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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_check_rbac, 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", @@ -1471,7 +1535,9 @@ 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", @@ -1486,7 +1552,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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1501,13 +1567,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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_map_user, patch.object( JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_validate_object, patch.object( JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_sync_user, ): # Set up mock return values mock_get_userinfo.return_value = userinfo_response @@ -1548,7 +1614,9 @@ 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() @@ -1572,14 +1640,18 @@ 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, # test-quality-ok: [TQ008] collaborator injected via its import site; there is no seam to patch otherwise + ) as mock_check_rbac, 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", @@ -1587,7 +1659,9 @@ 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", @@ -1600,7 +1674,9 @@ 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", @@ -1613,9 +1689,15 @@ 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 @@ -1661,7 +1743,9 @@ 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() @@ -1680,7 +1764,9 @@ 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), @@ -1721,7 +1807,9 @@ 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 @@ -1791,7 +1879,9 @@ 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, @@ -1802,7 +1892,9 @@ 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", @@ -1810,7 +1902,9 @@ 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", @@ -2027,7 +2121,9 @@ 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, @@ -2056,7 +2152,9 @@ 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 @@ -2150,7 +2248,9 @@ 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", @@ -2202,7 +2302,9 @@ 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", @@ -2256,7 +2358,9 @@ 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", @@ -2300,7 +2404,9 @@ 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) @@ -2308,7 +2414,9 @@ 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 @@ -2339,7 +2447,9 @@ 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 @@ -2371,7 +2481,9 @@ 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( @@ -2414,7 +2526,9 @@ 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 @@ -2424,7 +2538,9 @@ 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, @@ -2504,7 +2620,9 @@ 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 @@ -2512,7 +2630,9 @@ 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 @@ -2550,7 +2670,9 @@ 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 ( @@ -2628,7 +2750,9 @@ 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() @@ -2659,7 +2783,9 @@ 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() @@ -2803,7 +2929,9 @@ 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 @@ -2831,7 +2959,9 @@ 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 @@ -2941,7 +3071,9 @@ 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: @@ -3000,7 +3132,9 @@ 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, @@ -3017,7 +3151,9 @@ 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 @@ -3114,7 +3250,9 @@ 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, @@ -3126,7 +3264,9 @@ 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", @@ -3161,7 +3301,9 @@ 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) @@ -3172,7 +3314,9 @@ def test_build_decode_kwargs_no_env_disables_both_verifications(monkeypatch, _re 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) @@ -3184,7 +3328,9 @@ def test_build_decode_kwargs_audience_only_enables_aud_verification(monkeypatch, 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/") @@ -3195,7 +3341,9 @@ def test_build_decode_kwargs_issuer_only_enables_iss_verification(monkeypatch, _ 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/") @@ -3207,7 +3355,9 @@ def test_build_decode_kwargs_both_set_enables_full_verification(monkeypatch, _re 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 @@ -3223,12 +3373,17 @@ def test_build_decode_kwargs_warns_once_when_unscoped(monkeypatch, _reset_unscop 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") @@ -3237,7 +3392,11 @@ def test_build_decode_kwargs_no_warning_when_scoped(monkeypatch, _reset_unscoped 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 == [] @@ -3295,7 +3454,11 @@ 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) @@ -3367,7 +3530,9 @@ 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( @@ -3440,7 +3605,9 @@ 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() @@ -3505,7 +3672,11 @@ 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) @@ -3542,7 +3713,12 @@ 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(): @@ -3550,20 +3726,28 @@ 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(): @@ -3573,7 +3757,10 @@ 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" ) @@ -3589,7 +3776,9 @@ 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"}, @@ -3625,7 +3814,9 @@ 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"}, @@ -3639,7 +3830,9 @@ 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") @@ -3753,7 +3946,9 @@ 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] @@ -3998,7 +4193,9 @@ 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, @@ -4346,7 +4543,9 @@ 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 @@ -4370,7 +4569,9 @@ 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( @@ -4466,7 +4667,9 @@ 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 @@ -4569,7 +4772,9 @@ 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 @@ -4602,7 +4807,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) @@ -4676,7 +4881,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) @@ -4731,7 +4936,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=[ { @@ -4748,7 +4953,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=[ { @@ -4760,7 +4965,9 @@ 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 @@ -4813,15 +5020,21 @@ 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 @@ -4861,11 +5074,15 @@ 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) @@ -4915,7 +5132,11 @@ 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 @@ -5052,7 +5273,9 @@ 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"], @@ -5130,7 +5353,9 @@ 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=[]), @@ -5170,7 +5395,9 @@ 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, @@ -5274,7 +5501,9 @@ 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"]), } @@ -5416,7 +5645,9 @@ 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=[]), @@ -5435,7 +5666,9 @@ 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", @@ -5444,7 +5677,9 @@ 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, @@ -5473,7 +5708,9 @@ 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 @@ -5493,15 +5730,19 @@ 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 @@ -5698,7 +5939,9 @@ 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, @@ -5850,7 +6093,9 @@ 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, @@ -5932,7 +6177,9 @@ 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, @@ -5989,10 +6236,14 @@ 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, @@ -6039,7 +6290,9 @@ 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, @@ -6108,7 +6361,9 @@ 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, @@ -6178,14 +6433,18 @@ 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"] @@ -6244,7 +6503,9 @@ 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, @@ -6319,7 +6580,9 @@ 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, @@ -6369,7 +6632,9 @@ 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=[]), @@ -6397,7 +6662,9 @@ 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, @@ -6457,9 +6724,13 @@ 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 deleted file mode 100644 index c0e8626dba8..00000000000 --- a/tests/test_litellm/proxy/auth/test_team_grants.py +++ /dev/null @@ -1,130 +0,0 @@ -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, - tpd_limit=200000, - 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 f7bfea6c7ed..d44f96d95bf 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,7 +210,10 @@ 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 @@ -224,7 +227,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.""" @@ -248,7 +251,9 @@ 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 @@ -348,7 +353,9 @@ 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}, @@ -379,7 +386,9 @@ 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}, @@ -410,7 +419,9 @@ 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}, @@ -446,7 +457,9 @@ 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 { @@ -708,7 +721,9 @@ 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, @@ -765,7 +780,9 @@ def _assert_get_api_key_with_custom_litellm_key_header(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, @@ -859,7 +876,10 @@ def test_routing_selector_matches_claim_parametrized(selector_value, claim_value ], ) 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(): @@ -938,9 +958,12 @@ 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. Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}" + f"team_metadata not correctly mapped. " + f"Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}" ) # Specifically verify tags are present @@ -979,7 +1002,9 @@ 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 = [ @@ -988,7 +1013,9 @@ 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 = [ @@ -1008,7 +1035,9 @@ 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 = [ @@ -1018,7 +1047,9 @@ 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 = [ @@ -1028,7 +1059,9 @@ 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 = [ @@ -1043,7 +1076,9 @@ 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 = [ @@ -1054,7 +1089,9 @@ 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 = [ @@ -1073,7 +1110,9 @@ 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 = [ @@ -1085,9 +1124,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 @@ -1134,7 +1173,9 @@ 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) @@ -1171,7 +1212,9 @@ 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) @@ -1195,30 +1238,36 @@ 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(): @@ -1256,7 +1305,9 @@ 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() @@ -1277,7 +1328,9 @@ 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) @@ -1344,7 +1397,9 @@ 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 @@ -1363,7 +1418,9 @@ 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) @@ -1415,7 +1472,9 @@ 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 @@ -1434,7 +1493,9 @@ 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) @@ -1487,7 +1548,9 @@ 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() @@ -1508,7 +1571,9 @@ 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) @@ -1928,7 +1993,10 @@ 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.""" @@ -2037,7 +2105,10 @@ 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 @@ -2229,7 +2300,9 @@ 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" @@ -2307,7 +2380,10 @@ 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): @@ -2383,7 +2459,8 @@ 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 @@ -3188,7 +3265,9 @@ 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 @@ -3320,7 +3399,9 @@ 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 @@ -3370,9 +3451,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(): @@ -3697,7 +3778,9 @@ 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. @@ -3826,7 +3909,9 @@ 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(): @@ -4231,7 +4316,9 @@ 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 @@ -4246,7 +4333,9 @@ 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") @@ -5047,7 +5136,9 @@ 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", @@ -5342,7 +5433,9 @@ 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} @@ -5386,7 +5479,9 @@ 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={ @@ -5396,7 +5491,9 @@ 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} @@ -5794,7 +5891,9 @@ 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 @@ -5844,7 +5943,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", @@ -5923,7 +6022,9 @@ 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 @@ -6044,7 +6145,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}", @@ -6083,9 +6184,13 @@ 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() @@ -6132,7 +6237,9 @@ 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}} @@ -6245,7 +6352,9 @@ 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"} @@ -6270,7 +6379,9 @@ 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() @@ -6284,7 +6395,9 @@ 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). @@ -6293,7 +6406,9 @@ 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 @@ -6363,7 +6478,9 @@ 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() @@ -6431,7 +6548,9 @@ 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() @@ -6511,7 +6630,9 @@ 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 @@ -6677,7 +6798,9 @@ 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 @@ -6687,7 +6810,9 @@ 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 @@ -6716,7 +6841,9 @@ 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() @@ -6745,119 +6872,3 @@ 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 ca32b8198d4..b9a884e910c 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,10 +31,6 @@ 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, @@ -1031,7 +1027,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, slot=None): + async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client): return MagicMock(model_id=str(uuid.uuid4())) mock_team_model_add = AsyncMock() @@ -1055,7 +1051,7 @@ class TestTeamModelSiblingRouting: side_effect=mock_add_model_to_db, ), patch( - "litellm.proxy.management_endpoints.model_management_endpoints.append_team_models", + "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add", mock_team_model_add, ), ): @@ -1211,7 +1207,6 @@ 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_") @@ -1442,7 +1437,6 @@ 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) @@ -1703,7 +1697,6 @@ 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 @@ -4487,124 +4480,3 @@ class TestTeamMemberAutoRouterWrites: assert saved == expected assert row.litellm_params["complexity_router_config"] == stored_config assert request.litellm_params.complexity_router_config == config - - @pytest.mark.asyncio - @pytest.mark.parametrize("endpoint,change", [("patch", "config"), ("legacy", "strategy"), ("patch", "unrelated")]) - 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.append_team_models", 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 c2f29eb9bad..ffa6bc601e9 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -173,7 +173,9 @@ 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() @@ -216,17 +218,27 @@ 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 @@ -278,7 +290,9 @@ 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 @@ -330,7 +344,9 @@ 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) @@ -374,7 +390,10 @@ 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() @@ -437,7 +456,9 @@ 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 @@ -462,7 +483,9 @@ 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. @@ -492,7 +515,9 @@ async def test_new_team_rejects_a_duration_that_never_advances(mock_db_client, m @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 @@ -531,13 +556,17 @@ 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( @@ -554,7 +583,9 @@ 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() @@ -623,7 +654,9 @@ 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( @@ -635,10 +668,14 @@ 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() @@ -697,14 +734,20 @@ 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 @@ -719,7 +762,9 @@ async def test_new_team_disable_auto_add_proxy_admin_flag(mock_db_client, disabl 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() @@ -730,18 +775,17 @@ async def test_new_team_disable_auto_add_proxy_admin_flag(mock_db_client, disabl 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), @@ -790,12 +834,16 @@ 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 = { @@ -851,17 +899,21 @@ 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", } @@ -907,17 +959,21 @@ 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", } @@ -1023,10 +1079,14 @@ 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( @@ -1037,11 +1097,15 @@ 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 @@ -1061,10 +1125,14 @@ 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", @@ -1076,7 +1144,10 @@ 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( @@ -1105,7 +1176,9 @@ 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( @@ -1322,7 +1395,9 @@ 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 @@ -1761,7 +1836,9 @@ 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") @@ -1891,7 +1968,9 @@ 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) @@ -2027,7 +2106,9 @@ 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 = { @@ -2064,8 +2145,12 @@ 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": @@ -2107,58 +2192,18 @@ 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", @@ -2192,7 +2237,9 @@ 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" @@ -2225,15 +2272,9 @@ 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, @@ -2244,7 +2285,9 @@ 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) @@ -2275,7 +2318,9 @@ 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, @@ -2283,7 +2328,9 @@ 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, @@ -2295,14 +2342,20 @@ 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( @@ -2341,14 +2394,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() @@ -2374,9 +2427,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() @@ -2401,11 +2454,13 @@ 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}" - print("✅ All test cases passed: team_member_budget is properly excluded from database update operations") + print( + "✅ All test cases passed: team_member_budget is properly excluded from database update operations" + ) def test_clean_team_member_fields(): @@ -2468,7 +2523,9 @@ 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", @@ -2535,7 +2592,9 @@ 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 = { @@ -2579,7 +2638,9 @@ 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"} @@ -2638,7 +2699,9 @@ 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 = {} @@ -2686,12 +2749,16 @@ 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, @@ -2732,14 +2799,18 @@ 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, @@ -2773,14 +2844,18 @@ 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, @@ -2811,7 +2886,9 @@ 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, @@ -2819,7 +2896,9 @@ 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, @@ -2831,13 +2910,19 @@ 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, @@ -2905,7 +2990,9 @@ 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) @@ -2923,7 +3010,9 @@ 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( @@ -2976,7 +3065,9 @@ 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) @@ -3021,7 +3112,9 @@ 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) @@ -3210,7 +3303,9 @@ 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", @@ -3265,7 +3360,9 @@ 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) @@ -3275,7 +3372,9 @@ 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() @@ -3398,7 +3497,9 @@ 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) @@ -3446,7 +3547,9 @@ 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 @@ -3596,11 +3699,17 @@ 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) @@ -3789,7 +3898,9 @@ 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=[]) @@ -3885,7 +3996,10 @@ 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 @@ -4045,7 +4159,9 @@ 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() @@ -4090,7 +4206,9 @@ 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() @@ -4141,7 +4259,9 @@ 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", @@ -4326,7 +4446,9 @@ 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=[]) @@ -4359,7 +4481,9 @@ 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": [], @@ -4367,24 +4491,32 @@ 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) @@ -4401,7 +4533,9 @@ 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 @@ -4411,28 +4545,38 @@ async def test_team_member_delete_cleans_verification_tokens(mock_db_client, moc 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) @@ -4450,7 +4594,9 @@ async def test_team_member_delete_cleans_verification_tokens(mock_db_client, moc @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. @@ -4476,7 +4622,9 @@ async def test_team_member_delete_reads_on_the_lock_holding_transaction(mock_db_ "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 @@ -4508,14 +4656,18 @@ async def test_team_member_delete_reads_on_the_lock_holding_transaction(mock_db_ 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} ) @@ -4556,14 +4708,18 @@ 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() @@ -4575,21 +4731,29 @@ 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) @@ -4616,7 +4780,9 @@ 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 @@ -4637,19 +4803,25 @@ async def test_team_member_delete_is_atomic_across_its_four_writes(mock_db_clien 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") ) @@ -4713,7 +4885,9 @@ 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) @@ -4741,7 +4915,9 @@ 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) @@ -4778,7 +4954,9 @@ 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) @@ -4810,13 +4988,19 @@ 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() @@ -4839,7 +5023,9 @@ 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( @@ -4900,8 +5086,12 @@ 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) @@ -4941,13 +5131,19 @@ 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() @@ -4970,7 +5166,9 @@ 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( @@ -5021,7 +5219,9 @@ 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 ) @@ -5032,8 +5232,12 @@ 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) @@ -5075,13 +5279,19 @@ 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() @@ -5104,7 +5314,9 @@ 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( @@ -5164,7 +5376,9 @@ 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) @@ -5231,7 +5445,9 @@ 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) @@ -5256,7 +5472,9 @@ 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 @@ -5300,8 +5518,12 @@ 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) @@ -5375,8 +5597,12 @@ 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) @@ -5400,7 +5626,10 @@ 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 @@ -5444,7 +5673,9 @@ 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" @@ -5455,9 +5686,13 @@ 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: @@ -5506,7 +5741,9 @@ 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" @@ -5517,9 +5754,13 @@ 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() @@ -5534,7 +5775,9 @@ 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, @@ -5586,7 +5829,9 @@ 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" @@ -5597,9 +5842,13 @@ 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: @@ -5649,7 +5898,9 @@ 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" @@ -5660,9 +5911,13 @@ 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() @@ -5677,7 +5932,9 @@ 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, @@ -5735,7 +5992,9 @@ 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() @@ -5747,9 +6006,13 @@ 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. @@ -5771,7 +6034,9 @@ 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( @@ -5825,7 +6090,9 @@ 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" @@ -5836,9 +6103,13 @@ 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( @@ -5858,7 +6129,9 @@ 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, @@ -5919,7 +6192,9 @@ 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), @@ -5934,9 +6209,13 @@ 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: @@ -5948,7 +6227,10 @@ 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 @@ -5990,7 +6272,9 @@ 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() @@ -6002,9 +6286,13 @@ 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() @@ -6018,7 +6306,9 @@ 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, @@ -6081,7 +6371,9 @@ 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), @@ -6097,9 +6389,13 @@ 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 @@ -6108,7 +6404,9 @@ 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() @@ -6121,7 +6419,9 @@ 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( @@ -6182,7 +6482,9 @@ 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), @@ -6198,11 +6500,17 @@ 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() @@ -6215,7 +6523,9 @@ 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( @@ -6274,7 +6584,9 @@ 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), @@ -6289,9 +6601,13 @@ 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: @@ -6303,7 +6619,10 @@ 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 @@ -6354,7 +6673,9 @@ 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), @@ -6370,11 +6691,17 @@ 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() @@ -6395,7 +6722,9 @@ 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( @@ -6453,7 +6782,9 @@ 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() @@ -6467,7 +6798,9 @@ 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() @@ -6481,7 +6814,9 @@ 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, @@ -6529,7 +6864,9 @@ 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() @@ -6543,7 +6880,9 @@ 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() @@ -6557,7 +6896,9 @@ 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, @@ -6774,7 +7115,9 @@ 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), @@ -6805,9 +7148,13 @@ 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 @@ -6886,9 +7233,13 @@ 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: @@ -6968,9 +7319,13 @@ 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: @@ -7059,9 +7414,13 @@ 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() @@ -7076,7 +7435,9 @@ 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 @@ -7163,7 +7524,9 @@ 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 @@ -7188,14 +7551,20 @@ 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 @@ -7203,13 +7572,17 @@ 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 = { @@ -7217,7 +7590,9 @@ 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) @@ -7247,7 +7622,11 @@ 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 @@ -7298,7 +7677,9 @@ 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] @@ -7443,7 +7824,9 @@ 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 @@ -7561,7 +7944,9 @@ 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) @@ -7655,9 +8040,7 @@ 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) @@ -7975,7 +8358,9 @@ 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) @@ -8131,7 +8516,9 @@ 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) @@ -8165,13 +8552,19 @@ 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() @@ -8194,7 +8587,9 @@ 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 @@ -8323,7 +8718,9 @@ 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() @@ -8337,7 +8734,9 @@ 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( @@ -8347,8 +8746,14 @@ 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" @@ -8361,9 +8766,13 @@ 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 @@ -8409,7 +8818,9 @@ 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 @@ -8430,7 +8841,9 @@ 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( @@ -8445,7 +8858,9 @@ 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() @@ -8507,7 +8922,9 @@ 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( @@ -8592,7 +9009,9 @@ 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( @@ -8626,7 +9045,9 @@ 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( @@ -8762,7 +9183,9 @@ 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( @@ -8794,7 +9217,9 @@ 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( @@ -8831,7 +9256,9 @@ 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} @@ -8848,7 +9275,9 @@ 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( @@ -8973,7 +9402,9 @@ 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]) @@ -9029,7 +9460,9 @@ 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 @@ -9125,7 +9558,9 @@ 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, @@ -9222,7 +9657,9 @@ 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"]) @@ -9251,7 +9688,9 @@ 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"]) @@ -9273,7 +9712,9 @@ 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"]) @@ -9311,7 +9752,9 @@ 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 @@ -9348,7 +9791,9 @@ 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"] @@ -9364,7 +9809,10 @@ 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")), ] @@ -9457,7 +9905,9 @@ 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", @@ -9539,7 +9989,9 @@ 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, @@ -9551,7 +10003,9 @@ 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), @@ -9569,7 +10023,9 @@ 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. @@ -9604,9 +10060,13 @@ 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), @@ -9628,7 +10088,9 @@ 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, @@ -9647,7 +10109,9 @@ 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. @@ -9687,7 +10151,9 @@ 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, @@ -9733,12 +10199,18 @@ 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( @@ -9750,7 +10222,9 @@ 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 @@ -9765,7 +10239,9 @@ async def test_new_team_encrypts_callback_vars(mock_db_client, mock_admin_auth, # 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() @@ -9775,7 +10251,9 @@ async def test_new_team_encrypts_callback_vars(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 = 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()) @@ -9812,7 +10290,9 @@ async def test_new_team_encrypts_callback_vars(mock_db_client, mock_admin_auth, 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(): @@ -9828,8 +10308,12 @@ 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( @@ -9866,7 +10350,9 @@ 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(), ) @@ -9893,7 +10379,9 @@ 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(), ) @@ -10237,12 +10725,16 @@ 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), @@ -10280,7 +10772,9 @@ 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), @@ -10363,8 +10857,12 @@ 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={}), @@ -10412,7 +10910,9 @@ 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): @@ -10486,7 +10986,9 @@ 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" @@ -10564,7 +11066,9 @@ 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") ) @@ -10665,7 +11169,9 @@ _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) @@ -10793,7 +11299,9 @@ 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 @@ -10810,7 +11318,9 @@ 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 @@ -10820,7 +11330,9 @@ 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 @@ -11137,7 +11649,9 @@ 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 @@ -11452,7 +11966,9 @@ 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 @@ -11990,7 +12506,9 @@ 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, @@ -12016,7 +12534,9 @@ 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", @@ -12045,7 +12565,9 @@ 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" @@ -12084,7 +12606,9 @@ 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 @@ -12233,12 +12757,16 @@ 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] @@ -12485,9 +13013,7 @@ 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 a8c2c84b5e7..788a28b0772 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -28154,8 +28154,6 @@ 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 */ @@ -29479,8 +29477,6 @@ 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 */ @@ -31963,8 +31959,6 @@ 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 */ @@ -35929,8 +35923,6 @@ 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 */ @@ -36071,8 +36063,6 @@ 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 */ @@ -38044,10 +38034,6 @@ 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 [] @@ -38062,8 +38048,6 @@ 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 */ @@ -38553,11 +38537,6 @@ 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 */