mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): align team member add with existing user provisioning rules
Adding a team member by a user_id with no user row created that row as a side effect for any caller permitted to add members, while creating users directly is restricted to proxy admins. Restrict that path to proxy admins too; adding an existing user, and inviting a new one by user_email (where the user_id is allocated server-side), are unchanged. Also record the membership change, and any user row it creates, in the audit log, matching /team/update, /user/new and /key/*.
This commit is contained in:
parent
f2cfa86713
commit
e8e2e07ef6
2 changed files with 383 additions and 0 deletions
|
|
@ -13,6 +13,7 @@ import asyncio
|
|||
import json
|
||||
import math
|
||||
import traceback
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any, Dict, List, Mapping, Optional, Tuple, Union, cast
|
||||
|
||||
|
|
@ -2426,6 +2427,123 @@ def _emit_team_members_metric(team: LiteLLM_TeamTable) -> None:
|
|||
verbose_proxy_logger.debug("Prometheus: failed to emit team members metric: %s", str(e))
|
||||
|
||||
|
||||
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 frozenset(user.user_id for user in found if user is not None and user.user_id is not None)
|
||||
|
||||
|
||||
def _pre_existing_user_ids(
|
||||
members: Sequence[Member],
|
||||
caller_supplied_user_ids: frozenset[str],
|
||||
existing_user_ids: frozenset[str],
|
||||
) -> frozenset[str]:
|
||||
"""Return the user_ids that already had a user row before this request.
|
||||
|
||||
Combines the caller-supplied ids that resolved to a user with the ids
|
||||
``_validate_and_populate_member_user_info`` filled in, which it only does
|
||||
from a matched user row. Deriving it that way keeps this in step with the
|
||||
email matching that resolution performs, rather than repeating it here.
|
||||
"""
|
||||
populated_user_ids = frozenset(
|
||||
member.user_id
|
||||
for member in members
|
||||
if member.user_id is not None and member.user_id not in caller_supplied_user_ids
|
||||
)
|
||||
return existing_user_ids | populated_user_ids
|
||||
|
||||
|
||||
def _validate_member_user_id_provisioning(
|
||||
members: Sequence[Member],
|
||||
existing_user_ids: frozenset[str],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
"""Restrict adding a caller-chosen user_id that has no user row yet to proxy admins.
|
||||
|
||||
Team and org admins keep the ability to add users that already exist and to
|
||||
invite new ones by user_email, where the user_id is allocated server-side.
|
||||
"""
|
||||
if user_api_key_dict.user_role in (
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN.value,
|
||||
):
|
||||
return
|
||||
|
||||
unknown_user_ids = tuple(
|
||||
member.user_id for member in members if member.user_id is not None and member.user_id not in existing_user_ids
|
||||
)
|
||||
if not unknown_user_ids:
|
||||
return
|
||||
|
||||
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: {}. "
|
||||
"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))
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _members_audit_value(members: Sequence[Member]) -> str:
|
||||
"""Serialize a team's member list for an audit-log value.
|
||||
|
||||
The audit-log columns hold a JSON object, so the member list is nested
|
||||
under a key rather than serialized as a top-level array.
|
||||
"""
|
||||
return safe_dumps(
|
||||
{ # mutable-ok: the audit-log JSON column rejects a top-level array, so this value must be an object
|
||||
"members_with_roles": tuple(member.model_dump() for member in members)
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def _create_team_member_add_audit_logs(
|
||||
team_id: str,
|
||||
updated_users: Sequence[LiteLLM_UserTable],
|
||||
existing_user_ids: frozenset[str],
|
||||
before_members: Sequence[Member],
|
||||
after_members: Sequence[Member],
|
||||
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."""
|
||||
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(
|
||||
object_id=user.user_id,
|
||||
action="created",
|
||||
litellm_changed_by=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
table_name=LitellmTableNames.USER_TABLE_NAME,
|
||||
before_value=None,
|
||||
after_value=safe_dumps(user.model_dump(exclude_none=True)),
|
||||
)
|
||||
|
||||
await create_object_audit_log(
|
||||
object_id=team_id,
|
||||
action="updated",
|
||||
litellm_changed_by=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
table_name=LitellmTableNames.TEAM_TABLE_NAME,
|
||||
before_value=_members_audit_value(before_members),
|
||||
after_value=_members_audit_value(after_members),
|
||||
)
|
||||
|
||||
|
||||
async def _validate_and_populate_member_user_info(
|
||||
member: Member,
|
||||
prisma_client: PrismaClient,
|
||||
|
|
@ -2606,6 +2724,19 @@ async def team_member_add(
|
|||
data=data,
|
||||
)
|
||||
|
||||
requested_members = tuple(data.member) if isinstance(data.member, list) else (data.member,)
|
||||
caller_supplied_user_ids = frozenset(member.user_id for member in requested_members if member.user_id is not None)
|
||||
existing_user_ids = await _resolve_existing_member_user_ids(
|
||||
members=requested_members,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
_validate_member_user_id_provisioning(
|
||||
members=requested_members,
|
||||
existing_user_ids=existing_user_ids,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
members_before_add = tuple(complete_team_data.members_with_roles)
|
||||
|
||||
# Validate and populate user_email/user_id for members before processing
|
||||
if isinstance(data.member, Member):
|
||||
await _validate_and_populate_member_user_info(
|
||||
|
|
@ -2619,6 +2750,12 @@ async def team_member_add(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
pre_existing_user_ids = _pre_existing_user_ids(
|
||||
members=requested_members,
|
||||
caller_supplied_user_ids=caller_supplied_user_ids,
|
||||
existing_user_ids=existing_user_ids,
|
||||
)
|
||||
|
||||
(
|
||||
updated_team,
|
||||
updated_users,
|
||||
|
|
@ -2637,6 +2774,16 @@ async def team_member_add(
|
|||
|
||||
_emit_team_members_metric(complete_team_data)
|
||||
|
||||
await _create_team_member_add_audit_logs(
|
||||
team_id=data.team_id,
|
||||
updated_users=updated_users,
|
||||
existing_user_ids=pre_existing_user_ids,
|
||||
before_members=members_before_add,
|
||||
after_members=tuple(complete_team_data.members_with_roles),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
)
|
||||
|
||||
return TeamAddMemberResponse.model_validate(
|
||||
{
|
||||
**updated_team.model_dump(),
|
||||
|
|
|
|||
|
|
@ -10338,3 +10338,239 @@ async def test_list_available_teams_filters_joined_and_validates_rows(monkeypatc
|
|||
assert result[0].team_alias == "open team"
|
||||
find_many_kwargs = mock_prisma_client.db.litellm_teamtable.find_many.call_args.kwargs
|
||||
assert find_many_kwargs["where"] == {"team_id": {"in": ["team-open"]}}
|
||||
|
||||
|
||||
def _provisioning_caller(role: LitellmUserRoles) -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth(user_id="caller-1", user_role=role)
|
||||
|
||||
|
||||
def test_validate_member_user_id_provisioning_allows_proxy_admin():
|
||||
"""Proxy admins may add a user_id that has no user row yet."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
_validate_member_user_id_provisioning,
|
||||
)
|
||||
|
||||
_validate_member_user_id_provisioning(
|
||||
members=[Member(user_id="brand-new", role="user")],
|
||||
existing_user_ids=frozenset(),
|
||||
user_api_key_dict=_provisioning_caller(LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
|
||||
def test_validate_member_user_id_provisioning_rejects_unknown_user_id_for_non_proxy_admin():
|
||||
"""A non-proxy-admin cannot add a user_id that has no user row yet."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
_validate_member_user_id_provisioning,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_validate_member_user_id_provisioning(
|
||||
members=[Member(user_id="brand-new", role="user")],
|
||||
existing_user_ids=frozenset(),
|
||||
user_api_key_dict=_provisioning_caller(LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "brand-new" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
def test_validate_member_user_id_provisioning_allows_existing_user_id_for_non_proxy_admin():
|
||||
"""A non-proxy-admin may still add a user that already exists."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
_validate_member_user_id_provisioning,
|
||||
)
|
||||
|
||||
_validate_member_user_id_provisioning(
|
||||
members=[Member(user_id="already-here", role="user")],
|
||||
existing_user_ids=frozenset({"already-here"}),
|
||||
user_api_key_dict=_provisioning_caller(LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
|
||||
|
||||
def test_validate_member_user_id_provisioning_allows_email_only_member_for_non_proxy_admin():
|
||||
"""Inviting by user_email stays open to non-proxy-admins; the user_id is server-allocated."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
_validate_member_user_id_provisioning,
|
||||
)
|
||||
|
||||
_validate_member_user_id_provisioning(
|
||||
members=[Member(user_email="invitee@example.com", role="user")],
|
||||
existing_user_ids=frozenset(),
|
||||
user_api_key_dict=_provisioning_caller(LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
|
||||
|
||||
def test_validate_member_user_id_provisioning_rejects_unknown_user_id_paired_with_email():
|
||||
"""Supplying a user_email alongside an unknown user_id does not lift the restriction."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
_validate_member_user_id_provisioning,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_validate_member_user_id_provisioning(
|
||||
members=[Member(user_id="chosen-id", user_email="invitee@example.com", role="user")],
|
||||
existing_user_ids=frozenset(),
|
||||
user_api_key_dict=_provisioning_caller(LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
def test_validate_member_user_id_provisioning_reports_every_unknown_member():
|
||||
"""A bulk add names each unknown user_id rather than only the first."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
_validate_member_user_id_provisioning,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_validate_member_user_id_provisioning(
|
||||
members=[
|
||||
Member(user_id="known", role="user"),
|
||||
Member(user_id="unknown-a", role="user"),
|
||||
Member(user_id="unknown-b", role="user"),
|
||||
],
|
||||
existing_user_ids=frozenset({"known"}),
|
||||
user_api_key_dict=_provisioning_caller(LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
|
||||
detail = str(exc_info.value.detail)
|
||||
assert "unknown-a" in detail
|
||||
assert "unknown-b" in detail
|
||||
|
||||
|
||||
@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."""
|
||||
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
|
||||
|
||||
with patch("litellm.proxy.management_endpoints.team_endpoints.UserRepository") as repo:
|
||||
repo.return_value.find_by_id = AsyncMock(side_effect=find_by_id)
|
||||
|
||||
resolved = await _resolve_existing_member_user_ids(
|
||||
members=[
|
||||
Member(user_id="by-id", role="user"),
|
||||
Member(user_id="missing", role="user"),
|
||||
Member(user_email="someone@example.com", role="user"),
|
||||
],
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
assert resolved == frozenset({"by-id"})
|
||||
|
||||
|
||||
def test_pre_existing_user_ids_counts_ids_filled_in_by_member_resolution():
|
||||
"""An id the member-resolution step filled in came from a matched row, so it pre-existed.
|
||||
|
||||
This is what keeps a case-variant email invite of an existing user from being
|
||||
recorded as a newly created user.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import _pre_existing_user_ids
|
||||
|
||||
# member arrived email-only; resolution matched an existing row and filled in the id
|
||||
resolved_member = Member(user_id="matched-existing", user_email="Someone@Example.com", role="user")
|
||||
|
||||
assert _pre_existing_user_ids(
|
||||
members=[resolved_member],
|
||||
caller_supplied_user_ids=frozenset(),
|
||||
existing_user_ids=frozenset(),
|
||||
) == frozenset({"matched-existing"})
|
||||
|
||||
|
||||
def test_pre_existing_user_ids_excludes_caller_supplied_ids_that_do_not_exist():
|
||||
"""A caller-supplied id that resolved to nothing is genuinely new, so it stays out."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import _pre_existing_user_ids
|
||||
|
||||
assert _pre_existing_user_ids(
|
||||
members=[Member(user_id="brand-new", role="user"), Member(user_id="already-here", role="user")],
|
||||
caller_supplied_user_ids=frozenset({"brand-new", "already-here"}),
|
||||
existing_user_ids=frozenset({"already-here"}),
|
||||
) == frozenset({"already-here"})
|
||||
|
||||
|
||||
def test_members_audit_value_serializes_to_a_json_object():
|
||||
"""The audit-log columns hold a JSON object; a top-level array is rejected by the DB."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import _members_audit_value
|
||||
|
||||
payload = json.loads(_members_audit_value([Member(user_id="u1", role="admin"), Member(user_id="u2", role="user")]))
|
||||
|
||||
assert isinstance(payload, dict)
|
||||
assert [m["user_id"] for m in payload["members_with_roles"]] == ["u1", "u2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeypatch):
|
||||
"""A user created by a list payload must still be reported as newly created.
|
||||
|
||||
For a list payload the member-list reconciliation back-fills the caller's own
|
||||
Member objects with the ids of users this request just created. The set of
|
||||
pre-existing ids therefore has to be captured before that runs, otherwise a
|
||||
freshly created user looks like it was already there and no creation is recorded.
|
||||
"""
|
||||
from litellm.proxy._types import TeamMemberAddRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_add
|
||||
|
||||
team_id = "team-list-audit"
|
||||
created_user_id = "generated-uuid-for-new-invitee"
|
||||
member = Member(user_email="invitee@example.com", role="user")
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id")
|
||||
|
||||
team_row = LiteLLM_TeamTable(team_id=team_id, members_with_roles=[])
|
||||
created_user = LiteLLM_UserTable(
|
||||
user_id=created_user_id, user_email="invitee@example.com", max_budget=None, spend=0.0, models=[]
|
||||
)
|
||||
updated_team = MagicMock()
|
||||
updated_team.model_dump.return_value = {"team_id": team_id, "members_with_roles": []}
|
||||
|
||||
async def fake_add_team_members_to_team(**kwargs):
|
||||
# mirrors _update_team_members_list: the list branch mutates the caller's Member in place
|
||||
member.user_id = created_user_id
|
||||
return updated_team, [created_user], []
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=team_row,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._validate_team_member_add_permissions",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._validate_and_populate_member_user_info",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._resolve_existing_member_user_ids",
|
||||
new_callable=AsyncMock,
|
||||
return_value=frozenset(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team",
|
||||
side_effect=fake_add_team_members_to_team,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints._create_team_member_add_audit_logs",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_audit,
|
||||
):
|
||||
await team_member_add(
|
||||
data=TeamMemberAddRequest(team_id=team_id, member=[member]),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1"),
|
||||
)
|
||||
|
||||
mock_audit.assert_called_once()
|
||||
assert created_user_id not in mock_audit.call_args.kwargs["existing_user_ids"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue