fix(proxy): make /team/member_delete's four cleanups atomic (#37959)

The team roster update, the user.teams update, the team membership
delete, and the team-scoped verification token delete ran as four
sequential writes with no transaction around them, so a failure
between any two left the removal half applied. Thread a single
prisma transaction through all four writes, following the same
tx.<table> pattern /team/member_add and /team/member_update already
use, so either all four land or none do.
This commit is contained in:
Yassin Kortam 2026-08-22 14:24:57 -07:00 • committed by GitHub
parent 9349b22c64
commit 7ed91df836
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 162 additions and 45 deletions

View file

@ -20,7 +20,7 @@ import secrets
import traceback
from collections.abc import Awaitable, Callable, Mapping, Sequence
from datetime import datetime, timedelta, timezone
from typing import Any, Final, Literal, Optional, Protocol, TypeVar, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeVar, cast
import fastapi
import yaml
@ -148,6 +148,9 @@ from litellm.types.utils import (
TeamUIKeyGenerationConfig,
)
if TYPE_CHECKING:
from prisma import Prisma
_PrismaRowT = TypeVar("_PrismaRowT")
_RepositoryModelT = TypeVar("_RepositoryModelT", bound=BaseModel)
@ -4337,10 +4340,19 @@ def _transform_verification_tokens_to_deleted_records(
async def _save_deleted_verification_token_records(
records: Sequence[Mapping[str, object]],
prisma_client: PrismaClient,
tx: "Prisma | None" = None,
) -> None:
"""Save deleted verification token records to the database."""
"""Save deleted verification token records to the database.
``tx`` runs the write on that transaction's connection instead of a fresh
one, so a caller batching this with other writes gets one all-or-nothing
commit.
"""
if not records:
return
if tx is not None:
await tx.litellm_deletedverificationtoken.create_many(data=records)
return
await _deleted_verification_token_table(prisma_client).create_many(data=records)
@ -4349,6 +4361,7 @@ async def _persist_deleted_verification_tokens(
prisma_client: PrismaClient,
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: str | None = None,
tx: "Prisma | None" = None,
) -> None:
"""Persist deleted verification token records by transforming and saving them."""
records: Final = _transform_verification_tokens_to_deleted_records(
@ -4359,6 +4372,7 @@ async def _persist_deleted_verification_tokens(
await _save_deleted_verification_token_records(
records=records,
prisma_client=prisma_client,
tx=tx,
)

View file

@ -3286,15 +3286,6 @@ async def team_member_delete(
_db_new_team_members: Final[list[dict]] = [m.model_dump() for m in new_team_members]
_ = await _team_db(prisma_client).update(
where={
"team_id": data.team_id,
},
data={"members_with_roles": json.dumps(_db_new_team_members)},
)
_emit_team_members_metric(existing_team_row)
## DELETE TEAM ID from USER ROW, IF EXISTS ##
# get user row
removed_user_ids: Final = frozenset(m.user_id for m in removed_team_members if m.user_id is not None)
@ -3303,52 +3294,62 @@ async def team_member_delete(
)
existing_user_rows: Final[Sequence[LiteLLM_UserTable]] = await _user_db(prisma_client).find_many(where=key_val)
for existing_user in existing_user_rows:
if data.team_id in existing_user.teams:
await _user_db(prisma_client).update(
where={
"user_id": existing_user.user_id,
},
data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}},
)
# Also clean up any existing team membership rows for this user and team
user_ids_to_delete: Final = removed_user_ids.union(
(data.user_id,) if data.user_id is not None else (),
(user.user_id for user in existing_user_rows if user.user_id),
)
for _uid in sorted(user_ids_to_delete):
await _team_membership_db(prisma_client).delete_many(where={"team_id": data.team_id, "user_id": _uid})
## DELETE KEYS CREATED BY USER FOR THIS TEAM
if user_ids_to_delete:
from litellm.proxy.management_endpoints.key_management_endpoints import (
_persist_deleted_verification_tokens,
# Fetch keys before deletion so their audit records can be persisted alongside the delete.
# An empty user_ids_to_delete still resolves cleanly: prisma's "in": [] matches no rows.
keys_to_delete: Final[list[LiteLLM_VerificationToken]] = await _tokens_db(prisma_client).find_many(
where={
"user_id": {"in": sorted(user_ids_to_delete)},
"team_id": data.team_id,
}
)
# All four cleanups run on one connection so a failure between them leaves
# no partial removal: either every write below lands, or none of them do.
async with prisma_client.tx() as tx:
await tx.litellm_teamtable.update(
where={"team_id": data.team_id},
data={"members_with_roles": json.dumps(_db_new_team_members)},
)
# Fetch keys before deletion to persist them
keys_to_delete: Final[list[LiteLLM_VerificationToken]] = await _tokens_db(prisma_client).find_many(
where={
"user_id": {"in": sorted(user_ids_to_delete)},
"team_id": data.team_id,
}
)
for existing_user in existing_user_rows:
if data.team_id in existing_user.teams:
await tx.litellm_usertable.update(
where={"user_id": existing_user.user_id},
data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}},
)
if keys_to_delete:
await _persist_deleted_verification_tokens(
keys=keys_to_delete,
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
for _uid in sorted(user_ids_to_delete):
await tx.litellm_teammembership.delete_many(where={"team_id": data.team_id, "user_id": _uid})
if user_ids_to_delete:
if keys_to_delete:
from litellm.proxy.management_endpoints.key_management_endpoints import (
_persist_deleted_verification_tokens,
)
await _persist_deleted_verification_tokens(
keys=keys_to_delete,
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
tx=tx,
)
await tx.litellm_verificationtoken.delete_many(
where={
"user_id": {"in": sorted(user_ids_to_delete)},
"team_id": data.team_id,
}
)
await _tokens_db(prisma_client).delete_many(
where={
"user_id": {"in": sorted(user_ids_to_delete)},
"team_id": data.team_id,
}
)
_emit_team_members_metric(existing_team_row)
return existing_team_row

View file

@ -80,6 +80,23 @@ def _wire_team_create_tx(prisma_client):
prisma_client.db.tx = lambda *_args, **_kwargs: _tx()
def _wire_member_delete_tx(prisma_client):
"""/team/member_delete's four cleanups run inside one transaction, so a mocked
client has to hand back its own table mocks out of `tx()` for the existing
per-table assertions to keep seeing the calls."""
tx = SimpleNamespace(
litellm_teamtable=prisma_client.db.litellm_teamtable,
litellm_usertable=prisma_client.db.litellm_usertable,
litellm_teammembership=prisma_client.db.litellm_teammembership,
litellm_verificationtoken=prisma_client.db.litellm_verificationtoken,
litellm_deletedverificationtoken=prisma_client.db.litellm_deletedverificationtoken,
)
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)
# Mock prisma_client
mock_prisma_client = MagicMock()
# Set up async mock for db operations
@ -4147,6 +4164,8 @@ async def test_team_member_delete_cleans_membership(mock_db_client, mock_admin_a
return_value=MagicMock()
)
_wire_member_delete_tx(mock_db_client)
# Execute
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id),
@ -4205,6 +4224,8 @@ async def test_team_member_delete_cleans_verification_tokens(
return_value=MagicMock()
)
_wire_member_delete_tx(mock_db_client)
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id),
user_api_key_dict=mock_admin_auth,
@ -4300,6 +4321,8 @@ async def test_team_member_delete_by_email_the_user_row_does_not_carry(
return_value=MagicMock()
)
_wire_member_delete_tx(mock_db_client)
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=test_team_id, user_email=roster_email),
user_api_key_dict=mock_admin_auth,
@ -4318,6 +4341,83 @@ async def test_team_member_delete_by_email_the_user_row_does_not_carry(
)
class _InjectedMemberDeleteFailure(Exception):
pass
@pytest.mark.asyncio
async def test_team_member_delete_is_atomic_across_its_four_writes(
mock_db_client, mock_admin_auth
):
"""
/team/member_delete's four cleanups (team roster, user.teams, team
membership, verification tokens) run as one transaction, so a failure
partway through must not leave the removal half applied.
Failing the second write (the user's ``teams`` update) pins two things a
non-transactional implementation gets wrong: the roster write that already
ran has to land on the SAME transaction client the failure raises on (so a
real database rolls it back too), and the writes still queued behind the
failure (membership delete, token delete) must never be attempted at all.
"""
from litellm.proxy._types import TeamMemberDeleteRequest
from litellm.proxy.management_endpoints.team_endpoints import team_member_delete
test_team_id = "team-del-atomic-123"
test_user_id = "user-atomic@example.com"
mock_team_row = MagicMock()
mock_team_row.model_dump.return_value = {
"team_id": test_team_id,
"members_with_roles": [
{"user_id": test_user_id, "user_email": None, "role": "user"}
],
"team_member_permissions": [],
"metadata": {},
"models": [],
"spend": 0.0,
}
mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(
return_value=mock_team_row
)
mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row)
mock_user_row = MagicMock()
mock_user_row.user_id = test_user_id
mock_user_row.teams = [test_team_id]
mock_db_client.db.litellm_usertable.find_many = AsyncMock(
return_value=[mock_user_row]
)
mock_db_client.db.litellm_usertable.update = AsyncMock(
side_effect=_InjectedMemberDeleteFailure("boom between writes 1 and 2")
)
mock_db_client.db.litellm_teammembership = MagicMock()
mock_db_client.db.litellm_teammembership.delete_many = AsyncMock()
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()
_wire_member_delete_tx(mock_db_client)
with pytest.raises(_InjectedMemberDeleteFailure):
await team_member_delete(
data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id),
user_api_key_dict=mock_admin_auth,
)
# The roster write ran, but on the transaction the injected failure also raised on.
mock_db_client.db.litellm_teamtable.update.assert_awaited_once()
mock_db_client.tx.assert_called_once()
aexit_args = mock_db_client.tx.return_value.__aexit__.await_args.args
assert aexit_args[0] is _InjectedMemberDeleteFailure
# Writes queued behind the failure inside that same transaction never ran.
mock_db_client.db.litellm_teammembership.delete_many.assert_not_awaited()
mock_db_client.db.litellm_verificationtoken.delete_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_new_team_max_budget_exceeds_user_max_budget():
"""
@ -7806,6 +7906,8 @@ async def test_team_member_delete_persists_deleted_keys(monkeypatch):
mock_create_many_keys
)
_wire_member_delete_tx(mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
mock_prisma_client,