diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 6c9e1511b00..9c7dece6894 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1,7 +1,7 @@ import enum import json import os -from collections.abc import Callable, Mapping +from collections.abc import Callable, Mapping, Sequence from datetime import datetime from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple, TypeAlias @@ -2241,6 +2241,12 @@ class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase): class DeleteTeamRequest(LiteLLMPydanticObjectBase): team_ids: list[str] # required + @field_validator("team_ids") + @classmethod + def distinct_team_ids(cls, team_ids: Sequence[str]) -> list[str]: + """One delete per team: a repeated id would otherwise write its tombstone and audit row twice.""" + return list(dict.fromkeys(team_ids)) + class BlockTeamRequest(LiteLLMPydanticObjectBase): team_id: str # required diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 6cec3e714ec..f0e59389d48 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -4517,27 +4517,13 @@ async def delete_team( llm_router=llm_router, ) - # ## DELETE TEAM MEMBERSHIPS - for team_row in team_rows: - ### get all team members - team_members = team_row.members_with_roles - ### call team_member_delete for each team member - tasks = [] - for team_member in team_members: - tasks.append( - _team_member_delete( - data=TeamMemberDeleteRequest( - team_id=team_row.team_id, - user_id=team_member.user_id, - user_email=team_member.user_email, - ), - user_api_key_dict=user_api_key_dict, - ) - ) - await asyncio.gather(*tasks) - await _sweep_deleted_team_references(team_ids=data.team_ids, prisma_client=prisma_client) + member_ids_per_team: Final = await _resolve_deleted_team_member_user_ids( + teams=team_rows, + prisma_client=prisma_client, + ) + ## DELETE TEAMS # 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), @@ -4565,8 +4551,15 @@ async def delete_team( user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) + await _invalidate_deleted_team_member_cache( + member_ids_per_team=member_ids_per_team, + user_api_key_cache=user_api_key_cache, + ) for deleted_team in team_rows: + _emit_team_members_metric( + deleted_team.model_copy(update={"members_with_roles": ()}) # mutable-ok: pydantic update payload + ) await sync_team_access_group_membership(prisma_client=prisma_client, team_id=deleted_team.team_id) return deleted_teams @@ -4641,6 +4634,63 @@ async def _invalidate_deleted_team_cache( ) +async def _invalidate_deleted_team_member_cache( + member_ids_per_team: Sequence[tuple[str, Sequence[str]]], + user_api_key_cache: UserApiKeyCache, +) -> None: + for team_id, member_user_ids in member_ids_per_team: + await _evict_deleted_team_member_cache( + team_id=team_id, + member_user_ids=member_user_ids, + user_api_key_cache=user_api_key_cache, + ) + + +async def _evict_deleted_team_member_cache( + team_id: str, + member_user_ids: Sequence[str], + user_api_key_cache: UserApiKeyCache, +) -> None: + await evict_and_broadcast(cache_keys=tuple(member_user_ids), user_api_key_cache=user_api_key_cache) + await asyncio.gather( + *( + invalidate_team_member_spend_state( + user_id=user_id, + team_id=team_id, + user_api_key_cache=user_api_key_cache, + ) + for user_id in member_user_ids + ) + ) + + +async def _resolve_deleted_team_member_user_ids( + teams: Sequence[LiteLLM_TeamTable], + prisma_client: PrismaClient, +) -> tuple[tuple[str, tuple[str, ...]], ...]: + resolved: Final = await asyncio.gather( + *(_deleted_team_member_user_ids(team=team, prisma_client=prisma_client) for team in teams) + ) + return tuple(zip((team.team_id for team in teams), resolved)) + + +async def _deleted_team_member_user_ids(team: LiteLLM_TeamTable, prisma_client: PrismaClient) -> tuple[str, ...]: + roster_user_ids: Final = frozenset( + member.user_id for member in team.members_with_roles if member.user_id is not None + ) + email_only_member_emails: Final = frozenset( + member.user_email + for member in team.members_with_roles + if member.user_id is None and member.user_email is not None + ) + if not email_only_member_emails: + return tuple(sorted(roster_user_ids)) + # One case-insensitive lookup for the whole roster. A per-email fan-out would size the + # query count by team membership, the same shape as the P2028 fan-out this path removed. + email_only_users: Final = await UserRepository(prisma_client).find_by_emails(sorted(email_only_member_emails)) + return tuple(sorted(roster_user_ids.union(user.user_id for user in email_only_users))) + + def _transform_teams_to_deleted_records( teams: list[LiteLLM_TeamTable], user_api_key_dict: UserAPIKeyAuth, diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index 4a2aea46197..7bf516a2bc1 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -3,13 +3,15 @@ User repository for database operations on LiteLLM_UserTable. """ import json -from collections.abc import Mapping +from collections.abc import Mapping, Sequence +from itertools import chain from typing import TYPE_CHECKING, Final from pydantic import TypeAdapter from litellm.models.user import LiteLLM_UserTable, SCIMPlaceholder from litellm.repositories.base_repository import BaseRepository, DbRecord, record_to_dict +from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE from litellm.repositories.prisma_protocols import TableActions if TYPE_CHECKING: @@ -71,6 +73,31 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): records: Final = await self.find_many(where={"user_email": user_email}) return records[0] if records else None + async def find_by_emails(self, user_emails: Sequence[str]) -> Sequence[LiteLLM_UserTable]: + """Every user whose email matches one of ``user_emails``, ignoring case. + + A roster entry stored by email can differ in case from its user row (member_add + resolves emails case-insensitively), so an exact match would miss it. The list goes + out in slices of ``IN_LIST_CHUNK_SIZE`` so one statement stays under Postgres's + bind-parameter cap; ``chunked_in.find_many_in`` cannot carry the insensitive mode. + """ + unique: Final = sorted(frozenset(user_emails)) + pages: Final = tuple( + [ + await self.find_many( + where={ # mutable-ok: Prisma query filters are dict-shaped + "user_email": { # mutable-ok: Prisma query filters are dict-shaped + # bounded-ok: sliced to IN_LIST_CHUNK_SIZE values per statement + "in": unique[start : start + IN_LIST_CHUNK_SIZE], + "mode": "insensitive", + } + } + ) + for start in range(0, len(unique), IN_LIST_CHUNK_SIZE) + ] + ) + return tuple(chain.from_iterable(pages)) + async def find_by_sso_id(self, sso_user_id: str) -> LiteLLM_UserTable | None: """Find a user by SSO ID.""" return await self.find_by_id(sso_user_id, id_field="sso_user_id") diff --git a/tests/e2e/coverage_registry/mgmt.yaml b/tests/e2e/coverage_registry/mgmt.yaml index e1a840b1239..bc728efacd4 100644 --- a/tests/e2e/coverage_registry/mgmt.yaml +++ b/tests/e2e/coverage_registry/mgmt.yaml @@ -35,6 +35,7 @@ - {id: mgmt.team.update.team_admin_cannot_grow_budget, module: mgmt, tier: P0, surface: api, assertions: [team_admin_cannot_grow_budget], source: "team_endpoints.py:1203", fail_before_fix: proven, rationale: "With max_budget enabled, a team admin may keep or lower its team's budget; raising or removing it is 403 and writes nothing, also under an organization's larger cap"} - {id: mgmt.team.update.team_admin_resend_keeps_budget_reset, module: mgmt, tier: P1, surface: api, assertions: [team_admin_resend_keeps_budget_reset], source: "team_admin_field_permissions.py:147", fail_before_fix: proven, rationale: "A team admin resending unchanged budget settings with an enabled field must not push the team's budget reset times back"} - {id: mgmt.team.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py:1750", rationale: "Deletion prevents key access"} +- {id: mgmt.team.delete.membership_larger_than_db_pool, module: mgmt, tier: P0, surface: api, assertions: [membership_larger_than_db_pool], source: "team_endpoints.py:4362", rationale: "Deleting a team with more members than the Prisma connection pool still completes instead of exhausting the pool and answering 500", fail_before_fix: proven} - {id: mgmt.team.block.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py", rationale: "Block suspends all members"} - {id: mgmt.team.info.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "team_endpoints.py:2244", rationale: "Metadata+members+budgets"} - {id: mgmt.team.daily_activity.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "vendor testing strategy §9.20 / LIT-4778", rationale: "GET /team/daily/activity returns results+metadata for a valid date range"} diff --git a/tests/e2e/management/management_client.py b/tests/e2e/management/management_client.py index 7366695c0d1..3da9bea12a3 100644 --- a/tests/e2e/management/management_client.py +++ b/tests/e2e/management/management_client.py @@ -411,6 +411,27 @@ class ManagementClient: assert last is not None raise AssertionError(last) + def add_team_members(self, team_id: str, members: list[TeamMemberEntry]) -> None: + """Bulk form of /team/member_add: `member` accepts a list, so one call + seeds a whole roster the way an admin import does.""" + _ = unwrap( + self.proxy.transport.post( + "/team/member_add", + headers=self.proxy.management_headers(), + json=TeamMemberAddBody(team_id=team_id, member=members), + response_type=NoBody, + ) + ) + + def delete_team_status(self, team_id: str) -> StreamingResponse: + """POST /team/delete judged by HTTP outcome: the raw status and body, so a + test can assert on what a caller actually sees when the delete fails.""" + return self.proxy.transport.send( + "/team/delete", + headers=self.proxy.management_headers(), + json=TeamDeleteBody(team_ids=[team_id]), + ) + def delete_team_member(self, team_id: str, user_id: str) -> None: _ = unwrap( self.proxy.transport.post( diff --git a/tests/e2e/management/test_management_e2e.py b/tests/e2e/management/test_management_e2e.py index da0fc37aff8..908eb752611 100644 --- a/tests/e2e/management/test_management_e2e.py +++ b/tests/e2e/management/test_management_e2e.py @@ -37,6 +37,7 @@ from models import ( OrgUpdateBody, TagListEntry, TagNewBody, + TeamMemberEntry, TeamNewBody, TeamUpdateBody, UserNewBody, @@ -48,6 +49,7 @@ pytestmark = pytest.mark.e2e REGENERATE_GRACE_PERIOD = "15s" REGENERATE_GRACE_SECONDS = 15.0 +TEAM_DELETE_POOL_OVERFLOW_MEMBERS = 250 def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T: @@ -479,6 +481,43 @@ class TestTeamRoutes: client, rejected, "team-bound key was still accepted on chat (never rejected 401) after team deletion" ) + @pytest.mark.covers("mgmt.team.delete.membership_larger_than_db_pool") + def test_team_delete_succeeds_for_team_larger_than_db_pool( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """Customer repro: /team/delete fans one transaction per member out over a + Prisma pool of 10 connections, each queued on the team's advisory lock, + so a team bigger than the pool must still delete cleanly instead of + answering 500 P2028.""" + team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", []) + user_ids = tuple( + _create_user( + client, + resources, + UserNewBody( + user_email=f"e2e-mgmt-bulk-{i}-{unique_marker()}@example.com", + user_role="internal_user", + ), + ) + for i in range(TEAM_DELETE_POOL_OVERFLOW_MEMBERS) + ) + client.add_team_members(team_id, [TeamMemberEntry(role="user", user_id=user_id) for user_id in user_ids]) + seated = len(client.team_info(team_id).members_with_roles) + assert seated >= len(user_ids), ( + f"/team/info lists {seated} members after the bulk /team/member_add, expected at least {len(user_ids)}" + ) + + outcome = client.delete_team_status(team_id) + + assert outcome.status_code == 200, ( + f"/team/delete on a {len(user_ids)}-member team must succeed, got " + f"{outcome.status_code}: {outcome.body[:500]}" + ) + probe = client.team_info_status(team_id) + assert probe.status_code == 404, ( + f"deleted team {team_id} still resolves: /team/info returned {probe.status_code}: {probe.body[:300]}" + ) + @pytest.mark.covers("mgmt.team.member_add.persists") def test_member_add_and_delete_persist_to_team_info( self, client: ManagementClient, resources: ResourceManager diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 0383da48c43..4f8f6926095 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -1485,7 +1485,7 @@ class TeamInfoResponse(BaseModel): class TeamMemberAddBody(BaseModel): team_id: str - member: TeamMemberEntry + member: TeamMemberEntry | list[TeamMemberEntry] class TeamMemberDeleteBody(BaseModel): diff --git a/tests/integration/management/test_team_delete_chaos.py b/tests/integration/management/test_team_delete_chaos.py new file mode 100644 index 00000000000..ebe515f59ec --- /dev/null +++ b/tests/integration/management/test_team_delete_chaos.py @@ -0,0 +1,499 @@ +"""Chaos rows for ``/team/delete`` on an owned two-worker proxy: C1 worker kill, C2 Redis outage, C3 proxy restart. + +Each leg creates 24 teams through the owned proxy (two internal users per team in one bulk +``/team/member_add``, plus one team key), then deletes all 24 in a 24-thread burst and breaks the +infrastructure while a delete is provably in flight: the test holds the first team's advisory lock +from its own transaction, waits until that team's delete is queued behind it inside Postgres with +its request unanswered, applies the failure once the third of the other deletes has answered, and +only then releases the lock. The outage therefore overlaps a live delete on every run and both legs, +and the pinned delete finishes, or is dropped, under the failure: + +- C1 SIGKILLs one uvicorn worker child; the survivor still answers ``/health/readiness`` and uvicorn + respawns the worker. +- C2 shuts the owned Redis down; ``/cache/ping`` reports it, the deletes keep answering 200 because + cache eviction and the invalidation broadcast are best-effort, then Redis comes back. +- C3 SIGTERMs the owned proxy root and a fresh proxy starts on the same database. + +After recovery the burst outcomes (status or transport error per team) are recorded, every team whose +row survived is deleted once more, and the invariants must hold for every team: no ``LiteLLM_TeamTable`` +row, no ``LiteLLM_TeamMembership`` row, no ``LiteLLM_UserTable.teams`` entry naming it, its key gone +from ``LiteLLM_VerificationToken``, and one ``LiteLLM_DeletedTeamTable`` row per attempt that reached +the tombstone write. Both legs commit that tombstone before the locked transaction that removes the +team, so an attempt that died in between leaves a tombstone for a live team and the retry adds a +second; that count is pinned as observed (pre-existing, outside this PR's diff, recorded in the audit +report) and the affected teams are recorded as ``double_tombstones``. Teams found half-deleted before +the retry are recorded as ``partial_states_before_retry`` and named in any failure; the pinned team's +outcome is recorded as ``pinned_delete`` and the answers the outage interrupted as +``answered_before_outage``. + +Nothing sleeps, and only processes the test started are signalled. +""" + +from __future__ import annotations + +import os +import threading +import uuid +from collections import Counter +from collections.abc import Callable, Iterator, Sequence +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import httpx +import psutil +import psycopg +import pytest + +from tests.integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + delete_key_if_present, + eventually, + string_value, +) +from tests.integration._support.database import read_rows +from tests.integration._support.process import owned_proxy_process +from tests.integration._support.redis_process import owned_redis + +RecordProperty = Callable[[str, object], None] + +TEAMS: Final = 24 +MEMBERS_PER_TEAM: Final = 2 +CHAOS_AFTER_ANSWERS: Final = 3 +WORKERS: Final = 2 +DELETE_TIMEOUT_SECONDS: Final = 60 +REMOVE_FROM_ENVIRONMENT: Final = ("DATABASE_URL_READ_REPLICA",) + +TEAM_SQL: Final = 'SELECT team_id FROM "LiteLLM_TeamTable" WHERE team_id = %s' +TOMBSTONE_SQL: Final = 'SELECT id FROM "LiteLLM_DeletedTeamTable" WHERE team_id = %s' +MEMBERSHIP_SQL: Final = 'SELECT user_id FROM "LiteLLM_TeamMembership" WHERE team_id = %s' +REFERENCING_USERS_SQL: Final = 'SELECT user_id FROM "LiteLLM_UserTable" WHERE %s = ANY(teams)' +TOKEN_SQL: Final = 'SELECT token FROM "LiteLLM_VerificationToken" WHERE token = %s' +TAKE_TEAM_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext(%s))" +# Sessions blocked on an advisory lock the given backend holds: the pinned team's delete, on either leg. +WAITERS_ON_HELD_LOCK_SQL: Final = """ +SELECT count(*)::int AS waiting +FROM pg_locks waiter +JOIN pg_stat_activity session ON session.pid = waiter.pid +WHERE waiter.locktype = 'advisory' + AND NOT waiter.granted + AND session.wait_event_type = 'Lock' + AND session.query ILIKE %s + AND (waiter.classid, waiter.objid, waiter.objsubid) IN ( + SELECT held.classid, held.objid, held.objsubid + FROM pg_locks held + WHERE held.locktype = 'advisory' AND held.granted AND held.pid = %s::int + ) +""" + + +@dataclass(frozen=True, slots=True) +class Team: + team_id: str + members: tuple[str, ...] + hashed_key: str + + +@dataclass(frozen=True, slots=True) +class Outcome: + """One burst delete: the HTTP status, or ``None`` with the transport error's class and message.""" + + team_id: str + status: int | None + detail: str + + @property + def label(self) -> str: + return str(self.status) if self.status is not None else self.detail.split(":", 1)[0] + + @property + def answered_or_dropped(self) -> bool: + """200, a 5xx from a dying process, or a transport error; a 4xx would mean a wrong delete.""" + return self.status is None or self.status == 200 or self.status >= 500 + + +@dataclass(frozen=True, slots=True) +class TeamState: + team_id: str + row_present: bool + tombstones: int + memberships: tuple[str, ...] + referencing_users: tuple[str, ...] + key_present: bool + + @property + def clean(self) -> bool: + """Row, memberships, ``teams`` references and key all gone; tombstones are counted per attempt.""" + return not self.row_present and not self.memberships and not self.referencing_users and not self.key_present + + @property + def untouched(self) -> bool: + return self.row_present and self.tombstones == 0 and self.key_present + + @property + def partial(self) -> bool: + return not (self.clean and self.tombstones == 1) and not self.untouched + + def describe(self) -> str: + return ( + f"{self.team_id}: row={'present' if self.row_present else 'gone'} tombstones={self.tombstones} " + f"memberships={len(self.memberships)} referencing_users={len(self.referencing_users)} " + f"key={'present' if self.key_present else 'gone'}" + ) + + +def _state(team: Team) -> TeamState: + return TeamState( + team.team_id, + row_present=bool(read_rows(TEAM_SQL, (team.team_id,))), + tombstones=len(read_rows(TOMBSTONE_SQL, (team.team_id,))), + memberships=tuple(string_value(row["user_id"]) for row in read_rows(MEMBERSHIP_SQL, (team.team_id,))), + referencing_users=tuple( + string_value(row["user_id"]) for row in read_rows(REFERENCING_USERS_SQL, (team.team_id,)) + ), + key_present=bool(read_rows(TOKEN_SQL, (team.hashed_key,))), + ) + + +def _states(fleet: Sequence[Team]) -> tuple[TeamState, ...]: + return tuple(_state(team) for team in fleet) + + +def _overrides() -> dict[str, str]: + return {"DATABASE_URL": os.environ["DATABASE_URL"]} + + +def _user(candidate: Gateway, scenario: Scenario) -> str: + """An internal user created through ``candidate``; its removal is registered on the shared rig.""" + user_id: Final = f"integration-chaos-{uuid.uuid4().hex}" + candidate.post("/user/new", {"user_id": user_id, "auto_create_key": False, "user_role": "internal_user"}) + scenario.cleanups.callback(scenario.delete_user, user_id) + return user_id + + +def _delete_team_if_present(candidate: Gateway, team_id: str) -> None: + if read_rows(TEAM_SQL, (team_id,)): + candidate.post("/team/delete", {"team_ids": [team_id]}) + assert read_rows(TEAM_SQL, (team_id,)) == [] + + +def _team(candidate: Gateway, scenario: Scenario, index: int) -> Team: + alias: Final = f"integration-chaos-{index:02d}-{uuid.uuid4().hex}" + team_id: Final = string_value(candidate.post("/team/new", {"team_alias": alias})["team_id"]) + scenario.cleanups.callback(_delete_team_if_present, scenario.gateway, team_id) + members: Final = tuple(_user(candidate, scenario) for _ in range(MEMBERS_PER_TEAM)) + candidate.post( + "/team/member_add", + {"team_id": team_id, "member": [{"role": "user", "user_id": user_id} for user_id in members]}, + ) + key: Final = string_value(candidate.post("/key/generate", {"team_id": team_id, "key_alias": alias})["key"]) + scenario.cleanups.callback(delete_key_if_present, scenario.gateway, key) + return Team(team_id, members, sha256(key.encode()).hexdigest()) + + +def _fleet(candidate: Gateway, scenario: Scenario) -> tuple[Team, ...]: + """24 teams created through ``candidate``, each verified intact: row, key, both members' membership + rows and ``teams`` entries present, so the invariants after the burst have something to remove. + + A master-key ``/team/new`` also seats ``default_user_id`` as an admin (roster entry, membership row and + ``teams`` entry), so the checks are supersets. Cleanup is registered on the shared rig; the team + callback only acts when a run fails before its delete. + """ + fleet: Final = tuple(_team(candidate, scenario, index) for index in range(TEAMS)) + for team, state in zip(fleet, _states(fleet)): + assert state.untouched, state.describe() + assert set(state.memberships) >= set(team.members), state.describe() + assert set(state.referencing_users) >= set(team.members), state.describe() + return fleet + + +class Burst: + """One ``/team/delete`` per team on ``target``, all submitted at once; ``chaos_point`` is set once the + third delete has answered (or failed), so the leg breaks the infrastructure mid-burst.""" + + def __init__(self, target: Gateway) -> None: + self._target: Final = target + self._lock: Final = threading.Lock() + self._answers = 0 # rebind-ok: counter behind _lock + self._futures: dict[str, Future[Outcome]] = {} + self.chaos_point: Final = threading.Event() + + def start(self, pool: ThreadPoolExecutor, fleet: Sequence[Team]) -> None: + assert not self._futures, "burst already started" + self._futures.update((team.team_id, pool.submit(self._delete, team)) for team in fleet) + assert self.chaos_point.wait(DELETE_TIMEOUT_SECONDS), ( + f"fewer than {CHAOS_AFTER_ANSWERS} deletes answered within {DELETE_TIMEOUT_SECONDS}s" + ) + + def _delete(self, team: Team) -> Outcome: + try: + response: Final = self._target.client.request( + "POST", + "/team/delete", + json={"team_ids": [team.team_id]}, + headers={"Authorization": f"Bearer {self._target.key}"}, + timeout=DELETE_TIMEOUT_SECONDS, + ) + outcome = Outcome(team.team_id, response.status_code, response.text[:200]) + except httpx.HTTPError as error: # a killed worker or a stopped proxy drops the in-flight request + outcome = Outcome(team.team_id, None, f"{type(error).__name__}: {error}"[:200]) + with self._lock: + self._answers += 1 + if self._answers >= CHAOS_AFTER_ANSWERS: + self.chaos_point.set() + return outcome + + def answered(self) -> int: + with self._lock: + return self._answers + + def pending(self, team_id: str) -> bool: + return not self._futures[team_id].done() + + def outcomes(self) -> tuple[Outcome, ...]: + return tuple(future.result(timeout=DELETE_TIMEOUT_SECONDS + 30) for future in self._futures.values()) + + +def _waiters_on_lock_held_by(backend_pid: int) -> int: + rows: Final = read_rows(WAITERS_ON_HELD_LOCK_SQL, ("%pg_advisory_xact_lock%", str(backend_pid))) + waiting: Final = rows[0]["waiting"] + assert isinstance(waiting, int) + return waiting + + +@contextmanager +def _holding_team_lock(team_id: str) -> Iterator[int]: + """Hold ``team_id``'s advisory lock in a test-owned transaction and yield the holder's backend pid; + leaving the block commits, which releases the lock.""" + with psycopg.connect(os.environ["DATABASE_URL"]) as holder: + holder.execute(TAKE_TEAM_LOCK_SQL, (team_id,)) + yield holder.info.backend_pid + + +def _await_pinned_delete_blocked(burst: Burst, pinned: Team, holder_pid: int, record_property: RecordProperty) -> None: + """The pinned team's delete is queued behind the held lock inside Postgres with its request unanswered, + so the failure applied next lands on a live delete; records how many other deletes had answered.""" + eventually(lambda: _waiters_on_lock_held_by(holder_pid), lambda waiting: waiting >= 1, seconds=20) + assert burst.pending(pinned.team_id), f"{pinned.team_id}: delete answered while its team lock was held" + record_property("answered_before_outage", burst.answered()) + + +def _record_burst( + record_property: RecordProperty, outcomes: Sequence[Outcome], observed: Sequence[TeamState], pinned: Team +) -> None: + """Record the status split, the pinned team's outcome and the half-deleted teams seen before the retry.""" + split: Final = Counter(outcome.label for outcome in outcomes) + record_property("status_split", dict(sorted(split.items()))) + pinned_outcome: Final = next(outcome for outcome in outcomes if outcome.team_id == pinned.team_id) + record_property( + "pinned_delete", + {"team_id": pinned.team_id, "status": pinned_outcome.status, "detail": pinned_outcome.detail}, + ) + record_property("partial_states_before_retry", [state.describe() for state in observed if state.partial]) + record_property("rows_present_before_retry", sum(state.row_present for state in observed)) + + +def _retry_survivors(target: Gateway, fleet: Sequence[Team], observed: Sequence[TeamState]) -> tuple[str, ...]: + """Delete once more, through ``target``, every team whose row survived the burst; each must answer 200.""" + survivors: Final = tuple(team.team_id for team, state in zip(fleet, observed) if state.row_present) + for team_id in survivors: + assert target.post("/team/delete", {"team_ids": [team_id]}) == {"deleted_teams": [team_id]} + return survivors + + +def _expected_tombstones(before: TeamState, retried: bool) -> int: + """One ``LiteLLM_DeletedTeamTable`` row per attempt that reached the tombstone write. + + Both legs commit the tombstone before the locked transaction that removes the team, so a burst + attempt that died in between left one (``before.tombstones``, 0 or 1) for a team whose row + survived, and the retry adds one more. Pinned as observed: pre-existing on the merge base, + outside this PR's diff, recorded in the audit report. + """ + assert before.tombstones <= 1, before.describe() + return before.tombstones + (1 if retried else 0) + + +def _assert_every_team_fully_deleted( + record_property: RecordProperty, + before_retry: Sequence[TeamState], + final: Sequence[TeamState], + retried: Sequence[str], +) -> None: + """Every team: row, memberships, ``teams`` references and key gone; tombstones one per attempt.""" + expected: Final = {state.team_id: _expected_tombstones(state, state.team_id in retried) for state in before_retry} + record_property("double_tombstones", sorted(team_id for team_id, count in expected.items() if count == 2)) + violations: Final = tuple( + f"{state.describe()} expected tombstones={expected[state.team_id]}" + for state in final + if not state.clean or state.tombstones != expected[state.team_id] or expected[state.team_id] == 0 + ) + assert not violations, ( + f"{len(violations)} of {len(final)} teams are not fully deleted after the retry:\n " + + "\n ".join(violations) + + f"\nhalf-deleted before the retry ({sum(state.partial for state in before_retry)}):\n " + + "\n ".join(state.describe() for state in before_retry if state.partial) + + f"\nretried ({len(retried)}): {sorted(retried)}" + ) + + +def _workers(root: psutil.Process) -> tuple[psutil.Process, ...]: + """uvicorn's worker children of the owned proxy root, spawned through ``multiprocessing.spawn``. + + The root's other child is the multiprocessing resource tracker; each worker's prisma query engine + is a grandchild. A worker that just died shows as a zombie whose cmdline raises, so it is left out. + """ + workers: Final = [] + for child in root.children(): + try: + cmdline = child.cmdline() + except (psutil.NoSuchProcess, psutil.AccessDenied): + continue + if any("multiprocessing.spawn" in part for part in cmdline): + workers.append(child) + return tuple(sorted(workers, key=lambda process: process.pid)) + + +def _cache_ping(target: Gateway) -> httpx.Response: + return target.request("GET", "/cache/ping") + + +def _cache_status(response: httpx.Response) -> str: + assert response.status_code == 200, f"/cache/ping: {response.status_code} {response.text}" + return string_value(JSON_OBJECT.validate_json(response.content)["status"]) + + +@pytest.mark.timeout(240) # owned two-worker proxy boot plus a 24-team fleet and its cleanup +def test_worker_killed_mid_burst_leaves_every_team_fully_deleted_after_retry( + gateway: Gateway, tmp_path: Path, record_property: RecordProperty +) -> None: + with ( + gateway.scenario() as scenario, + owned_proxy_process( + gateway, tmp_path, _overrides(), remove_environment=REMOVE_FROM_ENVIRONMENT, workers=WORKERS + ) as owned, + ThreadPoolExecutor(TEAMS) as pool, + ): + root: Final = psutil.Process(owned.process.pid) + fleet: Final = _fleet(owned.gateway, scenario) + pinned: Final = fleet[0] + burst: Final = Burst(owned.gateway) + with _holding_team_lock(pinned.team_id) as holder_pid: + burst.start(pool, fleet) + _await_pinned_delete_blocked(burst, pinned, holder_pid, record_property) + before: Final = _workers(root) + assert len(before) == WORKERS, [process.pid for process in before] + victim: Final = before[0] + victim.kill() # SIGKILL with the pinned delete blocked: the worker cannot finish its in-flight deletes + victim.wait(timeout=10) + with httpx.Client(base_url=str(owned.gateway.client.base_url), timeout=15, trust_env=False) as fresh: + readiness: Final = fresh.get("/health/readiness") + assert readiness.status_code == 200, ( + f"/health/readiness with worker {victim.pid} dead: {readiness.status_code} {readiness.text}" + ) + # The lock is released: the pinned delete finishes on the survivor, or was dropped with the victim. + outcomes: Final = burst.outcomes() + respawned: Final = eventually( + lambda: tuple(process.pid for process in _workers(root)), + lambda pids: len(pids) == WORKERS and victim.pid not in pids, + seconds=60, + ) + record_property( + "worker_pids", {"before": [process.pid for process in before], "killed": victim.pid, "after": respawned} + ) + observed: Final = _states(fleet) + _record_burst(record_property, outcomes, observed, pinned) + assert all(outcome.answered_or_dropped for outcome in outcomes), [ + (outcome.team_id, outcome.status, outcome.detail) for outcome in outcomes if not outcome.answered_or_dropped + ] + retried: Final = _retry_survivors(owned.gateway, fleet, observed) + _assert_every_team_fully_deleted(record_property, observed, _states(fleet), retried) + + +@pytest.mark.timeout(240) # owned Redis, owned two-worker proxy boot, 24-team fleet, Redis restart +def test_redis_stopped_mid_burst_keeps_deletes_answering_200( + gateway: Gateway, tmp_path: Path, record_property: RecordProperty +) -> None: + with ( + gateway.scenario() as scenario, + owned_redis(tmp_path) as coordination, + owned_proxy_process( + gateway, + tmp_path, + { + **_overrides(), + "REDIS_HOST": coordination.host, + "REDIS_PORT": str(coordination.port), + # The breaker opens during the outage; the default 60 s before it probes again would + # keep /cache/ping (whose set_cache runs under the breaker) at 503 long after restart. + "REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "5", + }, + remove_environment=REMOVE_FROM_ENVIRONMENT, + workers=WORKERS, + ) as owned, + ThreadPoolExecutor(TEAMS) as pool, + ): + fleet: Final = _fleet(owned.gateway, scenario) + pinned: Final = fleet[0] + assert _cache_status(_cache_ping(owned.gateway)) == "healthy" + burst: Final = Burst(owned.gateway) + with _holding_team_lock(pinned.team_id) as holder_pid: + burst.start(pool, fleet) + _await_pinned_delete_blocked(burst, pinned, holder_pid, record_property) + coordination.stop() + down: Final = _cache_ping(owned.gateway) + assert down.status_code == 503, f"/cache/ping with Redis stopped: {down.status_code} {down.text}" + assert "Service Unhealthy" in down.text, down.text + assert burst.pending(pinned.team_id), f"{pinned.team_id}: delete answered while its team lock was held" + # The lock is released with Redis down: the pinned delete's cache eviction runs against the outage. + outcomes: Final = burst.outcomes() + coordination.start() + recovered: Final = eventually(lambda: _cache_ping(owned.gateway), lambda r: r.status_code == 200, seconds=60) + assert _cache_status(recovered) == "healthy" + + observed: Final = _states(fleet) + _record_burst(record_property, outcomes, observed, pinned) + assert all(outcome.status == 200 for outcome in outcomes), ( + "deletes not answered 200 while Redis was down: " + + str([(outcome.team_id, outcome.status, outcome.detail) for outcome in outcomes if outcome.status != 200]) + + f"; split {dict(Counter(outcome.label for outcome in outcomes))}" + ) + retried: Final = _retry_survivors(owned.gateway, fleet, observed) + _assert_every_team_fully_deleted(record_property, observed, _states(fleet), retried) + + +@pytest.mark.timeout(240) # two owned two-worker proxy boots (before and after SIGTERM) plus a 24-team fleet +def test_proxy_terminated_mid_burst_then_restarted_leaves_every_team_fully_deleted( + gateway: Gateway, tmp_path: Path, record_property: RecordProperty +) -> None: + with gateway.scenario() as scenario, ThreadPoolExecutor(TEAMS) as pool: + with owned_proxy_process( + gateway, tmp_path, _overrides(), remove_environment=REMOVE_FROM_ENVIRONMENT, workers=WORKERS + ) as doomed: + fleet: Final = _fleet(doomed.gateway, scenario) + pinned: Final = fleet[0] + burst: Final = Burst(doomed.gateway) + with _holding_team_lock(pinned.team_id) as holder_pid: + burst.start(pool, fleet) + _await_pinned_delete_blocked(burst, pinned, holder_pid, record_property) + doomed.process.terminate() # SIGTERM with the pinned delete blocked: uvicorn stops accepting and drains + assert burst.pending(pinned.team_id), f"{pinned.team_id}: delete answered while its team lock was held" + # The lock is released: the drain lets the pinned delete finish before the proxy exits. + doomed.process.wait(timeout=120) + outcomes: Final = burst.outcomes() + + at_restart: Final = _states(fleet) + _record_burst(record_property, outcomes, at_restart, pinned) + assert all(outcome.answered_or_dropped for outcome in outcomes), [ + (outcome.team_id, outcome.status, outcome.detail) for outcome in outcomes if not outcome.answered_or_dropped + ] + with owned_proxy_process( + gateway, tmp_path, _overrides(), remove_environment=REMOVE_FROM_ENVIRONMENT, workers=WORKERS + ) as fresh: + retried: Final = _retry_survivors(fresh.gateway, fleet, at_restart) + final: Final = _states(fleet) + _assert_every_team_fully_deleted(record_property, at_restart, final, retried) diff --git a/tests/integration/management/test_team_delete_inputs.py b/tests/integration/management/test_team_delete_inputs.py new file mode 100644 index 00000000000..d385c915696 --- /dev/null +++ b/tests/integration/management/test_team_delete_inputs.py @@ -0,0 +1,330 @@ +"""Sad inputs for /team/delete: malformed ids, callers without access, and rosters the API can no longer produce. + +Legacy roster shapes (email-only entries, entries with neither id nor email) are seeded straight into +``members_with_roles`` because ``/team/member_add`` backfills ``user_id`` and will not write them any more. +""" + +from __future__ import annotations + +import json +import os +import uuid +from collections.abc import Callable, Mapping +from hashlib import sha256 +from typing import Final + +import httpx +import pytest +from pydantic import JsonValue +from redis import Redis + +from tests.integration._support.client import ( + Gateway, + Scenario, + delete_key_if_present, + eventually, + object_value, + string_value, +) +from tests.integration._support.database import read_rows, write_rows + +RecordProperty = Callable[[str, object], None] + +TEAM_SQL: Final = 'SELECT team_id FROM "LiteLLM_TeamTable" WHERE team_id = %s' +ROSTER_READ_SQL: Final = 'SELECT members_with_roles FROM "LiteLLM_TeamTable" WHERE team_id = %s' +ROSTER_SQL: Final = 'UPDATE "LiteLLM_TeamTable" SET members_with_roles = %s::jsonb WHERE team_id = %s' +MEMBERSHIP_SQL: Final = 'SELECT user_id FROM "LiteLLM_TeamMembership" WHERE team_id = %s' +USER_SQL: Final = 'SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = %s' +USER_EMAIL_SQL: Final = 'UPDATE "LiteLLM_UserTable" SET user_email = %s WHERE user_id = %s' +TOKEN_SQL: Final = 'SELECT token FROM "LiteLLM_VerificationToken" WHERE token = %s' +TOMBSTONE_SQL: Final = 'SELECT id FROM "LiteLLM_DeletedTeamTable" WHERE team_id = %s' +AUDIT_SQL: Final = 'SELECT id, table_name, action FROM "LiteLLM_AuditLog" WHERE object_id = %s' + +NOT_FOUND: Final = "Team not found, passed team_id=" +# /team/delete sits on management_routes but on no internal-user route list, so the route gate in +# RouteChecks.non_proxy_admin_allowed_routes_check answers 401 before _verify_team_access ever runs +# (pinned by tests/integration/authorization/test_team_admin_gate.py as team_admin=401, others=401). +ROUTE_GATE_MESSAGE: Final = "Only proxy admin can be used" +UNKNOWN_TEAM: Final = f"integration-missing-{uuid.uuid4().hex}" +FIVE_KB_TEAM: Final = "t" * 5120 + + +def _team_rows(team_id: str) -> list[dict[str, JsonValue]]: + return read_rows(TEAM_SQL, (team_id,)) + + +def _delete_team_if_present(gateway: Gateway, team_id: str) -> None: + """Cleanup for a team the test deletes itself: a no-op once the row is gone.""" + if _team_rows(team_id): + gateway.post("/team/delete", {"team_ids": [team_id]}) + assert _team_rows(team_id) == [] + + +def _reset_roster_if_present(team_id: str) -> None: + """Cleanup for a seeded roster: put back a shape the delete path always accepts.""" + if _team_rows(team_id): + write_rows(ROSTER_SQL, ("[]", team_id)) + + +def _delete_user_if_present(gateway: Gateway, user_id: str) -> None: + if read_rows(USER_SQL, (user_id,)): + response: Final = gateway.request("POST", "/user/delete", {"user_ids": [user_id]}) + assert response.status_code == 200, response.text + assert read_rows(USER_SQL, (user_id,)) == [] + + +def _own_team(scenario: Scenario, **fields: JsonValue) -> str: + """A team the test deletes itself, so cleanup tolerates the row already being gone.""" + created: Final = scenario.gateway.post("/team/new", {"team_alias": f"integration-{uuid.uuid4().hex}", **fields}) + team_id: Final = string_value(created["team_id"]) + scenario.cleanups.callback(_delete_team_if_present, scenario.gateway, team_id) + return team_id + + +def _own_key(scenario: Scenario, **fields: JsonValue) -> str: + """A key the team delete is expected to remove, so cleanup tolerates it already being gone.""" + created: Final = scenario.gateway.post("/key/generate", fields) + token: Final = string_value(created["key"]) + scenario.cleanups.callback(delete_key_if_present, scenario.gateway, token) + return token + + +def _own_user(scenario: Scenario) -> str: + """An internal user the test may remove by SQL, so cleanup tolerates the row already being gone.""" + created: Final = scenario.gateway.post( + "/user/new", + {"user_id": f"integration-{uuid.uuid4().hex}", "auto_create_key": False, "user_role": "internal_user"}, + ) + user_id: Final = string_value(created["user_id"]) + scenario.cleanups.callback(_delete_user_if_present, scenario.gateway, user_id) + return user_id + + +def _seed_roster(scenario: Scenario, team_id: str, entries: list[dict[str, JsonValue]]) -> None: + write_rows(ROSTER_SQL, (json.dumps(entries), team_id)) + scenario.cleanups.callback(_reset_roster_if_present, team_id) + + +def _team_admin(scenario: Scenario, team_id: str) -> str: + """Add a member and flip their roster role to admin by SQL: the API gates that role behind a license.""" + user_id: Final = scenario.member(team_id) + rows: Final = read_rows(ROSTER_READ_SQL, (team_id,)) + assert len(rows) == 1, rows + roster: Final = rows[0]["members_with_roles"] + assert isinstance(roster, list), roster + promoted: Final = [ + {**object_value(entry), "role": "admin"} if object_value(entry).get("user_id") == user_id else entry + for entry in roster + ] + assert any(object_value(entry).get("user_id") == user_id for entry in promoted), promoted + write_rows(ROSTER_SQL, (json.dumps(promoted), team_id)) + return user_id + + +def _membership_user_ids(team_id: str) -> frozenset[str]: + return frozenset(string_value(row["user_id"]) for row in read_rows(MEMBERSHIP_SQL, (team_id,))) + + +def _delete(gateway: Gateway, team_ids: JsonValue, *, key: str | None = None) -> httpx.Response: + return gateway.request("POST", "/team/delete", {"team_ids": team_ids}, key=key) + + +def _hashed(token: str) -> str: + return sha256(token.encode()).hexdigest() + + +@pytest.mark.parametrize( + ("body", "status", "needle"), + [ + pytest.param({"team_ids": [UNKNOWN_TEAM]}, 404, f"{NOT_FOUND}{UNKNOWN_TEAM}", id="S1-unknown-id"), + pytest.param({"team_ids": "abc"}, 422, "list_type", id="S3-string-not-list"), + pytest.param({"team_ids": [123]}, 422, "string_type", id="S4-integer-item"), + pytest.param({"team_ids": [""]}, 404, NOT_FOUND, id="S5-empty-id"), + pytest.param({"team_ids": [FIVE_KB_TEAM]}, 404, NOT_FOUND, id="S6-5kb-id"), + ], +) +def test_rejects_malformed_team_ids(gateway: Gateway, body: Mapping[str, JsonValue], status: int, needle: str) -> None: + response: Final = gateway.request("POST", "/team/delete", body) + assert response.status_code == status, f"{response.status_code} {response.text}" + assert needle in response.text, response.text + + +def test_empty_list_deletes_nothing(gateway: Gateway) -> None: + response: Final = _delete(gateway, []) + assert response.status_code == 200, f"{response.status_code} {response.text}" + assert response.json() == {"deleted_teams": []}, response.text + + +def test_duplicate_ids_delete_once(gateway: Gateway, record_property: RecordProperty) -> None: + """Repeated ids collapse to one delete: the body names the team once and exactly one tombstone row lands.""" + with gateway.scenario() as scenario: + team: Final = _own_team(scenario) + first: Final = scenario.member(team) + second: Final = scenario.member(team) + key: Final = _own_key(scenario, team_id=team) + # The master key's /team/new also seats the proxy admin, so the table holds more than these two. + assert {first, second} <= _membership_user_ids(team), _membership_user_ids(team) + response: Final = _delete(gateway, [team, team]) + # Read every table before the first assert so a red cell carries the partial state with it. + present: Final = _team_rows(team) + memberships: Final = _membership_user_ids(team) + key_rows: Final = read_rows(TOKEN_SQL, (_hashed(key),)) + tombstones: Final = read_rows(TOMBSTONE_SQL, (team,)) + audit: Final = read_rows(AUDIT_SQL, (team,)) + state: Final = ( + f"team_present={bool(present)} membership_rows={len(memberships)} key_present={bool(key_rows)} " + f"tombstone_rows={len(tombstones)} audit_rows={len(audit)}" + ) + record_property("status", response.status_code) + record_property("body", response.text) + record_property("state_after", state) + record_property("audit_rows", len(audit)) # recorded only: the shared rigs cannot enable audit logging + assert response.status_code == 200, f"{response.status_code} {response.text}; {state}" + assert response.json() == {"deleted_teams": [team]}, response.text + assert present == [], state + assert memberships == frozenset(), state + assert key_rows == [], state + assert len(tombstones) == 1, f"tombstone rows for {team}: {len(tombstones)}; {state}" + + +def test_missing_authorization_is_401(gateway: Gateway) -> None: + response: Final = gateway.client.post("/team/delete", json={"team_ids": [UNKNOWN_TEAM]}) + assert response.status_code == 401, f"{response.status_code} {response.text}" + assert "error" in response.text.lower(), response.text + + +def test_internal_user_outside_team_is_refused_by_the_route_gate(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + outsider: Final = scenario.user(user_role="internal_user") + key: Final = scenario.key(user_id=outsider) + response: Final = _delete(gateway, [team], key=key) + assert response.status_code == 401, f"{response.status_code} {response.text}" + assert ROUTE_GATE_MESSAGE in response.text, response.text + assert len(_team_rows(team)) == 1 + + +def test_admin_of_another_team_is_refused_by_the_route_gate(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + target: Final = scenario.team() + other: Final = scenario.team() + admin: Final = _team_admin(scenario, other) + key: Final = scenario.key(team_id=other, user_id=admin) + response: Final = _delete(gateway, [target], key=key) + assert response.status_code == 401, f"{response.status_code} {response.text}" + assert ROUTE_GATE_MESSAGE in response.text, response.text + assert len(_team_rows(target)) == 1 + assert len(_team_rows(other)) == 1 + + +def test_team_admin_of_own_team_is_refused_by_the_route_gate(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + admin: Final = _team_admin(scenario, team) + key: Final = scenario.key(team_id=team, user_id=admin) + response: Final = _delete(gateway, [team], key=key) + assert response.status_code == 401, f"{response.status_code} {response.text}" + assert ROUTE_GATE_MESSAGE in response.text, response.text + assert len(_team_rows(team)) == 1 + + +def test_roster_user_whose_row_was_removed(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = _own_team(scenario) + ghost: Final = _own_user(scenario) + gateway.post("/team/member_add", {"team_id": team, "member": {"role": "user", "user_id": ghost}}) + assert ghost in _membership_user_ids(team), _membership_user_ids(team) + write_rows('DELETE FROM "LiteLLM_UserTable" WHERE user_id = %s', (ghost,)) + assert read_rows(USER_SQL, (ghost,)) == [] + response: Final = _delete(gateway, [team]) + assert response.status_code == 200, f"{response.status_code} {response.text}" + assert _team_rows(team) == [] + assert read_rows(MEMBERSHIP_SQL, (team,)) == [] + + +def test_email_only_roster_entry_matching_no_user(gateway: Gateway, record_property: RecordProperty) -> None: + with gateway.scenario() as scenario: + team: Final = _own_team(scenario) + _seed_roster( + scenario, + team, + [{"role": "user", "user_id": None, "user_email": f"nobody-{uuid.uuid4().hex}@example.com"}], + ) + response: Final = _delete(gateway, [team]) + record_property("status", response.status_code) + record_property("body", response.text) + assert response.status_code == 200, f"{response.status_code} {response.text}" + assert _team_rows(team) == [] + + +def test_email_only_roster_entry_matching_two_case_variants(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + tag: Final = uuid.uuid4().hex + upper: Final = scenario.user(user_role="internal_user") + lower: Final = scenario.user(user_role="internal_user") + # /user/new rejects a second email that matches case-insensitively, so the pair is seeded by SQL. + write_rows(USER_EMAIL_SQL, (f"Case-{tag}@example.com", upper)) + write_rows(USER_EMAIL_SQL, (f"case-{tag}@example.com", lower)) + for user in (upper, lower): + gateway.chat(model, key=scenario.key(user_id=user, models=[model])) + with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache: + + def cached() -> dict[str, int]: + return {user: int(cache.exists(user)) for user in (upper, lower)} + + assert eventually(cached, lambda seen: seen == {upper: 1, lower: 1}, seconds=10) == {upper: 1, lower: 1} + team: Final = _own_team(scenario) + _seed_roster(scenario, team, [{"role": "user", "user_id": None, "user_email": f"CASE-{tag}@EXAMPLE.COM"}]) + response: Final = _delete(gateway, [team]) + assert response.status_code == 200, f"{response.status_code} {response.text}" + assert _team_rows(team) == [] + remaining: Final = eventually( + cached, lambda seen: seen == {upper: 0, lower: 0}, seconds=10, return_last_on_timeout=True + ) + assert remaining == {upper: 0, lower: 0}, ( + f"user cache entries still present after /team/delete: " + f"{upper} exists={remaining[upper]}, {lower} exists={remaining[lower]}" + ) + + +def test_roster_entry_without_id_or_email_pins_the_500(gateway: Gateway, record_property: RecordProperty) -> None: + """Pins a pre-existing defect outside this PR's diff until it gets its own ticket: for a roster entry with neither + id nor email, LiteLLM_TeamTable.model_validate raises outside delete_team's 404 try/except, so the call is a 500 + that writes nothing (team row intact, no tombstone, membership rows untouched).""" + with gateway.scenario() as scenario: + team: Final = _own_team(scenario) + before: Final = _membership_user_ids(team) + _seed_roster(scenario, team, [{"role": "user", "user_id": None, "user_email": None}]) + response: Final = _delete(gateway, [team]) + present: Final = _team_rows(team) + tombstones: Final = read_rows(TOMBSTONE_SQL, (team,)) + after: Final = _membership_user_ids(team) + state: Final = f"team_present={bool(present)} tombstone_rows={len(tombstones)} membership_rows={len(after)}" + record_property("status", response.status_code) + record_property("body", response.text) + record_property("state_after", state) + assert response.status_code == 500, f"{response.status_code} {response.text}; {state}" + assert "Internal server error" in response.text, response.text + assert len(present) == 1, state + assert tombstones == [], state + assert after == before, f"membership rows changed: before={sorted(before)} after={sorted(after)}" + + +def test_failed_delete_leaves_unrelated_key_serving(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = scenario.key(models=[model]) + assert object_value(gateway.chat(model, key=key)["usage"])["total_tokens"] == 40 + missing: Final = f"integration-missing-{uuid.uuid4().hex}" + response: Final = _delete(gateway, [missing]) + assert response.status_code == 404, f"{response.status_code} {response.text}" + assert f"{NOT_FOUND}{missing}" in response.text, response.text + completion: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "after failed delete"}]}, + key=key, + ) + assert completion.status_code == 200, f"{completion.status_code} {completion.text}" + assert object_value(object_value(completion.json())["usage"])["total_tokens"] == 40 diff --git a/tests/integration/management/test_team_delete_large_membership.py b/tests/integration/management/test_team_delete_large_membership.py new file mode 100644 index 00000000000..5912256e5cd --- /dev/null +++ b/tests/integration/management/test_team_delete_large_membership.py @@ -0,0 +1,634 @@ +"""`/team/delete` as one locked transaction, whatever the roster size. + +The delete removes the team row, its membership rows, every member's `teams` reference and every +team key in one pass, writes one tombstone per team, evicts the cached team object and takes the +team's advisory lock (the one `/team/member_add` takes) before it writes. A roster larger than the +Prisma pool used to fail with P2028 because each member got its own transaction. +""" + +from __future__ import annotations + +import json +import os +import uuid +from collections.abc import Callable, Mapping, Sequence +from concurrent.futures import Future, ThreadPoolExecutor +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import httpx +import psycopg +import pytest +import yaml +from pydantic import JsonValue +from redis import Redis + +from tests.integration._support.client import ( + Gateway, + Scenario, + delete_key_if_present, + eventually, + object_value, + string_value, +) +from tests.integration._support.database import read_rows, write_rows +from tests.integration._support.process import owned_proxy_process + +LARGE_ROSTER: Final = 250 +POOL_LIMIT: Final = 5 +# One statement seeds the whole roster: 250 individual /user/new calls would dominate the runtime. +SEED_USERS_SQL: Final = """ +INSERT INTO "LiteLLM_UserTable" (user_id, user_role, teams, models) +SELECT %s || '-' || lpad(n::text, 3, '0'), 'internal_user', '{}'::text[], '{}'::text[] +FROM generate_series(1, %s::int) AS n +""" +# The master key's user id. `/team/new` appends the creator to the roster as an admin, so every team +# created here carries this member alongside the ones the test adds. +PROXY_ADMIN: Final = "default_user_id" +TAKE_TEAM_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext(%s))" +# Sessions blocked on the advisory lock the given backend holds, and nothing else on the shared rig. +WAITERS_ON_HELD_LOCK_SQL: Final = """ +SELECT count(*)::int AS waiting +FROM pg_locks waiter +JOIN pg_stat_activity session ON session.pid = waiter.pid +WHERE waiter.locktype = 'advisory' + AND NOT waiter.granted + AND session.wait_event_type = 'Lock' + AND session.query ILIKE %s + AND (waiter.classid, waiter.objid, waiter.objsubid) IN ( + SELECT held.classid, held.objid, held.objsubid + FROM pg_locks held + WHERE held.locktype = 'advisory' AND held.granted AND held.pid = %s::int + ) +""" + + +def _hashed(key: str) -> str: + return sha256(key.encode()).hexdigest() + + +def _team_rows(team_id: str) -> list[dict[str, JsonValue]]: + return read_rows('SELECT team_id FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team_id,)) + + +def _membership_user_ids(team_id: str) -> list[JsonValue]: + rows: Final = read_rows( + 'SELECT user_id FROM "LiteLLM_TeamMembership" WHERE team_id = %s ORDER BY user_id', (team_id,) + ) + return [row["user_id"] for row in rows] + + +def _user_teams(user_id: str) -> JsonValue: + rows: Final = read_rows('SELECT teams FROM "LiteLLM_UserTable" WHERE user_id = %s', (user_id,)) + assert len(rows) == 1, f"user row for {user_id}: {rows}" + return rows[0]["teams"] + + +def _users_referencing(team_id: str) -> list[JsonValue]: + rows: Final = read_rows( + 'SELECT user_id FROM "LiteLLM_UserTable" WHERE %s = ANY(teams) ORDER BY user_id', (team_id,) + ) + return [row["user_id"] for row in rows] + + +def _user_ids_with_prefix(prefix: str) -> list[JsonValue]: + rows: Final = read_rows( + 'SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id LIKE %s ORDER BY user_id', (f"{prefix}-%",) + ) + return [row["user_id"] for row in rows] + + +def _live_token(hashed: str) -> list[dict[str, JsonValue]]: + return read_rows('SELECT token, team_id FROM "LiteLLM_VerificationToken" WHERE token = %s', (hashed,)) + + +def _deleted_token(hashed: str) -> list[dict[str, JsonValue]]: + return read_rows('SELECT token, team_id FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s', (hashed,)) + + +def _tombstones(team_id: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT team_id, members_with_roles FROM "LiteLLM_DeletedTeamTable" WHERE team_id = %s', (team_id,) + ) + + +def _roster_user_ids(roster: JsonValue) -> list[str]: + assert isinstance(roster, list), f"roster is not a list: {roster!r}" + return sorted(string_value(object_value(member)["user_id"]) for member in roster) + + +def _waiters_on_lock_held_by(backend_pid: int) -> int: + rows: Final = read_rows(WAITERS_ON_HELD_LOCK_SQL, ("%pg_advisory_xact_lock%", str(backend_pid))) + waiting: Final = rows[0]["waiting"] + assert isinstance(waiting, int) + return waiting + + +def _remove_team_by_sql(team_id: str) -> None: + """Cleanup for a team the test expects to have deleted itself. Whatever a failed delete left behind + (row, memberships, `teams` references) goes by SQL so the shared rig stays clean without sending + another request through the proxy under test.""" + if not _team_rows(team_id): + return + write_rows('DELETE FROM "LiteLLM_TeamMembership" WHERE team_id = %s', (team_id,)) + write_rows( + 'UPDATE "LiteLLM_UserTable" SET teams = array_remove(teams, %s) WHERE %s = ANY(teams)', (team_id, team_id) + ) + write_rows('DELETE FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team_id,)) + assert _team_rows(team_id) == [] + + +def _remove_team_if_present(gateway: Gateway, team_id: str) -> None: + """Cleanup for teams the test deletes itself: the API delete first, SQL for anything it leaves.""" + if not _team_rows(team_id): + return + gateway.request("POST", "/team/delete", {"team_ids": [team_id]}) + _remove_team_by_sql(team_id) + + +def _create_team(scenario: Scenario, **fields: JsonValue) -> str: + """A team the test deletes itself; cleanup only removes it if the test left it behind.""" + created: Final = scenario.gateway.post("/team/new", {"team_alias": f"integration-{uuid.uuid4().hex}", **fields}) + team_id: Final = string_value(created["team_id"]) + scenario.cleanups.callback(_remove_team_if_present, scenario.gateway, team_id) + return team_id + + +def _generate_key(scenario: Scenario, **fields: JsonValue) -> str: + """A key the team delete is expected to remove; cleanup only deletes it if it is still live.""" + key: Final = string_value(scenario.gateway.post("/key/generate", fields)["key"]) + scenario.cleanups.callback(delete_key_if_present, scenario.gateway, key) + return key + + +def _delete_seeded_users(prefix: str) -> None: + write_rows('DELETE FROM "LiteLLM_UserTable" WHERE user_id LIKE %s', (f"{prefix}-%",)) + assert _user_ids_with_prefix(prefix) == [] + + +def _seed_users(scenario: Scenario, prefix: str, count: int) -> tuple[str, ...]: + """Insert `count` user rows in one statement; ids are `-001` … `-`.""" + users: Final = tuple(f"{prefix}-{index:03d}" for index in range(1, count + 1)) + write_rows(SEED_USERS_SQL, (prefix, str(count))) + scenario.cleanups.callback(_delete_seeded_users, prefix) + assert _user_ids_with_prefix(prefix) == list(users) + return users + + +def _bulk_member_add(gateway: Gateway, team_id: str, users: Sequence[str]) -> None: + gateway.post( + "/team/member_add", + {"team_id": team_id, "member": [{"role": "user", "user_id": user_id} for user_id in users]}, + ) + + +def _delete_teams(gateway: Gateway, team_ids: Sequence[str]) -> httpx.Response: + return gateway.request("POST", "/team/delete", {"team_ids": list(team_ids)}) + + +def _team_info(gateway: Gateway, team_id: str) -> httpx.Response: + return gateway.request("GET", "/team/info", params={"team_id": team_id}) + + +def _team_not_found_body(team_id: str) -> dict[str, JsonValue]: + """The proxy's exception handler wraps the 404 detail as an `error` object with the detail stringified.""" + return { + "error": { + "message": f"{{'message': 'Team not found, passed team id: {team_id}.'}}", + "type": "auth_error", + "param": "None", + "code": "404", + } + } + + +def _chat(gateway: Gateway, model: str, key: str) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"team delete {uuid.uuid4().hex}"}]}, + key=key, + ) + + +def _post_with_timeout(gateway: Gateway, path: str, body: Mapping[str, JsonValue], timeout: float) -> httpx.Response: + """Like `Gateway.request` with a per-call timeout longer than the client's default 15 s.""" + return gateway.client.request( + "POST", path, json=body, headers={"Authorization": f"Bearer {gateway.key}"}, timeout=timeout + ) + + +def _post_in_background( + pool: ThreadPoolExecutor, gateway: Gateway, path: str, body: Mapping[str, JsonValue] +) -> Future[httpx.Response]: + return pool.submit(_post_with_timeout, gateway, path, body, 60) + + +def _config_with_pool_limit(tmp_path: Path, pool_limit: int) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"]["database_connection_pool_limit"] = pool_limit + config["general_settings"]["database_connection_pool_timeout"] = 60 + path: Final = tmp_path / f"pool-{pool_limit}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _redis() -> Redis: + return Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) + + +def test_delete_small_team_removes_rows_keys_tombstone_and_cache( + gateway: Gateway, record_property: Callable[[str, object], None] +) -> None: + with gateway.scenario() as scenario, _redis() as cache: + model: Final = scenario.model() + users: Final = sorted(scenario.user() for _ in range(3)) + team: Final = _create_team(scenario) + _bulk_member_add(gateway, team, users) + keys: Final = tuple(_generate_key(scenario, team_id=team) for _ in range(2)) + hashed: Final = tuple(_hashed(key) for key in keys) + roster: Final = sorted([PROXY_ADMIN, *users]) + assert _membership_user_ids(team) == roster + assert _users_referencing(team) == roster + assert all(_user_teams(user) == [team] for user in users), [_user_teams(user) for user in users] + assert all(len(_live_token(digest)) == 1 for digest in hashed), hashed + + warm: Final = _chat(gateway, model, keys[0]) + assert warm.status_code == 200, warm.text + team_cache_key: Final = f"team_id:{team}" + eventually(lambda: cache.exists(team_cache_key), lambda present: present == 1, seconds=10) + record_property("redis_keys_before_delete", sorted(entry.decode() for entry in cache.keys(f"*{team}*"))) + + response: Final = _delete_teams(gateway, [team]) + assert response.status_code == 200, response.text + assert response.json() == {"deleted_teams": [team]} + + assert _team_rows(team) == [] + assert _membership_user_ids(team) == [] + assert _users_referencing(team) == [] + assert [_user_teams(user) for user in users] == [[], [], []] + assert [_live_token(digest) for digest in hashed] == [[], []] + assert [_deleted_token(digest) for digest in hashed] == [ + [{"token": hashed[0], "team_id": team}], + [{"token": hashed[1], "team_id": team}], + ] + tombstones: Final = _tombstones(team) + assert len(tombstones) == 1, tombstones + assert tombstones[0]["team_id"] == team + assert _roster_user_ids(tombstones[0]["members_with_roles"]) == roster + + info: Final = _team_info(gateway, team) + assert info.status_code == 404, info.text + assert info.json() == _team_not_found_body(team) + + assert cache.exists(team_cache_key) == 0 + record_property("redis_keys_after_delete", sorted(entry.decode() for entry in cache.keys(f"*{team}*"))) + + +@pytest.mark.timeout(240) # owned two-worker proxy boot plus a 250-member roster +def test_delete_250_member_team_succeeds_with_pool_limit_five_on_two_workers( + gateway: Gateway, tmp_path: Path, record_property: Callable[[str, object], None] +) -> None: + prefix: Final = f"integration-roster-{uuid.uuid4().hex}" + # The scenario is bound to the shared gateway and its cleanups are SQL, so a failed delete on the + # owned proxy (and whatever it does to that proxy's workers) cannot mask the assertion below with a + # second failure during cleanup. The owned proxy is stopped before the cleanups run. + with ( + gateway.scenario() as scenario, + owned_proxy_process( + gateway, + tmp_path, + {"DATABASE_URL": os.environ["DATABASE_URL"]}, + config=_config_with_pool_limit(tmp_path, POOL_LIMIT), + remove_environment=("DATABASE_URL_READ_REPLICA",), + workers=2, + ) as owned, + ): + team: Final = string_value( + owned.gateway.post("/team/new", {"team_alias": f"integration-{uuid.uuid4().hex}"})["team_id"] + ) + scenario.cleanups.callback(_remove_team_by_sql, team) + users: Final = _seed_users(scenario, prefix, LARGE_ROSTER) + added: Final = _post_with_timeout( + owned.gateway, + "/team/member_add", + {"team_id": team, "member": [{"role": "user", "user_id": user_id} for user_id in users]}, + timeout=120, + ) + assert added.status_code == 200, added.text + roster: Final = sorted([PROXY_ADMIN, *users]) + assert _membership_user_ids(team) == roster + assert _users_referencing(team) == roster + + response: Final = _post_with_timeout(owned.gateway, "/team/delete", {"team_ids": [team]}, timeout=120) + record_property("h2_delete_response", f"{response.status_code} {response.text[:300]}") + assert response.status_code == 200, ( + f"/team/delete of a {LARGE_ROSTER}-member team with database_connection_pool_limit={POOL_LIMIT}: " + f"{response.status_code} {response.text}" + ) + assert response.json() == {"deleted_teams": [team]} + assert _team_rows(team) == [] + assert _membership_user_ids(team) == [] + assert _users_referencing(team) == [] + assert _user_ids_with_prefix(prefix) == list(users) + tombstones: Final = _tombstones(team) + assert len(tombstones) == 1, tombstones + assert _roster_user_ids(tombstones[0]["members_with_roles"]) == roster + + +def test_delete_waits_for_the_team_advisory_lock_and_completes_after_release( + gateway: Gateway, record_property: Callable[[str, object], None] +) -> None: + with gateway.scenario() as scenario: + user: Final = scenario.user() + team: Final = _create_team(scenario) + _bulk_member_add(gateway, team, [user]) + hashed: Final = _hashed(_generate_key(scenario, team_id=team)) + with ThreadPoolExecutor(max_workers=1) as pool: + with psycopg.connect(os.environ["DATABASE_URL"]) as holder: + holder.execute(TAKE_TEAM_LOCK_SQL, (team,)) + pending: Final = _post_in_background(pool, gateway, "/team/delete", {"team_ids": [team]}) + eventually( + lambda: _waiters_on_lock_held_by(holder.info.backend_pid), + lambda waiting: waiting >= 1, + seconds=20, + ) + assert not pending.done(), "delete returned while the team lock was still held" + assert _team_rows(team) == [{"team_id": team}], "team row deleted while the team lock was held" + # Recorded before the count assertion so both legs document what the delete had already + # written by the time it reached the lock. + record_property( + "state_while_blocked", + json.dumps( + { + "membership_user_ids": _membership_user_ids(team), + "user_teams": _user_teams(user), + "live_token_rows": len(_live_token(hashed)), + "tombstones": len(_tombstones(team)), + } + ), + ) + # One transaction per delete: a per-member fan-out would queue one waiter per roster entry. + eventually( + lambda: _waiters_on_lock_held_by(holder.info.backend_pid), + lambda waiting: waiting == 1, + seconds=10, + ) + # leaving the holder block commits its transaction, which releases the advisory lock + response: Final = pending.result(timeout=60) + assert response.status_code == 200, response.text + assert response.json() == {"deleted_teams": [team]} + assert _team_rows(team) == [] + assert _membership_user_ids(team) == [] + assert _user_teams(user) == [] + assert _live_token(hashed) == [] + assert len(_tombstones(team)) == 1, _tombstones(team) + + +def test_deleting_two_teams_sharing_a_member_in_one_call_clears_both_from_the_member(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + user: Final = scenario.user() + first: Final = _create_team(scenario) + second: Final = _create_team(scenario) + _bulk_member_add(gateway, first, [user]) + _bulk_member_add(gateway, second, [user]) + assert _user_teams(user) == [first, second] + + response: Final = _delete_teams(gateway, [first, second]) + assert response.status_code == 200, response.text + assert response.json() == {"deleted_teams": [first, second]} + + assert _team_rows(first) == [] + assert _team_rows(second) == [] + assert _membership_user_ids(first) == [] + assert _membership_user_ids(second) == [] + assert _user_teams(user) == [] + assert [row["team_id"] for row in _tombstones(first)] == [first] + assert [row["team_id"] for row in _tombstones(second)] == [second] + + +def test_deleting_one_team_leaves_the_members_other_team_and_key_intact(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + user: Final = scenario.user() + deleted: Final = _create_team(scenario) + kept: Final = scenario.team() + _bulk_member_add(gateway, deleted, [user]) + _bulk_member_add(gateway, kept, [user]) + kept_key: Final = scenario.key(team_id=kept, user_id=user) + before: Final = _chat(gateway, model, kept_key) + assert before.status_code == 200, before.text + + response: Final = _delete_teams(gateway, [deleted]) + assert response.status_code == 200, response.text + assert response.json() == {"deleted_teams": [deleted]} + + assert _team_rows(deleted) == [] + assert _team_rows(kept) == [{"team_id": kept}] + assert _membership_user_ids(deleted) == [] + assert _membership_user_ids(kept) == [PROXY_ADMIN, user] + assert _user_teams(user) == [kept] + assert _live_token(_hashed(kept_key)) == [{"token": _hashed(kept_key), "team_id": kept}] + after: Final = _chat(gateway, model, kept_key) + assert after.status_code == 200, after.text + + +@pytest.mark.timeout(240) # owned proxy boot +def test_delete_writes_one_audit_row_for_the_team_and_one_per_key(gateway: Gateway, tmp_path: Path) -> None: + with ( + owned_proxy_process( + gateway, + tmp_path, + {"DATABASE_URL": os.environ["DATABASE_URL"], "LITELLM_STORE_AUDIT_LOGS": "true"}, + remove_environment=("DATABASE_URL_READ_REPLICA",), + ) as owned, + owned.gateway.scenario() as scenario, + ): + user: Final = scenario.user() + team: Final = _create_team(scenario) + _bulk_member_add(owned.gateway, team, [user]) + hashed: Final = _hashed(_generate_key(scenario, team_id=team)) + + response: Final = _delete_teams(owned.gateway, [team]) + assert response.status_code == 200, response.text + assert response.json() == {"deleted_teams": [team]} + assert _team_rows(team) == [] + + def deleted_audit_rows() -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT table_name, action, object_id FROM "LiteLLM_AuditLog" ' + "WHERE object_id IN (%s, %s) AND action = 'deleted' ORDER BY table_name", + (team, hashed), + ) + + rows: Final = eventually(deleted_audit_rows, lambda found: len(found) >= 2, seconds=30) + assert rows == [ + {"table_name": "LiteLLM_TeamTable", "action": "deleted", "object_id": team}, + {"table_name": "LiteLLM_VerificationToken", "action": "deleted", "object_id": hashed}, + ] + + +def test_second_delete_of_the_same_team_is_404_with_one_tombstone(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + user: Final = scenario.user() + team: Final = _create_team(scenario) + _bulk_member_add(gateway, team, [user]) + + first: Final = _delete_teams(gateway, [team]) + assert first.status_code == 200, first.text + assert first.json() == {"deleted_teams": [team]} + + second: Final = _delete_teams(gateway, [team]) + assert second.status_code == 404, second.text + assert second.json() == {"detail": {"error": f"Team not found, passed team_id={team}"}} + + assert _team_rows(team) == [] + assert [row["team_id"] for row in _tombstones(team)] == [team] + assert _user_teams(user) == [] + + +def test_delete_empty_team_writes_tombstone_and_team_info_is_404(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = _create_team(scenario) + present: Final = _team_info(gateway, team) + assert present.status_code == 200, present.text + + response: Final = _delete_teams(gateway, [team]) + assert response.status_code == 200, response.text + assert response.json() == {"deleted_teams": [team]} + + assert _team_rows(team) == [] + assert _membership_user_ids(team) == [] + assert _tombstones(team) == [ + {"team_id": team, "members_with_roles": [{"role": "admin", "user_id": PROXY_ADMIN, "user_email": None}]} + ] + info: Final = _team_info(gateway, team) + assert info.status_code == 404, info.text + assert info.json() == _team_not_found_body(team) + + +def test_delete_keys_only_team_removes_keys_and_revokes_them(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = _create_team(scenario) + keys: Final = tuple(_generate_key(scenario, team_id=team) for _ in range(2)) + hashed: Final = tuple(_hashed(key) for key in keys) + assert _membership_user_ids(team) == [PROXY_ADMIN] + for key in keys: + warm = _chat(gateway, model, key) + assert warm.status_code == 200, warm.text + + response: Final = _delete_teams(gateway, [team]) + assert response.status_code == 200, response.text + assert response.json() == {"deleted_teams": [team]} + + assert _team_rows(team) == [] + assert [_live_token(digest) for digest in hashed] == [[], []] + assert [_deleted_token(digest) for digest in hashed] == [ + [{"token": hashed[0], "team_id": team}], + [{"token": hashed[1], "team_id": team}], + ] + for key in keys: + revoked = _chat(gateway, model, key) + assert revoked.status_code == 401, f"{revoked.status_code} {revoked.text}" + + +def test_delete_three_teams_in_one_call_lists_all_and_tombstones_each_once(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + teams: Final = tuple(_create_team(scenario) for _ in range(3)) + for team in teams: + _bulk_member_add(gateway, team, [scenario.user()]) + + response: Final = _delete_teams(gateway, teams) + assert response.status_code == 200, response.text + assert response.json() == {"deleted_teams": list(teams)} + + for team in teams: + assert _team_rows(team) == [] + assert _membership_user_ids(team) == [] + assert _users_referencing(team) == [] + assert [row["team_id"] for row in _tombstones(team)] == [team] + + +def test_recreating_the_same_team_id_after_delete_serves_the_fresh_team(gateway: Gateway) -> None: + with gateway.scenario() as scenario, _redis() as cache: + model: Final = scenario.model() + original_member: Final = scenario.user() + replacement_member: Final = scenario.user() + team: Final = _create_team(scenario) + _bulk_member_add(gateway, team, [original_member]) + original_key: Final = _generate_key(scenario, team_id=team) + warm: Final = _chat(gateway, model, original_key) + assert warm.status_code == 200, warm.text + team_cache_key: Final = f"team_id:{team}" + eventually(lambda: cache.exists(team_cache_key), lambda present: present == 1, seconds=10) + + response: Final = _delete_teams(gateway, [team]) + assert response.status_code == 200, response.text + assert _team_rows(team) == [] + assert cache.exists(team_cache_key) == 0 + + fresh_alias: Final = f"integration-recreated-{uuid.uuid4().hex}" + recreated: Final = gateway.request( + "POST", + "/team/new", + { + "team_id": team, + "team_alias": fresh_alias, + "members_with_roles": [{"role": "user", "user_id": replacement_member}], + }, + ) + assert recreated.status_code == 200, recreated.text + assert recreated.json()["team_id"] == team + + info: Final = _team_info(gateway, team) + assert info.status_code == 200, info.text + team_info: Final = object_value(info.json()["team_info"]) + assert team_info["team_alias"] == fresh_alias + assert _roster_user_ids(team_info["members_with_roles"]) == [PROXY_ADMIN, replacement_member] + assert _membership_user_ids(team) == [PROXY_ADMIN, replacement_member] + assert _user_teams(replacement_member) == [team] + assert _user_teams(original_member) == [] + + fresh_key: Final = scenario.key(team_id=team) + served: Final = _chat(gateway, model, fresh_key) + assert served.status_code == 200, served.text + cached: Final = eventually(lambda: cache.get(team_cache_key), lambda value: value is not None, seconds=10) + assert isinstance(cached, bytes), cached + assert json.loads(cached)["team_alias"] == fresh_alias, cached + revoked: Final = _chat(gateway, model, original_key) + assert revoked.status_code == 401, f"{revoked.status_code} {revoked.text}" + + +def test_member_add_and_delete_released_together_leave_no_team_reference(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + newcomer: Final = scenario.user() + team: Final = _create_team(scenario) + with ThreadPoolExecutor(max_workers=2) as pool: + with psycopg.connect(os.environ["DATABASE_URL"]) as holder: + holder.execute(TAKE_TEAM_LOCK_SQL, (team,)) + pending_delete: Final = _post_in_background(pool, gateway, "/team/delete", {"team_ids": [team]}) + pending_add: Final = _post_in_background( + pool, + gateway, + "/team/member_add", + {"team_id": team, "member": {"role": "user", "user_id": newcomer}}, + ) + eventually( + lambda: _waiters_on_lock_held_by(holder.info.backend_pid), + lambda waiting: waiting == 2, + seconds=20, + ) + assert not pending_delete.done() and not pending_add.done() + # leaving the holder block commits its transaction, which releases the advisory lock + deleted: Final = pending_delete.result(timeout=60) + added: Final = pending_add.result(timeout=60) + assert deleted.status_code == 200, deleted.text + assert deleted.json() == {"deleted_teams": [team]} + assert added.status_code in (200, 404), f"{added.status_code} {added.text}" + assert _team_rows(team) == [] + assert _membership_user_ids(team) == [] + assert _user_teams(newcomer) == [] + assert _users_referencing(team) == [] + assert [row["team_id"] for row in _tombstones(team)] == [team] diff --git a/tests/integration/management/test_team_delete_member_cache_eviction.py b/tests/integration/management/test_team_delete_member_cache_eviction.py new file mode 100644 index 00000000000..0b1daa1f527 --- /dev/null +++ b/tests/integration/management/test_team_delete_member_cache_eviction.py @@ -0,0 +1,446 @@ +""" +`/team/delete` cache eviction across both proxies: member user objects, the team object and the +team's keys must stop being served by every worker once the team rows are gone. + +Auth caches the user object under the Redis key ``, the team under `team_id:` +and the key under its sha256; `enable_redis_auth_cache` is on, so Redis is the observable and +the pubsub channel carries the in-memory eviction to the peer proxy. +""" + +import asyncio +import json +import os +import uuid +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from hashlib import sha256 +from typing import Final + +import anthropic +import httpx +import openai +import pytest +from pydantic import JsonValue +from redis import Redis + +from tests.integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + delete_key_if_present, + eventually, + string_value, +) +from tests.integration._support.database import read_rows, write_rows +from tests.integration._support.wire import Reply, Request, wire_server + +_USAGE: Final = {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15} +_CACHE_KEY_HEADER: Final = "x-litellm-cache-key" + + +def _redis() -> Redis: + return Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) + + +def _cached_user(cache: Redis, user_id: str) -> dict[str, JsonValue] | None: + raw: Final = cache.get(user_id) + if raw is None: + return None + assert isinstance(raw, bytes), raw + return JSON_OBJECT.validate_json(raw) + + +def _warmed_user(cache: Redis, user_id: str) -> dict[str, JsonValue]: + """The cached user once its Redis SET has landed: auth writes memory at once but sends the Redis + SET on the request's pipeline, so the entry can trail the response that warmed it.""" + cached: Final = eventually(lambda: _cached_user(cache, user_id), lambda value: value is not None, seconds=10) + assert cached is not None + return cached + + +def _delete_team_if_present(gateway: Gateway, team_id: str) -> None: + if read_rows('SELECT team_id FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team_id,)): + gateway.post("/team/delete", {"team_ids": [team_id]}) + + +def _team(gateway: Gateway, scenario: Scenario) -> str: + """A team the test deletes itself; cleanup removes it only if the test failed before that delete.""" + created: Final = gateway.post("/team/new", {"team_alias": f"integration-{uuid.uuid4().hex}"}) + team_id: Final = string_value(created["team_id"]) + scenario.cleanups.callback(_delete_team_if_present, gateway, team_id) + return team_id + + +def _team_key(gateway: Gateway, scenario: Scenario, team_id: str, model: str) -> str: + """A key `/team/delete` removes; cleanup deletes it only if the team delete never ran.""" + created: Final = gateway.post("/key/generate", {"team_id": team_id, "models": [model]}) + token: Final = string_value(created["key"]) + scenario.cleanups.callback(delete_key_if_present, gateway, token) + return token + + +def _delete_team(gateway: Gateway, team_id: str) -> None: + deleted: Final = gateway.post("/team/delete", {"team_ids": [team_id]}) + assert deleted == {"deleted_teams": [team_id]}, deleted + assert read_rows('SELECT team_id FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team_id,)) == [] + + +def _chat_body(model: str, text: str, stream: bool = False) -> dict[str, JsonValue]: + body: dict[str, JsonValue] = {"model": model, "messages": [{"role": "user", "content": text}]} + if stream: + body["stream"] = True + return body + + +def _chat(proxy: Gateway, model: str, key: str, text: str) -> httpx.Response: + return proxy.request("POST", "/v1/chat/completions", _chat_body(model, text), key=key) + + +def _team_info(proxy: Gateway, team_id: str) -> httpx.Response: + return proxy.request("GET", "/team/info", params={"team_id": team_id}) + + +@pytest.mark.parametrize("roster_case", ("exact", "lower"), ids=("exact-case", "different-case")) +def test_team_delete_evicts_legacy_email_only_member_from_redis(gateway: Gateway, roster_case: str) -> None: + """A roster entry carrying only an email (pre-backfill legacy shape) still names a cached user; the + delete has to resolve it, in whatever case the roster stored it, and drop that user's cache entry.""" + with gateway.scenario() as scenario, _redis() as cache: + model: Final = scenario.model() + email: Final = f"Legacy-{uuid.uuid4().hex[:12]}@Example.com" + user: Final = scenario.user(user_email=email) + key: Final = scenario.key(user_id=user, models=[model]) + team: Final = _team(gateway, scenario) + roster_email: Final = email if roster_case == "exact" else email.lower() + assert (roster_email == email) is (roster_case == "exact"), (email, roster_email) + write_rows( + 'UPDATE "LiteLLM_TeamTable" SET members_with_roles = %s::jsonb WHERE team_id = %s', + (json.dumps([{"role": "user", "user_id": None, "user_email": roster_email}]), team), + ) + write_rows('UPDATE "LiteLLM_UserTable" SET teams = array_append(teams, %s) WHERE user_id = %s', (team, user)) + warm: Final = _chat(gateway, model, key, "warm legacy member " + uuid.uuid4().hex) + assert warm.status_code == 200, warm.text + warmed: Final = _warmed_user(cache, user) + assert warmed["teams"] == [team], warmed + + _delete_team(gateway, team) + + eventually(lambda: _cached_user(cache, user), lambda cached: cached is None, seconds=10) + rows: Final = read_rows('SELECT teams FROM "LiteLLM_UserTable" WHERE user_id = %s', (user,)) + assert rows == [{"teams": []}], rows + + +def test_team_delete_evicts_member_cached_on_peer_and_peer_rehydrates_without_the_team( + gateway: Gateway, peer: Gateway +) -> None: + """The peer's in-memory copy of the member is evicted over pubsub: its next request misses locally + and re-caches the user from the db, whose `teams` no longer holds the deleted team.""" + with gateway.scenario() as scenario, _redis() as cache: + model: Final = scenario.model() + user: Final = scenario.user(user_role="internal_user") + team: Final = _team(gateway, scenario) + gateway.post("/team/member_add", {"team_id": team, "member": {"role": "user", "user_id": user}}) + key: Final = scenario.key(user_id=user, models=[model]) + warm: Final = _chat(peer, model, key, "warm member on peer " + uuid.uuid4().hex) + assert warm.status_code == 200, warm.text + warmed: Final = _warmed_user(cache, user) + assert warmed["teams"] == [team], warmed + + _delete_team(gateway, team) + + eventually(lambda: _cached_user(cache, user), lambda cached: cached is None, seconds=10) + + def rehydrate() -> dict[str, JsonValue] | None: + # A peer worker still holding the stale in-memory copy answers from it and never + # rewrites Redis, so each poll issues a fresh request rather than re-reading Redis alone. + response: Final = _chat(peer, model, key, "rehydrate member on peer " + uuid.uuid4().hex) + assert response.status_code == 200, response.text + return _cached_user(cache, user) + + rehydrated: Final = eventually(rehydrate, lambda cached: cached is not None, seconds=10) + assert rehydrated is not None and rehydrated["teams"] == [], rehydrated + + +def _sse(events: Sequence[object]) -> tuple[bytes, ...]: + return tuple(b"data: " + json.dumps(event).encode() + b"\n\n" for event in events) + (b"data: [DONE]\n\n",) + + +def _chat_reply(stream: bool) -> Reply: + identity: Final = "chatcmpl-" + uuid.uuid4().hex + if not stream: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "team probe"}, "finish_reason": "stop"} + ], + "usage": _USAGE, + } + ).encode() + ) + head: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + return Reply( + content_type="text/event-stream", + chunks=_sse( + ( + {**head, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "team "}}]}, + {**head, "choices": [{"index": 0, "delta": {"content": "probe"}}]}, + {**head, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + {**head, "choices": [], "usage": _USAGE}, + ) + ), + ) + + +def _responses_reply(stream: bool) -> Reply: + identity: Final = uuid.uuid4().hex + completed: Final = { + "id": "resp_" + identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": "msg_" + identity, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "team probe", "annotations": []}], + } + ], + "parallel_tool_calls": False, + "tool_choice": "auto", + "tools": [], + "usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15}, + } + if not stream: + return Reply(body=json.dumps(completed).encode()) + events: Final = ( + {"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}}, + { + "type": "response.output_text.delta", + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": "team probe", + }, + {"type": "response.completed", "response": completed}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _upstream(request: Request) -> Reply: + stream: Final = json.loads(request.body).get("stream") is True + if request.target.endswith("/responses"): + return _responses_reply(stream) + return _chat_reply(stream) + + +def _v1(proxy: Gateway) -> str: + return str(proxy.client.base_url).rstrip("/") + "/v1" + + +def _sdk_status(error: openai.APIStatusError | anthropic.APIStatusError) -> int | str: + if isinstance(error, (openai.AuthenticationError, anthropic.AuthenticationError)): + return error.status_code + return f"{type(error).__name__}:{error.status_code}" + + +def _httpx_chat(proxy: Gateway, model: str, key: str, stream: bool, text: str) -> int | str: + return proxy.request("POST", "/v1/chat/completions", _chat_body(model, text, stream), key=key).status_code + + +def _httpx_messages(proxy: Gateway, model: str, key: str, stream: bool, text: str) -> int | str: + body: Final = {"model": model, "max_tokens": 64, "messages": [{"role": "user", "content": text}], "stream": stream} + return proxy.request("POST", "/v1/messages", body, key=key).status_code + + +def _httpx_responses(proxy: Gateway, model: str, key: str, stream: bool, text: str) -> int | str: + return proxy.request( + "POST", "/v1/responses", {"model": model, "input": text, "stream": stream}, key=key + ).status_code + + +def _openai_sync(proxy: Gateway, model: str, key: str, stream: bool, text: str) -> int | str: + with openai.OpenAI( + api_key=key, base_url=_v1(proxy), max_retries=0, http_client=httpx.Client(timeout=15, trust_env=False) + ) as client: + try: + if stream: + for _ in client.chat.completions.create( + model=model, messages=[{"role": "user", "content": text}], stream=True + ): + pass + else: + client.chat.completions.create(model=model, messages=[{"role": "user", "content": text}]) + except openai.APIStatusError as error: + return _sdk_status(error) + return 200 + + +def _openai_async(proxy: Gateway, model: str, key: str, stream: bool, text: str) -> int | str: + async def call() -> int | str: + async with openai.AsyncOpenAI( + api_key=key, base_url=_v1(proxy), max_retries=0, http_client=httpx.AsyncClient(timeout=15, trust_env=False) + ) as client: + try: + if stream: + async for _ in await client.chat.completions.create( + model=model, messages=[{"role": "user", "content": text}], stream=True + ): + pass + else: + await client.chat.completions.create(model=model, messages=[{"role": "user", "content": text}]) + except openai.APIStatusError as error: + return _sdk_status(error) + return 200 + + return asyncio.run(call()) + + +def _anthropic_sync(proxy: Gateway, model: str, key: str, stream: bool, text: str) -> int | str: + with anthropic.Anthropic( + api_key=key, + base_url=str(proxy.client.base_url), + max_retries=0, + http_client=httpx.Client(timeout=15, trust_env=False), + ) as client: + try: + if stream: + for _ in client.messages.create( + model=model, max_tokens=64, messages=[{"role": "user", "content": text}], stream=True + ): + pass + else: + client.messages.create(model=model, max_tokens=64, messages=[{"role": "user", "content": text}]) + except anthropic.APIStatusError as error: + return _sdk_status(error) + return 200 + + +@dataclass(frozen=True, slots=True) +class _Client: + name: str + call: Callable[[Gateway, str, str, bool, str], int | str] + stream: bool + + +_CLIENTS: Final = ( + _Client("httpx-chat", _httpx_chat, False), + _Client("httpx-chat-stream", _httpx_chat, True), + _Client("httpx-messages", _httpx_messages, False), + _Client("httpx-messages-stream", _httpx_messages, True), + _Client("httpx-responses", _httpx_responses, False), + _Client("httpx-responses-stream", _httpx_responses, True), + _Client("openai-sync", _openai_sync, False), + _Client("openai-sync-stream", _openai_sync, True), + _Client("openai-async", _openai_async, False), + _Client("openai-async-stream", _openai_async, True), + _Client("anthropic-sync", _anthropic_sync, False), + _Client("anthropic-sync-stream", _anthropic_sync, True), +) + + +def _observe(proxies: Mapping[str, Gateway], model: str, key: str) -> dict[str, int | str]: + """One cell per proxy and client; unique text per cell keeps the response cache out of the picture.""" + return { + f"{proxy_name}/{client.name}": client.call( + proxy, model, key, client.stream, f"{client.name} {uuid.uuid4().hex}" + ) + for proxy_name, proxy in proxies.items() + for client in _CLIENTS + } + + +def _off(observed: Mapping[str, int | str], expected: int) -> dict[str, int | str]: + return {cell: status for cell, status in observed.items() if status != expected} + + +def test_team_delete_refuses_the_team_key_for_every_client_on_both_proxies(gateway: Gateway, peer: Gateway) -> None: + """Every surface a deleted team's key can reach, on the primary and on the peer, answers 401 + once the team is gone; every cell is checked and every failing cell is reported at once.""" + proxies: Final = {"primary": gateway, "peer": peer} + with wire_server(_upstream) as upstream, gateway.scenario() as scenario: + model: Final = scenario.model(api_base=upstream.url + "/v1") + team: Final = _team(gateway, scenario) + key: Final = _team_key(gateway, scenario, team, model) + before: Final = _observe(proxies, model, key) + assert _off(before, 200) == {}, _off(before, 200) + + _delete_team(gateway, team) + + eventually( + lambda: _httpx_chat(peer, model, key, False, "deleted team key on peer"), + lambda status: status == 401, + seconds=10, + ) + after: Final = _observe(proxies, model, key) + assert _off(after, 401) == {}, _off(after, 401) + + +def test_team_delete_rejects_the_deleted_key_before_the_response_cache(gateway: Gateway) -> None: + """A request the response cache already answers for this key is refused at auth after the delete: + 401, and the upstream never sees it, so the cache-hit path cannot outlive the key.""" + with wire_server(_upstream) as upstream, gateway.scenario() as scenario: + model: Final = scenario.model(api_base=upstream.url + "/v1") + team: Final = _team(gateway, scenario) + key: Final = _team_key(gateway, scenario, team, model) + marker: Final = "cache twin " + uuid.uuid4().hex + body: Final = _chat_body(model, marker) + first: Final = gateway.request("POST", "/v1/chat/completions", body, key=key) + assert first.status_code == 200, first.text + assert first.headers.get(_CACHE_KEY_HEADER) is None, dict(first.headers) + second: Final = gateway.request("POST", "/v1/chat/completions", body, key=key) + assert second.status_code == 200, second.text + assert second.headers.get(_CACHE_KEY_HEADER), dict(second.headers) + assert second.json()["id"] == first.json()["id"], (first.text, second.text) + received: Final = upstream.drain() + assert len(received) == 1 and marker.encode() in received[0].body, received + + _delete_team(gateway, team) + + third: Final = gateway.request("POST", "/v1/chat/completions", body, key=key) + assert third.status_code == 401, third.text + assert "token_not_found_in_db" in third.text, third.text + assert upstream.drain() == (), "upstream saw a request for the deleted key" + + +def test_team_delete_evicts_team_object_and_key_on_both_proxies(gateway: Gateway, peer: Gateway) -> None: + """Team object and key warm on both proxies before the delete: `/team/info` is 404 and the key is + 401 on both afterwards, and neither the team nor the key entry is left in Redis.""" + with gateway.scenario() as scenario, _redis() as cache: + model: Final = scenario.model() + team: Final = _team(gateway, scenario) + key: Final = _team_key(gateway, scenario, team, model) + hashed: Final = sha256(key.encode()).hexdigest() + for proxy in (gateway, peer): + info: httpx.Response = _team_info(proxy, team) + assert info.status_code == 200 and info.json()["team_id"] == team, info.text + warm: httpx.Response = _chat(proxy, model, key, "warm team key " + uuid.uuid4().hex) + assert warm.status_code == 200, warm.text + # Both SETs ride the warming request's Redis pipeline and can land after its response. + eventually(lambda: cache.exists(f"team_id:{team}"), lambda present: present == 1, seconds=10) + eventually(lambda: cache.exists(hashed), lambda present: present == 1, seconds=10) + + _delete_team(gateway, team) + + eventually(lambda: _team_info(peer, team).status_code, lambda status: status == 404, seconds=10) + eventually( + lambda: _chat(peer, model, key, "deleted team key on peer").status_code, + lambda status: status == 401, + seconds=10, + ) + for proxy in (gateway, peer): + gone: httpx.Response = _team_info(proxy, team) + assert gone.status_code == 404 and "Team not found" in gone.text, gone.text + refused: httpx.Response = _chat(proxy, model, key, "deleted team key " + uuid.uuid4().hex) + assert refused.status_code == 401 and "token_not_found_in_db" in refused.text, refused.text + assert cache.exists(f"team_id:{team}") == 0, cache.keys(f"*{team}*") + assert cache.exists(hashed) == 0, cache.keys(f"*{hashed}*") diff --git a/tests/integration/management/test_team_delete_prometheus.py b/tests/integration/management/test_team_delete_prometheus.py new file mode 100644 index 00000000000..c5e383131f2 --- /dev/null +++ b/tests/integration/management/test_team_delete_prometheus.py @@ -0,0 +1,122 @@ +"""H7: the Prometheus team members gauge follows ``/team/member_add`` and ``/team/delete``. + +An owned single-worker proxy registers the ``prometheus`` callback, so ``GET /metrics/`` serves the +in-process registry (one worker, so no ``PROMETHEUS_MULTIPROC_DIR``). A team with an alias takes three +users in one bulk ``/team/member_add``; the ``litellm_team_members_metric`` series carrying that team's +id then reads 3.0. ``/team/delete`` re-emits the gauge with an empty roster instead of dropping the +series, so the same series afterwards reads 0.0. + +``disable_auto_add_proxy_admin_to_teams`` is on for the owned proxy: a master-key ``/team/new`` +otherwise seeds the roster with ``default_user_id`` and the gauge would read 4.0 after three adds. +""" + +from __future__ import annotations + +import os +import uuid +from pathlib import Path +from typing import Final + +import pytest +import yaml + +from tests.integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.process import owned_proxy_process + +METRIC: Final = "litellm_team_members_metric" +METRICS_ROUTE: Final = "/metrics/" +MEMBERS: Final = 3 +TEAM_SQL: Final = 'SELECT team_id, members_with_roles FROM "LiteLLM_TeamTable" WHERE team_id = %s' + + +def _prometheus_config(tmp_path: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["callbacks"] = ["prometheus"] + config["general_settings"]["disable_auto_add_proxy_admin_to_teams"] = True + path: Final = tmp_path / "prometheus.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _labels(text: str) -> dict[str, str]: + """``team="a",team_alias="b"`` to ``{"team": "a", "team_alias": "b"}``; ids and aliases carry no commas or quotes.""" + return {name: value.strip('"') for name, _, value in (pair.partition("=") for pair in text.split(","))} + + +def _team_members_series(scrape: str, team_id: str) -> tuple[dict[str, str], float] | None: + """The one ``litellm_team_members_metric`` sample whose ``team`` label is ``team_id``, as (labels, value).""" + samples: Final = tuple( + (labels, float(value)) + for line in scrape.splitlines() + if line.startswith(METRIC + "{") + for label_text, _, value in (line[len(METRIC) + 1 :].partition("} "),) + for labels in (_labels(label_text),) + if labels.get("team") == team_id + ) + assert len(samples) <= 1, f"{METRIC} exported more than one series for team {team_id}: {samples}" + return samples[0] if samples else None + + +def _scrape(candidate: Gateway) -> str: + response: Final = candidate.request("GET", METRICS_ROUTE) + assert response.status_code == 200, f"GET {METRICS_ROUTE}: {response.status_code} {response.text}" + return response.text + + +def _user(candidate: Gateway, scenario: Scenario) -> str: + """An internal user created through ``candidate``; its removal is registered on the shared rig.""" + user_id: Final = f"integration-h7-{uuid.uuid4().hex}" + candidate.post("/user/new", {"user_id": user_id, "auto_create_key": False, "user_role": "internal_user"}) + scenario.cleanups.callback(scenario.delete_user, user_id) + return user_id + + +def _delete_team_if_present(candidate: Gateway, team_id: str) -> None: + if read_rows(TEAM_SQL, (team_id,)): + candidate.post("/team/delete", {"team_ids": [team_id]}) + assert read_rows(TEAM_SQL, (team_id,)) == [] + + +@pytest.mark.timeout(240) # owned proxy boot (prisma db push + readiness) takes 20-40 s +def test_team_members_gauge_reads_roster_size_then_zero_after_delete(gateway: Gateway, tmp_path: Path) -> None: + with ( + gateway.scenario() as scenario, + owned_proxy_process( + gateway, + tmp_path, + {"DATABASE_URL": os.environ["DATABASE_URL"]}, + config=_prometheus_config(tmp_path), + remove_environment=("DATABASE_URL_READ_REPLICA",), + ) as owned, + ): + candidate: Final = owned.gateway + alias: Final = f"integration-h7-{uuid.uuid4().hex}" + team_id: Final = string_value(candidate.post("/team/new", {"team_alias": alias})["team_id"]) + scenario.cleanups.callback(_delete_team_if_present, gateway, team_id) + users: Final = tuple(_user(candidate, scenario) for _ in range(MEMBERS)) + candidate.post( + "/team/member_add", + {"team_id": team_id, "member": [{"role": "user", "user_id": user_id} for user_id in users]}, + ) + rows: Final = read_rows(TEAM_SQL, (team_id,)) + assert len(rows) == 1, rows + roster: Final = rows[0]["members_with_roles"] + assert isinstance(roster, list), roster + assert sorted(string_value(object_value(member)["user_id"]) for member in roster) == sorted(users), roster + + before: Final = eventually( + lambda: _team_members_series(_scrape(candidate), team_id), + lambda sample: sample is not None, + seconds=30, + ) + assert before == ({"team": team_id, "team_alias": alias}, 3.0), before + + assert candidate.post("/team/delete", {"team_ids": [team_id]}) == {"deleted_teams": [team_id]} + assert read_rows(TEAM_SQL, (team_id,)) == [] + after: Final = eventually( + lambda: _team_members_series(_scrape(candidate), team_id), + lambda sample: sample is not None and sample[1] == 0.0, + seconds=30, + ) + assert after == ({"team": team_id, "team_alias": alias}, 0.0), after 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 a53894fcd1b..f2ce01f899e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -8869,10 +8869,6 @@ async def test_delete_team_persists_deleted_teams( "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin", ) - monkeypatch.setattr( - "litellm.proxy.management_endpoints.team_endpoints._team_member_delete", - AsyncMock(return_value=(team1, (), ())), - ) data = DeleteTeamRequest(team_ids=["team-1"]) @@ -9015,6 +9011,113 @@ async def test_delete_team_sweeps_references_outside_members_with_roles( assert cache_state_when_rows_deleted["doomed_still_cached"] is True +def test_delete_team_request_collapses_repeated_ids_in_order(): + """`[T, T, U]` deletes T once and U once: one tombstone, one audit row and one eviction per team.""" + from litellm.proxy._types import DeleteTeamRequest + + assert DeleteTeamRequest(team_ids=["team-a", "team-b", "team-a", "team-b", "team-c"]).team_ids == [ + "team-a", + "team-b", + "team-c", + ] + + +@pytest.mark.asyncio +async def test_delete_team_evicts_member_caches_with_one_transaction( + monkeypatch, + disable_audit_logging_for_mocked_team, +): + """ + Regression pin for LIT-8533: `delete_team` used to fan out one + `_team_member_delete` per roster entry via `asyncio.gather`, and each opened + its own `prisma_client.tx()` and queued on the team's advisory lock, so a + team larger than the Prisma pool exhausted it and the late transactions died + on P2028. Every member-side db effect is already covered by the key delete + and the locked sweep, so the only work left is evicting each member's cache + entries, which needs no transaction at all. + """ + from litellm.proxy._types import DeleteTeamRequest + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + member_user_ids = tuple(f"member-{i}" for i in range(3)) + team = LiteLLM_TeamTable( + team_id="team-doomed", + team_alias="doomed-team", + members_with_roles=[Member(user_id=user_id, role="user") for user_id in member_user_ids] + + [ + Member(user_id=None, user_email="invitee@example.com", role="user"), + Member(user_id=None, user_email="Second.Invitee@Example.com", role="user"), + ], + metadata={}, + model_max_budget={}, + model_spend={}, + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) + mock_prisma_client.delete_data = AsyncMock(return_value={"deleted_keys": 0}) + mock_prisma_client.db.litellm_deletedteamtable.create_many = AsyncMock() + mock_prisma_client.db.litellm_deletedverificationtoken.create_many = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.execute_raw = AsyncMock() + mock_prisma_client.db.litellm_teammembership.delete_many = AsyncMock() + mock_prisma_client.db.litellm_usertable.find_many = AsyncMock( + return_value=[ + LiteLLM_UserTable(user_id="invited-user", user_email="invitee@example.com"), + LiteLLM_UserTable(user_id="second-invited-user", user_email="second.invitee@example.com"), + ] + ) + + mock_tx = AsyncMock() + mock_tx.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_tx_cm = MagicMock() + mock_tx_cm.__aenter__ = AsyncMock(return_value=mock_tx) + mock_tx_cm.__aexit__ = AsyncMock(return_value=False) + mock_prisma_client.db.tx = MagicMock(return_value=mock_tx_cm) + _wire_team_delete_tx(mock_prisma_client) + + fresh_cache = UserApiKeyCache() + for user_id in member_user_ids: + fresh_cache.set_cache(key=user_id, value=UserAPIKeyAuth(user_id=user_id)) + fresh_cache.set_cache(key="invited-user", value=UserAPIKeyAuth(user_id="invited-user")) + fresh_cache.set_cache(key="second-invited-user", value=UserAPIKeyAuth(user_id="second-invited-user")) + fresh_cache.set_cache(key="bystander-user", value=UserAPIKeyAuth(user_id="bystander-user")) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", fresh_cache) + monkeypatch.setattr("litellm.proxy.proxy_server.create_audit_log_for_update", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") + + await delete_team( + data=DeleteTeamRequest(team_ids=["team-doomed"]), + http_request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth( + user_id="admin-user", + api_key="sk-admin", + user_role=LitellmUserRoles.PROXY_ADMIN.value, + ), + litellm_changed_by="admin-user", + ) + + assert mock_prisma_client.tx.call_count == 1, ( + f"delete_team must run a single locked transaction for the whole delete, not one per member; " + f"prisma_client.tx() was entered {mock_prisma_client.tx.call_count} times for " + f"{len(member_user_ids)} members" + ) + for user_id in member_user_ids: + assert fresh_cache.get_cache(key=user_id) is None, ( + f"member {user_id}'s cached user object survived the team delete" + ) + for user_id in ("invited-user", "second-invited-user"): + assert fresh_cache.get_cache(key=user_id) is None, ( + f"the email-only roster entry resolving to {user_id} must have its cached user object evicted too" + ) + assert fresh_cache.get_cache(key="bystander-user") is not None + assert mock_prisma_client.db.litellm_usertable.find_many.await_count == 1, ( + "email-only roster entries must resolve in one lookup, not one query per email" + ) + + @pytest.mark.asyncio async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes( monkeypatch, @@ -14146,12 +14249,6 @@ async def test_delete_team_emits_only_the_deleted_audit_event(monkeypatch): monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") - removals = [(team, members, members[1:]), (team, members[1:], ())] - monkeypatch.setattr( - "litellm.proxy.management_endpoints.team_endpoints._team_member_delete", - AsyncMock(side_effect=lambda **_kwargs: removals.pop(0)), - ) - await delete_team( data=DeleteTeamRequest(team_ids=["team-gone"]), http_request=MagicMock(), diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py index e185d95ffb8..bd0f194b326 100644 --- a/tests/unit/repositories/test_repositories.py +++ b/tests/unit/repositories/test_repositories.py @@ -18,6 +18,7 @@ from litellm.models.credentials import CredentialItem from litellm.models.team import LiteLLM_TeamTable from litellm.repositories.base_repository import BaseRepository from litellm.repositories.budget_repository import BudgetRepository +from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.credentials_repository import CredentialsRepository from litellm.repositories.model_repository import ModelRepository @@ -891,6 +892,32 @@ class TestUserRepository: user = await repo.find_by_email("test@example.com") assert user is not None + @pytest.mark.asyncio + async def test_find_by_emails_is_one_case_insensitive_query(self, repo): + repo._prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[]) + await repo.find_by_emails(["B@Example.com", "a@example.com", "B@Example.com"]) + repo._prisma_client.db.litellm_usertable.find_many.assert_awaited_once() + where = repo._prisma_client.db.litellm_usertable.find_many.await_args.kwargs["where"] + assert where["user_email"] == {"in": ["B@Example.com", "a@example.com"], "mode": "insensitive"} + + @pytest.mark.asyncio + async def test_find_by_emails_slices_the_list_into_bounded_statements(self, repo): + repo._prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[]) + emails = [f"user{index}@example.com" for index in range(IN_LIST_CHUNK_SIZE + 1)] + await repo.find_by_emails(emails) + assert repo._prisma_client.db.litellm_usertable.find_many.await_count == 2 + sizes = [ + len(call.kwargs["where"]["user_email"]["in"]) + for call in repo._prisma_client.db.litellm_usertable.find_many.await_args_list + ] + assert sizes == [IN_LIST_CHUNK_SIZE, 1] + + @pytest.mark.asyncio + async def test_find_by_emails_skips_the_query_for_no_emails(self, repo): + repo._prisma_client.db.litellm_usertable.find_many = AsyncMock() + assert await repo.find_by_emails(()) == () + repo._prisma_client.db.litellm_usertable.find_many.assert_not_awaited() + @pytest.mark.asyncio async def test_find_by_sso_id(self, repo): repo._prisma_client.db.litellm_usertable._records["sso-123"] = {