mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
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:
parent
595aba3cb0
commit
82872c9627
3 changed files with 178 additions and 140 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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"]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue