mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
9349b22c64
commit
7ed91df836
3 changed files with 162 additions and 45 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue