diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 49461d7841d..ba4fb81323f 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -16,11 +16,12 @@ import traceback from collections.abc import Mapping, Sequence from datetime import datetime, timezone from types import MappingProxyType -from typing import Annotated, Final, NamedTuple, NoReturn, Protocol, TypedDict, TypeVar, cast +from typing import Annotated, Final, NamedTuple, NoReturn, Protocol, TypeVar, cast import fastapi from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from pydantic import BaseModel, JsonValue +from typing_extensions import ReadOnly, TypedDict import litellm from litellm._logging import verbose_proxy_logger @@ -116,6 +117,7 @@ from litellm.proxy.management_endpoints.tag_management_endpoints import ( get_daily_activity, ) from litellm.proxy.management_helpers.access_group_team_sync import ( + TEAM_ADVISORY_LOCK_SQL, AccessGroupSyncTx, invalidate_access_group_caches, reconcile_team_access_group_membership, @@ -134,6 +136,7 @@ from litellm.proxy.management_helpers.team_metadata_validation import ( validate_team_metadata_if_configured, ) from litellm.proxy.management_helpers.utils import ( + MemberWriteTx, add_new_member, management_endpoint_wrapper, ) @@ -330,11 +333,44 @@ class _TeamIdInFilter(TypedDict, total=False): team_id: Mapping[str, Sequence[str]] +class _DeletedTeamsResult(TypedDict): + deleted_teams: ReadOnly[Sequence[str]] + + +class _ErrorDetail(TypedDict): + error: ReadOnly[str] + + class _TeamCreateTx(AccessGroupSyncTx, Protocol): @property def litellm_teamtable(self) -> "_PrismaTableActions[LiteLLM_TeamTable]": ... +class _MemberDeleteTx(Protocol): + """The tables `/team/member_delete` reads while it holds the team's advisory lock. + + Reading them off the transaction keeps the whole endpoint on the one pooled connection + it already checked out: a request that has the lock but still needs another connection + can be starved by the lock waiters, which is a deadlock rather than a wait when enough + of them hold the rest of the pool.""" + + @property + def litellm_usertable(self) -> "_PrismaTableActions[LiteLLM_UserTable]": ... + + @property + def litellm_verificationtoken(self) -> "_PrismaTableActions[LiteLLM_VerificationToken]": ... + + +class _TeamDeleteTx(AccessGroupSyncTx, Protocol): + async def execute_raw(self, query: str, *args: object) -> int: ... + + @property + def litellm_teamtable(self) -> "_PrismaTableActions[LiteLLM_TeamTable]": ... + + @property + def litellm_teammembership(self) -> "_PrismaTableActions[LiteLLM_TeamMembership]": ... + + _STRIP_DELETED_TEAM_FROM_USERS_SQL: Final = """ UPDATE "LiteLLM_UserTable" SET teams = array_remove(teams, $1) WHERE $1 = ANY(teams) """ @@ -2580,8 +2616,13 @@ async def _process_team_members( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, + tx: MemberWriteTx | None = None, ) -> tuple[list[LiteLLM_UserTable], list[LiteLLM_TeamMembership]]: - """Process and add new team members.""" + """Process and add new team members. + + ``tx`` is the caller's open transaction, when it has one, so the member writes run on the + connection it already holds instead of checking out a second one. + """ updated_users: Final[list[LiteLLM_UserTable]] = [] updated_team_memberships: Final[list[LiteLLM_TeamMembership]] = [] @@ -2607,6 +2648,7 @@ async def _process_team_members( default_team_budget_id=default_team_budget_id, allowed_models=member_allowed_models, budget_duration=data.budget_duration, + tx=tx, ) except Exception as e: raise HTTPException( @@ -2629,6 +2671,7 @@ async def _process_team_members( default_team_budget_id=default_team_budget_id, allowed_models=member_allowed_models, budget_duration=data.budget_duration, + tx=tx, ) except Exception as e: raise HTTPException( @@ -2708,65 +2751,40 @@ async def _add_team_members_to_team( user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, ) -> tuple[LiteLLM_TeamTable, list[LiteLLM_UserTable], list[LiteLLM_TeamMembership]]: - """Add team members to the team. + """Add team members to the team, under the team's advisory lock. - The members_with_roles reconciliation runs inside a transaction that locks - the team row with ``SELECT ... FOR UPDATE`` before reading the current - membership. Concurrent /team/member_add calls for the same team therefore - serialize on the row lock and each appends onto the other's committed - result, instead of both rewriting the whole JSON array from a stale - snapshot (which silently drops one member on the losing write). + The lock (``TEAM_ADVISORY_LOCK_SQL``, keyed on the team id) is taken first, and the + team is re-read under it before any write, so a delete that already committed is + visible here before this call writes anything: the user and membership writes only + happen once the re-read proves the team is still live. /team/delete takes the same + lock around its own sweep-and-delete, so the two can never interleave; whichever + acquires the lock first runs to completion before the other's re-read can proceed. - The same lock serializes this against /team/delete: the delete cannot remove - the row while the reconcile holds it, and a reconcile that finds the row - already gone cleans up after itself rather than leaving the member pointing - at a deleted team id. - """ - # Process and add new members - updated_users, updated_team_memberships = await _process_team_members( - data=data, - complete_team_data=complete_team_data, - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ) - - updated_team: Final = await _write_members_with_roles_locked( - data=data, - complete_team_data=complete_team_data, - prisma_client=prisma_client, - updated_users=updated_users, - ) - if updated_team is None: - await _sweep_deleted_team_references(team_ids=(data.team_id,), prisma_client=prisma_client) - raise HTTPException( - status_code=404, - detail={"error": f"Team={data.team_id} was deleted while this member add was running"}, - ) - - return updated_team, updated_users, updated_team_memberships - - -async def _write_members_with_roles_locked( - data: TeamMemberAddRequest, - complete_team_data: LiteLLM_TeamTable, - prisma_client: PrismaClient, - updated_users: list[LiteLLM_UserTable], -) -> LiteLLM_TeamTable | None: - """Reconcile members_with_roles under the team row lock. None when the team row is gone. - - That read is at least as recent as the user and membership writes the caller - already made, so a missing row means /team/delete committed after them. Its - post-delete sweep can have run before those writes landed, which is why the - caller sweeps this team id again rather than only reporting the 404. + The user and membership writes run on this transaction too, not on a second + connection from the pool: a lock waiter that needs a connection it hasn't got yet is + a waiter that can deadlock the pool, since enough concurrent adds for one team would + hold every connection waiting on the lock while the holder waits for a free one. """ async with prisma_client.tx() as tx: + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, data.team_id) + locked_members: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, data.team_id) if locked_members is None: - return None - + gone_detail: Final[_ErrorDetail] = { + "error": f"Team={data.team_id} was deleted while this member add was running" + } + raise HTTPException(status_code=404, detail=gone_detail) complete_team_data.members_with_roles = locked_members + updated_users, updated_team_memberships = await _process_team_members( + data=data, + complete_team_data=complete_team_data, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, + tx=tx, + ) + await _update_team_members_list( data=data, complete_team_data=complete_team_data, @@ -2774,11 +2792,13 @@ async def _write_members_with_roles_locked( ) _db_team_members: Final = [m.model_dump() for m in complete_team_data.members_with_roles] - return await tx.litellm_teamtable.update( + updated_team: Final = await tx.litellm_teamtable.update( where={"team_id": data.team_id}, data={"members_with_roles": json.dumps(_db_team_members)}, ) + return updated_team, updated_users, updated_team_memberships + def _emit_team_members_metric(team: LiteLLM_TeamTable) -> None: """Update the Prometheus team members gauge after a membership change. @@ -3159,10 +3179,6 @@ async def team_member_add( litellm_proxy_admin_name=litellm_proxy_admin_name, ) - # Check if updated_team is None - if updated_team is None: - raise HTTPException(status_code=404, detail={"error": f"Team with id {data.team_id} not found"}) - _emit_team_members_metric(complete_team_data) await _create_team_member_add_audit_logs( @@ -3276,45 +3292,63 @@ async def team_member_delete( ) ## DELETE MEMBER FROM TEAM - removed_team_members, new_team_members = _cleanup_members_with_roles( - existing_team_row=existing_team_row, - data=data, - ) - - if not removed_team_members: - raise HTTPException(status_code=400, detail={"error": "User not found in team"}) - - existing_team_row.members_with_roles = new_team_members - - _db_new_team_members: Final[list[dict]] = [m.model_dump() for m in new_team_members] - - ## 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) - key_val: Final[Mapping[str, object]] = ( - {"user_id": {"in": sorted(removed_user_ids)}} if removed_user_ids else {"user_email": data.user_email} - ) - existing_user_rows: Final[Sequence[LiteLLM_UserTable]] = await _user_db(prisma_client).find_many(where=key_val) - - # 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), - ) - - ## DELETE KEYS CREATED BY USER FOR THIS TEAM - # 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. + # Everything from here on runs under the team's advisory lock, the same one + # /team/member_add and /team/delete take: without it, this endpoint's own row-level + # update lock used to be the only thing serializing it against a concurrent member_add, + # and only by accident (their SELECT ... FOR UPDATE contended for the same row lock this + # UPDATE takes). Now that member_add reads under the advisory lock instead, this has to + # take it too, and re-read the roster under it rather than off the snapshot validated + # above, or a member_add that commits in between can have its addition silently + # overwritten by this delete computing from stale data. async with prisma_client.tx() as tx: + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, data.team_id) + + fresh_members: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, data.team_id) + if fresh_members is None: + raise HTTPException( + status_code=400, + detail={"error": f"Team id={data.team_id} does not exist in db"}, + ) + + removed_team_members, new_team_members = _cleanup_members_with_roles( + existing_team_row=LiteLLM_TeamTable(team_id=data.team_id, members_with_roles=fresh_members), + data=data, + ) + + if not removed_team_members: + raise HTTPException(status_code=400, detail={"error": "User not found in team"}) + + existing_team_row.members_with_roles = new_team_members + + _db_new_team_members: Final[list[dict]] = [m.model_dump() for m in new_team_members] + + ## 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) + key_val: Final[Mapping[str, object]] = ( + {"user_id": {"in": sorted(removed_user_ids)}} if removed_user_ids else {"user_email": data.user_email} + ) + member_tx: Final[_MemberDeleteTx] = tx + existing_user_rows: Final[Sequence[LiteLLM_UserTable]] = await member_tx.litellm_usertable.find_many( + where=key_val + ) + + # 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), + ) + + ## DELETE KEYS CREATED BY USER FOR THIS TEAM + # 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 member_tx.litellm_verificationtoken.find_many( + where={ + "user_id": {"in": sorted(user_ids_to_delete)}, + "team_id": data.team_id, + } + ) + await tx.litellm_teamtable.update( where={"team_id": data.team_id}, data={"members_with_roles": json.dumps(_db_new_team_members)}, @@ -4009,7 +4043,21 @@ async def delete_team( await _sweep_deleted_team_references(team_ids=data.team_ids, prisma_client=prisma_client) ## DELETE TEAMS - deleted_teams: Final = await prisma_client.delete_data(team_id_list=data.team_ids, table_name="team") + # Both the delete and the reconcile sweep run under every team's advisory lock + # (TEAM_ADVISORY_LOCK_SQL, the same one /team/member_add takes before its own writes), + # sorted so two overlapping batch deletes always request their locks in the same order. + # A member_add mid-flight for one of these teams either finishes its write and releases + # the lock before this transaction starts, in which case this sweep reaches what it wrote, + # or is still waiting on the lock, in which case its own re-read happens after this commits + # and sees the row gone before it writes anything. + delete_filter: Final[_TeamIdInFilter] = {"team_id": {"in": data.team_ids}} + async with prisma_client.tx() as tx: + for team_id in sorted(data.team_ids): + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id) + await tx.litellm_teamtable.delete_many(where=delete_filter) + await _sweep_deleted_team_references_tx(team_ids=data.team_ids, tx=tx) + + deleted_teams: Final[_DeletedTeamsResult] = {"deleted_teams": data.team_ids} # Evict AFTER the rows are gone. Both writers of these keys (`_cache_team_object` and # `get_team_object_by_alias`) hydrate from the db, so evicting first leaves a window where a @@ -4022,12 +4070,6 @@ async def delete_team( proxy_logging_obj=proxy_logging_obj, ) - # Sweep again now the team is gone. A `/team/member_add` that landed between the first sweep - # and the delete would have re-appended the reference; an add still in flight sees the row - # missing under its own row lock and sweeps what it wrote. Both passes are idempotent, and - # keeping the first one means a failure here still leaves a team the admin can retry deleting. - await _sweep_deleted_team_references(team_ids=data.team_ids, prisma_client=prisma_client) - for deleted_team in team_rows: await sync_team_access_group_membership(prisma_client=prisma_client, team_id=deleted_team.team_id) @@ -4056,6 +4098,16 @@ async def _sweep_deleted_team_references(team_ids: Sequence[str], prisma_client: _ = await _team_membership_db(prisma_client).delete_many(where=_TeamIdInFilter(team_id={"in": tuple(team_ids)})) +async def _sweep_deleted_team_references_tx(team_ids: Sequence[str], tx: _TeamDeleteTx) -> None: + """Same sweep as `_sweep_deleted_team_references`, run on the transaction that holds + every id's advisory lock and deletes the team rows, so it commits or rolls back with them.""" + for team_id in team_ids: + _ = await tx.execute_raw(_STRIP_DELETED_TEAM_FROM_USERS_SQL, team_id) + + membership_filter: Final[_TeamIdInFilter] = {"team_id": {"in": tuple(team_ids)}} + _ = await tx.litellm_teammembership.delete_many(where=membership_filter) + + async def _invalidate_deleted_key_cache( keys: Sequence[LiteLLM_VerificationToken], user_api_key_cache: UserApiKeyCache, diff --git a/litellm/proxy/management_helpers/access_group_team_sync.py b/litellm/proxy/management_helpers/access_group_team_sync.py index 55c0346e375..664e36c9f10 100644 --- a/litellm/proxy/management_helpers/access_group_team_sync.py +++ b/litellm/proxy/management_helpers/access_group_team_sync.py @@ -23,9 +23,11 @@ from pydantic import BaseModel, TypeAdapter from litellm.proxy.auth.auth_checks import _delete_cache_access_object # hashtext collisions only cost two unrelated teams a little serialization, and the -# lock is never taken by the access-group endpoints, so it cannot join their -# access-group-then-team lock order to form a cycle. -_LOCK_TEAM_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" +# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, +# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints +# reuses this exact statement to serialize /team/member_add and /team/delete against each +# other and against this mirror, rather than defining a second, divergent lock on the same key. +TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" _READ_TEAM_SQL: Final = 'SELECT access_group_ids FROM "LiteLLM_TeamTable" WHERE team_id = $1' @@ -138,7 +140,7 @@ async def reconcile_team_access_group_membership(tx: AccessGroupSyncTx, team_id: concurrent write for a different team cannot be lost the way a read-modify-write of the whole array can, and the pair commits together or not at all. """ - await tx.query_raw(_LOCK_TEAM_SQL, team_id) + await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id) team_rows: Final = _TeamRows.validate_python(await tx.query_raw(_READ_TEAM_SQL, team_id)) desired: Final = (team_rows[0].access_group_ids or ()) if team_rows else () affected: Final = _AffectedGroups.validate_python(await tx.query_raw(_AFFECTED_SQL, team_id, desired)) diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index cb30ce90c7f..e2d7262fb69 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -34,7 +34,7 @@ from litellm.proxy._types import ( # key request types; user request types; tea ) from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time -from litellm.proxy.utils import PrismaClient +from litellm.proxy.utils import PrismaClient, jsonify_object from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.table_repositories import TeamMembershipRepository from litellm.repositories.user_repository import UserRepository @@ -79,6 +79,8 @@ class _PrismaUserTable(Protocol): self, *, where: Mapping[str, object], data: Mapping[str, Mapping[str, object]] ) -> _PrismaUserRecord | None: ... + async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_PrismaUserRecord]: ... + class _PrismaTeamMembershipTable(Protocol): """Team membership table actions the management helpers issue.""" @@ -86,6 +88,73 @@ class _PrismaTeamMembershipTable(Protocol): async def create(self, *, data: Mapping[str, object], include: Mapping[str, bool]) -> _PrismaRecord: ... +class MemberWriteTx(Protocol): + """Transaction surface `add_new_member` writes through when the caller owns one. + + A caller already holding a transaction, and with it a pooled connection plus that + transaction's locks, passes it here so these writes reuse that connection rather than + checking out another one that lock waiters may already have drained from the pool. + """ + + @property + def litellm_usertable(self) -> _PrismaUserTable: ... + + @property + def litellm_budgettable(self) -> _PrismaBudgetTable: ... + + @property + def litellm_teammembership(self) -> _PrismaTeamMembershipTable: ... + + +def _user_table(prisma_client: PrismaClient, tx: MemberWriteTx | None) -> _PrismaUserTable: + return tx.litellm_usertable if tx is not None else UserRepository(prisma_client).table + + +def _budget_table(prisma_client: PrismaClient, tx: MemberWriteTx | None) -> _PrismaBudgetTable: + return tx.litellm_budgettable if tx is not None else BudgetRepository(prisma_client).table + + +def _team_membership_table(prisma_client: PrismaClient, tx: MemberWriteTx | None) -> _PrismaTeamMembershipTable: + return tx.litellm_teammembership if tx is not None else TeamMembershipRepository(prisma_client).table + + +async def _find_users_by_email( + prisma_client: PrismaClient, tx: MemberWriteTx | None, user_email: str +) -> Sequence[_PrismaUserRecord]: + if tx is not None: + return await tx.litellm_usertable.find_many(where={"user_email": user_email}) + rows: Final[Sequence[_PrismaUserRecord] | None] = await prisma_client.get_data( + key_val={"user_email": user_email}, + table_name="user", + query_type="find_all", + ) + return rows if rows is not None else () + + +async def _upsert_user_row( + user_table: _PrismaUserTable, user_id: str, create_data: Mapping[str, object] +) -> _PrismaUserRecord | None: + """Insert the user row if it is absent, leaving an existing row as it is. + + Upserting keeps concurrent provisioning of the same new user from racing on create. + The update branch re-states user_id rather than being empty because Prisma only + compiles an upsert down to INSERT ... ON CONFLICT when the update is non-empty, and + otherwise falls back to a racy SELECT-then-INSERT. + """ + return await user_table.upsert( + where={"user_id": user_id}, + data={"create": create_data, "update": {"user_id": user_id}}, + ) + + +async def _create_user_row( + prisma_client: PrismaClient, tx: MemberWriteTx | None, user_data: dict[str, object] +) -> _PrismaUserRecord | None: + if tx is not None: + return await _upsert_user_row(tx.litellm_usertable, str(user_data["user_id"]), jsonify_object(user_data)) + return await prisma_client.insert_data(data=user_data, table_name="user") + + def get_new_internal_user_defaults(user_id: str, user_email: str | None = None) -> dict[str, object]: user_info: Final = litellm.default_internal_user_params or {} @@ -206,6 +275,7 @@ async def _clone_team_default_budget_for_member( user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, budget_duration_override: str | None = None, + tx: MemberWriteTx | None = None, ) -> str | None: """ Create a new budget row that copies the values from the team's default @@ -220,7 +290,7 @@ async def _clone_team_default_budget_for_member( member while keeping the default's other limits, so an admin can set a member's reset cadence without discarding the team default's max_budget. """ - budget_table: Final[_PrismaBudgetTable] = BudgetRepository(prisma_client).table + budget_table: Final[_PrismaBudgetTable] = _budget_table(prisma_client, tx) default_budget: Final = await budget_table.find_unique(where={"budget_id": default_team_budget_id}) if default_budget is None: return None @@ -248,7 +318,7 @@ async def _clone_team_default_budget_for_member( if cloned_data.get("budget_duration"): cloned_data["budget_reset_at"] = get_budget_reset_time(cloned_data["budget_duration"]) - new_budget: Final[_PrismaBudgetRecord] = await BudgetRepository(prisma_client).table.create(data=cloned_data) + new_budget: Final[_PrismaBudgetRecord] = await budget_table.create(data=cloned_data) return new_budget.budget_id @@ -260,6 +330,7 @@ async def _resolve_member_budget_id( allowed_models: list[str] | None, budget_duration: str | None, default_team_budget_id: str | None, + tx: MemberWriteTx | None = None, ) -> str | None: """ Resolve the budget a new team member should be linked to. @@ -279,6 +350,7 @@ async def _resolve_member_budget_id( user_api_key_dict=user_api_key_dict, litellm_proxy_admin_name=litellm_proxy_admin_name, budget_duration_override=budget_duration, + tx=tx, ) if not has_explicit_limit and budget_duration is None: @@ -295,12 +367,14 @@ async def _resolve_member_budget_id( if budget_duration is not None: budget_data["budget_duration"] = budget_duration budget_data["budget_reset_at"] = get_budget_reset_time(budget_duration=budget_duration) - budget_table: Final[_PrismaBudgetTable] = BudgetRepository(prisma_client).table + budget_table: Final[_PrismaBudgetTable] = _budget_table(prisma_client, tx) response: Final = await budget_table.create(data=budget_data) return response.budget_id -async def _append_team_id_if_absent(prisma_client: PrismaClient, user_id: str, team_id: str) -> None: +async def _append_team_id_if_absent( + prisma_client: PrismaClient, user_id: str, team_id: str, tx: MemberWriteTx | None = None +) -> None: """Append team_id to a user's teams array, only if it is not already present. The row-level filter makes the append a no-op once the team is present, so @@ -309,7 +383,7 @@ async def _append_team_id_if_absent(prisma_client: PrismaClient, user_id: str, t number of teams a user belongs to). Teams added concurrently for a different team id are unaffected, since each update filters on its own team id. """ - user_table: Final[_PrismaUserTable] = UserRepository(prisma_client).table + user_table: Final[_PrismaUserTable] = _user_table(prisma_client, tx) await user_table.update_many( where={"user_id": user_id, "NOT": {"teams": {"has": team_id}}}, data={"teams": {"push": [team_id]}}, @@ -326,6 +400,7 @@ async def add_new_member( default_team_budget_id: str | None = None, allowed_models: list[str] | None = None, budget_duration: str | None = None, + tx: MemberWriteTx | None = None, ) -> tuple[LiteLLM_UserTable, LiteLLM_TeamMembership | None]: """ Add a new member to a team @@ -334,49 +409,41 @@ async def add_new_member( - add team member w/ budget to team member table Returns created/existing user + team membership w/ budget id + + Callers already inside a transaction pass it as ``tx`` so every write here runs on that + connection instead of borrowing more from the pool while the caller's locks are held. """ returned_user: LiteLLM_UserTable | None = None returned_team_membership: LiteLLM_TeamMembership | None = None ## ADD TEAM ID, to USER TABLE IF NEW ## if new_member.user_id is not None: new_user_defaults = get_new_internal_user_defaults(user_id=new_member.user_id) - # Upsert ensures the user row exists atomically (no create race when the - # same new user is provisioned concurrently), seeding teams on create. - # The teams append lives in the filtered update below rather than the - # upsert's update branch so an already-existing user does not get a - # duplicate team id. The update branch still has to write something: - # Prisma only compiles an upsert down to INSERT ... ON CONFLICT when it - # is non-empty, and falls back to a racy SELECT-then-INSERT when it is - # not, so this re-states user_id as a no-op rather than being empty. - user_table: Final[_PrismaUserTable] = UserRepository(prisma_client).table - _returned_user: _PrismaUserRecord | None = await user_table.upsert( - where={"user_id": new_member.user_id}, - data={ - "create": {"teams": [team_id], **new_user_defaults}, - "update": {"user_id": new_member.user_id}, - }, + # The teams append lives in the filtered update below rather than the upsert's + # update branch so an already-existing user does not get a duplicate team id. + _returned_user: _PrismaUserRecord | None = await _upsert_user_row( + _user_table(prisma_client, tx), + new_member.user_id, + {"teams": [team_id], **new_user_defaults}, ) - await _append_team_id_if_absent(prisma_client, new_member.user_id, team_id) + await _append_team_id_if_absent(prisma_client, new_member.user_id, team_id, tx) if _returned_user is not None: returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) elif new_member.user_email is not None: new_user_defaults = get_new_internal_user_defaults(user_id=str(uuid.uuid4()), user_email=new_member.user_email) ## user email is not unique acc. to prisma schema -> future improvement ### for now: check if it exists in db, if not - insert it - existing_user_row: Final[list[_PrismaUserRecord] | None] = await prisma_client.get_data( - key_val={"user_email": new_member.user_email}, - table_name="user", - query_type="find_all", + existing_user_row: Final[Sequence[_PrismaUserRecord]] = await _find_users_by_email( + prisma_client, tx, new_member.user_email ) - if existing_user_row is None or (isinstance(existing_user_row, list) and len(existing_user_row) == 0): + if len(existing_user_row) == 0: new_user_defaults["teams"] = [team_id] - _returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") + _returned_user = await _create_user_row(prisma_client, tx, new_user_defaults) if _returned_user is not None: returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) elif len(existing_user_row) == 1: user_info: Final = existing_user_row[0] - await _append_team_id_if_absent(prisma_client, user_info.user_id, team_id) + await _append_team_id_if_absent(prisma_client, user_info.user_id, team_id, tx) returned_user = LiteLLM_UserTable.model_validate(user_info.model_dump()) elif len(existing_user_row) > 1: raise HTTPException( @@ -392,10 +459,11 @@ async def add_new_member( allowed_models=allowed_models, budget_duration=budget_duration, default_team_budget_id=default_team_budget_id, + tx=tx, ) if _budget_id and returned_user is not None and returned_user.user_id is not None: - membership_table: Final[_PrismaTeamMembershipTable] = TeamMembershipRepository(prisma_client).table + membership_table: Final[_PrismaTeamMembershipTable] = _team_membership_table(prisma_client, tx) _returned_team_membership: Final = await membership_table.create( data={ "team_id": team_id, diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 7efd32288e4..d636592f925 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -58,19 +58,22 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return LiteLLM_TeamTable.model_validate(data) async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> list[Member] | None: - """Return the team's members_with_roles, locking the row FOR UPDATE. + """Return the team's members_with_roles. The caller must already hold + ``TEAM_ADVISORY_LOCK_SQL`` for this team_id on ``tx`` before calling this. - ``None`` when the team row is gone, which a caller holding the lock can - only see if a delete committed under it, as opposed to ``[]`` for a team - that simply has no members. + ``None`` when the team row is gone, which is only possible under that lock if + a delete committed before this read, as opposed to ``[]`` for a team that + simply has no members. - Must be called inside a transaction so the row lock is held until - commit. This serializes concurrent membership writers on the team row - so the losing writer appends onto the winner's committed result instead - of overwriting it from a stale snapshot. + A plain read is enough here because the advisory lock, not a row lock, is what + serializes this against a concurrent writer: ``SELECT ... FOR UPDATE`` would + additionally take a row lock on ``LiteLLM_TeamTable``, and the access-group + endpoints lock an access group and then a team row, so a team-row-first lock + here can deadlock with them. The advisory lock cannot, since those endpoints + never take it. """ rows: Final = await tx.query_raw( - 'SELECT members_with_roles FROM "LiteLLM_TeamTable" WHERE team_id = $1 FOR UPDATE', + 'SELECT members_with_roles FROM "LiteLLM_TeamTable" WHERE team_id = $1', team_id, ) if not rows: diff --git a/tests/proxy_admin_ui_tests/test_team_delete_member_add_race.py b/tests/proxy_admin_ui_tests/test_team_delete_member_add_race.py new file mode 100644 index 00000000000..30544a8bb81 --- /dev/null +++ b/tests/proxy_admin_ui_tests/test_team_delete_member_add_race.py @@ -0,0 +1,307 @@ +""" +Real-Postgres coverage for the /team/member_add vs /team/delete race (LIT-5544), and for +/team/member_delete's participation in the same lock. + +A member_add that validated the team before a delete began could previously still commit +its writes after the delete's reference sweeps had already run, leaving a user record and +a membership row pointing at a team id that no longer exists. Neither side of that race can +be forced by a sequential script: it needs one request to be genuinely mid-flight while the +other commits. A mocked prisma cannot arbitrate that either, since the property under test +is whether Postgres's own advisory lock actually serializes the two requests. + +These tests pin the interleaving the same way test_access_group_team_sync.py does: a second +real connection holds the team's advisory lock in its own transaction, so the function under +test is provably blocked on it rather than hoping a sleep lands in the right gap. +""" + +import asyncio +import json +import os +from contextlib import asynccontextmanager +from datetime import timedelta +from unittest.mock import MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import ( + DeleteTeamRequest, + LitellmUserRoles, + Member, + TeamMemberAddRequest, + UserAPIKeyAuth, +) +from litellm.caching.caching import DualCache +from litellm.proxy.utils import PrismaClient, ProxyLogging + +TEAM = "lit5544-race-team" +USER = "lit5544-race-user" +_DELETE_SEEDED = 'DELETE FROM "LiteLLM_TeamMembership" WHERE team_id = $1' +_DELETE_USER = 'DELETE FROM "LiteLLM_UserTable" WHERE user_id = $1' +_DELETE_TEAM = 'DELETE FROM "LiteLLM_TeamTable" WHERE team_id = $1' +_LOCK_SQL = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" + + +@asynccontextmanager +async def _clean_db(): + """Connects inside the running test's loop: an async fixture would be torn up on a + different loop than the test body, which prisma's engine lock refuses outright.""" + from prisma import Prisma + + if not os.getenv("DATABASE_URL"): + pytest.fail("DATABASE_URL is required; these tests must not silently skip") + + db = Prisma() + await db.connect() + try: + await db.execute_raw(_DELETE_SEEDED, TEAM) + await db.execute_raw(_DELETE_USER, USER) + await db.execute_raw(_DELETE_TEAM, TEAM) + yield db + finally: + await db.execute_raw(_DELETE_SEEDED, TEAM) + await db.execute_raw(_DELETE_USER, USER) + await db.execute_raw(_DELETE_TEAM, TEAM) + await db.disconnect() + + +@asynccontextmanager +async def _real_prisma_client(): + """The full app-level PrismaClient, not the raw generated client: add_new_member reads + and writes through PrismaClient.get_data/insert_data, which the raw client doesn't have.""" + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + client = PrismaClient(database_url=os.environ["DATABASE_URL"], proxy_logging_obj=proxy_logging_obj) + await client.connect() + try: + yield client + finally: + await client.db.disconnect() + + +def _admin_auth(): + return UserAPIKeyAuth(user_id="lit5544-admin", api_key="sk-lit5544", user_role=LitellmUserRoles.PROXY_ADMIN.value) + + +@pytest.mark.asyncio +async def test_member_add_blocked_by_delete_writes_no_dangling_reference(): + """ + member_add re-reads the team under the advisory lock before writing anything. When a + delete already holds that lock and then removes the row, member_add's re-read must see + the row gone and raise, without ever calling the write that appends the user/membership + references, which is the only way this leaves zero trace after the delete wins. + """ + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.management_endpoints.team_endpoints import ( + _add_team_members_to_team, + ) + + async with _clean_db() as db: + await db.litellm_teamtable.create(data={"team_id": TEAM, "team_alias": TEAM, "members_with_roles": "[]"}) + + async with _real_prisma_client() as prisma_client: + from prisma import Prisma + + blocker = Prisma() + await blocker.connect() + lock_acquired = asyncio.Event() + + async def add_member(): + lock_acquired.set() + await _add_team_members_to_team( + data=TeamMemberAddRequest( + team_id=TEAM, + member=Member(user_id=USER, role="user"), + max_budget_in_team=5.0, + ), + complete_team_data=LiteLLM_TeamTable(team_id=TEAM, members_with_roles=[]), + prisma_client=prisma_client, + user_api_key_dict=_admin_auth(), + litellm_proxy_admin_name="lit5544-admin", + ) + + try: + async with blocker.tx(timeout=timedelta(seconds=30)) as held: + await held.query_raw(_LOCK_SQL, TEAM) + task = asyncio.create_task(add_member()) + await lock_acquired.wait() + await asyncio.sleep(0.2) + assert not task.done(), "member_add did not wait on the team's advisory lock" + + # the delete wins the race: strip the team row while the lock is held + await held.execute_raw(_DELETE_TEAM, TEAM) + + with pytest.raises(HTTPException) as exc_info: + await asyncio.wait_for(task, timeout=30) + assert exc_info.value.status_code == 404 + finally: + await blocker.disconnect() + + user_row = await db.litellm_usertable.find_unique(where={"user_id": USER}) + assert user_row is None, "member_add must not have written a user row for a team that was gone under its lock" + + membership_row = await db.litellm_teammembership.find_first(where={"team_id": TEAM, "user_id": USER}) + assert membership_row is None + + +@pytest.mark.asyncio +async def test_member_delete_blocked_by_member_add_removes_from_the_fresh_roster(): + """ + team_member_delete takes the same advisory lock and re-reads the roster under it, so a + member_add that committed while member_delete was waiting on the lock is not silently + undone. Without the re-read, member_delete would compute its new roster from the stale + snapshot it validated against before the lock, and its write would overwrite the + member_add's addition right back out even though member_add's request already succeeded. + """ + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import TeamMemberDeleteRequest + from litellm.proxy.management_endpoints.team_endpoints import team_member_delete + + other_user = f"{USER}-other" + seeded_roster = '[{"user_id": "%s", "user_email": null, "role": "user"}]' % USER + winning_add_roster = ( + '[{"user_id": "%s", "user_email": null, "role": "user"}, ' + '{"user_id": "%s", "user_email": null, "role": "user"}]' % (USER, other_user) + ) + + async with _clean_db() as db: + await db.litellm_teamtable.create( + data={"team_id": TEAM, "team_alias": TEAM, "members_with_roles": seeded_roster} + ) + + async with _real_prisma_client() as prisma_client: + original_prisma_client = proxy_server_module.prisma_client + proxy_server_module.prisma_client = prisma_client + + try: + from prisma import Prisma + + blocker = Prisma() + await blocker.connect() + lock_acquired = asyncio.Event() + + async def run_delete(): + lock_acquired.set() + return await team_member_delete( + data=TeamMemberDeleteRequest(team_id=TEAM, user_id=USER), + user_api_key_dict=_admin_auth(), + ) + + try: + async with blocker.tx(timeout=timedelta(seconds=30)) as held: + await held.query_raw(_LOCK_SQL, TEAM) + task = asyncio.create_task(run_delete()) + await lock_acquired.wait() + await asyncio.sleep(0.2) + assert not task.done(), "member_delete did not wait on the team's advisory lock" + + # member_add wins the race: it adds `other_user` while holding the lock + await held.litellm_teamtable.update( + where={"team_id": TEAM}, + data={"members_with_roles": winning_add_roster}, + ) + + await asyncio.wait_for(task, timeout=30) + finally: + await blocker.disconnect() + finally: + proxy_server_module.prisma_client = original_prisma_client + + team_row = await db.litellm_teamtable.find_unique(where={"team_id": TEAM}) + raw_roster = team_row.members_with_roles + parsed_roster = json.loads(raw_roster) if isinstance(raw_roster, str) else raw_roster + remaining_ids = {m["user_id"] for m in parsed_roster} + assert remaining_ids == {other_user}, ( + "member_delete must remove only the user it targeted from the roster it actually " + "committed to, not silently drop the member the winning add just committed" + ) + + +@pytest.mark.asyncio +async def test_delete_blocked_by_member_add_sweeps_the_fresh_reference(): + """ + A member_add that wins the lock race writes its reference and releases the lock; the + delete that was waiting on it must then run its locked sweep against the row as it + actually is, not a stale snapshot, and reap that reference rather than leaving it + stranded on a team id the delete is about to remove. + """ + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.management_endpoints.team_endpoints import delete_team + + async with _clean_db() as db: + await db.litellm_teamtable.create(data={"team_id": TEAM, "team_alias": TEAM, "members_with_roles": "[]"}) + + async with _real_prisma_client() as prisma_client: + proxy_logging_obj = prisma_client.proxy_logging_obj + original_prisma_client = proxy_server_module.prisma_client + original_admin_name = proxy_server_module.litellm_proxy_admin_name + original_proxy_logging_obj = proxy_server_module.proxy_logging_obj + original_cache = proxy_server_module.user_api_key_cache + original_router = proxy_server_module.llm_router + proxy_server_module.prisma_client = prisma_client + proxy_server_module.litellm_proxy_admin_name = "lit5544-admin" + proxy_server_module.proxy_logging_obj = proxy_logging_obj + proxy_server_module.user_api_key_cache = original_cache or proxy_logging_obj.internal_usage_cache + proxy_server_module.llm_router = None + + async def restore(): + proxy_server_module.prisma_client = original_prisma_client + proxy_server_module.litellm_proxy_admin_name = original_admin_name + proxy_server_module.proxy_logging_obj = original_proxy_logging_obj + proxy_server_module.user_api_key_cache = original_cache + proxy_server_module.llm_router = original_router + + try: + from prisma import Prisma + + blocker = Prisma() + await blocker.connect() + lock_acquired = asyncio.Event() + + async def run_delete(): + lock_acquired.set() + return await delete_team( + data=DeleteTeamRequest(team_ids=[TEAM]), + http_request=MagicMock(), + user_api_key_dict=_admin_auth(), + litellm_changed_by="lit5544-admin", + ) + + try: + async with blocker.tx(timeout=timedelta(seconds=30)) as held: + await held.query_raw(_LOCK_SQL, TEAM) + task = asyncio.create_task(run_delete()) + await lock_acquired.wait() + await asyncio.sleep(0.3) + assert not task.done(), "delete_team did not wait on the team's advisory lock" + + # member_add wins the race: write the reference while holding the lock + await held.litellm_usertable.upsert( + where={"user_id": USER}, + data={ + "create": {"user_id": USER, "teams": [TEAM]}, + "update": {"teams": {"push": [TEAM]}}, + }, + ) + await held.litellm_teammembership.create(data={"team_id": TEAM, "user_id": USER}) + await held.litellm_teamtable.update( + where={"team_id": TEAM}, + data={"members_with_roles": '[{"user_id": "%s", "role": "user"}]' % USER}, + ) + + await asyncio.wait_for(task, timeout=30) + finally: + await blocker.disconnect() + finally: + await restore() + + team_row = await db.litellm_teamtable.find_unique(where={"team_id": TEAM}) + assert team_row is None + + user_row = await db.litellm_usertable.find_unique(where={"user_id": USER}) + assert user_row is not None and TEAM not in user_row.teams, ( + "delete_team's locked sweep must reap the reference member_add wrote just before losing the lock" + ) + + membership_row = await db.litellm_teammembership.find_first(where={"team_id": TEAM, "user_id": USER}) + assert membership_row is None diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 21dbf3e090f..ceaabf0a70f 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -1169,6 +1169,22 @@ async def test_create_user_default_budget(prisma_client, user_role): # noqa: F8 assert mock_client.call_args.kwargs["data"]["budget_duration"] is None +def _member_add_tx_cm(team_table): + """Transaction whose member writes land on whatever tables are mocked on `prisma_client.db`""" + + class _Tx: + query_raw = AsyncMock(return_value=[{"members_with_roles": []}]) + litellm_teamtable = team_table + + def __getattr__(self, table_name): + return getattr(litellm.proxy.proxy_server.prisma_client.db, table_name) + + tx_cm = MagicMock() + tx_cm.__aenter__ = AsyncMock(return_value=_Tx()) + tx_cm.__aexit__ = AsyncMock(return_value=None) + return tx_cm + + @pytest.mark.parametrize("new_member_method", ["user_id", "user_email"]) @pytest.mark.asyncio @pytest.mark.skip(reason="Requires reliable external DB connection (prisma).") @@ -1230,7 +1246,7 @@ async def test_create_team_member_add(prisma_client, new_member_method): # noqa ) ) mock_litellm_usertable.upsert = mock_client - mock_litellm_usertable.find_many = AsyncMock(return_value=None) + mock_litellm_usertable.find_many = AsyncMock(return_value=[]) # Mock find_first for user_email validation (returns None for new users) mock_litellm_usertable.find_first = AsyncMock(return_value=None) # Mock find_unique for user_id validation (returns None for new users) @@ -1245,12 +1261,7 @@ async def test_create_team_member_add(prisma_client, new_member_method): # noqa return_value=LiteLLM_TeamTableCachedObj(team_id="1234") ) - tx_mock = AsyncMock() - tx_mock.query_raw = AsyncMock(return_value=[{"members_with_roles": []}]) - tx_mock.litellm_teamtable = team_mock_client - tx_cm = MagicMock() - tx_cm.__aenter__ = AsyncMock(return_value=tx_mock) - tx_cm.__aexit__ = AsyncMock(return_value=None) + tx_cm = _member_add_tx_cm(team_mock_client) original_tx = litellm.proxy.proxy_server.prisma_client.tx litellm.proxy.proxy_server.prisma_client.tx = MagicMock( return_value=tx_cm @@ -1432,7 +1443,7 @@ async def test_create_team_member_add_team_admin( ) ) mock_litellm_usertable.upsert = mock_client - mock_litellm_usertable.find_many = AsyncMock(return_value=None) + mock_litellm_usertable.find_many = AsyncMock(return_value=[]) # Mock find_first for user_email validation (returns None for new users) mock_litellm_usertable.find_first = AsyncMock(return_value=None) # Mock find_unique for user_id validation (returns None for new users) @@ -1443,12 +1454,7 @@ async def test_create_team_member_add_team_admin( return_value=LiteLLM_TeamTableCachedObj(team_id="1234") ) - tx_mock = AsyncMock() - tx_mock.query_raw = AsyncMock(return_value=[{"members_with_roles": []}]) - tx_mock.litellm_teamtable = team_mock_client - tx_cm = MagicMock() - tx_cm.__aenter__ = AsyncMock(return_value=tx_mock) - tx_cm.__aexit__ = AsyncMock(return_value=None) + tx_cm = _member_add_tx_cm(team_mock_client) with ( patch.object( 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 7f5d3eb0a14..f33854eeb81 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -4,7 +4,7 @@ from contextlib import asynccontextmanager from datetime import datetime, timezone from types import SimpleNamespace from typing import Optional, cast -from unittest.mock import AsyncMock, MagicMock, call, patch +from unittest.mock import AsyncMock, MagicMock, PropertyMock, call, patch import pytest from fastapi import HTTPException @@ -58,6 +58,9 @@ from litellm.proxy.management_endpoints.team_endpoints import ( update_team, validate_team_org_change, ) +from litellm.proxy.management_helpers.access_group_team_sync import ( + TEAM_ADVISORY_LOCK_SQL, +) from litellm.proxy.management_helpers.team_member_permission_checks import ( TeamMemberPermissionChecks, ) @@ -75,7 +78,11 @@ client = TestClient(app) def _wire_team_create_tx(prisma_client): """`/team/new` inserts the team and mirrors it onto the access groups in one transaction, - so a mocked client has to hand its team table back out of `db.tx()`.""" + so a mocked client has to hand its team table back out of `db.tx()`. + + A `/team/new` carrying members then adds them under the team's advisory lock, and those + writes run on that lock's transaction, so `tx()` has to hand back the mocked tables too + for the per-table assertions on `prisma_client.db.*` to keep seeing them.""" @asynccontextmanager async def _tx(): @@ -85,18 +92,67 @@ def _wire_team_create_tx(prisma_client): ) prisma_client.db.tx = lambda *_args, **_kwargs: _tx() + _wire_member_add_tx(prisma_client) + + +def _wire_member_add_tx(prisma_client): + """/team/member_add takes the team's advisory lock, re-reads the roster under it, and runs + the user, budget, and membership writes on that same transaction, so a mocked client has + to hand its own table mocks back out of `tx()`. + + Tables resolve on access, not here, since tests routinely replace `db.