diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 28c12173ea7..8e5d62976fa 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) @@ -3160,6 +3163,27 @@ def _members_audit_value(members: Sequence[Member]) -> str: ) +async def _create_team_membership_audit_log( + team_id: str, + 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 + + 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 _create_team_member_add_audit_logs( team_id: str, updated_users: Sequence[LiteLLM_UserTable], @@ -3191,15 +3215,12 @@ async def _create_team_member_add_audit_logs( if user.user_id is not None and user.user_id not in existing_user_ids ) - membership_entry: Final = create_object_audit_log( - object_id=team_id, - action="updated", - litellm_changed_by=None, + membership_entry: Final = _create_team_membership_audit_log( + team_id=team_id, + 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) @@ -3508,7 +3529,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 + ) + + if before_members != after_members: + await _create_team_membership_audit_log( + team_id=existing_team_row.team_id, + 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"}) @@ -3672,7 +3719,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( @@ -3692,7 +3739,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"}) @@ -3800,8 +3852,12 @@ async def team_member_update( ### update team member role if data.role is not None: + members_before_role_update: Final = tuple( + Member(user_id=member.user_id, user_email=member.user_email, role=member.role) + for member in team_table.members_with_roles + ) team_members: Final[list[Member]] = [] - for member in team_table.members_with_roles: + for member in members_before_role_update: if member.user_id == received_user_id: team_members.append( Member( @@ -3820,6 +3876,14 @@ async def team_member_update( where={"team_id": data.team_id}, data={"members_with_roles": json.dumps(_db_team_members)}, ) + if members_before_role_update != tuple(team_members): + await _create_team_membership_audit_log( + team_id=data.team_id, + 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( team_id=data.team_id, @@ -4298,7 +4362,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 690b5ae80b6..72c854f1d6c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -12,6 +12,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 +24,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 +72,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, @@ -8696,8 +8701,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"]) @@ -13316,6 +13321,301 @@ async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeyp assert created_user_id not in mock_audit.call_args.kwargs["existing_user_ids"] +class _RecordingAuditLogger(CustomLogger): + """An audit_log_callbacks sink that keeps every payload it is handed.""" + + 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: + """Turn audit logging on and register one recording callback, the way an operator's + `litellm_settings.audit_log_callbacks` entry would be.""" + 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: + """Audit callbacks run on `asyncio.create_task`, so give the loop a few turns.""" + 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"]} + + +@pytest.mark.asyncio +async def test_new_team_created_audit_event_carries_the_final_roster(monkeypatch): + """The `created` event a `/team/new` hands to audit_log_callbacks must list the members + the team was created with. The team row is inserted empty and the members attached + afterwards, so a snapshot taken from the pre-insert object reports no members and a + downstream consumer syncing membership from the event has nothing to sync.""" + 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.create = 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): + """Removing a member must reach audit_log_callbacks as a TEAM_TABLE `updated` event whose + before and after rosters differ by exactly the removed user, the same shape + `/team/member_add` already emits, so one consumer can diff both directions.""" + 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", + "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"} + + 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): + """Changing a member's role must reach audit_log_callbacks as a TEAM_TABLE `updated` + event whose before roster carries the old role and whose after roster carries the new one.""" + audit_logger = _wire_audit_log_callback(monkeypatch) + + mock_prisma_client = MagicMock() + team_row = LiteLLM_TeamTable( + team_id="team-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", + 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(return_value=team_row) + mock_prisma_client.db.litellm_auditlog.create = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_tx = AsyncMock() + mock_prisma_client.tx.return_value.__aenter__ = AsyncMock(return_value=mock_tx) + mock_prisma_client.tx.return_value.__aexit__ = AsyncMock(return_value=None) + + 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"} + + 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" + ) + + +@pytest.mark.asyncio +async def test_delete_team_emits_only_the_deleted_audit_event(monkeypatch): + """`/team/delete` removes every member on its way out through the same code path + `/team/member_delete` uses. Those removals must not surface as TEAM_TABLE `updated` + roster events trailing the `deleted` one: the team is gone, and the `deleted` event + already carries the roster it went out with.""" + 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")] + + 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 (