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.
This commit is contained in:
Yuneng Jiang 2026-07-31 22:13:57 -07:00
parent e8e2e07ef6
commit 2a13bbe1cb
No known key found for this signature in database
3 changed files with 87 additions and 21 deletions

View file

@ -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,

View file

@ -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:

View file

@ -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