diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 9ff00922de4..dbc709a1742 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1784,7 +1784,10 @@ async def new_team( ) if is_audit_logging_enabled(): - _updated_values = complete_team_data.json(exclude_none=True) + created_team_snapshot: Final = complete_team_data.model_copy( + update={"members_with_roles": list(team_row.members_with_roles)} + ) + _updated_values = created_team_snapshot.json(exclude_none=True) _updated_values = json.dumps(_updated_values, default=str) @@ -3050,6 +3053,38 @@ async def _add_team_members_to_team( return updated_team, updated_users, updated_team_memberships +async def _update_team_member_role( + tx: "Prisma", + prisma_client: PrismaClient, + team_id: str, + user_id: str, + role: Literal["admin", "user"], + user_email: str | None, +) -> tuple[tuple[Member, ...], tuple[Member, ...]]: + """Rewrite one member's role from the roster read under the team lock; returns (before, after).""" + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id) + + locked_members: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, team_id) + if locked_members is None: + raise HTTPException(status_code=404, detail={"error": f"Team id={team_id} does not exist in db"}) + + before: Final = tuple(locked_members) + if all(member.user_id != user_id for member in before): + raise HTTPException(status_code=404, detail={"error": f"User {user_id} is not a member of team {team_id}"}) + + after: Final = tuple( + Member(user_id=member.user_id, role=role, user_email=user_email or member.user_email) + if member.user_id == user_id + else member + for member in before + ) + await _team_tx_db(tx).update( + where={"team_id": team_id}, + data={"members_with_roles": json.dumps([m.model_dump() for m in after])}, + ) + return before, after + + def _emit_team_members_metric(team: LiteLLM_TeamTable) -> None: """Update the Prometheus team members gauge after a membership change. @@ -3147,7 +3182,7 @@ def _validate_member_user_id_provisioning( ) -def _members_audit_value(members: Sequence[Member]) -> str: +def _members_audit_value(team_alias: str | None, 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 @@ -3155,13 +3190,45 @@ def _members_audit_value(members: Sequence[Member]) -> str: """ 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) + "team_alias": team_alias, + "members_with_roles": tuple(member.model_dump() for member in members), } ) -async def _create_team_member_add_audit_logs( +def _schedule_team_membership_audit_log( team_id: str, + team_alias: str | None, + before_members: Sequence[Member], + after_members: Sequence[Member], + user_api_key_dict: UserAPIKeyAuth, + litellm_proxy_admin_name: str, +) -> None: + from litellm.proxy.management_helpers.audit_logs import ( + create_object_audit_log, + is_audit_logging_enabled, + ) + + if not is_audit_logging_enabled() or tuple(before_members) == tuple(after_members): + return + + asyncio.create_task( + 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(team_alias, before_members), + after_value=_members_audit_value(team_alias, after_members), + ) + ) + + +def _schedule_team_member_add_audit_logs( + team_id: str, + team_alias: str | None, updated_users: Sequence[LiteLLM_UserTable], existing_user_ids: frozenset[str], before_members: Sequence[Member], @@ -3169,41 +3236,40 @@ 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. - - 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 - - created_user_entries: Final = tuple( - 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)), - ) - for user in updated_users - if user.user_id is not None and user.user_id not in existing_user_ids + """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, + is_audit_logging_enabled, ) - membership_entry: Final = create_object_audit_log( - object_id=team_id, - action="updated", - litellm_changed_by=None, + if not is_audit_logging_enabled(): + return + + for user in updated_users: + if user.user_id in existing_user_ids: + continue + asyncio.create_task( + 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)), + ) + ) + + _schedule_team_membership_audit_log( + team_id=team_id, + team_alias=team_alias, + before_members=before_members, + after_members=after_members, 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), ) - await asyncio.gather(*created_user_entries, membership_entry) - async def _validate_and_populate_member_user_info( member: Member, @@ -3442,8 +3508,9 @@ async def team_member_add( _emit_team_members_metric(complete_team_data) - await _create_team_member_add_audit_logs( + _schedule_team_member_add_audit_logs( team_id=data.team_id, + team_alias=complete_team_data.team_alias, updated_users=updated_users, existing_user_ids=pre_existing_user_ids, before_members=members_before_add, @@ -3513,7 +3580,33 @@ async def team_member_delete( }' ``` """ - from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache + from litellm.proxy.proxy_server import litellm_proxy_admin_name + + existing_team_row, before_members, after_members = await _team_member_delete( + data=data, user_api_key_dict=user_api_key_dict + ) + + _schedule_team_membership_audit_log( + team_id=existing_team_row.team_id, + team_alias=existing_team_row.team_alias, + before_members=before_members, + after_members=after_members, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + ) + + return existing_team_row + + +async def _team_member_delete( + data: TeamMemberDeleteRequest, + user_api_key_dict: UserAPIKeyAuth, +) -> tuple[LiteLLM_TeamTable, tuple[Member, ...], tuple[Member, ...]]: + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) if prisma_client is None: raise HTTPException(status_code=500, detail={"error": "No db connected"}) @@ -3677,7 +3770,7 @@ async def team_member_delete( _emit_team_members_metric(existing_team_row) - return existing_team_row + return existing_team_row, tuple(fresh_members), tuple(new_team_members) @router.post( @@ -3697,7 +3790,12 @@ async def team_member_update( Update team member budgets and team member role """ - from litellm.proxy.proxy_server import premium_user, prisma_client, user_api_key_cache + from litellm.proxy.proxy_server import ( + litellm_proxy_admin_name, + premium_user, + prisma_client, + user_api_key_cache, + ) if prisma_client is None: raise HTTPException(status_code=500, detail={"error": "No db connected"}) @@ -3787,6 +3885,18 @@ async def team_member_update( ### upsert new budget budget_patch: Final = member_budget_patch(data) async with prisma_client.tx() as tx: + role_change: Final = ( + await _update_team_member_role( + tx=tx, + prisma_client=prisma_client, + team_id=data.team_id, + user_id=received_user_id, + role=data.role, + user_email=data.user_email, + ) + if data.role is not None + else None + ) await _upsert_budget_and_membership( tx=tx, team_id=data.team_id, @@ -3803,27 +3913,16 @@ async def team_member_update( user_api_key_cache=user_api_key_cache, ) - ### update team member role - if data.role is not None: - team_members: Final[list[Member]] = [] - for member in team_table.members_with_roles: - if member.user_id == received_user_id: - team_members.append( - Member( - user_id=member.user_id, - role=data.role, - user_email=data.user_email or member.user_email, - ) - ) - else: - team_members.append(member) - - team_table.members_with_roles = team_members - - _db_team_members: Final[list[dict]] = [m.model_dump() for m in team_members] - await _team_db(prisma_client).update( - where={"team_id": data.team_id}, - data={"members_with_roles": json.dumps(_db_team_members)}, + if role_change is not None: + members_before_role_update, team_members = role_change + team_table.members_with_roles = list(team_members) + _schedule_team_membership_audit_log( + team_id=data.team_id, + team_alias=team_table.team_alias, + before_members=members_before_role_update, + after_members=team_members, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, ) return TeamMemberUpdateResponse( @@ -4303,7 +4402,7 @@ async def delete_team( tasks = [] for team_member in team_members: tasks.append( - team_member_delete( + _team_member_delete( data=TeamMemberDeleteRequest( team_id=team_row.team_id, user_id=team_member.user_id, 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 9fd388887f4..7cb62a8da11 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -3,6 +3,7 @@ import json from contextlib import asynccontextmanager, contextmanager from datetime import datetime, timezone from types import SimpleNamespace +from collections.abc import Sequence from typing import Final, Optional, cast from unittest.mock import AsyncMock, MagicMock, PropertyMock, call, patch @@ -12,6 +13,7 @@ from fastapi.testclient import TestClient from pydantic import ValidationError from litellm._uuid import uuid +from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import ( LiteLLM_BudgetTable, LiteLLM_BudgetTableFull, @@ -23,11 +25,14 @@ from litellm.proxy._types import ( LiteLLM_TeamTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, + LitellmTableNames, LitellmUserRoles, Member, ProxyErrorTypes, ProxyException, ResetSpendRequest, + TeamInfoMember, + TeamInfoResponseObjectTeamTable, TeamMemberAddRequest, TeamMemberUpdateRequest, UpdateTeamRequest, @@ -68,6 +73,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import ( BulkTeamMemberAddResponse, TeamMemberAddResult, ) +from litellm.types.utils import StandardAuditLogPayload from tests.test_litellm.proxy.management_endpoints.jwt_key_mapping_doubles import ( CascadingJWTMappingTable, JWTMappingRow, @@ -143,11 +149,11 @@ def _wire_member_add_tx(prisma_client): def _wire_member_delete_tx(prisma_client): - """/team/member_delete's four cleanups, plus the advisory-lock re-read that now guards - them, run inside one transaction, so a mocked client has to hand back its own table - mocks (and a `query_raw` that answers the locked re-read from the same team row the - test already configured on `find_unique`) out of `tx()` for the existing per-table - assertions to keep seeing the calls.""" + """/team/member_delete's four cleanups and /team/member_update's role rewrite, plus the + advisory-lock re-read that guards them, run inside one transaction, so a mocked client + has to hand back its own table mocks (and a `query_raw` that answers the locked re-read + from the same team row the test already configured on `find_unique`) out of `tx()` for + the existing per-table assertions to keep seeing the calls.""" async def _query_raw(sql, team_id): if sql != TEAM_ADVISORY_LOCK_SQL: @@ -163,10 +169,12 @@ def _wire_member_delete_tx(prisma_client): return getattr(prisma_client.db, table_name) tx = _Tx() + tx.query_raw = AsyncMock(side_effect=_query_raw) tx_cm = MagicMock() tx_cm.__aenter__ = AsyncMock(return_value=tx) tx_cm.__aexit__ = AsyncMock(return_value=None) prisma_client.tx = MagicMock(return_value=tx_cm) + return tx def _wire_team_delete_tx(prisma_client): @@ -8696,8 +8704,8 @@ async def test_delete_team_persists_deleted_teams( "admin", ) monkeypatch.setattr( - "litellm.proxy.management_endpoints.team_endpoints.team_member_delete", - AsyncMock(return_value=team1), + "litellm.proxy.management_endpoints.team_endpoints._team_member_delete", + AsyncMock(return_value=(team1, (), ())), ) data = DeleteTeamRequest(team_ids=["team-1"]) @@ -13240,9 +13248,12 @@ 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")])) + payload = json.loads( + _members_audit_value("my-team", [Member(user_id="u1", role="admin"), Member(user_id="u2", role="user")]) + ) assert isinstance(payload, dict) + assert payload["team_alias"] == "my-team" assert [m["user_id"] for m in payload["members_with_roles"]] == ["u1", "u2"] @@ -13267,7 +13278,7 @@ async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeyp 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=[]) + team_row = LiteLLM_TeamTable(team_id=team_id, team_alias="list-audit", members_with_roles=[]) created_user = LiteLLM_UserTable( user_id=created_user_id, user_email="invitee@example.com", max_budget=None, spend=0.0, models=[] ) @@ -13303,8 +13314,7 @@ async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeyp side_effect=fake_add_team_members_to_team, ), patch( - "litellm.proxy.management_endpoints.team_endpoints._create_team_member_add_audit_logs", - new_callable=AsyncMock, + "litellm.proxy.management_endpoints.team_endpoints._schedule_team_member_add_audit_logs", ) as mock_audit, ): await team_member_add( @@ -13314,6 +13324,603 @@ 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"] + assert mock_audit.call_args.kwargs["team_alias"] == "list-audit" + + +class _RecordingAuditLogger(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: list[StandardAuditLogPayload] = [] + + async def async_log_audit_log_event(self, audit_log_payload: StandardAuditLogPayload) -> None: + self.payloads.append(audit_log_payload) + + +def _wire_audit_log_callback(monkeypatch: pytest.MonkeyPatch) -> _RecordingAuditLogger: + audit_logger = _RecordingAuditLogger() + monkeypatch.setattr("litellm.store_audit_logs", True) + monkeypatch.setattr("litellm.audit_log_callbacks", [audit_logger]) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + return audit_logger + + +async def _settle_audit_log_tasks() -> None: + for _ in range(5): + await asyncio.sleep(0) + + +def _team_roster_events(audit_logger: _RecordingAuditLogger, action: str) -> list[StandardAuditLogPayload]: + return [ + p + for p in audit_logger.payloads + if p["table_name"] == LitellmTableNames.TEAM_TABLE_NAME and p["action"] == action + ] + + +def _roster_user_roles(members_json: str | None) -> dict[str, str]: + assert members_json is not None + return {m["user_id"]: m["role"] for m in json.loads(members_json)["members_with_roles"]} + + +def _roster_team_alias(members_json: str | None) -> str | None: + assert members_json is not None + return json.loads(members_json)["team_alias"] + + +@pytest.mark.asyncio +async def test_new_team_created_audit_event_carries_the_final_roster(monkeypatch): + from fastapi import Request + + from litellm.proxy._types import NewTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import new_team + + audit_logger = _wire_audit_log_callback(monkeypatch) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) + mock_prisma.jsonify_team_object = lambda db_data: db_data + mock_prisma.get_data = AsyncMock(return_value=None) + mock_prisma.update_data = AsyncMock() + created_team = MagicMock() + created_team.team_id = "team-audit-roster" + created_team.members_with_roles = [] + created_team.metadata = None + created_team.default_team_member_models = None + created_team.model_dump.return_value = {"team_id": "team-audit-roster", "members_with_roles": []} + mock_prisma.db.litellm_teamtable.create = AsyncMock(return_value=created_team) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=created_team) + mock_prisma.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model-1")) + user_row = MagicMock() + user_row.user_id = "alice" + user_row.model_dump.return_value = {"user_id": "alice", "teams": ["team-audit-roster"]} + mock_prisma.db.litellm_usertable.upsert = AsyncMock(return_value=user_row) + mock_prisma.db.litellm_usertable.update_many = AsyncMock() + mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_usertable.update = AsyncMock(return_value=user_row) + membership_row = MagicMock() + membership_row.model_dump.return_value = {"team_id": "team-audit-roster", "user_id": "alice", "budget_id": None} + mock_prisma.db.litellm_teammembership.upsert = AsyncMock(return_value=membership_row) + mock_prisma.db.litellm_auditlog.create = AsyncMock() + _wire_team_create_tx(mock_prisma) + + mock_license = MagicMock() + mock_license.is_team_count_over_limit.return_value = False + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr("litellm.proxy.proxy_server._license_check", mock_license) + monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") + + await new_team( + data=NewTeamRequest( + team_id="team-audit-roster", + team_alias="audit-roster", + members_with_roles=[Member(user_id="alice", role="admin"), Member(user_id="bob", role="user")], + ), + http_request=MagicMock(spec=Request), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1", api_key="sk-a"), + ) + await _settle_audit_log_tasks() + + created_events = _team_roster_events(audit_logger, "created") + assert [e["object_id"] for e in created_events] == ["team-audit-roster"] + assert _roster_user_roles(created_events[0]["updated_values"]) == { + "admin-1": "admin", + "alice": "admin", + "bob": "user", + } + + +@pytest.mark.asyncio +async def test_team_member_delete_emits_a_roster_audit_event(monkeypatch, mock_db_client, mock_admin_auth): + from litellm.proxy._types import TeamMemberDeleteRequest + + audit_logger = _wire_audit_log_callback(monkeypatch) + + team_row = MagicMock() + team_row.model_dump.return_value = { + "team_id": "team-del-audit", + "team_alias": "del-audit", + "members_with_roles": [ + {"user_id": "alice", "user_email": None, "role": "admin"}, + {"user_id": "bob", "user_email": None, "role": "user"}, + ], + "team_member_permissions": [], + "metadata": {}, + "models": [], + "spend": 0.0, + } + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=team_row) + user_row = MagicMock() + user_row.user_id = "bob" + user_row.teams = ["team-del-audit"] + mock_db_client.db.litellm_usertable.find_many = AsyncMock(return_value=[user_row]) + mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock()) + mock_db_client.db.litellm_teammembership = MagicMock() + mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(return_value=MagicMock()) + mock_db_client.db.litellm_verificationtoken = MagicMock() + mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(return_value=MagicMock()) + _wire_member_delete_tx(mock_db_client) + + await team_member_delete( + data=TeamMemberDeleteRequest(team_id="team-del-audit", user_id="bob"), + user_api_key_dict=mock_admin_auth, + ) + await _settle_audit_log_tasks() + + updated_events = _team_roster_events(audit_logger, "updated") + assert [e["object_id"] for e in updated_events] == ["team-del-audit"] + assert _roster_user_roles(updated_events[0]["before_value"]) == {"alice": "admin", "bob": "user"} + assert _roster_user_roles(updated_events[0]["updated_values"]) == {"alice": "admin"} + assert _roster_team_alias(updated_events[0]["before_value"]) == "del-audit" + assert _roster_team_alias(updated_events[0]["updated_values"]) == "del-audit" + + stale_user_row = MagicMock() + stale_user_row.user_id = "carol" + stale_user_row.teams = ["team-del-audit"] + mock_db_client.db.litellm_usertable.find_many = AsyncMock(return_value=[stale_user_row]) + + await team_member_delete( + data=TeamMemberDeleteRequest(team_id="team-del-audit", user_id="carol"), + user_api_key_dict=mock_admin_auth, + ) + await _settle_audit_log_tasks() + + assert len(_team_roster_events(audit_logger, "updated")) == 1, ( + "scrubbing a stale team reference off a user row leaves the roster as it was, so no roster event" + ) + + +@pytest.mark.asyncio +async def test_team_member_update_role_change_emits_a_roster_audit_event(monkeypatch): + audit_logger = _wire_audit_log_callback(monkeypatch) + + mock_prisma_client = MagicMock() + team_row = LiteLLM_TeamTable( + team_id="team-role-audit", + team_alias="role-audit", + metadata={}, + members_with_roles=[Member(user_id="alice", role="admin"), Member(user_id="bob", role="user")], + ) + + def _team_info_as_read_from_db(bob_role: str): + return { + "team_info": TeamInfoResponseObjectTeamTable( + team_id="team-role-audit", + team_alias="role-audit", + metadata={}, + members_with_roles=( + TeamInfoMember(user_id="alice", role="admin", user_alias="Alice"), + TeamInfoMember(user_id="bob", role=bob_role, user_alias="Bob"), + ), + ), + "team_memberships": [LiteLLM_TeamMembership(user_id="bob", team_id="team-role-audit", budget_id=None)], + } + + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(side_effect=_roster_writer(team_row)) + mock_prisma_client.db.litellm_auditlog.create = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + _wire_member_delete_tx(mock_prisma_client) + + with ( + patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.team_info", + AsyncMock(side_effect=[_team_info_as_read_from_db("user"), _team_info_as_read_from_db("admin")]), + ), + patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + AsyncMock(), + ), + ): + await team_member_update( + data=TeamMemberUpdateRequest(team_id="team-role-audit", user_id="bob", role="admin"), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + await _settle_audit_log_tasks() + + updated_events = _team_roster_events(audit_logger, "updated") + assert [e["object_id"] for e in updated_events] == ["team-role-audit"] + assert _roster_user_roles(updated_events[0]["before_value"]) == {"alice": "admin", "bob": "user"} + assert _roster_user_roles(updated_events[0]["updated_values"]) == {"alice": "admin", "bob": "admin"} + assert _roster_team_alias(updated_events[0]["updated_values"]) == "role-audit" + + await team_member_update( + data=TeamMemberUpdateRequest(team_id="team-role-audit", user_id="bob", role="admin"), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + await _settle_audit_log_tasks() + + assert len(_team_roster_events(audit_logger, "updated")) == 1, ( + "re-sending the role a member already holds leaves the roster as it was, so no roster event" + ) + + +def _roster_writer(team_row: LiteLLM_TeamTable): + """An `update` side effect that lands `members_with_roles` on the team row later reads see.""" + + async def _update(where, data): + team_row.members_with_roles = [Member(**m) for m in json.loads(data["members_with_roles"])] + return team_row + + return _update + + +def _member_update_patches(team_snapshot: LiteLLM_TeamTable): + return ( + patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints.team_info", + AsyncMock( + return_value={ + "team_info": TeamInfoResponseObjectTeamTable(**team_snapshot.model_dump()), + "team_memberships": [ + LiteLLM_TeamMembership(user_id="bob", team_id=team_snapshot.team_id, budget_id=None) + ], + } + ), + ), + patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests + "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + AsyncMock(), + ), + ) + + +@pytest.mark.asyncio +async def test_team_member_update_role_change_rewrites_the_roster_it_read_under_the_lock(monkeypatch): + """Regression: a member added between /team/member_update's permission checks and its write + was dropped, because the new roster was built from the pre-check snapshot.""" + audit_logger = _wire_audit_log_callback(monkeypatch) + + stale_snapshot = LiteLLM_TeamTable( + team_id="team-race", + team_alias="race", + metadata={}, + members_with_roles=[Member(user_id="alice", role="admin"), Member(user_id="bob", role="user")], + ) + team_row = LiteLLM_TeamTable( + **{ + **stale_snapshot.model_dump(), + "members_with_roles": [*stale_snapshot.members_with_roles, Member(user_id="carol", role="user")], + } + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(side_effect=_roster_writer(team_row)) + mock_prisma_client.db.litellm_auditlog.create = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + tx = _wire_member_delete_tx(mock_prisma_client) + + team_info_patch, upsert_patch = _member_update_patches(stale_snapshot) + with team_info_patch, upsert_patch: + response = await team_member_update( + data=TeamMemberUpdateRequest(team_id="team-race", user_id="bob", role="admin"), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + await _settle_audit_log_tasks() + + assert {m.user_id: m.role for m in team_row.members_with_roles} == { + "alice": "admin", + "bob": "admin", + "carol": "user", + } + assert response.team_id == "team-race" and response.user_id == "bob" + assert tx.query_raw.await_args_list[0].args == (TEAM_ADVISORY_LOCK_SQL, "team-race"), ( + "the roster must be read only after the team advisory lock is held" + ) + + [event] = _team_roster_events(audit_logger, "updated") + assert _roster_user_roles(event["before_value"]) == {"alice": "admin", "bob": "user", "carol": "user"} + assert _roster_user_roles(event["updated_values"]) == {"alice": "admin", "bob": "admin", "carol": "user"} + assert _roster_team_alias(event["updated_values"]) == "race" + + +@pytest.mark.asyncio +async def test_team_member_update_role_change_404s_when_the_team_is_gone_under_the_lock(monkeypatch): + _wire_audit_log_callback(monkeypatch) + snapshot = LiteLLM_TeamTable( + team_id="team-gone-race", + team_alias="gone-race", + metadata={}, + members_with_roles=[Member(user_id="bob", role="user")], + ) + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=[snapshot, None]) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + _wire_member_delete_tx(mock_prisma_client) + + team_info_patch, upsert_patch = _member_update_patches(snapshot) + with team_info_patch, upsert_patch, pytest.raises(HTTPException) as exc_info: + await team_member_update( + data=TeamMemberUpdateRequest(team_id="team-gone-race", user_id="bob", role="admin"), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + + assert exc_info.value.status_code == 404 + mock_prisma_client.db.litellm_teamtable.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_team_member_update_role_change_404s_when_the_member_left_before_the_locked_read(monkeypatch): + """Regression: a member removed between the pre-lock read and the locked read was reported as updated.""" + audit_logger = _wire_audit_log_callback(monkeypatch) + snapshot = LiteLLM_TeamTable( + team_id="team-member-gone-race", + team_alias="member-gone-race", + metadata={}, + members_with_roles=[Member(user_id="alice", role="admin"), Member(user_id="bob", role="user")], + ) + locked_row = LiteLLM_TeamTable( + team_id="team-member-gone-race", + team_alias="member-gone-race", + metadata={}, + members_with_roles=[Member(user_id="alice", role="admin")], + ) + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=[snapshot, locked_row]) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + _wire_member_delete_tx(mock_prisma_client) + + team_info_patch, upsert_patch = _member_update_patches(snapshot) + with team_info_patch, upsert_patch as upsert_budget, pytest.raises(HTTPException) as exc_info: + await team_member_update( + data=TeamMemberUpdateRequest( + team_id="team-member-gone-race", user_id="bob", role="admin", max_budget_in_team=5.0 + ), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ), + ) + + assert exc_info.value.status_code == 404 + assert "bob" in str(exc_info.value.detail) + mock_prisma_client.db.litellm_teamtable.update.assert_not_awaited() + upsert_budget.assert_not_awaited() + await _settle_audit_log_tasks() + assert audit_logger.payloads == [] + + +@pytest.mark.asyncio +async def test_team_member_delete_response_does_not_wait_for_the_audit_insert( + monkeypatch, mock_db_client, mock_admin_auth +): + """Regression: the roster audit row was awaited on the request path, so a slow audit table + held every /team/member_delete response.""" + from litellm.proxy._types import TeamMemberDeleteRequest + + audit_logger = _wire_audit_log_callback(monkeypatch) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + + team_row = LiteLLM_TeamTable( + team_id="team-slow-audit", + team_alias="slow-audit", + metadata={}, + members_with_roles=[Member(user_id="alice", role="admin"), Member(user_id="bob", role="user")], + ) + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=team_row) + mock_db_client.db.litellm_usertable.find_many = AsyncMock(return_value=[]) + mock_db_client.db.litellm_teammembership.delete_many = AsyncMock(return_value=MagicMock()) + mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock(return_value=MagicMock()) + _wire_member_delete_tx(mock_db_client) + + audit_table_answers = asyncio.Event() + audit_rows = [] + + async def _blocked_create(data): + await audit_table_answers.wait() + audit_rows.append(data) + + mock_db_client.db.litellm_auditlog.create = AsyncMock(side_effect=_blocked_create) + + await asyncio.wait_for( + team_member_delete( + data=TeamMemberDeleteRequest(team_id="team-slow-audit", user_id="bob"), + user_api_key_dict=mock_admin_auth, + ), + timeout=1, + ) + assert audit_rows == [], "the response returned while the audit table was still blocked" + + audit_table_answers.set() + await _settle_audit_log_tasks() + + assert [row["object_id"] for row in audit_rows] == ["team-slow-audit"] + [event] = _team_roster_events(audit_logger, "updated") + assert _roster_user_roles(event["before_value"]) == {"alice": "admin", "bob": "user"} + assert _roster_user_roles(event["updated_values"]) == {"alice": "admin"} + + +class _UntouchableRoster(Sequence[Member]): + """A roster that fails the test the moment anything reads it.""" + + def __getitem__(self, index): + raise AssertionError("the roster was read while audit logging is off") + + def __len__(self) -> int: + raise AssertionError("the roster was read while audit logging is off") + + +class _UntouchableUsers(Sequence[LiteLLM_UserTable]): + def __getitem__(self, index): + raise AssertionError("the created users were read while audit logging is off") + + def __len__(self) -> int: + raise AssertionError("the created users were read while audit logging is off") + + +def _schedule_membership_audit_with_untouchable_roster() -> None: + from litellm.proxy.management_endpoints.team_endpoints import _schedule_team_membership_audit_log + + _schedule_team_membership_audit_log( + team_id="team-quiet", + team_alias="quiet", + before_members=_UntouchableRoster(), + after_members=_UntouchableRoster(), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user"), + litellm_proxy_admin_name="admin", + ) + + +def _schedule_member_add_audit_with_untouchable_roster() -> None: + from litellm.proxy.management_endpoints.team_endpoints import _schedule_team_member_add_audit_logs + + _schedule_team_member_add_audit_logs( + team_id="team-quiet", + team_alias="quiet", + updated_users=_UntouchableUsers(), + existing_user_ids=frozenset(), + before_members=_UntouchableRoster(), + after_members=_UntouchableRoster(), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user"), + litellm_proxy_admin_name="admin", + ) + + +@pytest.mark.parametrize( + "schedule", + [_schedule_membership_audit_with_untouchable_roster, _schedule_member_add_audit_with_untouchable_roster], +) +def test_membership_audit_scheduling_skips_the_roster_entirely_when_audit_logging_is_off(monkeypatch, schedule): + """Regression: the before/after rosters were serialized on every membership change, even + when audit logs are not stored.""" + monkeypatch.setattr("litellm.store_audit_logs", False) + + schedule() + + +@pytest.mark.asyncio +async def test_member_add_audit_reports_only_the_users_it_created_plus_the_roster_change(monkeypatch): + from litellm.proxy.management_endpoints.team_endpoints import _schedule_team_member_add_audit_logs + + audit_logger = _wire_audit_log_callback(monkeypatch) + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_auditlog.create = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + before = (Member(user_id="alice", role="admin"),) + after = (*before, Member(user_id="bob", role="user"), Member(user_id="carol", role="user")) + + _schedule_team_member_add_audit_logs( + team_id="team-add-audit", + team_alias="add-audit", + updated_users=[ + LiteLLM_UserTable(user_id="bob", user_email="bob@example.com", teams=["team-add-audit"]), + LiteLLM_UserTable(user_id="carol", teams=["team-add-audit"]), + ], + existing_user_ids=frozenset({"bob"}), + before_members=before, + after_members=after, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-1"), + litellm_proxy_admin_name="admin", + ) + await _settle_audit_log_tasks() + + created_users = [ + p + for p in audit_logger.payloads + if p["table_name"] == LitellmTableNames.USER_TABLE_NAME and p["action"] == "created" + ] + assert [p["object_id"] for p in created_users] == ["carol"], "only the user this request created is audited" + assert json.loads(created_users[0]["updated_values"])["teams"] == ["team-add-audit"] + + [roster_event] = _team_roster_events(audit_logger, "updated") + assert roster_event["object_id"] == "team-add-audit" + assert _roster_user_roles(roster_event["before_value"]) == {"alice": "admin"} + assert _roster_user_roles(roster_event["updated_values"]) == {"alice": "admin", "bob": "user", "carol": "user"} + assert _roster_team_alias(roster_event["updated_values"]) == "add-audit" + + +@pytest.mark.asyncio +async def test_delete_team_emits_only_the_deleted_audit_event(monkeypatch): + from litellm.proxy._types import DeleteTeamRequest + + audit_logger = _wire_audit_log_callback(monkeypatch) + + members = (Member(user_id="alice", role="admin"), Member(user_id="bob", role="user")) + team = LiteLLM_TeamTable( + team_id="team-gone", + team_alias="gone", + members_with_roles=list(members), + metadata={}, + model_max_budget={}, + model_spend={}, + ) + mock_prisma = AsyncMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) + mock_prisma.get_data = AsyncMock( + return_value=SimpleNamespace(json=lambda **_kwargs: team.model_dump_json(exclude_none=True)) + ) + mock_prisma.delete_data = AsyncMock(return_value={"deleted_keys": 0}) + mock_prisma.db.litellm_deletedteamtable.create_many = AsyncMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_auditlog.create = AsyncMock() + mock_tx = AsyncMock() + mock_tx.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_tx_cm = MagicMock() + mock_tx_cm.__aenter__ = AsyncMock(return_value=mock_tx) + mock_tx_cm.__aexit__ = AsyncMock(return_value=False) + mock_prisma.db.tx = MagicMock(return_value=mock_tx_cm) + _wire_team_delete_tx(mock_prisma) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") + + removals = [(team, members, members[1:]), (team, members[1:], ())] + monkeypatch.setattr( + "litellm.proxy.management_endpoints.team_endpoints._team_member_delete", + AsyncMock(side_effect=lambda **_kwargs: removals.pop(0)), + ) + + await delete_team( + data=DeleteTeamRequest(team_ids=["team-gone"]), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1", api_key="sk-a"), + litellm_changed_by=None, + ) + await _settle_audit_log_tasks() + + team_events = [ + (p["object_id"], p["action"]) for p in audit_logger.payloads if p["table_name"] == "LiteLLM_TeamTable" + ] + assert team_events == [("team-gone", "deleted")] @pytest.mark.asyncio @@ -13374,8 +13981,7 @@ async def test_team_member_add_evicts_the_new_members_cached_user_row_on_every_w side_effect=fake_add_team_members_to_team, ), patch( # test-quality-ok: team_member_add has no injection seam for its prisma-backed helpers - "litellm.proxy.management_endpoints.team_endpoints._create_team_member_add_audit_logs", - new_callable=AsyncMock, + "litellm.proxy.management_endpoints.team_endpoints._schedule_team_member_add_audit_logs", ), ): await team_member_add( diff --git a/tests/test_litellm/proxy/test_team_member_update.py b/tests/test_litellm/proxy/test_team_member_update.py index 352c68d491c..ace4c4e65af 100644 --- a/tests/test_litellm/proxy/test_team_member_update.py +++ b/tests/test_litellm/proxy/test_team_member_update.py @@ -14,7 +14,10 @@ from litellm.proxy._types import ( TeamMemberUpdateRequest, UserAPIKeyAuth, ) -from litellm.proxy.management_endpoints.team_endpoints import team_member_update +from litellm.proxy.management_endpoints.team_endpoints import ( + TEAM_ADVISORY_LOCK_SQL, + team_member_update, +) @pytest.mark.asyncio @@ -65,13 +68,20 @@ def happy_path_upsert(monkeypatch): prisma_client.db.litellm_teamtable.update = AsyncMock() class _FakeTx: + litellm_teamtable = prisma_client.db.litellm_teamtable + async def __aenter__(self): return self async def __aexit__(self, *args): return False - prisma_client.db.tx = MagicMock(return_value=_FakeTx()) + async def query_raw(self, sql, team_id): + if sql == TEAM_ADVISORY_LOCK_SQL: + return [] + return [{"members_with_roles": team_row.model_dump()["members_with_roles"]}] + + prisma_client.tx = MagicMock(return_value=_FakeTx()) monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "premium_user", False)