From de4b520153718864dddfa53a37dd6b56825c958b Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:12:56 +0000 Subject: [PATCH 01/11] fix(proxy): track team member spend when the member has no budget add_new_member only wrote a LiteLLM_TeamMembership row when a budget id resolved, and the spend writer used update_many so a missing row failed silently. Members without a budget therefore never accrued per-member spend. The membership row is now always upserted (budget_id NULL when no budget applies) and the spend write is an upsert so members added before this fix start accruing on their next request. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 16 +++- litellm/proxy/management_helpers/utils.py | 23 +++--- litellm/repositories/prisma_protocols.py | 2 + .../proxy/db/test_db_spend_update_writer.py | 37 ++++++--- .../test_team_endpoints.py | 10 +-- .../test_management_helpers_utils.py | 75 ++++++++++++++----- 6 files changed, 114 insertions(+), 49 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index c13b852484e..913675705e3 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1693,11 +1693,19 @@ class DBSpendUpdateWriter: team_id = key.split("::")[1] user_id = key.split("::")[3] - batcher.litellm_teammembership.update_many( # 'update_many' prevents error from being raised if no row exists - where={"team_id": team_id, "user_id": user_id}, + batcher.litellm_teammembership.upsert( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, data={ - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, + "create": { + "team_id": team_id, + "user_id": user_id, + "spend": response_cost, + "total_spend": response_cost, + }, + "update": { + "spend": {"increment": response_cost}, + "total_spend": {"increment": response_cost}, + }, }, ) # Transaction succeeded, break out of retry loop diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index f3bd4b0f6dd..d940eba86d3 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -86,7 +86,9 @@ class _PrismaUserTable(Protocol): class _PrismaTeamMembershipTable(Protocol): """Team membership table actions the management helpers issue.""" - async def create(self, *, data: Mapping[str, object], include: Mapping[str, bool]) -> _PrismaRecord: ... + async def upsert( + self, *, where: Mapping[str, object], data: Mapping[str, Mapping[str, object]], include: Mapping[str, bool] + ) -> _PrismaRecord: ... class MemberWriteTx(Protocol): @@ -348,7 +350,7 @@ async def _resolve_member_budget_id( default member budget is cloned (with ``budget_duration`` overriding its reset window while keeping its other limits). A lone ``budget_duration`` with no team default creates a window-only budget. With nothing set the - member gets no budget. + member gets no budget, though ``add_new_member`` still writes its membership row. """ has_explicit_limit: Final = max_budget_in_team is not None or allowed_models is not None @@ -415,9 +417,9 @@ async def add_new_member( Add a new member to a team - add team id to user table - - add team member w/ budget to team member table + - add team member to team member table, linked to a budget when one resolves - Returns created/existing user + team membership w/ budget id + Returns created/existing user + team membership (``budget_id`` is ``None`` when no budget applies) Callers already inside a transaction pass it as ``tx`` so every write here runs on that connection instead of borrowing more from the pool while the caller's locks are held. @@ -471,14 +473,13 @@ async def add_new_member( tx=tx, ) - if _budget_id and returned_user is not None and returned_user.user_id is not None: + if returned_user is not None and returned_user.user_id is not None: membership_table: Final[_PrismaTeamMembershipTable] = _team_membership_table(prisma_client, tx) - _returned_team_membership: Final = await membership_table.create( - data={ - "team_id": team_id, - "user_id": returned_user.user_id, - "budget_id": _budget_id, - }, + membership_key: Final[Mapping[str, object]] = {"user_id": returned_user.user_id, "team_id": team_id} + budget_link: Final[Mapping[str, object]] = {"budget_id": _budget_id} if _budget_id is not None else {} + _returned_team_membership: Final = await membership_table.upsert( + where={"user_id_team_id": membership_key}, + data={"create": {**membership_key, **budget_link}, "update": {**budget_link}}, include={"litellm_budget_table": True}, ) diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 93b8c5c7cd7..0919ae9f808 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -123,6 +123,8 @@ class BatchTable(Protocol): def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... + def upsert(self, *, where: Mapping[str, object], data: Mapping[str, Mapping[str, object]]) -> None: ... + class PrismaBatch(Protocol): @property diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index b3f5a60877d..bed903d7a1f 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -918,7 +918,11 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total """ Verify that _commit_spend_updates_to_db increments BOTH spend (cycle-scoped) and total_spend (non-resetting) on LiteLLM_TeamMembership in a single - update_many call, using the same response_cost. + upsert call, using the same response_cost. + + Regression (LIT-5502): members added without a budget had no membership row, and + the previous update_many matched zero rows, so their spend was silently dropped. + The upsert has to create the row seeded with this call's cost in that case. """ db_writer = DBSpendUpdateWriter() @@ -930,7 +934,7 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total mock_batcher.litellm_teamtable = MagicMock() mock_batcher.litellm_teamtable.update_many = MagicMock() mock_batcher.litellm_teammembership = MagicMock() - mock_batcher.litellm_teammembership.update_many = MagicMock() + mock_batcher.litellm_teammembership.upsert = MagicMock() mock_batcher.litellm_organizationtable = MagicMock() mock_batcher.litellm_organizationtable.update_many = MagicMock() mock_batcher.litellm_tagtable = MagicMock() @@ -979,12 +983,21 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total db_spend_update_transactions=db_spend_update_transactions, ) - mock_batcher.litellm_teammembership.update_many.assert_called_once() - call_kwargs = mock_batcher.litellm_teammembership.update_many.call_args[1] - assert call_kwargs["where"] == {"team_id": team_id, "user_id": user_id} + mock_batcher.litellm_teammembership.upsert.assert_called_once() + mock_batcher.litellm_teammembership.update_many.assert_not_called() + call_kwargs = mock_batcher.litellm_teammembership.upsert.call_args.kwargs + assert call_kwargs["where"] == {"user_id_team_id": {"user_id": user_id, "team_id": team_id}} assert call_kwargs["data"] == { - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, + "create": { + "team_id": team_id, + "user_id": user_id, + "spend": response_cost, + "total_spend": response_cost, + }, + "update": { + "spend": {"increment": response_cost}, + "total_spend": {"increment": response_cost}, + }, } @@ -2219,9 +2232,13 @@ async def test_commit_daily_tag_spend_no_requeue_on_success(): "team_id::team_b::user_id::user_x": 0.3, }, "litellm_teammembership", - "update_many", - "team_id", - ["team_a", "team_b", "team_c"], + "upsert", + "user_id_team_id", + [ + {"user_id": "user_x", "team_id": "team_a"}, + {"user_id": "user_x", "team_id": "team_b"}, + {"user_id": "user_x", "team_id": "team_c"}, + ], id="team_member", ), pytest.param( 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 4f5c5066367..f7157da3100 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -2086,7 +2086,7 @@ async def test_add_team_members_runs_member_writes_on_the_lock_holding_transacti tx.litellm_usertable.upsert = AsyncMock(return_value=added_user) tx.litellm_usertable.update_many = AsyncMock() tx.litellm_budgettable.create = AsyncMock(return_value=created_budget) - tx.litellm_teammembership.create = AsyncMock(return_value=membership) + tx.litellm_teammembership.upsert = AsyncMock(return_value=membership) tx_cm = MagicMock() tx_cm.__aenter__ = AsyncMock(return_value=tx) @@ -5772,7 +5772,7 @@ async def test_new_team_max_budget_within_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( + mock_prisma.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_membership ) @@ -5915,7 +5915,7 @@ async def test_new_team_org_scoped_budget_bypasses_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( + mock_prisma.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_membership ) @@ -6063,7 +6063,7 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( + mock_prisma.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_membership ) @@ -9525,7 +9525,7 @@ async def test_new_team_soft_budget_validation( "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( + mock_prisma.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_membership ) diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index a6b1fc32eda..0dcf7e1b8ad 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -234,7 +234,7 @@ async def test_add_new_member_clones_default_team_budget_id(): "budget_id": test_cloned_budget_id, "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -257,7 +257,7 @@ async def test_add_new_member_clones_default_team_budget_id(): assert result_team_membership.budget_id != test_default_budget_id mock_prisma_client.db.litellm_usertable.upsert.assert_called_once() - mock_prisma_client.db.litellm_teammembership.create.assert_called_once() + mock_prisma_client.db.litellm_teammembership.upsert.assert_called_once() # The clone must have happened: find_unique on the default, create for the clone. mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with( @@ -274,9 +274,9 @@ async def test_add_new_member_clones_default_team_budget_id(): assert cloned_create_data["created_by"] == user_api_key_dict.user_id team_membership_call_args = ( - mock_prisma_client.db.litellm_teammembership.create.call_args + mock_prisma_client.db.litellm_teammembership.upsert.call_args ) - create_data = team_membership_call_args.kwargs["data"] + create_data = team_membership_call_args.kwargs["data"]["create"] assert create_data["budget_id"] == test_cloned_budget_id @@ -332,7 +332,7 @@ async def test_add_new_member_budget_duration_only_clones_default_max_budget(): "budget_id": "cloned-dc", "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -362,7 +362,8 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget(): Test that add_new_member links no budget to the team membership when neither max_budget_in_team nor default_team_budget_id is provided. - When the team has no default member budget, new members get nothing. + When the team has no default member budget, no budget row is created, but the + membership row still is, otherwise the member's spend has nowhere to accrue. """ from litellm.proxy._types import LitellmUserRoles @@ -393,7 +394,19 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget(): # Even though we mock these, they must NOT be called on the no-budget path. mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock() mock_prisma_client.db.litellm_budgettable.create = AsyncMock() - mock_prisma_client.db.litellm_teammembership.create = AsyncMock() + + mock_team_membership_response = MagicMock() + mock_team_membership_response.model_dump.return_value = { + "team_id": test_team_id, + "user_id": test_user_id, + "budget_id": None, + "spend": 0.0, + "total_spend": 0.0, + "litellm_budget_table": None, + } + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( + return_value=mock_team_membership_response + ) result_user, result_team_membership = await add_new_member( new_member=new_member, @@ -408,11 +421,20 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget(): assert result_user is not None assert result_user.user_id == test_user_id - # No budget id, so no team membership row is created. - assert result_team_membership is None mock_prisma_client.db.litellm_budgettable.find_unique.assert_not_called() mock_prisma_client.db.litellm_budgettable.create.assert_not_called() - mock_prisma_client.db.litellm_teammembership.create.assert_not_called() + + # Regression (LIT-5502): the membership row is what per-member spend increments land on, + # so it has to exist even when the member has no budget. Skipping it silently dropped spend. + assert result_team_membership is not None + assert result_team_membership.budget_id is None + mock_prisma_client.db.litellm_teammembership.upsert.assert_awaited_once() + upsert_kwargs = mock_prisma_client.db.litellm_teammembership.upsert.call_args.kwargs + assert upsert_kwargs["where"] == { + "user_id_team_id": {"user_id": test_user_id, "team_id": test_team_id} + } + assert upsert_kwargs["data"]["create"] == {"user_id": test_user_id, "team_id": test_team_id} + assert "budget_id" not in upsert_kwargs["data"]["update"] @pytest.mark.asyncio @@ -473,7 +495,7 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): "budget_id": test_new_budget_id, "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -502,10 +524,10 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): # Verify the team membership was created with the correct budget_id team_membership_call_args = ( - mock_prisma_client.db.litellm_teammembership.create.call_args + mock_prisma_client.db.litellm_teammembership.upsert.call_args ) assert team_membership_call_args is not None - create_data = team_membership_call_args.kwargs["data"] + create_data = team_membership_call_args.kwargs["data"]["create"] assert create_data["budget_id"] == test_new_budget_id @@ -546,7 +568,7 @@ async def test_add_new_member_persists_budget_duration(): "budget_id": "budget-dur", "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -610,7 +632,7 @@ async def test_add_new_member_persists_budget_duration_without_max_budget(): "budget_id": "budget-dur2", "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -700,7 +722,7 @@ async def test_add_new_member_with_user_email_clones_default_budget(): "budget_id": test_cloned_budget_id, "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -1031,8 +1053,15 @@ async def test_add_new_member_appends_team_only_if_absent_for_existing_user(): } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_after) mock_prisma_client.db.litellm_usertable.update_many = AsyncMock() - # no team default budget and no explicit budget -> no team membership row mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) + mock_membership = MagicMock() + mock_membership.model_dump.return_value = { + "team_id": "team-1", + "user_id": "existing-user", + "budget_id": None, + "litellm_budget_table": None, + } + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(return_value=mock_membership) result_user, _ = await add_new_member( new_member=new_member, @@ -1099,6 +1128,14 @@ async def test_add_new_member_creates_missing_user_atomically_via_upsert(): mock_prisma_client.db.litellm_usertable.update_many = AsyncMock() mock_prisma_client.db.litellm_usertable.create = AsyncMock() mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) + mock_membership = MagicMock() + mock_membership.model_dump.return_value = { + "team_id": "team-1", + "user_id": "brand-new-user", + "budget_id": None, + "litellm_budget_table": None, + } + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(return_value=mock_membership) result_user, _ = await add_new_member( new_member=new_member, @@ -1147,7 +1184,7 @@ def _member_write_tx() -> MagicMock: tx.litellm_usertable.find_many = AsyncMock(return_value=[]) tx.litellm_budgettable.find_unique = AsyncMock(return_value=None) tx.litellm_budgettable.create = AsyncMock(return_value=created_budget) - tx.litellm_teammembership.create = AsyncMock(return_value=membership) + tx.litellm_teammembership.upsert = AsyncMock(return_value=membership) return tx @@ -1192,7 +1229,7 @@ async def test_add_new_member_runs_every_write_on_the_caller_transaction(new_mem assert result_membership.budget_id == "budget-pool" assert tx.litellm_budgettable.create.await_count == 1 - assert tx.litellm_teammembership.create.await_count == 1 + assert tx.litellm_teammembership.upsert.await_count == 1 assert tx.litellm_usertable.upsert.await_count + tx.litellm_usertable.create.await_count == 1 prisma_client.db.assert_not_called() From d9b48ac9414adbfbcf5c5f15521a7cb4293d278d Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:35:53 +0000 Subject: [PATCH 02/11] test(proxy): mock team membership upsert in team admin member add test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/proxy_unit_tests/test_proxy_server.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 1fcdaa67143..659097abbb9 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -1372,6 +1372,7 @@ async def test_create_team_member_add_team_admin( from fastapi import Request from litellm.proxy._types import ( + LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, Member, @@ -1454,6 +1455,10 @@ async def test_create_team_member_add_team_admin( team_mock_client.update = AsyncMock( return_value=LiteLLM_TeamTableCachedObj(team_id="1234") ) + membership_mock_client = AsyncMock() + membership_mock_client.upsert = AsyncMock( + return_value=LiteLLM_TeamMembership(user_id="1234", team_id=_team_id) + ) tx_cm = _member_add_tx_cm(team_mock_client) @@ -1463,6 +1468,11 @@ async def test_create_team_member_add_team_admin( "litellm_teamtable", team_mock_client, ), + patch.object( # test-quality-ok: legacy test swaps the prisma table on the module-level client + litellm.proxy.proxy_server.prisma_client.db, + "litellm_teammembership", + membership_mock_client, + ), patch.object( litellm.proxy.proxy_server.prisma_client, "tx", From 6dce85c7285e7e7a6ea1b0b450ff4041815470ba Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 02:10:29 +0000 Subject: [PATCH 03/11] fix(proxy): skip recreating membership rows for members removed before a spend flush Take the team advisory lock in the spend flush transaction and read the roster through it, so a delayed flush after /team/member_delete cannot recreate the deleted LiteLLM_TeamMembership row. TEAM_ADVISORY_LOCK_SQL moves to team_repository so the spend writer can import it without a circular import Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 72 +++++++--- .../access_group_team_sync.py | 8 +- litellm/repositories/prisma_protocols.py | 8 +- litellm/repositories/team_repository.py | 12 +- .../proxy/db/test_db_spend_update_writer.py | 134 ++++++++++++------ 5 files changed, 160 insertions(+), 74 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 913675705e3..5c86c52c539 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -75,7 +75,8 @@ from litellm.proxy.spend_tracking.savings import ( marks_gateway_injection, ) from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error -from litellm.repositories.prisma_protocols import BatchTable +from litellm.repositories.prisma_protocols import BatchTable, RawQueryTransaction +from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL, TeamRepository from litellm.types.utils import CallTypes if TYPE_CHECKING: @@ -133,7 +134,7 @@ class _SpendBatchManager(Protocol): async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... -class _SpendTransaction(Protocol): +class _SpendTransaction(RawQueryTransaction, Protocol): def batch_(self) -> _SpendBatchManager: ... @@ -161,6 +162,50 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager: return tx +async def _lock_and_read_rosters( + prisma_client: PrismaClient, transaction: _SpendTransaction, team_ids: Sequence[str] +) -> frozenset[tuple[str, str]]: + """Take each team's advisory lock on ``transaction`` and return its rostered (user_id, team_id) pairs. + + A spend flush can land after ``/team/member_delete`` removed the member. Holding the same + lock that endpoint takes, until this transaction commits, means a member read here is on + the team for the whole flush, so only they may have a missing membership row created. + """ + repository: Final = TeamRepository(prisma_client) + + async def locked_roster(team_id: str) -> tuple[tuple[str, str], ...]: + await transaction.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id) + roster: Final = await repository.get_members_with_roles_locked(transaction, team_id) + return tuple((member.user_id, team_id) for member in roster or () if member.user_id is not None) + + rosters: Final = tuple([await locked_roster(team_id) for team_id in sorted(frozenset(team_ids))]) + return frozenset(pair for roster in rosters for pair in roster) + + +def _queue_team_member_spend( + memberships: BatchTable, user_id: str, team_id: str, response_cost: float, rostered: bool +) -> None: + increments: Final = { + "spend": {"increment": response_cost}, + "total_spend": {"increment": response_cost}, + } + if not rostered: + memberships.update_many(where={"team_id": team_id, "user_id": user_id}, data=increments) + return + memberships.upsert( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, + data={ + "create": { + "team_id": team_id, + "user_id": user_id, + "spend": response_cost, + "total_spend": response_cost, + }, + "update": increments, + }, + ) + + def get_llm_router(): """The proxy's router, or None outside a running proxy. @@ -1685,6 +1730,9 @@ class DBSpendUpdateWriter: start_time = time.time() try: async with _spend_update_tx(prisma_client) as transaction: + rostered_members = await _lock_and_read_rosters( + prisma_client, transaction, tuple(team_id for _, team_id in team_memberships_to_invalidate) + ) async with transaction.batch_() as batcher: # Sort by composite key for consistent lock ordering across pods to prevent deadlocks. # Key format "team_id::::user_id::" makes the string sort equivalent to sorting by (team_id, user_id). @@ -1693,20 +1741,12 @@ class DBSpendUpdateWriter: team_id = key.split("::")[1] user_id = key.split("::")[3] - batcher.litellm_teammembership.upsert( - where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, - data={ - "create": { - "team_id": team_id, - "user_id": user_id, - "spend": response_cost, - "total_spend": response_cost, - }, - "update": { - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, - }, - }, + _queue_team_member_spend( + batcher.litellm_teammembership, + user_id, + team_id, + response_cost, + (user_id, team_id) in rostered_members, ) # Transaction succeeded, break out of retry loop break diff --git a/litellm/proxy/management_helpers/access_group_team_sync.py b/litellm/proxy/management_helpers/access_group_team_sync.py index 664e36c9f10..cfe207e66ea 100644 --- a/litellm/proxy/management_helpers/access_group_team_sync.py +++ b/litellm/proxy/management_helpers/access_group_team_sync.py @@ -21,13 +21,7 @@ from typing import Final, Protocol from pydantic import BaseModel, TypeAdapter from litellm.proxy.auth.auth_checks import _delete_cache_access_object - -# hashtext collisions only cost two unrelated teams a little serialization, and the -# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, -# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints -# reuses this exact statement to serialize /team/member_add and /team/delete against each -# other and against this mirror, rather than defining a second, divergent lock on the same key. -TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" +from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL _READ_TEAM_SQL: Final = 'SELECT access_group_ids FROM "LiteLLM_TeamTable" WHERE team_id = $1' diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 0919ae9f808..c742f3f5f3f 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -7,7 +7,7 @@ private ones per file. """ from collections.abc import Mapping, Sequence -from typing import Protocol, TypeVar +from typing import LiteralString, Protocol, TypeVar RowT_co = TypeVar("RowT_co", covariant=True) @@ -108,6 +108,12 @@ class PrismaRecord(Protocol): def dict(self) -> Mapping[str, object]: ... +class RawQueryTransaction(Protocol): + """A prisma transaction handle that can run raw SQL, e.g. an advisory lock or a locked read.""" + + async def query_raw(self, query: LiteralString, *args: str) -> Sequence[Mapping[str, object]]: ... + + class ReadOnlyTable(Protocol): async def find_many(self, *, where: Mapping[str, object]) -> Sequence[PrismaRecord]: ... diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 5ff07d76b5d..1b9cba48e41 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -15,10 +15,9 @@ from litellm.repositories.base_repository import ( DbRecord, record_to_dict, ) -from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.prisma_protocols import RawQueryTransaction, TableActions if TYPE_CHECKING: - from prisma import Prisma from prisma import models as prisma_models @@ -40,6 +39,13 @@ def _team_arrays(team: LiteLLM_TeamTable) -> _TeamArrays: return team +# hashtext collisions only cost two unrelated teams a little serialization, and the +# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, +# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints, +# the access-group mirror and the team member spend flush all reuse this exact statement to +# serialize against each other, rather than defining a second, divergent lock on the same key. +TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" + _MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member]) _JSON_ENCODED_TEAM_FIELDS: Final = ( "metadata", @@ -78,7 +84,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return LiteLLM_TeamTable.model_validate(data) - async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> list[Member] | None: + async def get_members_with_roles_locked(self, tx: RawQueryTransaction, team_id: str) -> list[Member] | None: """Return the team's members_with_roles. The caller must already hold ``TEAM_ADVISORY_LOCK_SQL`` for this team_id on ``tx`` before calling this. diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index bed903d7a1f..b12dc68b7b1 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -20,6 +20,7 @@ from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( build_window_spend_transaction, ) +from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL @pytest.mark.asyncio @@ -913,38 +914,22 @@ async def test_commit_spend_updates_to_db_increments_agent_spend(): assert call_kwargs["data"] == {"spend": {"increment": response_cost}} -@pytest.mark.asyncio -async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total_spend(): - """ - Verify that _commit_spend_updates_to_db increments BOTH spend (cycle-scoped) - and total_spend (non-resetting) on LiteLLM_TeamMembership in a single - upsert call, using the same response_cost. - - Regression (LIT-5502): members added without a budget had no membership row, and - the previous update_many matched zero rows, so their spend was silently dropped. - The upsert has to create the row seeded with this call's cost in that case. - """ - db_writer = DBSpendUpdateWriter() - +def _team_member_flush_fixtures(team_id: str, rostered_user_ids: list[str]) -> tuple[MagicMock, AsyncMock, MagicMock]: + """A batcher, transaction and prisma client whose locked roster read for `team_id` lists `rostered_user_ids`.""" mock_batcher = MagicMock() - mock_batcher.litellm_verificationtoken = MagicMock() - mock_batcher.litellm_verificationtoken.update_many = MagicMock() - mock_batcher.litellm_usertable = MagicMock() - mock_batcher.litellm_usertable.update_many = MagicMock() - mock_batcher.litellm_teamtable = MagicMock() - mock_batcher.litellm_teamtable.update_many = MagicMock() mock_batcher.litellm_teammembership = MagicMock() mock_batcher.litellm_teammembership.upsert = MagicMock() - mock_batcher.litellm_organizationtable = MagicMock() - mock_batcher.litellm_organizationtable.update_many = MagicMock() - mock_batcher.litellm_tagtable = MagicMock() - mock_batcher.litellm_tagtable.update_many = MagicMock() - mock_batcher.litellm_agentstable = MagicMock() - mock_batcher.litellm_agentstable.update_many = MagicMock() + mock_batcher.litellm_teammembership.update_many = MagicMock() + + roster_row = {"members_with_roles": json.dumps([{"user_id": uid, "role": "user"} for uid in rostered_user_ids])} + + async def query_raw(query: str, *args: str) -> list[dict[str, object]]: + return [] if query == TEAM_ADVISORY_LOCK_SQL else [roster_row] mock_transaction = AsyncMock() mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction) mock_transaction.__aexit__ = AsyncMock(return_value=False) + mock_transaction.query_raw = AsyncMock(side_effect=query_raw) mock_transaction.batch_ = MagicMock( return_value=AsyncMock( __aenter__=AsyncMock(return_value=mock_batcher), @@ -955,16 +940,11 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total mock_prisma_client = MagicMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction) + return mock_batcher, mock_transaction, mock_prisma_client - mock_proxy_logging = MagicMock() - # Skip team-membership cache invalidation — out of scope for this test. - mock_proxy_logging.call_details.get = MagicMock(return_value=None) - team_id = "team-abc" - user_id = "user-xyz" - response_cost = 0.75 - entity_id = f"team_id::{team_id}::user_id::{user_id}" - db_spend_update_transactions = { +def _team_member_only_transactions(entity_id: str, response_cost: float) -> dict[str, dict[str, float]]: + return { "user_list_transactions": {}, "end_user_list_transactions": {}, "key_list_transactions": {}, @@ -975,14 +955,38 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total "agent_list_transactions": {}, } - with patch("litellm.proxy.utils._raise_failed_update_spend_exception"): - await db_writer._commit_spend_updates_to_db( - prisma_client=mock_prisma_client, - n_retry_times=0, - proxy_logging_obj=mock_proxy_logging, - db_spend_update_transactions=db_spend_update_transactions, - ) +@pytest.mark.asyncio +async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total_spend(): + """ + Verify that _commit_spend_updates_to_db increments BOTH spend (cycle-scoped) + and total_spend (non-resetting) on LiteLLM_TeamMembership in a single + upsert call, using the same response_cost. + + Regression (LIT-5502): members added without a budget had no membership row, and + the previous update_many matched zero rows, so their spend was silently dropped. + For a member still on the team roster the upsert has to create the row seeded + with this call's cost in that case. + """ + db_writer = DBSpendUpdateWriter() + team_id = "team-abc" + user_id = "user-xyz" + response_cost = 0.75 + mock_batcher, mock_transaction, mock_prisma_client = _team_member_flush_fixtures(team_id, [user_id]) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.call_details.get = MagicMock(return_value=None) + + await db_writer._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=0, + proxy_logging_obj=mock_proxy_logging, + db_spend_update_transactions=_team_member_only_transactions( + f"team_id::{team_id}::user_id::{user_id}", response_cost + ), + ) + + assert mock_transaction.query_raw.await_args_list[0] == call(TEAM_ADVISORY_LOCK_SQL, team_id) mock_batcher.litellm_teammembership.upsert.assert_called_once() mock_batcher.litellm_teammembership.update_many.assert_not_called() call_kwargs = mock_batcher.litellm_teammembership.upsert.call_args.kwargs @@ -1001,6 +1005,43 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total } +@pytest.mark.asyncio +async def test_commit_spend_updates_to_db_does_not_recreate_membership_of_removed_team_member(): + """ + A spend flush that lands after /team/member_delete must not resurrect the deleted + membership row: a user missing from the team roster, read under the team's advisory + lock inside the flush transaction, only gets an increment on whatever row still + exists, never a create. + """ + db_writer = DBSpendUpdateWriter() + team_id = "team-abc" + removed_user_id = "user-removed" + response_cost = 0.75 + mock_batcher, mock_transaction, mock_prisma_client = _team_member_flush_fixtures(team_id, ["user-still-here"]) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.call_details.get = MagicMock(return_value=None) + + await db_writer._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=0, + proxy_logging_obj=mock_proxy_logging, + db_spend_update_transactions=_team_member_only_transactions( + f"team_id::{team_id}::user_id::{removed_user_id}", response_cost + ), + ) + + assert mock_transaction.query_raw.await_args_list[0] == call(TEAM_ADVISORY_LOCK_SQL, team_id) + mock_batcher.litellm_teammembership.upsert.assert_not_called() + mock_batcher.litellm_teammembership.update_many.assert_called_once_with( + where={"team_id": team_id, "user_id": removed_user_id}, + data={ + "spend": {"increment": response_cost}, + "total_spend": {"increment": response_cost}, + }, + ) + + @pytest.mark.asyncio async def test_org_spend_increments_organization_membership_row_for_the_calling_user(): """A request made with a user_id inside an org must increment that user's @@ -2232,13 +2273,9 @@ async def test_commit_daily_tag_spend_no_requeue_on_success(): "team_id::team_b::user_id::user_x": 0.3, }, "litellm_teammembership", - "upsert", - "user_id_team_id", - [ - {"user_id": "user_x", "team_id": "team_a"}, - {"user_id": "user_x", "team_id": "team_b"}, - {"user_id": "user_x", "team_id": "team_c"}, - ], + "update_many", + "team_id", + ["team_a", "team_b", "team_c"], id="team_member", ), pytest.param( @@ -2312,6 +2349,8 @@ async def test_commit_spend_updates_iterates_in_sorted_order( ) ) + mock_transaction.query_raw = AsyncMock(return_value=[]) + mock_prisma_client = MagicMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction) @@ -3071,6 +3110,7 @@ def _good_tx(mock_batcher): tx = AsyncMock() tx.__aenter__ = AsyncMock(return_value=tx) tx.__aexit__ = AsyncMock(return_value=False) + tx.query_raw = AsyncMock(return_value=[]) tx.batch_ = MagicMock( return_value=AsyncMock( __aenter__=AsyncMock(return_value=mock_batcher), From 7c068a4cf7772b45246568346e36cef7e4f27ec2 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 02:16:11 +0000 Subject: [PATCH 04/11] fix(repositories): import LiteralString from typing_extensions for python 3.10 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/repositories/prisma_protocols.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index c742f3f5f3f..c7421c9284a 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -7,7 +7,9 @@ private ones per file. """ from collections.abc import Mapping, Sequence -from typing import LiteralString, Protocol, TypeVar +from typing import Protocol, TypeVar + +from typing_extensions import LiteralString RowT_co = TypeVar("RowT_co", covariant=True) From efef6ab68491299a550739286990cf922330dd89 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 02:53:47 +0000 Subject: [PATCH 05/11] fix(proxy): write team member spend as one roster checked upsert statement Replaces the per team advisory lock and Pydantic roster parse in the spend flush with a single INSERT ... ON CONFLICT statement that checks the stored roster in SQL, so malformed roster JSON cannot fail the whole flush and large batches no longer issue two queries per team inside the fixed transaction deadline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 95 +++++-------- .../access_group_team_sync.py | 8 +- litellm/repositories/prisma_protocols.py | 10 -- litellm/repositories/team_repository.py | 12 +- .../proxy/db/test_db_spend_update_writer.py | 126 ++++++------------ 5 files changed, 87 insertions(+), 164 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 5c86c52c539..ecf107ef4be 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -18,7 +18,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload from urllib.parse import quote, unquote -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import LiteralString, ReadOnly, TypedDict import litellm from litellm._logging import verbose_proxy_logger @@ -75,8 +75,7 @@ from litellm.proxy.spend_tracking.savings import ( marks_gateway_injection, ) from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error -from litellm.repositories.prisma_protocols import BatchTable, RawQueryTransaction -from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL, TeamRepository +from litellm.repositories.prisma_protocols import BatchTable from litellm.types.utils import CallTypes if TYPE_CHECKING: @@ -134,9 +133,11 @@ class _SpendBatchManager(Protocol): async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... -class _SpendTransaction(RawQueryTransaction, Protocol): +class _SpendTransaction(Protocol): def batch_(self) -> _SpendBatchManager: ... + async def execute_raw(self, query: LiteralString, *args: object) -> int: ... + class _SpendTransactionManager(Protocol): async def __aenter__(self) -> _SpendTransaction: ... @@ -162,47 +163,36 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager: return tx -async def _lock_and_read_rosters( - prisma_client: PrismaClient, transaction: _SpendTransaction, team_ids: Sequence[str] -) -> frozenset[tuple[str, str]]: - """Take each team's advisory lock on ``transaction`` and return its rostered (user_id, team_id) pairs. - - A spend flush can land after ``/team/member_delete`` removed the member. Holding the same - lock that endpoint takes, until this transaction commits, means a member read here is on - the team for the whole flush, so only they may have a missing membership row created. - """ - repository: Final = TeamRepository(prisma_client) - - async def locked_roster(team_id: str) -> tuple[tuple[str, str], ...]: - await transaction.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id) - roster: Final = await repository.get_members_with_roles_locked(transaction, team_id) - return tuple((member.user_id, team_id) for member in roster or () if member.user_id is not None) - - rosters: Final = tuple([await locked_roster(team_id) for team_id in sorted(frozenset(team_ids))]) - return frozenset(pair for roster in rosters for pair in roster) +# One statement adds every member's cost to their membership row. A missing row is created only +# while the user is still on the team's roster; FOR SHARE on the team row makes that check and +# the insert atomic against /team/member_delete and /team/delete, which update or delete that row +# before removing memberships, so a spend flush landing after a removal never recreates the member. +_TEAM_MEMBER_SPEND_SQL: Final = """ +INSERT INTO "LiteLLM_TeamMembership" (user_id, team_id, spend, total_spend) +SELECT p.user_id, p.team_id, p.cost, p.cost +FROM unnest($1::text[], $2::text[], $3::float8[]) AS p(user_id, team_id, cost) +WHERE EXISTS ( + SELECT 1 FROM "LiteLLM_TeamTable" t + WHERE t.team_id = p.team_id + AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id)) + FOR SHARE +) + OR EXISTS (SELECT 1 FROM "LiteLLM_TeamMembership" m WHERE m.user_id = p.user_id AND m.team_id = p.team_id) +ON CONFLICT (user_id, team_id) DO UPDATE +SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend, + total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend +""" -def _queue_team_member_spend( - memberships: BatchTable, user_id: str, team_id: str, response_cost: float, rostered: bool -) -> None: - increments: Final = { - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, - } - if not rostered: - memberships.update_many(where={"team_id": team_id, "user_id": user_id}, data=increments) - return - memberships.upsert( - where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, - data={ - "create": { - "team_id": team_id, - "user_id": user_id, - "spend": response_cost, - "total_spend": response_cost, - }, - "update": increments, - }, +async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None: + # key is "team_id::::user_id::"; the string sort orders rows by (team_id, user_id), + # keeping lock order consistent across pods to prevent deadlocks + keys: Final = tuple(sorted(spend_by_member_key)) + _ = await transaction.execute_raw( + _TEAM_MEMBER_SPEND_SQL, + [key.split("::")[3] for key in keys], + [key.split("::")[1] for key in keys], + [spend_by_member_key[key] for key in keys], ) @@ -1730,24 +1720,7 @@ class DBSpendUpdateWriter: start_time = time.time() try: async with _spend_update_tx(prisma_client) as transaction: - rostered_members = await _lock_and_read_rosters( - prisma_client, transaction, tuple(team_id for _, team_id in team_memberships_to_invalidate) - ) - async with transaction.batch_() as batcher: - # Sort by composite key for consistent lock ordering across pods to prevent deadlocks. - # Key format "team_id::::user_id::" makes the string sort equivalent to sorting by (team_id, user_id). - for key, response_cost in sorted(team_member_list_transactions.items()): - # key is "team_id::::user_id::" - team_id = key.split("::")[1] - user_id = key.split("::")[3] - - _queue_team_member_spend( - batcher.litellm_teammembership, - user_id, - team_id, - response_cost, - (user_id, team_id) in rostered_members, - ) + await _write_team_member_spend(transaction, team_member_list_transactions) # Transaction succeeded, break out of retry loop break except Exception as e: diff --git a/litellm/proxy/management_helpers/access_group_team_sync.py b/litellm/proxy/management_helpers/access_group_team_sync.py index cfe207e66ea..664e36c9f10 100644 --- a/litellm/proxy/management_helpers/access_group_team_sync.py +++ b/litellm/proxy/management_helpers/access_group_team_sync.py @@ -21,7 +21,13 @@ from typing import Final, Protocol from pydantic import BaseModel, TypeAdapter from litellm.proxy.auth.auth_checks import _delete_cache_access_object -from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL + +# hashtext collisions only cost two unrelated teams a little serialization, and the +# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, +# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints +# reuses this exact statement to serialize /team/member_add and /team/delete against each +# other and against this mirror, rather than defining a second, divergent lock on the same key. +TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" _READ_TEAM_SQL: Final = 'SELECT access_group_ids FROM "LiteLLM_TeamTable" WHERE team_id = $1' diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index c7421c9284a..93b8c5c7cd7 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -9,8 +9,6 @@ private ones per file. from collections.abc import Mapping, Sequence from typing import Protocol, TypeVar -from typing_extensions import LiteralString - RowT_co = TypeVar("RowT_co", covariant=True) @@ -110,12 +108,6 @@ class PrismaRecord(Protocol): def dict(self) -> Mapping[str, object]: ... -class RawQueryTransaction(Protocol): - """A prisma transaction handle that can run raw SQL, e.g. an advisory lock or a locked read.""" - - async def query_raw(self, query: LiteralString, *args: str) -> Sequence[Mapping[str, object]]: ... - - class ReadOnlyTable(Protocol): async def find_many(self, *, where: Mapping[str, object]) -> Sequence[PrismaRecord]: ... @@ -131,8 +123,6 @@ class BatchTable(Protocol): def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... - def upsert(self, *, where: Mapping[str, object], data: Mapping[str, Mapping[str, object]]) -> None: ... - class PrismaBatch(Protocol): @property diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 1b9cba48e41..5ff07d76b5d 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -15,9 +15,10 @@ from litellm.repositories.base_repository import ( DbRecord, record_to_dict, ) -from litellm.repositories.prisma_protocols import RawQueryTransaction, TableActions +from litellm.repositories.prisma_protocols import TableActions if TYPE_CHECKING: + from prisma import Prisma from prisma import models as prisma_models @@ -39,13 +40,6 @@ def _team_arrays(team: LiteLLM_TeamTable) -> _TeamArrays: return team -# hashtext collisions only cost two unrelated teams a little serialization, and the -# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, -# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints, -# the access-group mirror and the team member spend flush all reuse this exact statement to -# serialize against each other, rather than defining a second, divergent lock on the same key. -TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" - _MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member]) _JSON_ENCODED_TEAM_FIELDS: Final = ( "metadata", @@ -84,7 +78,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return LiteLLM_TeamTable.model_validate(data) - async def get_members_with_roles_locked(self, tx: RawQueryTransaction, team_id: str) -> list[Member] | None: + async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> list[Member] | None: """Return the team's members_with_roles. The caller must already hold ``TEAM_ADVISORY_LOCK_SQL`` for this team_id on ``tx`` before calling this. diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index b12dc68b7b1..05769eea343 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -16,11 +16,10 @@ from redis.exceptions import DataError import litellm from litellm.proxy._types import Litellm_EntityType -from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter +from litellm.proxy.db.db_spend_update_writer import _TEAM_MEMBER_SPEND_SQL, DBSpendUpdateWriter from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( build_window_spend_transaction, ) -from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL @pytest.mark.asyncio @@ -914,42 +913,26 @@ async def test_commit_spend_updates_to_db_increments_agent_spend(): assert call_kwargs["data"] == {"spend": {"increment": response_cost}} -def _team_member_flush_fixtures(team_id: str, rostered_user_ids: list[str]) -> tuple[MagicMock, AsyncMock, MagicMock]: - """A batcher, transaction and prisma client whose locked roster read for `team_id` lists `rostered_user_ids`.""" - mock_batcher = MagicMock() - mock_batcher.litellm_teammembership = MagicMock() - mock_batcher.litellm_teammembership.upsert = MagicMock() - mock_batcher.litellm_teammembership.update_many = MagicMock() - - roster_row = {"members_with_roles": json.dumps([{"user_id": uid, "role": "user"} for uid in rostered_user_ids])} - - async def query_raw(query: str, *args: str) -> list[dict[str, object]]: - return [] if query == TEAM_ADVISORY_LOCK_SQL else [roster_row] - +def _team_member_flush_fixtures() -> tuple[AsyncMock, MagicMock]: + """A transaction and prisma client that record the raw statement the member spend flush runs.""" mock_transaction = AsyncMock() mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction) mock_transaction.__aexit__ = AsyncMock(return_value=False) - mock_transaction.query_raw = AsyncMock(side_effect=query_raw) - mock_transaction.batch_ = MagicMock( - return_value=AsyncMock( - __aenter__=AsyncMock(return_value=mock_batcher), - __aexit__=AsyncMock(return_value=False), - ) - ) + mock_transaction.execute_raw = AsyncMock(return_value=1) mock_prisma_client = MagicMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction) - return mock_batcher, mock_transaction, mock_prisma_client + return mock_transaction, mock_prisma_client -def _team_member_only_transactions(entity_id: str, response_cost: float) -> dict[str, dict[str, float]]: +def _team_member_only_transactions(spend_by_member_key: dict[str, float]) -> dict[str, dict[str, float]]: return { "user_list_transactions": {}, "end_user_list_transactions": {}, "key_list_transactions": {}, "team_list_transactions": {}, - "team_member_list_transactions": {entity_id: response_cost}, + "team_member_list_transactions": spend_by_member_key, "org_list_transactions": {}, "tag_list_transactions": {}, "agent_list_transactions": {}, @@ -957,22 +940,20 @@ def _team_member_only_transactions(entity_id: str, response_cost: float) -> dict @pytest.mark.asyncio -async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total_spend(): +async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster_checked_upsert(): """ - Verify that _commit_spend_updates_to_db increments BOTH spend (cycle-scoped) - and total_spend (non-resetting) on LiteLLM_TeamMembership in a single - upsert call, using the same response_cost. + Regression (LIT-5502): members added without a budget had no membership row, and the + previous update_many matched zero rows, so their spend was silently dropped. - Regression (LIT-5502): members added without a budget had no membership row, and - the previous update_many matched zero rows, so their spend was silently dropped. - For a member still on the team roster the upsert has to create the row seeded - with this call's cost in that case. + The flush now runs one INSERT ... ON CONFLICT statement for the whole batch that adds + the cost to both spend and total_spend and creates the missing row for a user still on + the team roster, so no per-team read can fail or time out ahead of the writes. """ db_writer = DBSpendUpdateWriter() team_id = "team-abc" user_id = "user-xyz" response_cost = 0.75 - mock_batcher, mock_transaction, mock_prisma_client = _team_member_flush_fixtures(team_id, [user_id]) + mock_transaction, mock_prisma_client = _team_member_flush_fixtures() mock_proxy_logging = MagicMock() mock_proxy_logging.call_details.get = MagicMock(return_value=None) @@ -982,42 +963,31 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total n_retry_times=0, proxy_logging_obj=mock_proxy_logging, db_spend_update_transactions=_team_member_only_transactions( - f"team_id::{team_id}::user_id::{user_id}", response_cost + {f"team_id::{team_id}::user_id::{user_id}": response_cost} ), ) - assert mock_transaction.query_raw.await_args_list[0] == call(TEAM_ADVISORY_LOCK_SQL, team_id) - mock_batcher.litellm_teammembership.upsert.assert_called_once() - mock_batcher.litellm_teammembership.update_many.assert_not_called() - call_kwargs = mock_batcher.litellm_teammembership.upsert.call_args.kwargs - assert call_kwargs["where"] == {"user_id_team_id": {"user_id": user_id, "team_id": team_id}} - assert call_kwargs["data"] == { - "create": { - "team_id": team_id, - "user_id": user_id, - "spend": response_cost, - "total_spend": response_cost, - }, - "update": { - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, - }, - } + mock_transaction.execute_raw.assert_awaited_once() + statement, user_ids, team_ids, costs = mock_transaction.execute_raw.await_args.args + assert statement is _TEAM_MEMBER_SPEND_SQL + assert (user_ids, team_ids, costs) == ([user_id], [team_id], [response_cost]) + assert 'INSERT INTO "LiteLLM_TeamMembership"' in statement + assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))" in statement + assert "FOR SHARE" in statement + assert "ON CONFLICT (user_id, team_id) DO UPDATE" in statement + assert 'spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend' in statement + assert 'total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend' in statement @pytest.mark.asyncio -async def test_commit_spend_updates_to_db_does_not_recreate_membership_of_removed_team_member(): +async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_user(): """ - A spend flush that lands after /team/member_delete must not resurrect the deleted - membership row: a user missing from the team roster, read under the team's advisory - lock inside the flush transaction, only gets an increment on whatever row still - exists, never a create. + The single member spend statement locks rows in the order of its input arrays, so the + batch is handed over sorted by (team_id, user_id), with each cost kept next to its + member, so concurrent pods lock in the same order and cannot deadlock. """ db_writer = DBSpendUpdateWriter() - team_id = "team-abc" - removed_user_id = "user-removed" - response_cost = 0.75 - mock_batcher, mock_transaction, mock_prisma_client = _team_member_flush_fixtures(team_id, ["user-still-here"]) + mock_transaction, mock_prisma_client = _team_member_flush_fixtures() mock_proxy_logging = MagicMock() mock_proxy_logging.call_details.get = MagicMock(return_value=None) @@ -1027,19 +997,22 @@ async def test_commit_spend_updates_to_db_does_not_recreate_membership_of_remove n_retry_times=0, proxy_logging_obj=mock_proxy_logging, db_spend_update_transactions=_team_member_only_transactions( - f"team_id::{team_id}::user_id::{removed_user_id}", response_cost + { + "team_id::team_c::user_id::user_x": 0.1, + "team_id::team_a::user_id::user_y": 0.2, + "team_id::team_a::user_id::user_x": 0.3, + "team_id::team_b::user_id::user_x": 0.4, + } ), ) - assert mock_transaction.query_raw.await_args_list[0] == call(TEAM_ADVISORY_LOCK_SQL, team_id) - mock_batcher.litellm_teammembership.upsert.assert_not_called() - mock_batcher.litellm_teammembership.update_many.assert_called_once_with( - where={"team_id": team_id, "user_id": removed_user_id}, - data={ - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, - }, - ) + _statement, user_ids, team_ids, costs = mock_transaction.execute_raw.await_args.args + assert list(zip(team_ids, user_ids, costs)) == [ + ("team_a", "user_x", 0.3), + ("team_a", "user_y", 0.2), + ("team_b", "user_x", 0.4), + ("team_c", "user_x", 0.1), + ] @pytest.mark.asyncio @@ -2265,19 +2238,6 @@ async def test_commit_daily_tag_spend_no_requeue_on_success(): ["team_a", "team_b", "team_c"], id="team", ), - pytest.param( - "team_member_list_transactions", - { - "team_id::team_c::user_id::user_x": 0.1, - "team_id::team_a::user_id::user_x": 0.2, - "team_id::team_b::user_id::user_x": 0.3, - }, - "litellm_teammembership", - "update_many", - "team_id", - ["team_a", "team_b", "team_c"], - id="team_member", - ), pytest.param( "org_list_transactions", {"org_c": 0.1, "org_a": 0.2, "org_b": 0.3}, From 60077e90aabfcb52c9babeb19c28efb3ad0c1fdd Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 03:16:54 +0000 Subject: [PATCH 06/11] fix(proxy): take the team advisory lock before the member spend flush Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 18 +++++++--- .../proxy/db/test_db_spend_update_writer.py | 34 +++++++++++++------ 2 files changed, 36 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ecf107ef4be..f617302d8ad 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -163,10 +163,17 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager: return tx +# The per-team advisory lock the team endpoints hold while changing a roster (TEAM_ADVISORY_LOCK_SQL), +# taken in sorted order so the roster check below cannot interleave with their writes. A row lock would +# deadlock with the access-group endpoints, which lock a team row after an access-group lock. +_TEAM_ADVISORY_LOCKS_SQL: Final = """ +SELECT pg_advisory_xact_lock(hashtext(teams.team_id)) +FROM (SELECT DISTINCT team_id FROM unnest($1::text[]) AS team_id ORDER BY team_id) AS teams +""" + # One statement adds every member's cost to their membership row. A missing row is created only -# while the user is still on the team's roster; FOR SHARE on the team row makes that check and -# the insert atomic against /team/member_delete and /team/delete, which update or delete that row -# before removing memberships, so a spend flush landing after a removal never recreates the member. +# while the user is still on the team's roster, so a spend flush landing after a removal never +# recreates the member. _TEAM_MEMBER_SPEND_SQL: Final = """ INSERT INTO "LiteLLM_TeamMembership" (user_id, team_id, spend, total_spend) SELECT p.user_id, p.team_id, p.cost, p.cost @@ -175,7 +182,6 @@ WHERE EXISTS ( SELECT 1 FROM "LiteLLM_TeamTable" t WHERE t.team_id = p.team_id AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id)) - FOR SHARE ) OR EXISTS (SELECT 1 FROM "LiteLLM_TeamMembership" m WHERE m.user_id = p.user_id AND m.team_id = p.team_id) ON CONFLICT (user_id, team_id) DO UPDATE @@ -188,10 +194,12 @@ async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_memb # key is "team_id::::user_id::"; the string sort orders rows by (team_id, user_id), # keeping lock order consistent across pods to prevent deadlocks keys: Final = tuple(sorted(spend_by_member_key)) + team_ids: Final = [key.split("::")[1] for key in keys] + _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCKS_SQL, team_ids) _ = await transaction.execute_raw( _TEAM_MEMBER_SPEND_SQL, [key.split("::")[3] for key in keys], - [key.split("::")[1] for key in keys], + team_ids, [spend_by_member_key[key] for key in keys], ) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 05769eea343..9f3bc5acd92 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -16,7 +16,11 @@ from redis.exceptions import DataError import litellm from litellm.proxy._types import Litellm_EntityType -from litellm.proxy.db.db_spend_update_writer import _TEAM_MEMBER_SPEND_SQL, DBSpendUpdateWriter +from litellm.proxy.db.db_spend_update_writer import ( + _TEAM_ADVISORY_LOCKS_SQL, + _TEAM_MEMBER_SPEND_SQL, + DBSpendUpdateWriter, +) from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( build_window_spend_transaction, ) @@ -945,9 +949,10 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster Regression (LIT-5502): members added without a budget had no membership row, and the previous update_many matched zero rows, so their spend was silently dropped. - The flush now runs one INSERT ... ON CONFLICT statement for the whole batch that adds - the cost to both spend and total_spend and creates the missing row for a user still on - the team roster, so no per-team read can fail or time out ahead of the writes. + The flush now takes the same per-team advisory lock the team endpoints hold, then runs + one INSERT ... ON CONFLICT statement for the whole batch that adds the cost to both spend + and total_spend and creates the missing row for a user still on the team roster, so no + per-team read can fail or time out ahead of the writes. """ db_writer = DBSpendUpdateWriter() team_id = "team-abc" @@ -967,13 +972,17 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster ), ) - mock_transaction.execute_raw.assert_awaited_once() - statement, user_ids, team_ids, costs = mock_transaction.execute_raw.await_args.args + lock_call, spend_call = mock_transaction.execute_raw.await_args_list + lock_statement, locked_team_ids = lock_call.args + assert lock_statement is _TEAM_ADVISORY_LOCKS_SQL + assert locked_team_ids == [team_id] + assert "pg_advisory_xact_lock(hashtext(teams.team_id))" in lock_statement + assert "ORDER BY team_id" in lock_statement + statement, user_ids, team_ids, costs = spend_call.args assert statement is _TEAM_MEMBER_SPEND_SQL assert (user_ids, team_ids, costs) == ([user_id], [team_id], [response_cost]) assert 'INSERT INTO "LiteLLM_TeamMembership"' in statement assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))" in statement - assert "FOR SHARE" in statement assert "ON CONFLICT (user_id, team_id) DO UPDATE" in statement assert 'spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend' in statement assert 'total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend' in statement @@ -982,9 +991,10 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster @pytest.mark.asyncio async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_user(): """ - The single member spend statement locks rows in the order of its input arrays, so the - batch is handed over sorted by (team_id, user_id), with each cost kept next to its - member, so concurrent pods lock in the same order and cannot deadlock. + The member spend statement touches rows in the order of its input arrays, so the batch + is handed over sorted by (team_id, user_id), with each cost kept next to its member, and + the advisory lock statement receives the same team ids, so concurrent pods lock in the + same order and cannot deadlock. """ db_writer = DBSpendUpdateWriter() mock_transaction, mock_prisma_client = _team_member_flush_fixtures() @@ -1006,7 +1016,9 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u ), ) - _statement, user_ids, team_ids, costs = mock_transaction.execute_raw.await_args.args + lock_call, spend_call = mock_transaction.execute_raw.await_args_list + _statement, user_ids, team_ids, costs = spend_call.args + assert lock_call.args == (_TEAM_ADVISORY_LOCKS_SQL, ["team_a", "team_a", "team_b", "team_c"]) assert list(zip(team_ids, user_ids, costs)) == [ ("team_a", "user_x", 0.3), ("team_a", "user_y", 0.2), From 499c334fce491fa67fc1db92e5e1c160cd00b33d Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 03:35:34 +0000 Subject: [PATCH 07/11] fix(proxy): lock each team in sorted order before the member spend upsert Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 12 +++++------ .../proxy/db/test_db_spend_update_writer.py | 21 +++++++++++-------- 2 files changed, 17 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index f617302d8ad..09ce6b6f451 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -164,12 +164,9 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager: # The per-team advisory lock the team endpoints hold while changing a roster (TEAM_ADVISORY_LOCK_SQL), -# taken in sorted order so the roster check below cannot interleave with their writes. A row lock would -# deadlock with the access-group endpoints, which lock a team row after an access-group lock. -_TEAM_ADVISORY_LOCKS_SQL: Final = """ -SELECT pg_advisory_xact_lock(hashtext(teams.team_id)) -FROM (SELECT DISTINCT team_id FROM unnest($1::text[]) AS team_id ORDER BY team_id) AS teams -""" +# so the roster check below cannot interleave with their writes. A row lock would deadlock with the +# access-group endpoints, which lock a team row after an access-group lock. +_TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" # One statement adds every member's cost to their membership row. A missing row is created only # while the user is still on the team's roster, so a spend flush landing after a removal never @@ -195,7 +192,8 @@ async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_memb # keeping lock order consistent across pods to prevent deadlocks keys: Final = tuple(sorted(spend_by_member_key)) team_ids: Final = [key.split("::")[1] for key in keys] - _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCKS_SQL, team_ids) + for team_id in dict.fromkeys(team_ids): + _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id) _ = await transaction.execute_raw( _TEAM_MEMBER_SPEND_SQL, [key.split("::")[3] for key in keys], diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 9f3bc5acd92..2ac9fc8eccd 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -17,7 +17,7 @@ from redis.exceptions import DataError import litellm from litellm.proxy._types import Litellm_EntityType from litellm.proxy.db.db_spend_update_writer import ( - _TEAM_ADVISORY_LOCKS_SQL, + _TEAM_ADVISORY_LOCK_SQL, _TEAM_MEMBER_SPEND_SQL, DBSpendUpdateWriter, ) @@ -973,11 +973,10 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster ) lock_call, spend_call = mock_transaction.execute_raw.await_args_list - lock_statement, locked_team_ids = lock_call.args - assert lock_statement is _TEAM_ADVISORY_LOCKS_SQL - assert locked_team_ids == [team_id] - assert "pg_advisory_xact_lock(hashtext(teams.team_id))" in lock_statement - assert "ORDER BY team_id" in lock_statement + lock_statement, locked_team_id = lock_call.args + assert lock_statement is _TEAM_ADVISORY_LOCK_SQL + assert locked_team_id == team_id + assert "pg_advisory_xact_lock(hashtext($1))" in lock_statement statement, user_ids, team_ids, costs = spend_call.args assert statement is _TEAM_MEMBER_SPEND_SQL assert (user_ids, team_ids, costs) == ([user_id], [team_id], [response_cost]) @@ -993,7 +992,7 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u """ The member spend statement touches rows in the order of its input arrays, so the batch is handed over sorted by (team_id, user_id), with each cost kept next to its member, and - the advisory lock statement receives the same team ids, so concurrent pods lock in the + each distinct team is locked once, in that same order, so concurrent pods lock in the same order and cannot deadlock. """ db_writer = DBSpendUpdateWriter() @@ -1016,9 +1015,13 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u ), ) - lock_call, spend_call = mock_transaction.execute_raw.await_args_list + *lock_calls, spend_call = mock_transaction.execute_raw.await_args_list _statement, user_ids, team_ids, costs = spend_call.args - assert lock_call.args == (_TEAM_ADVISORY_LOCKS_SQL, ["team_a", "team_a", "team_b", "team_c"]) + assert [lock_call.args for lock_call in lock_calls] == [ + (_TEAM_ADVISORY_LOCK_SQL, "team_a"), + (_TEAM_ADVISORY_LOCK_SQL, "team_b"), + (_TEAM_ADVISORY_LOCK_SQL, "team_c"), + ] assert list(zip(team_ids, user_ids, costs)) == [ ("team_a", "user_x", 0.3), ("team_a", "user_y", 0.2), From 2d61fa66b165d8e5e9b4295bba2587346709aab9 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 00:37:03 +0000 Subject: [PATCH 08/11] fix(proxy): lock teams in sorted team id order and keep existing member budgets on re-add Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 12 ++++----- litellm/proxy/management_helpers/utils.py | 2 +- .../proxy/db/test_db_spend_update_writer.py | 27 ++++++++++--------- .../test_management_helpers_utils.py | 3 +++ 4 files changed, 24 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 09ce6b6f451..77d16e90706 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -188,17 +188,17 @@ SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend, async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None: - # key is "team_id::::user_id::"; the string sort orders rows by (team_id, user_id), - # keeping lock order consistent across pods to prevent deadlocks - keys: Final = tuple(sorted(spend_by_member_key)) - team_ids: Final = [key.split("::")[1] for key in keys] + # key is "team_id::::user_id::"; rows are sorted by (team_id, user_id) so the teams are + # locked in the same `sorted(team_ids)` order the team endpoints use, preventing deadlocks + rows: Final = sorted((key.split("::")[1], key.split("::")[3], cost) for key, cost in spend_by_member_key.items()) + team_ids: Final = [team_id for team_id, _user_id, _cost in rows] for team_id in dict.fromkeys(team_ids): _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id) _ = await transaction.execute_raw( _TEAM_MEMBER_SPEND_SQL, - [key.split("::")[3] for key in keys], + [user_id for _team_id, user_id, _cost in rows], team_ids, - [spend_by_member_key[key] for key in keys], + [cost for _team_id, _user_id, cost in rows], ) diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index d940eba86d3..d84d32431ec 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -479,7 +479,7 @@ async def add_new_member( budget_link: Final[Mapping[str, object]] = {"budget_id": _budget_id} if _budget_id is not None else {} _returned_team_membership: Final = await membership_table.upsert( where={"user_id_team_id": membership_key}, - data={"create": {**membership_key, **budget_link}, "update": {**budget_link}}, + data={"create": {**membership_key, **budget_link}, "update": {}}, include={"litellm_budget_table": True}, ) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 2ac9fc8eccd..70199ee0c73 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -992,8 +992,9 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u """ The member spend statement touches rows in the order of its input arrays, so the batch is handed over sorted by (team_id, user_id), with each cost kept next to its member, and - each distinct team is locked once, in that same order, so concurrent pods lock in the - same order and cannot deadlock. + each distinct team is locked once, in `sorted(team_ids)` order, the order /team/delete + locks in, so a concurrent flush and delete cannot deadlock. `eng` and `eng2` pin that: + sorting the composite keys instead would lock `eng2` first because `2` < `:`. """ db_writer = DBSpendUpdateWriter() mock_transaction, mock_prisma_client = _team_member_flush_fixtures() @@ -1007,10 +1008,10 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u proxy_logging_obj=mock_proxy_logging, db_spend_update_transactions=_team_member_only_transactions( { - "team_id::team_c::user_id::user_x": 0.1, - "team_id::team_a::user_id::user_y": 0.2, - "team_id::team_a::user_id::user_x": 0.3, - "team_id::team_b::user_id::user_x": 0.4, + "team_id::eng2::user_id::user_x": 0.1, + "team_id::eng::user_id::user_y": 0.2, + "team_id::eng::user_id::user_x": 0.3, + "team_id::eng-b::user_id::user_x": 0.4, } ), ) @@ -1018,15 +1019,15 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u *lock_calls, spend_call = mock_transaction.execute_raw.await_args_list _statement, user_ids, team_ids, costs = spend_call.args assert [lock_call.args for lock_call in lock_calls] == [ - (_TEAM_ADVISORY_LOCK_SQL, "team_a"), - (_TEAM_ADVISORY_LOCK_SQL, "team_b"), - (_TEAM_ADVISORY_LOCK_SQL, "team_c"), + (_TEAM_ADVISORY_LOCK_SQL, "eng"), + (_TEAM_ADVISORY_LOCK_SQL, "eng-b"), + (_TEAM_ADVISORY_LOCK_SQL, "eng2"), ] assert list(zip(team_ids, user_ids, costs)) == [ - ("team_a", "user_x", 0.3), - ("team_a", "user_y", 0.2), - ("team_b", "user_x", 0.4), - ("team_c", "user_x", 0.1), + ("eng", "user_x", 0.3), + ("eng", "user_y", 0.2), + ("eng-b", "user_x", 0.4), + ("eng2", "user_x", 0.1), ] diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index 0dcf7e1b8ad..09d5f684f3d 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -446,6 +446,8 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): 1. When max_budget_in_team is provided 2. A new budget is created in the litellm_budgettable 3. The new budget_id is used for the team membership + 4. The upsert's update branch stays empty, so a bulk /team/member_add that names a member + already on the team does not replace the budget_id (and the spend) their existing row carries """ from litellm.proxy._types import LitellmUserRoles @@ -529,6 +531,7 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): assert team_membership_call_args is not None create_data = team_membership_call_args.kwargs["data"]["create"] assert create_data["budget_id"] == test_new_budget_id + assert team_membership_call_args.kwargs["data"]["update"] == {} @pytest.mark.asyncio From b6f4ad190e00a21935479091a423e16e94530778 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 00:49:48 +0000 Subject: [PATCH 09/11] refactor(proxy): freeze the member spend arrays and budget link to stay within the type discipline budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 9 ++++----- litellm/proxy/management_helpers/utils.py | 4 +++- .../test_litellm/proxy/db/test_db_spend_update_writer.py | 2 +- 3 files changed, 8 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 77d16e90706..8bcfe28488e 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -188,17 +188,16 @@ SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend, async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None: - # key is "team_id::::user_id::"; rows are sorted by (team_id, user_id) so the teams are - # locked in the same `sorted(team_ids)` order the team endpoints use, preventing deadlocks + # key is "team_id::::user_id::"; locks are taken in sorted team_id order like the team endpoints rows: Final = sorted((key.split("::")[1], key.split("::")[3], cost) for key, cost in spend_by_member_key.items()) - team_ids: Final = [team_id for team_id, _user_id, _cost in rows] + team_ids: Final = tuple(team_id for team_id, _user_id, _cost in rows) for team_id in dict.fromkeys(team_ids): _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id) _ = await transaction.execute_raw( _TEAM_MEMBER_SPEND_SQL, - [user_id for _team_id, user_id, _cost in rows], + tuple(user_id for _team_id, user_id, _cost in rows), team_ids, - [cost for _team_id, _user_id, cost in rows], + tuple(cost for _team_id, _user_id, cost in rows), ) diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index d84d32431ec..c10e5f9b23d 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -476,7 +476,9 @@ async def add_new_member( if returned_user is not None and returned_user.user_id is not None: membership_table: Final[_PrismaTeamMembershipTable] = _team_membership_table(prisma_client, tx) membership_key: Final[Mapping[str, object]] = {"user_id": returned_user.user_id, "team_id": team_id} - budget_link: Final[Mapping[str, object]] = {"budget_id": _budget_id} if _budget_id is not None else {} + budget_link: Final[Mapping[str, str]] = ( + MappingProxyType({"budget_id": _budget_id}) if _budget_id is not None else MappingProxyType({}) + ) _returned_team_membership: Final = await membership_table.upsert( where={"user_id_team_id": membership_key}, data={"create": {**membership_key, **budget_link}, "update": {}}, diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 70199ee0c73..e38a65f2b12 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -979,7 +979,7 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster assert "pg_advisory_xact_lock(hashtext($1))" in lock_statement statement, user_ids, team_ids, costs = spend_call.args assert statement is _TEAM_MEMBER_SPEND_SQL - assert (user_ids, team_ids, costs) == ([user_id], [team_id], [response_cost]) + assert (list(user_ids), list(team_ids), list(costs)) == ([user_id], [team_id], [response_cost]) assert 'INSERT INTO "LiteLLM_TeamMembership"' in statement assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))" in statement assert "ON CONFLICT (user_id, team_id) DO UPDATE" in statement From 61d4c5b9b5b427929f2d97d216f3aafb611f0abb Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 01:12:54 +0000 Subject: [PATCH 10/11] fix(proxy): skip members already on the team before resolving a per-member budget A mixed /team/member_add list that names an existing member used to run add_new_member for them, which created or cloned a budget that the empty upsert update branch never linked to their membership row. Filter the requested members against the freshly locked roster first so budgets and membership rows are only written for members who are actually new Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/team_endpoints.py | 34 +++------ .../test_team_endpoints.py | 71 +++++++++++++++++++ 2 files changed, 79 insertions(+), 26 deletions(-) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 4957b3bd925..28c12173ea7 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -2904,10 +2904,15 @@ async def _process_team_members( if member_allowed_models is None and team_default_member_models: member_allowed_models = team_default_member_models - if isinstance(data.member, Member): + requested_members: Final[Sequence[Member]] = ( + (data.member,) if isinstance(data.member, Member) else tuple(data.member) + ) + for m in requested_members: + if _member_already_in_team(m, complete_team_data): + continue try: updated_user, updated_tm = await add_new_member( - new_member=data.member, + new_member=m, max_budget_in_team=data.max_budget_in_team, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -2921,34 +2926,11 @@ async def _process_team_members( except Exception as e: raise HTTPException( status_code=500, - detail={"error": f"Unable to add user - {data.member}, to team - {data.team_id}, for reason - {e}"}, + detail={"error": f"Unable to add user - {m}, to team - {data.team_id}, for reason - {e}"}, ) updated_users.append(updated_user) if updated_tm is not None: updated_team_memberships.append(updated_tm) - elif isinstance(data.member, list): - for m in data.member: - try: - updated_user, updated_tm = await add_new_member( - new_member=m, - max_budget_in_team=data.max_budget_in_team, - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - team_id=data.team_id, - default_team_budget_id=default_team_budget_id, - allowed_models=member_allowed_models, - budget_duration=data.budget_duration, - tx=tx, - ) - except Exception as e: - raise HTTPException( - status_code=500, - detail={"error": f"Unable to add user - {m}, to team - {data.team_id}, for reason - {e}"}, - ) - updated_users.append(updated_user) - if updated_tm is not None: - updated_team_memberships.append(updated_tm) return updated_users, updated_team_memberships 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 f7157da3100..690b5ae80b6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -1794,6 +1794,7 @@ async def test_process_team_members_single_member(): mock_team = MagicMock(spec=LiteLLM_TeamTable) mock_team.metadata = {"team_member_budget_id": "budget-123"} mock_team.default_team_member_models = None + mock_team.members_with_roles = [] # Mock user and membership objects mock_user = MagicMock(spec=LiteLLM_UserTable) @@ -1854,6 +1855,7 @@ async def test_process_team_members_multiple_members(): mock_team = MagicMock(spec=LiteLLM_TeamTable) mock_team.metadata = None mock_team.default_team_member_models = None + mock_team.members_with_roles = [] # Create multiple members as dictionaries (they will be converted to Member objects) members = [ @@ -2114,6 +2116,75 @@ async def test_add_team_members_runs_member_writes_on_the_lock_holding_transacti assert [tm.budget_id for tm in updated_team_memberships] == ["budget-pool"] +@pytest.mark.asyncio +async def test_add_team_members_skips_budget_and_membership_writes_for_members_already_on_the_roster(): + """ + Regression pin for orphaned budgets on a mixed /team/member_add list. + + A list naming one member already on the team and one new member must only create a + budget and membership row for the new member. Running add_new_member for the existing + member would create a per-member budget that nothing links to, since their membership + row (and the budget it already carries) is left untouched. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + _add_team_members_to_team, + ) + + added_user = MagicMock() + added_user.user_id = "bob" + added_user.model_dump.return_value = {"user_id": "bob", "teams": ["team-mixed"]} + created_budget = MagicMock() + created_budget.budget_id = "budget-bob" + membership = MagicMock() + membership.model_dump.return_value = { + "team_id": "team-mixed", + "user_id": "bob", + "budget_id": "budget-bob", + "litellm_budget_table": None, + } + + tx = MagicMock() + tx.query_raw = AsyncMock( + return_value=[{"members_with_roles": [{"user_id": "alice", "user_email": None, "role": "user"}]}] + ) + tx.litellm_teamtable.update = AsyncMock( + return_value=LiteLLM_TeamTable(team_id="team-mixed", members_with_roles=[]) + ) + tx.litellm_usertable.upsert = AsyncMock(return_value=added_user) + tx.litellm_usertable.update_many = AsyncMock() + tx.litellm_budgettable.create = AsyncMock(return_value=created_budget) + tx.litellm_teammembership.upsert = AsyncMock(return_value=membership) + + tx_cm = MagicMock() + tx_cm.__aenter__ = AsyncMock(return_value=tx) + tx_cm.__aexit__ = AsyncMock(return_value=None) + + prisma_client = MagicMock() + prisma_client.tx = MagicMock(return_value=tx_cm) + + _, updated_users, updated_team_memberships = await _add_team_members_to_team( + data=TeamMemberAddRequest( + team_id="team-mixed", + member=[Member(user_id="alice", role="user"), Member(user_id="bob", role="user")], + max_budget_in_team=50.0, + ), + complete_team_data=LiteLLM_TeamTable(team_id="team-mixed", members_with_roles=[]), + prisma_client=cast(object, prisma_client), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + litellm_proxy_admin_name="admin", + ) + + tx.litellm_budgettable.create.assert_awaited_once() + tx.litellm_teammembership.upsert.assert_awaited_once() + assert tx.litellm_teammembership.upsert.call_args.kwargs["where"] == { + "user_id_team_id": {"user_id": "bob", "team_id": "team-mixed"} + } + assert [user.user_id for user in updated_users] == ["bob"] + assert [tm.user_id for tm in updated_team_memberships] == ["bob"] + written_ids = [m["user_id"] for m in json.loads(tx.litellm_teamtable.update.call_args.kwargs["data"]["members_with_roles"])] + assert written_ids == ["alice", "bob"] + + @pytest.mark.asyncio async def test_add_team_members_writes_nothing_when_the_team_is_deleted_mid_request(): """ From c21e86a4436abdd48e0ad83cd3284adc35a50ea7 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 18 Sep 2026 12:27:39 -0700 Subject: [PATCH 11/11] feat(ui): reset a team member's spend from the Members tab Every member now carries a membership row, so a member who spent with no budget is over budget the moment a member budget is added later. The only fix was POST /team/{team_id}/member/{user_id}/reset_spend, which had no UI. The Members tab gets a Reset spend action on rows that have current cycle spend. It confirms in a dialog, posts reset_to 0 through the typed client, and refreshes the team without remounting the page so the tab stays open. A team admin does not see it on their own row because the backend rejects that reset --- .../hooks/teams/useResetTeamMemberSpend.ts | 17 +++ .../TableIconActionButton.tsx | 1 + .../common_components/MemberTable.tsx | 18 +++ .../src/components/team/TeamInfo.tsx | 10 ++ .../components/team/TeamMemberTab.test.tsx | 114 +++++++++++++++++- .../src/components/team/TeamMemberTab.tsx | 94 ++++++++++++--- 6 files changed, 234 insertions(+), 20 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useResetTeamMemberSpend.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useResetTeamMemberSpend.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useResetTeamMemberSpend.ts new file mode 100644 index 00000000000..1fdf8fbd98d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useResetTeamMemberSpend.ts @@ -0,0 +1,17 @@ +import { useMutation } from "@tanstack/react-query"; +import { fetchClient } from "@/lib/http/api"; + +export interface ResetTeamMemberSpendParams { + teamId: string; + userId: string; +} + +export const resetTeamMemberSpend = async ({ teamId, userId }: ResetTeamMemberSpendParams): Promise => { + await fetchClient.POST("/team/{team_id}/member/{user_id}/reset_spend", { + params: { path: { team_id: teamId, user_id: userId } }, + body: { reset_to: 0 }, + }); +}; + +export const useResetTeamMemberSpend = () => + useMutation({ mutationFn: resetTeamMemberSpend }); diff --git a/ui/litellm-dashboard/src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx b/ui/litellm-dashboard/src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx index 396355e1c0b..9da9c2dc702 100644 --- a/ui/litellm-dashboard/src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx +++ b/ui/litellm-dashboard/src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx @@ -30,6 +30,7 @@ export const TableIconActionButtonMap: Record boolean; + onResetSpend?: (member: Member) => void; + showResetSpendForMember?: (member: Member) => boolean; emptyText?: string; } @@ -73,6 +75,8 @@ interface MemberColumnDeps { roleTooltip?: string; extraColumns: MemberTableColumn[]; showDeleteForMember?: (member: Member) => boolean; + onResetSpend?: (member: Member) => void; + showResetSpendForMember?: (member: Member) => boolean; } const extraColumnDef = (column: MemberTableColumn): ColumnDef => { @@ -105,6 +109,8 @@ const buildColumns = ({ roleTooltip, extraColumns, showDeleteForMember, + onResetSpend, + showResetSpendForMember, }: MemberColumnDeps): ColumnDef[] => [ { id: "user_alias", @@ -173,6 +179,14 @@ const buildColumns = ({ dataTestId="edit-member" onClick={() => onEdit(row.original)} /> + {onResetSpend && (showResetSpendForMember?.(row.original) ?? true) && ( + onResetSpend(row.original)} + /> + )} {(!showDeleteForMember || showDeleteForMember(row.original)) && ( = ({ } }; + const refreshTeamData = async () => { + if (!accessToken) return; + try { + setTeamData(await teamInfoCall(accessToken, teamId)); + } catch { + toast.fromError("Failed to load team information"); + } + }; + useEffect(() => { fetchTeamInfo(); }, [teamId, accessToken]); @@ -1351,6 +1360,7 @@ const TeamInfoView: React.FC = ({ teamData={teamData} canEditTeam={canEditTeam} handleMemberDelete={handleMemberDelete} + onMemberSpendReset={refreshTeamData} setSelectedEditMember={setSelectedEditMember} setIsEditMemberModalVisible={setIsEditMemberModalVisible} setIsAddMemberModalVisible={setIsAddMemberModalVisible} diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx index 760074d5dc9..52cba1e6330 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx @@ -1,4 +1,4 @@ -import { fireEvent, screen, within } from "@testing-library/react"; +import { fireEvent, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders } from "../../../tests/test-utils"; @@ -13,6 +13,9 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: vi.fn(), })); +const { POST } = vi.hoisted(() => ({ POST: vi.fn() })); +vi.mock("@/lib/http/api", () => ({ fetchClient: { POST } })); + vi.mock("@/utils/roles", () => ({ isUserTeamAdminForSingleTeam: vi.fn(() => false), isProxyAdminRole: vi.fn(() => false), @@ -26,6 +29,7 @@ const mockHandleMemberDelete = vi.fn(); const mockSetSelectedEditMember = vi.fn(); const mockSetIsEditMemberModalVisible = vi.fn(); const mockSetIsAddMemberModalVisible = vi.fn(); +const mockOnMemberSpendReset = vi.fn(); const budgetResetIso = new Date(2026, 6, 15, 12, 0, 0).toISOString(); @@ -121,6 +125,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -136,6 +141,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -154,6 +160,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -172,6 +179,7 @@ describe("TeamMembersComponent", () => { const props = { canEditTeam: false, handleMemberDelete: mockHandleMemberDelete, + onMemberSpendReset: mockOnMemberSpendReset, setSelectedEditMember: mockSetSelectedEditMember, setIsEditMemberModalVisible: mockSetIsEditMemberModalVisible, setIsAddMemberModalVisible: mockSetIsAddMemberModalVisible, @@ -195,6 +203,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -221,6 +230,7 @@ describe("TeamMembersComponent", () => { })} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -247,6 +257,7 @@ describe("TeamMembersComponent", () => { })} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -262,6 +273,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -280,6 +292,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -295,6 +308,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -311,6 +325,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -330,6 +345,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -364,6 +380,7 @@ describe("TeamMembersComponent", () => { teamData={teamData} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -417,6 +434,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -447,6 +465,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -466,6 +485,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -482,6 +502,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -491,4 +512,95 @@ describe("TeamMembersComponent", () => { expect(screen.queryByTestId("edit-member")).not.toBeInTheDocument(); expect(screen.queryByTestId("delete-member")).not.toBeInTheDocument(); }); + + describe("reset spend", () => { + const renderEditableTab = () => + renderWithProviders( + , + ); + + it("resets the member's current cycle spend to $0 after confirming, then refreshes the team", async () => { + const user = userEvent.setup(); + POST.mockResolvedValue({ data: {} }); + renderEditableTab(); + + const memberRow = screen.getByRole("row", { name: /user1@test\.com/ }); + await user.click(within(memberRow).getByTestId("reset-member-spend")); + + const dialog = await screen.findByRole("dialog", { name: "Reset Team Member Spend" }); + expect(dialog).toHaveTextContent("user1@test.com"); + expect(dialog).toHaveTextContent("$100.5000"); + expect(POST).not.toHaveBeenCalled(); + + await user.click(within(dialog).getByRole("button", { name: "Reset" })); + + await waitFor(() => expect(mockOnMemberSpendReset).toHaveBeenCalledTimes(1)); + expect(POST).toHaveBeenCalledExactlyOnceWith("/team/{team_id}/member/{user_id}/reset_spend", { + params: { path: { team_id: "team-123", user_id: "user1@test.com" } }, + body: { reset_to: 0 }, + }); + expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); + }); + + it("keeps the dialog open and does not refresh the team when the reset fails", async () => { + const user = userEvent.setup(); + POST.mockRejectedValue(new Error("Cannot reset your own spend. Ask a proxy admin.")); + renderEditableTab(); + + await user.click(screen.getByTestId("reset-member-spend")); + const dialog = await screen.findByRole("dialog", { name: "Reset Team Member Spend" }); + await user.click(within(dialog).getByRole("button", { name: "Reset" })); + + await waitFor(() => expect(POST).toHaveBeenCalledTimes(1)); + expect(mockOnMemberSpendReset).not.toHaveBeenCalled(); + expect(screen.getByRole("dialog", { name: "Reset Team Member Spend" })).toBeInTheDocument(); + }); + + it("does not call the API when the dialog is cancelled", async () => { + const user = userEvent.setup(); + renderEditableTab(); + + await user.click(screen.getByTestId("reset-member-spend")); + const dialog = await screen.findByRole("dialog", { name: "Reset Team Member Spend" }); + await user.click(within(dialog).getByRole("button", { name: "Cancel" })); + + await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); + expect(POST).not.toHaveBeenCalled(); + }); + + it("only offers the reset on members that have current cycle spend", () => { + renderEditableTab(); + + expect( + within(screen.getByRole("row", { name: /user1@test\.com/ })).getByTestId("reset-member-spend"), + ).toBeVisible(); + expect( + within(screen.getByRole("row", { name: /user2@test\.com/ })).queryByTestId("reset-member-spend"), + ).not.toBeInTheDocument(); + }); + + it("hides the reset on the caller's own row for a team admin, since the backend rejects it", () => { + vi.mocked(useAuthorized).mockReturnValue({ userId: "user1@test.com", userRole: "Internal User" } as never); + vi.mocked(isProxyAdminRole).mockReturnValue(false); + renderEditableTab(); + + expect(screen.queryByTestId("reset-member-spend")).not.toBeInTheDocument(); + }); + + it("shows the reset on the caller's own row for a proxy admin", () => { + vi.mocked(useAuthorized).mockReturnValue({ userId: "user1@test.com", userRole: "Admin" } as never); + vi.mocked(isProxyAdminRole).mockReturnValue(true); + renderEditableTab(); + + expect(screen.getByTestId("reset-member-spend")).toBeVisible(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx index 16f3d12d71c..a869c1ad624 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx @@ -1,13 +1,18 @@ +import { useResetTeamMemberSpend } from "@/app/(dashboard)/hooks/teams/useResetTeamMemberSpend"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { Button } from "@/components/ui/button"; +import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog"; import { SimpleTooltip } from "@/components/ui/tooltip"; import MemberTable from "@/components/common_components/MemberTable"; import { Member } from "@/components/networking"; +import { parseErrorMessage } from "@/components/shared/errorUtils"; import { DateCell, MoneyCell } from "@/components/shared/table_cells"; +import { toast } from "@/lib/toast"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { isProxyAdminRole, isUserTeamAdminForSingleTeam } from "@/utils/roles"; import { CircleHelp } from "lucide-react"; -import type { ComponentProps } from "react"; +import { useState, type ComponentProps } from "react"; import { TeamData, TeamMembership } from "./TeamInfo"; export const seedMemberBudgetFields = ( @@ -31,6 +36,7 @@ interface TeamMemberTabProps { setSelectedEditMember: (member: Member) => void; setIsEditMemberModalVisible: (visible: boolean) => void; setIsAddMemberModalVisible: (visible: boolean) => void; + onMemberSpendReset: () => void; } export default function TeamMemberTab({ @@ -40,7 +46,11 @@ export default function TeamMemberTab({ setSelectedEditMember, setIsEditMemberModalVisible, setIsAddMemberModalVisible, + onMemberSpendReset, }: TeamMemberTabProps) { + const [memberToResetSpend, setMemberToResetSpend] = useState(null); + const { mutate: resetMemberSpend, isPending: isResettingSpend } = useResetTeamMemberSpend(); + const formatNumber = (value: number | null): string => { if (value === null || value === undefined) return "0"; @@ -199,24 +209,70 @@ export default function TeamMemberTab({ }, ]; + const handleResetSpend = () => { + if (!memberToResetSpend?.user_id) return; + resetMemberSpend( + { teamId: teamData.team_id, userId: memberToResetSpend.user_id }, + { + onSuccess: () => { + toast.success("Team member spend reset to $0"); + setMemberToResetSpend(null); + onMemberSpendReset(); + }, + onError: (error) => toast.fromError(parseErrorMessage(error)), + }, + ); + }; + return ( - { - const membership = teamData.team_memberships.find((tm) => tm.user_id === record.user_id); - setSelectedEditMember(seedMemberBudgetFields(record, membership?.litellm_budget_table)); - setIsEditMemberModalVisible(true); - }} - onDelete={handleMemberDelete} - onAddMember={() => setIsAddMemberModalVisible(true)} - roleColumnTitle="Team Role" - roleTooltip="This role applies only to this team and is independent from the user's proxy-level role." - extraColumns={extraColumns} - showDeleteForMember={() => - isProxyAdmin || (canEditTeam && !isUserTeamAdmin) || (isUserTeamAdmin && !disableTeamAdminDeleteTeamUser) - } - /> + <> + { + const membership = teamData.team_memberships.find((tm) => tm.user_id === record.user_id); + setSelectedEditMember(seedMemberBudgetFields(record, membership?.litellm_budget_table)); + setIsEditMemberModalVisible(true); + }} + onDelete={handleMemberDelete} + onAddMember={() => setIsAddMemberModalVisible(true)} + roleColumnTitle="Team Role" + roleTooltip="This role applies only to this team and is independent from the user's proxy-level role." + extraColumns={extraColumns} + showDeleteForMember={() => + isProxyAdmin || (canEditTeam && !isUserTeamAdmin) || (isUserTeamAdmin && !disableTeamAdminDeleteTeamUser) + } + onResetSpend={setMemberToResetSpend} + showResetSpendForMember={(record) => + getUserCurrentCycleSpend(record.user_id) > 0 && (isProxyAdmin || record.user_id !== userId) + } + /> + !open && setMemberToResetSpend(null)}> + + + Reset Team Member Spend + +

+ Reset current cycle spend for{" "} + {memberToResetSpend?.user_email || memberToResetSpend?.user_id} in this team to{" "} + $0? +

+

+ Current cycle spend:{" "} + ${formatNumberWithCommas(getUserCurrentCycleSpend(memberToResetSpend?.user_id ?? null), 4)} + . This is the value checked against the member's budget. Total spend and logs are preserved. +

+ + + + +
+
+ ); }