mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
e8e2e07ef6
commit
2a13bbe1cb
3 changed files with 87 additions and 21 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue