From 2a13bbe1cb502ba472ad41fecac95363b077c68f Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 31 Jul 2026 22:13:57 -0700 Subject: [PATCH] refactor(proxy): resolve team member lookups in one query and cap the rejection message Resolve the requested member user_ids with a single find_many instead of one lookup per member, so a large member list no longer turns into that many round-trips before the permission check runs. Write the member-add audit entries concurrently rather than one after another, and list at most a few ids in the rejection message instead of echoing the whole request back. Update the team-admin member-add case that covered adding a user_id with no user row, which the endpoint now leaves to proxy admins. --- .../management_endpoints/team_endpoints.py | 49 +++++++++++----- tests/proxy_unit_tests/test_proxy_server.py | 2 +- .../test_team_endpoints.py | 57 ++++++++++++++++--- 3 files changed, 87 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 7b5ce33cb8b..a8c5fcea4df 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -2431,12 +2431,23 @@ async def _resolve_existing_member_user_ids( members: Sequence[Member], prisma_client: PrismaClient, ) -> frozenset[str]: - """Return the caller-supplied user_ids that already have a user row.""" - user_repository = UserRepository(prisma_client) - found = await asyncio.gather( - *(user_repository.find_by_id(member.user_id) for member in members if member.user_id is not None) + """Return the caller-supplied user_ids that already have a user row. + + Resolved with a single query so the number of members in the request does + not translate into that many concurrent connections. + """ + requested_user_ids = frozenset(member.user_id for member in members if member.user_id is not None) + if not requested_user_ids: + return frozenset() + + found = await UserRepository(prisma_client).table.find_many( + where={ # mutable-ok: Prisma query filters are dict-shaped + "user_id": { # mutable-ok: Prisma query filters are dict-shaped + "in": sorted(requested_user_ids) + } + } ) - return frozenset(user.user_id for user in found if user is not None and user.user_id is not None) + return frozenset(user.user_id for user in found or () if user.user_id is not None) def _pre_existing_user_ids( @@ -2459,6 +2470,9 @@ def _pre_existing_user_ids( return existing_user_ids | populated_user_ids +_MAX_REPORTED_UNKNOWN_USER_IDS = 10 + + def _validate_member_user_id_provisioning( members: Sequence[Member], existing_user_ids: frozenset[str], @@ -2481,13 +2495,15 @@ def _validate_member_user_id_provisioning( if not unknown_user_ids: return + listed = ", ".join(unknown_user_ids[:_MAX_REPORTED_UNKNOWN_USER_IDS]) + remaining = len(unknown_user_ids) - _MAX_REPORTED_UNKNOWN_USER_IDS raise HTTPException( status_code=403, detail={ # mutable-ok: HTTPException detail must be a plain mapping to keep this route's {"error": ...} response shape "error": ( - "Only proxy admins can add a user_id that does not exist yet: {}. " + "Only proxy admins can add a user_id that does not exist yet: {}{}. " "Add the member by user_email to invite a new user, or ask a proxy admin " - "to create the user first.".format(", ".join(unknown_user_ids)) + "to create the user first.".format(listed, " and {} more".format(remaining) if remaining > 0 else "") ) }, ) @@ -2515,13 +2531,15 @@ async def _create_team_member_add_audit_logs( user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, ) -> None: - """Record the membership change, and any user row it created, in the audit log.""" + """Record the membership change, and any user row it created, in the audit log. + + The entries are written concurrently so a request adding many members does + not pay for them one after another. + """ from litellm.proxy.management_helpers.audit_logs import create_object_audit_log - for user in updated_users: - if user.user_id is None or user.user_id in existing_user_ids: - continue - await create_object_audit_log( + created_user_entries = tuple( + create_object_audit_log( object_id=user.user_id, action="created", litellm_changed_by=None, @@ -2531,8 +2549,11 @@ async def _create_team_member_add_audit_logs( before_value=None, after_value=safe_dumps(user.model_dump(exclude_none=True)), ) + for user in updated_users + if user.user_id is not None and user.user_id not in existing_user_ids + ) - await create_object_audit_log( + membership_entry = create_object_audit_log( object_id=team_id, action="updated", litellm_changed_by=None, @@ -2543,6 +2564,8 @@ async def _create_team_member_add_audit_logs( after_value=_members_audit_value(after_members), ) + await asyncio.gather(*created_user_entries, membership_entry) + async def _validate_and_populate_member_user_info( member: Member, diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index bedd4dd1838..f64994cb3b1 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -1475,7 +1475,7 @@ async def test_create_team_member_add_team_admin( user_api_key_dict=valid_token, ) except HTTPException as e: - if user_role == "user": + if user_role == "user" or new_member_method == "user_id": assert e.status_code == 403 return else: 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 b1438447e4f..0de2c1ac71d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -10440,20 +10440,18 @@ def test_validate_member_user_id_provisioning_reports_every_unknown_member(): @pytest.mark.asyncio async def test_resolve_existing_member_user_ids_matches_caller_supplied_user_ids(): - """Only caller-supplied user_ids are looked up; unknown ones resolve to nothing.""" + """Caller-supplied user_ids resolve in one query; unknown ones resolve to nothing.""" from litellm.proxy.management_endpoints.team_endpoints import ( _resolve_existing_member_user_ids, ) prisma_client = MagicMock() - - async def find_by_id(user_id): - if user_id == "by-id": - return LiteLLM_UserTable(user_id="by-id", max_budget=None, spend=0.0, user_email=None, models=[]) - return None + find_many = AsyncMock( + return_value=[LiteLLM_UserTable(user_id="by-id", max_budget=None, spend=0.0, user_email=None, models=[])] + ) with patch("litellm.proxy.management_endpoints.team_endpoints.UserRepository") as repo: - repo.return_value.find_by_id = AsyncMock(side_effect=find_by_id) + repo.return_value.table.find_many = find_many resolved = await _resolve_existing_member_user_ids( members=[ @@ -10465,6 +10463,28 @@ async def test_resolve_existing_member_user_ids_matches_caller_supplied_user_ids ) assert resolved == frozenset({"by-id"}) + # one round-trip, and email-only members contribute no id to look up + find_many.assert_awaited_once() + assert find_many.await_args.kwargs["where"] == {"user_id": {"in": ["by-id", "missing"]}} + + +@pytest.mark.asyncio +async def test_resolve_existing_member_user_ids_skips_the_query_when_no_user_ids(): + """An all-email payload must not hit the database at all.""" + from litellm.proxy.management_endpoints.team_endpoints import ( + _resolve_existing_member_user_ids, + ) + + with patch("litellm.proxy.management_endpoints.team_endpoints.UserRepository") as repo: + repo.return_value.table.find_many = AsyncMock() + + resolved = await _resolve_existing_member_user_ids( + members=[Member(user_email="a@example.com", role="user")], + prisma_client=MagicMock(), + ) + + assert resolved == frozenset() + repo.return_value.table.find_many.assert_not_awaited() def test_pre_existing_user_ids_counts_ids_filled_in_by_member_resolution(): @@ -10574,3 +10594,26 @@ async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeyp mock_audit.assert_called_once() assert created_user_id not in mock_audit.call_args.kwargs["existing_user_ids"] + + +def test_validate_member_user_id_provisioning_caps_the_ids_it_echoes_back(): + """A large member list must not echo every id back in the error body.""" + from litellm.proxy.management_endpoints.team_endpoints import ( + _MAX_REPORTED_UNKNOWN_USER_IDS, + _validate_member_user_id_provisioning, + ) + + members = [Member(user_id=f"u{i}", role="user") for i in range(500)] + + with pytest.raises(HTTPException) as exc_info: + _validate_member_user_id_provisioning( + members=members, + existing_user_ids=frozenset(), + user_api_key_dict=_provisioning_caller(LitellmUserRoles.INTERNAL_USER), + ) + + detail = str(exc_info.value.detail) + assert "u0" in detail + assert f"u{_MAX_REPORTED_UNKNOWN_USER_IDS}" not in detail + assert f"and {500 - _MAX_REPORTED_UNKNOWN_USER_IDS} more" in detail + assert len(detail) < 1000