fix(proxy): run /user/bulk_delete team rewrites and user deletes in one transaction

Lock affected teams in sorted order inside a single 60s transaction so a
failure on any team rolls back every rewrite and every user row delete.
PrismaClient.tx() gains an optional timeout for the larger batch.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
ryan 2026-09-14 05:35:55 +00:00
parent 595aba3cb0
commit 82872c9627
3 changed files with 178 additions and 140 deletions

View file

@ -2,13 +2,16 @@
Each team a batch touches is rewritten exactly once, under the same advisory lock
`/team/member_delete` takes and from a roster re-read under that lock, so a concurrent
member_add on the team is never overwritten from a stale read.
member_add on the team is never overwritten from a stale read. A user batch runs in one
transaction, taking its team locks in sorted order, so either every team rewrite and every
user row delete lands or none of them does.
"""
import asyncio
import json
from collections.abc import Awaitable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import timedelta
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
@ -60,7 +63,8 @@ if TYPE_CHECKING:
from litellm.repositories.prisma_protocols import TableActions
_TEAM_WRITE_CONCURRENCY: Final = 10
_AUDIT_LOG_CONCURRENCY: Final = 10
_USER_BATCH_TX_TIMEOUT: Final = timedelta(seconds=60)
class _ErrorDetail(TypedDict):
@ -95,6 +99,12 @@ class _TeamRemoval:
deleted_key_tokens: tuple[str, ...]
@dataclass(frozen=True, slots=True)
class _UserBatchDeletion:
removals: Mapping[str, _TeamRemoval]
deleted_key_tokens: tuple[str, ...]
def _http_error(status_code: int, message: str) -> HTTPException:
detail: Final[_ErrorDetail] = {"error": message}
return HTTPException(status_code=status_code, detail=detail)
@ -161,7 +171,7 @@ def _error_message(exc: BaseException) -> str:
async def _bounded(awaitables: Iterable[Awaitable[object]]) -> tuple[object | BaseException, ...]:
semaphore: Final = asyncio.Semaphore(_TEAM_WRITE_CONCURRENCY)
semaphore: Final = asyncio.Semaphore(_AUDIT_LOG_CONCURRENCY)
async def run(awaitable: Awaitable[object]) -> object:
async with semaphore:
@ -172,54 +182,54 @@ async def _bounded(awaitables: Iterable[Awaitable[object]]) -> tuple[object | Ba
async def _remove_members_from_team(
prisma_client: PrismaClient,
tx: "Prisma",
team_id: str,
members: Sequence[MemberDeleteRequest],
user_api_key_dict: UserAPIKeyAuth,
) -> _TeamRemoval:
async with prisma_client.tx() as tx:
await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id)
roster: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, team_id)
if roster is None:
raise _http_error(400, f"Team id={team_id} does not exist in db")
await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id)
roster: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, team_id)
if roster is None:
raise _http_error(400, f"Team id={team_id} does not exist in db")
removed_members: Final = tuple(m for m in roster if any(_addresses_member(m, r) for r in members))
kept_members: Final = tuple(m for m in roster if not any(_addresses_member(m, r) for r in members))
removed_ids: Final = frozenset(m.user_id for m in removed_members if m.user_id is not None)
requested_ids: Final = frozenset(r.user_id for r in members if r.user_id is not None)
requested_emails: Final = frozenset(r.user_email for r in members if r.user_id is None and r.user_email)
user_rows: Final = await _user_tx_db(tx).find_many(
where=_any_filter(
_in_filter("user_id", removed_ids | requested_ids),
_in_filter("user_email", requested_emails),
)
removed_members: Final = tuple(m for m in roster if any(_addresses_member(m, r) for r in members))
kept_members: Final = tuple(m for m in roster if not any(_addresses_member(m, r) for r in members))
removed_ids: Final = frozenset(m.user_id for m in removed_members if m.user_id is not None)
requested_ids: Final = frozenset(r.user_id for r in members if r.user_id is not None)
requested_emails: Final = frozenset(r.user_email for r in members if r.user_id is None and r.user_email)
user_rows: Final = await _user_tx_db(tx).find_many(
where=_any_filter(
_in_filter("user_id", removed_ids | requested_ids),
_in_filter("user_email", requested_emails),
)
stale_rows: Final = tuple(u for u in user_rows if team_id in u.teams)
cleanup_ids: Final = removed_ids | frozenset(u.user_id for u in stale_rows)
matched: Final = frozenset(
i
for i, r in enumerate(members)
if any(_addresses_member(m, r) for m in removed_members) or any(_addresses_user(u, r) for u in stale_rows)
)
keys: Final = await _token_tx_db(tx).find_many(where=_team_users_filter(team_id, cleanup_ids))
)
stale_rows: Final = tuple(u for u in user_rows if team_id in u.teams)
cleanup_ids: Final = removed_ids | frozenset(u.user_id for u in stale_rows)
matched: Final = frozenset(
i
for i, r in enumerate(members)
if any(_addresses_member(m, r) for m in removed_members) or any(_addresses_user(u, r) for u in stale_rows)
)
keys: Final = await _token_tx_db(tx).find_many(where=_team_users_filter(team_id, cleanup_ids))
if removed_members:
roster_data: Final[_RosterData] = {
"members_with_roles": json.dumps(tuple(m.model_dump() for m in kept_members))
}
await _team_tx_db(tx).update(where=_eq_filter("team_id", team_id), data=roster_data)
for row in stale_rows:
teams_data: _TeamsData = {"teams": {"set": tuple(t for t in row.teams if t != team_id)}}
await _user_tx_db(tx).update(where=_eq_filter("user_id", row.user_id), data=teams_data)
await _membership_tx_db(tx).delete_many(where=_team_users_filter(team_id, cleanup_ids))
if keys:
await _persist_deleted_verification_tokens(
keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
tx=tx,
)
await _token_tx_db(tx).delete_many(where=_team_users_filter(team_id, cleanup_ids))
if removed_members:
roster_data: Final[_RosterData] = {
"members_with_roles": json.dumps(tuple(m.model_dump() for m in kept_members))
}
await _team_tx_db(tx).update(where=_eq_filter("team_id", team_id), data=roster_data)
for row in stale_rows:
teams_data: _TeamsData = {"teams": {"set": tuple(t for t in row.teams if t != team_id)}}
await _user_tx_db(tx).update(where=_eq_filter("user_id", row.user_id), data=teams_data)
await _membership_tx_db(tx).delete_many(where=_team_users_filter(team_id, cleanup_ids))
if keys:
await _persist_deleted_verification_tokens(
keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
tx=tx,
)
await _token_tx_db(tx).delete_many(where=_team_users_filter(team_id, cleanup_ids))
return _TeamRemoval(
team=LiteLLM_TeamTable(
@ -279,7 +289,8 @@ async def bulk_remove_team_members(
duplicates: Final = _duplicate_member_indexes(data.members)
kept_indexes: Final = tuple(i for i in range(len(data.members)) if i not in duplicates)
members: Final = tuple(data.members[i] for i in kept_indexes)
removal: Final = await _remove_members_from_team(prisma_client, data.team_id, members, user_api_key_dict)
async with prisma_client.tx() as tx:
removal: Final = await _remove_members_from_team(prisma_client, tx, data.team_id, members, user_api_key_dict)
await delete_cache_key_objects(
hashed_tokens=removal.deleted_key_tokens,
user_api_key_cache=user_api_key_cache,
@ -333,60 +344,101 @@ def _scope_error(user_id: str, target_org_ids: frozenset[str], caller_admin_org_
)
async def _delete_user_rows_tx(
async def _delete_user_rows(
prisma_client: PrismaClient,
tx: "Prisma",
user_ids: frozenset[str],
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: str | None,
) -> tuple[str, ...]:
async with prisma_client.tx() as tx:
keys: Final = await _token_tx_db(tx).find_many(where=_in_filter("user_id", user_ids))
if keys:
await _persist_deleted_verification_tokens(
keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
tx=tx,
)
await _token_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _invitation_tx_db(tx).delete_many(
where=_any_filter(
_in_filter("user_id", user_ids),
_in_filter("created_by", user_ids),
_in_filter("updated_by", user_ids),
)
keys: Final = await _token_tx_db(tx).find_many(where=_in_filter("user_id", user_ids))
if keys:
await _persist_deleted_verification_tokens(
keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
tx=tx,
)
await _org_membership_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _membership_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _user_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _token_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _invitation_tx_db(tx).delete_many(
where=_any_filter(
_in_filter("user_id", user_ids),
_in_filter("created_by", user_ids),
_in_filter("updated_by", user_ids),
)
)
await _org_membership_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _membership_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
await _user_tx_db(tx).delete_many(where=_in_filter("user_id", user_ids))
return tuple(k.token for k in keys)
async def _delete_user_rows(
async def _delete_users_tx(
prisma_client: PrismaClient,
users: Sequence["prisma_models.LiteLLM_UserTable"],
teams_of: Mapping[str, frozenset[str]],
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: str | None,
) -> _UserBatchDeletion:
"""Rewrites every team the users belong to and deletes their rows in one transaction, so a
failure anywhere rolls back the whole batch. Teams a user still names but which no longer exist
are skipped; the user row goes away regardless."""
async with prisma_client.tx(timeout=_USER_BATCH_TX_TIMEOUT) as tx:
team_rows: Final = await _team_tx_db(tx).find_many(
where=_in_filter("team_id", frozenset(t for teams in teams_of.values() for t in teams))
)
team_ids: Final = tuple(sorted(t.team_id for t in team_rows))
removals: Final = MappingProxyType(
{
tid: await _remove_members_from_team(
prisma_client,
tx,
tid,
tuple(
MemberDeleteRequest(user_id=u.user_id, user_email=u.user_email)
for u in users
if tid in teams_of[u.user_id]
),
user_api_key_dict,
)
for tid in team_ids
}
)
deleted_key_tokens: Final = await _delete_user_rows(
prisma_client, tx, frozenset(u.user_id for u in users), user_api_key_dict, litellm_changed_by
)
return _UserBatchDeletion(
removals=removals,
deleted_key_tokens=deleted_key_tokens + tuple(t for r in removals.values() for t in r.deleted_key_tokens),
)
async def _delete_users(
prisma_client: PrismaClient,
users: Sequence["prisma_models.LiteLLM_UserTable"],
teams_of: Mapping[str, frozenset[str]],
user_api_key_dict: UserAPIKeyAuth,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None,
litellm_proxy_admin_name: str | None,
litellm_changed_by: str | None,
) -> str | None:
) -> _UserBatchDeletion | str:
"""Returns the error message when the transaction rolled back, in which case no row was touched."""
user_ids: Final = frozenset(u.user_id for u in users)
try:
deleted_key_tokens: Final = await _delete_user_rows_tx(
prisma_client, user_ids, user_api_key_dict, litellm_changed_by
)
deletion: Final = await _delete_users_tx(prisma_client, users, teams_of, user_api_key_dict, litellm_changed_by)
except Exception as e: # noqa: BLE001 # the rolled-back batch is reported per row, not as a request failure
verbose_proxy_logger.error("/user/bulk_delete: failed to delete users %s: %s", sorted(user_ids), e)
return _error_message(e)
await delete_cache_key_objects(
hashed_tokens=deleted_key_tokens,
hashed_tokens=deletion.deleted_key_tokens,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await evict_and_broadcast(cache_keys=sorted(user_ids), user_api_key_cache=user_api_key_cache)
for removal in deletion.removals.values():
_emit_team_members_metric(removal.team)
audit_outcomes: Final = await _bounded(
UserManagementEventHooks.create_internal_user_audit_log(
user_id=u.user_id,
@ -401,7 +453,7 @@ async def _delete_user_rows(
for u, outcome in zip(users, audit_outcomes, strict=True):
if isinstance(outcome, BaseException):
verbose_proxy_logger.warning("Failed to create audit log for user %s: %s", u.user_id, outcome)
return None
return deletion
async def bulk_delete_users(
@ -452,54 +504,19 @@ async def bulk_delete_users(
for u in candidates
}
)
team_ids: Final = tuple(sorted(frozenset(t for u in candidates for t in teams_of[u.user_id])))
members_by_team: Final = MappingProxyType(
{
tid: tuple(
MemberDeleteRequest(user_id=u.user_id, user_email=u.user_email)
for u in candidates
if tid in teams_of[u.user_id]
)
for tid in team_ids
}
)
outcomes: Final = await _bounded(
_remove_members_from_team(prisma_client, tid, members_by_team[tid], user_api_key_dict) for tid in team_ids
)
removals: Final = MappingProxyType(
{tid: o for tid, o in zip(team_ids, outcomes, strict=True) if isinstance(o, _TeamRemoval)}
)
team_failures: Final = MappingProxyType(
{tid: _error_message(o) for tid, o in zip(team_ids, outcomes, strict=True) if isinstance(o, BaseException)}
)
for tid, err in team_failures.items():
verbose_proxy_logger.error("/user/bulk_delete: failed to remove users from team %s: %s", tid, err)
await delete_cache_key_objects(
hashed_tokens=tuple(t for removal in removals.values() for t in removal.deleted_key_tokens),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
for removal in removals.values():
_emit_team_members_metric(removal.team)
def team_errors(user_id: str) -> tuple[str, ...]:
return tuple(
f"Failed to remove from team {tid}: {err}" for tid, err in team_failures.items() if tid in teams_of[user_id]
)
deletable: Final = tuple(u for u in candidates if not team_errors(u.user_id))
delete_error: Final = (
await _delete_user_rows(
deletion: Final = (
await _delete_users(
prisma_client,
deletable,
candidates,
teams_of,
user_api_key_dict,
user_api_key_cache,
proxy_logging_obj,
litellm_proxy_admin_name,
litellm_changed_by,
)
if deletable
else None
if candidates
else _UserBatchDeletion(removals=MappingProxyType({}), deleted_key_tokens=())
)
def result(index: int, user_id: str) -> UserDeleteResult:
@ -508,15 +525,18 @@ async def bulk_delete_users(
error: Final = precheck_errors[user_id]
if error is not None:
return UserDeleteResult(user_id=user_id, success=False, error=error)
errors: Final = team_errors(user_id) or (
(f"Failed to delete user: {delete_error}",) if delete_error is not None else ()
)
if isinstance(deletion, str):
return UserDeleteResult(
user_id=user_id,
user_email=rows_by_id[user_id].user_email,
success=False,
error=f"Failed to delete user: {deletion}",
)
return UserDeleteResult(
user_id=user_id,
user_email=rows_by_id[user_id].user_email,
success=not errors,
teams_removed=tuple(tid for tid in team_ids if tid in removals and user_id in removals[tid].removed),
error="; ".join(errors) or None,
success=True,
teams_removed=tuple(tid for tid, r in deletion.removals.items() if user_id in r.removed),
)
results: Final = tuple(result(i, uid) for i, uid in enumerate(data.user_ids))

View file

@ -3792,6 +3792,7 @@ def jsonify_object(data: dict) -> dict:
# Bounded to prevent memory leaks from accumulated rotations.
_deprecated_key_cache: Final[LimitedSizeOrderedDict] = LimitedSizeOrderedDict(max_size=1000)
_DEPRECATED_KEY_CACHE_TTL_SECONDS: Final = 60
_PRISMA_DEFAULT_TX_TIMEOUT: Final = timedelta(seconds=5)
async def _lookup_deprecated_key(
@ -4170,13 +4171,13 @@ class PrismaClient:
return self.db.read_target
return self.db
def tx(self) -> "TransactionManager":
def tx(self, *, timeout: timedelta = _PRISMA_DEFAULT_TX_TIMEOUT) -> "TransactionManager":
"""Open an interactive transaction on the writer.
Callers go through this instead of reaching into ``self.db`` so writer
selection and read-replica routing stay encapsulated in the wrapper.
"""
return cast("TransactionManager", self.db.tx()) # cast-ok: wrappers delegate tx via __getattr__ (untyped)
return cast("TransactionManager", self.db.tx(timeout=timeout)) # cast-ok: untyped __getattr__ delegate
def get_request_status(self, payload: dict | SpendLogsPayload) -> Literal["success", "failure"]:
"""

View file

@ -95,6 +95,9 @@ class _TeamTable:
async def find_unique(self, where: Mapping[str, str]) -> LiteLLM_TeamTable | None:
return self.rows.get(where["team_id"])
async def find_many(self, where: Mapping[str, object]) -> list[LiteLLM_TeamTable]:
return [t for t in self.rows.values() if _matches({"team_id": t.team_id}, where)]
async def update(self, where: Mapping[str, str], data: Mapping[str, str]) -> LiteLLM_TeamTable:
self.update_calls += 1
team = self.rows[where["team_id"]]
@ -143,7 +146,7 @@ class _Tx:
self.locks.append(team_id)
self._on_lock(team_id)
return []
assert self.locks == [team_id], "roster must be read under this team's advisory lock"
assert team_id in self.locks, "roster must be read under this team's advisory lock"
self.roster_reads.append(team_id)
team = self.litellm_teamtable.rows.get(team_id)
if team is None:
@ -162,22 +165,22 @@ class _FakePrisma:
org_memberships: Sequence[Mapping[str, object]] = (),
on_lock: Callable[[str], None] = lambda _: None,
fail_locks: frozenset[str] = frozenset(),
fail_user_delete: bool = False,
fail_commit: bool = False,
) -> None:
self.db = _Db(users, teams, memberships, tokens, invitations, org_memberships)
self._on_lock = on_lock
self._fail_locks = fail_locks
self._fail_user_delete = fail_user_delete
self._fail_commit = fail_commit
self.locks: list[str] = []
self.roster_reads: list[str] = []
@asynccontextmanager
async def tx(self):
async def tx(self, *, timeout: object = None):
snapshot = copy.deepcopy(self.db)
tx = _Tx(self.db, self._on_lock, self._fail_locks)
try:
yield tx
if self._fail_user_delete and tx.locks == []:
if self._fail_commit:
raise RuntimeError("connection reset")
except BaseException:
self.db.__dict__.update(snapshot.__dict__)
@ -271,7 +274,7 @@ async def test_bulk_delete_removes_users_from_every_team_and_store():
assert [t["token"] for t in prisma.db.litellm_deletedverificationtoken.rows] == ["k1"]
assert [i["id"] for i in prisma.db.litellm_invitationlink.rows] == ["i3"]
assert prisma.db.litellm_organizationmembership.rows == []
assert sorted(prisma.locks) == ["t1", "t2"] and sorted(prisma.roster_reads) == ["t1", "t2"]
assert prisma.locks == ["t1", "t2"] and prisma.roster_reads == ["t1", "t2"]
@pytest.mark.asyncio
@ -320,22 +323,36 @@ async def test_bulk_delete_reports_missing_and_duplicate_ids_per_item_and_still_
@pytest.mark.asyncio
async def test_bulk_delete_keeps_user_when_a_team_rewrite_fails_and_deletes_the_others():
async def test_bulk_delete_rolls_back_every_team_and_user_when_one_team_rewrite_fails():
prisma = _FakePrisma(
users=[_user("u1", "bad", "good"), _user("u2", "good")],
teams=[_team("bad", "u1"), _team("good", "u1", "u2")],
fail_locks=frozenset({"bad"}),
users=[_user("u1", "a-good", "z-bad"), _user("u2", "a-good")],
teams=[_team("a-good", "u1", "u2"), _team("z-bad", "u1")],
tokens=[{"token": "k1", "user_id": "u1", "team_id": "a-good"}],
fail_locks=frozenset({"z-bad"}),
)
cache = _cache_with("k1")
response = await _delete(prisma, ["u1", "u2"])
response = await _delete(prisma, ["u1", "u2"], cache=cache)
assert [(r.user_id, r.success, r.teams_removed) for r in response.results] == [
("u1", False, ("good",)),
("u2", True, ("good",)),
assert [(r.user_id, r.success, r.teams_removed, r.error) for r in response.results] == [
("u1", False, (), "Failed to delete user: lock timeout"),
("u2", False, (), "Failed to delete user: lock timeout"),
]
assert response.results[0].error == "Failed to remove from team bad: lock timeout"
assert set(prisma.db.litellm_usertable.rows) == {"u1"}
assert _roster(prisma, "bad") == ["u1"] and _roster(prisma, "good") == []
assert set(prisma.db.litellm_usertable.rows) == {"u1", "u2"}
assert _roster(prisma, "a-good") == ["u1", "u2"] and _roster(prisma, "z-bad") == ["u1"]
assert [t["token"] for t in prisma.db.litellm_verificationtoken.rows] == ["k1"]
assert cache.get_cache(key="k1") is not None
@pytest.mark.asyncio
async def test_bulk_delete_skips_teams_the_user_still_names_but_which_no_longer_exist():
prisma = _FakePrisma(users=[_user("u1", "gone", "t1")], teams=[_team("t1", "u1", "keep")])
response = await _delete(prisma, ["u1"])
assert [(r.success, r.teams_removed) for r in response.results] == [(True, ("t1",))]
assert prisma.db.litellm_usertable.rows == {} and _roster(prisma, "t1") == ["keep"]
assert prisma.locks == ["t1"]
@pytest.mark.asyncio
@ -344,7 +361,7 @@ async def test_bulk_delete_rolls_back_every_user_row_and_reports_it_per_row_when
users=[_user("u1", "t1"), _user("u2")],
teams=[_team("t1", "u1")],
tokens=[{"token": "k1", "user_id": "u1"}],
fail_user_delete=True,
fail_commit=True,
)
cache = _cache_with("k1")