diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 5ed6a17608c..abe37b6ed7a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3706,6 +3706,13 @@ class TeamMemberUpdateResponse(MemberUpdateResponse): class TeamMemberBulkUpdateFields(LiteLLMPydanticObjectBase): + """Fields to apply to every selected member in a bulk update. Matches the + merge-patch semantics of the single-member endpoint: an omitted field is + left untouched, while an explicitly-passed ``None`` clears that field for + each member. ``require_at_least_one_field`` only enforces that the caller + set at least one field, so ``max_budget_in_team=None`` is a valid request + that clears the per-member budget.""" + max_budget_in_team: float | None = None role: Literal["admin", "user"] | None = None tpm_limit: int | None = None diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 21c7dfdcd1b..5953149ce6b 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -2924,6 +2924,23 @@ async def team_member_update( user_api_key_dict=user_api_key_dict, ) + return await _apply_team_member_update( + data=data, + returned_team_info=returned_team_info, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + + +async def _apply_team_member_update( + data: TeamMemberUpdateRequest, + returned_team_info: TeamInfoResponseObject, + prisma_client: PrismaClient, + user_api_key_dict: UserAPIKeyAuth, +) -> TeamMemberUpdateResponse: + """Apply a single member update against already-fetched team info so a bulk + caller can resolve team_info once and reuse it for every member, instead of + re-scanning the team, its keys, and all memberships per member.""" team_table = returned_team_info["team_info"] ## get user id @@ -2931,7 +2948,7 @@ async def team_member_update( if data.user_id is not None: received_user_id = data.user_id elif data.user_email is not None: - for member in returned_team_info["team_info"].members_with_roles: + for member in team_table.members_with_roles: if member.user_email is not None and member.user_email == data.user_email: received_user_id = member.user_id break @@ -3017,11 +3034,19 @@ async def bulk_update_team_members( http_request: Request, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], ): - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.proxy_server import premium_user, prisma_client if prisma_client is None: raise HTTPException(status_code=500, detail={"error": "No db connected"}) + if data.update_fields.role == "admin" and not premium_user: + raise HTTPException( + status_code=400, + detail="Assigning team admins is a premium feature. You must be a LiteLLM Enterprise user to use this feature. If you have a license please set `LITELLM_LICENSE` in your env. Get a 7 day trial key here: https://www.litellm.ai/#trial. Pricing: https://www.litellm.ai/#pricing", + ) + + _validate_budget_duration(data.update_fields.budget_duration) + existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id}) if existing_team_row is None: raise HTTPException( @@ -3059,19 +3084,29 @@ async def bulk_update_team_members( }, ) + returned_team_info: TeamInfoResponseObject = await team_info( + http_request=http_request, + team_id=data.team_id, + key_limit=None, + user_api_key_dict=user_api_key_dict, + ) + update_fields = data.update_fields.model_dump(exclude_unset=True) successful_updates: list[TeamMemberUpdateResponse] = [] failed_updates: list[FailedTeamMemberUpdate] = [] for user_id in user_ids: try: - response = await team_member_update( + response = await _apply_team_member_update( data=TeamMemberUpdateRequest(team_id=data.team_id, user_id=user_id, **update_fields), - http_request=http_request, + returned_team_info=returned_team_info, + prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, ) successful_updates.append(response) except HTTPException as exc: - failed_updates.append(FailedTeamMemberUpdate(user_id=user_id, failed_reason=str(exc.detail))) + detail = exc.detail + failed_reason = detail.get("error", str(detail)) if isinstance(detail, dict) else str(detail) + failed_updates.append(FailedTeamMemberUpdate(user_id=user_id, failed_reason=failed_reason)) except Exception as exc: verbose_proxy_logger.exception("Failed to bulk update team member %s in team %s", user_id, data.team_id) failed_updates.append(FailedTeamMemberUpdate(user_id=user_id, failed_reason=str(exc))) diff --git a/tests/test_litellm/proxy/test_team_member_update.py b/tests/test_litellm/proxy/test_team_member_update.py index bc6e5e69553..ab2702da926 100644 --- a/tests/test_litellm/proxy/test_team_member_update.py +++ b/tests/test_litellm/proxy/test_team_member_update.py @@ -184,13 +184,18 @@ async def test_bulk_team_member_update_applies_patch_and_returns_member_failures prisma_client = MagicMock() prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr( + team_endpoints, + "team_info", + AsyncMock(return_value={"team_info": team_row, "team_memberships": []}), + ) update_mock = AsyncMock( side_effect=[ team_endpoints.TeamMemberUpdateResponse(team_id="team-1234", user_id="user-1", tpm_limit=42), HTTPException(status_code=404, detail={"error": "User is not a team member"}), ] ) - monkeypatch.setattr(team_endpoints, "team_member_update", update_mock) + monkeypatch.setattr(team_endpoints, "_apply_team_member_update", update_mock) response = await bulk_update_team_members( data=BulkTeamMemberUpdateRequest( @@ -205,7 +210,9 @@ async def test_bulk_team_member_update_applies_patch_and_returns_member_failures assert response.total_requested == 2 assert [member.user_id for member in response.successful_updates] == ["user-1"] assert response.failed_updates[0].user_id == "user-2" - assert "not a team member" in response.failed_updates[0].failed_reason + # a dict HTTPException detail must surface the nested error string, not a + # python dict repr like "{'error': 'User is not a team member'}" + assert response.failed_updates[0].failed_reason == "User is not a team member" assert update_mock.await_args_list[0].kwargs["data"].model_dump(exclude_unset=True) == { "team_id": "team-1234", "user_id": "user-1", @@ -223,7 +230,12 @@ async def test_bulk_team_member_update_returns_unexpected_member_failure(monkeyp prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr( - team_endpoints, "team_member_update", AsyncMock(side_effect=RuntimeError("database unavailable")) + team_endpoints, + "team_info", + AsyncMock(return_value={"team_info": team_row, "team_memberships": []}), + ) + monkeypatch.setattr( + team_endpoints, "_apply_team_member_update", AsyncMock(side_effect=RuntimeError("database unavailable")) ) response = await bulk_update_team_members( @@ -242,6 +254,63 @@ async def test_bulk_team_member_update_returns_unexpected_member_failure(monkeyp ] +@pytest.mark.asyncio +async def test_bulk_team_member_update_resolves_team_info_once(monkeypatch): + """The whole batch must resolve team_info a single time and still upsert every + member; resolving it per member re-scans the team, its keys, and all + memberships on each iteration, which times out large teams.""" + team_row = LiteLLM_TeamTable( + team_id="team-1234", + members_with_roles=[ + Member(user_id="user-1", role="user"), + Member(user_id="user-2", role="user"), + Member(user_id="user-3", role="user"), + ], + metadata={}, + ) + prisma_client = MagicMock() + prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + + class _FakeTx: + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + prisma_client.db.tx = MagicMock(return_value=_FakeTx()) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "premium_user", False) + + team_info_mock = AsyncMock( + return_value={ + "team_info": team_row, + "team_memberships": [ + types.SimpleNamespace(user_id="user-1", budget_id="bud-1"), + types.SimpleNamespace(user_id="user-2", budget_id="bud-2"), + types.SimpleNamespace(user_id="user-3", budget_id="bud-3"), + ], + } + ) + monkeypatch.setattr(team_endpoints, "team_info", team_info_mock) + upsert_mock = AsyncMock() + monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock) + + response = await bulk_update_team_members( + data=BulkTeamMemberUpdateRequest( + team_id="team-1234", + all_members_in_team=True, + update_fields=TeamMemberBulkUpdateFields(tpm_limit=42), + ), + http_request=Request({"type": "http", "method": "POST", "path": "/team/member/bulk_update"}), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin"), + ) + + assert team_info_mock.await_count == 1 + assert upsert_mock.await_count == 3 + assert [member.user_id for member in response.successful_updates] == ["user-1", "user-2", "user-3"] + + def test_bulk_team_member_update_requires_exactly_one_member_selector(): with pytest.raises(ValueError, match="either user_ids or all_members_in_team"): BulkTeamMemberUpdateRequest(