Merge pull request #41349 from BerriAI/litellm_team_member_spend_without_budget

fix(proxy): track team member spend when the member has no budget
This commit is contained in:
ryan-crabbe-berri 2026-09-18 14:04:39 -07:00 committed by GitHub
commit 006080ea6d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 544 additions and 165 deletions

View file

@ -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
@ -136,6 +136,8 @@ class _SpendBatchManager(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: ...
@ -161,6 +163,44 @@ 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),
# 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
# 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))
)
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
"""
async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None:
# key is "team_id::<value>::user_id::<value>"; 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 = 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,
tuple(user_id for _team_id, user_id, _cost in rows),
team_ids,
tuple(cost for _team_id, _user_id, cost in rows),
)
def get_llm_router():
"""The proxy's router, or None outside a running proxy.
@ -1685,21 +1725,7 @@ class DBSpendUpdateWriter:
start_time = time.time()
try:
async with _spend_update_tx(prisma_client) as transaction:
async with transaction.batch_() as batcher:
# Sort by composite key for consistent lock ordering across pods to prevent deadlocks.
# Key format "team_id::<v>::user_id::<v>" 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::<value>::user_id::<value>"
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},
data={
"spend": {"increment": response_cost},
"total_spend": {"increment": response_cost},
},
)
await _write_team_member_spend(transaction, team_member_list_transactions)
# Transaction succeeded, break out of retry loop
break
except Exception as e:

View file

@ -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

View file

@ -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,15 @@ 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, 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": {}},
include={"litellm_budget_table": True},
)

View file

@ -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",

View file

@ -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 DBSpendUpdateWriter
from litellm.proxy.db.db_spend_update_writer import (
_TEAM_ADVISORY_LOCK_SQL,
_TEAM_MEMBER_SPEND_SQL,
DBSpendUpdateWriter,
)
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
build_window_spend_transaction,
)
@ -913,79 +917,118 @@ 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
update_many call, using the same response_cost.
"""
db_writer = DBSpendUpdateWriter()
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.update_many = 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()
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.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_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(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": {},
}
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,
)
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}
assert call_kwargs["data"] == {
"spend": {"increment": response_cost},
"total_spend": {"increment": response_cost},
}
@pytest.mark.asyncio
async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster_checked_upsert():
"""
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 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"
user_id = "user-xyz"
response_cost = 0.75
mock_transaction, mock_prisma_client = _team_member_flush_fixtures()
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}
),
)
lock_call, spend_call = mock_transaction.execute_raw.await_args_list
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 (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
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_orders_team_member_rows_by_team_then_user():
"""
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 `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()
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(
{
"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,
}
),
)
*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, "eng"),
(_TEAM_ADVISORY_LOCK_SQL, "eng-b"),
(_TEAM_ADVISORY_LOCK_SQL, "eng2"),
]
assert list(zip(team_ids, user_ids, costs)) == [
("eng", "user_x", 0.3),
("eng", "user_y", 0.2),
("eng-b", "user_x", 0.4),
("eng2", "user_x", 0.1),
]
@pytest.mark.asyncio
@ -2211,19 +2254,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},
@ -2295,6 +2325,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)
@ -3054,6 +3086,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),

View file

@ -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 = [
@ -2086,7 +2088,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)
@ -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():
"""
@ -5772,7 +5843,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 +5986,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 +6134,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 +9596,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
)

View file

@ -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
@ -424,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
@ -473,7 +497,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,11 +526,12 @@ 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
assert team_membership_call_args.kwargs["data"]["update"] == {}
@pytest.mark.asyncio
@ -546,7 +571,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 +635,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 +725,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 +1056,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 +1131,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 +1187,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 +1232,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()

View file

@ -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<void> => {
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<void, Error, ResetTeamMemberSpendParams>({ mutationFn: resetTeamMemberSpend });

View file

@ -30,6 +30,7 @@ export const TableIconActionButtonMap: Record<string, TableIconActionButtonBaseP
Delete: { icon: TrashIcon, className: "hover:text-destructive" },
Test: { icon: PlayIcon, className: "hover:text-info" },
Regenerate: { icon: RefreshIcon, className: "hover:text-success" },
Reset: { icon: RefreshIcon, className: "hover:text-info" },
Up: { icon: ChevronUpIcon, className: "hover:text-info" },
Down: { icon: ChevronDownIcon, className: "hover:text-info" },
Open: { icon: ExternalLinkIcon, className: "hover:text-success" },

View file

@ -36,6 +36,8 @@ export interface MemberTableProps {
roleTooltip?: string;
extraColumns?: MemberTableColumn[];
showDeleteForMember?: (member: Member) => 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<Member> => {
@ -105,6 +109,8 @@ const buildColumns = ({
roleTooltip,
extraColumns,
showDeleteForMember,
onResetSpend,
showResetSpendForMember,
}: MemberColumnDeps): ColumnDef<Member>[] => [
{
id: "user_alias",
@ -173,6 +179,14 @@ const buildColumns = ({
dataTestId="edit-member"
onClick={() => onEdit(row.original)}
/>
{onResetSpend && (showResetSpendForMember?.(row.original) ?? true) && (
<TableIconActionButton
variant="Reset"
tooltipText="Reset spend"
dataTestId="reset-member-spend"
onClick={() => onResetSpend(row.original)}
/>
)}
{(!showDeleteForMember || showDeleteForMember(row.original)) && (
<TableIconActionButton
variant="Delete"
@ -196,6 +210,8 @@ export default function MemberTable({
roleTooltip,
extraColumns = [],
showDeleteForMember,
onResetSpend,
showResetSpendForMember,
emptyText,
}: MemberTableProps) {
const [globalFilter, setGlobalFilter] = useState("");
@ -210,6 +226,8 @@ export default function MemberTable({
roleTooltip,
extraColumns,
showDeleteForMember,
onResetSpend,
showResetSpendForMember,
};
const columns = buildColumns(columnDeps);
const roleFilterItems = [

View file

@ -678,6 +678,15 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
}
};
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<TeamInfoProps> = ({
teamData={teamData}
canEditTeam={canEditTeam}
handleMemberDelete={handleMemberDelete}
onMemberSpendReset={refreshTeamData}
setSelectedEditMember={setSelectedEditMember}
setIsEditMemberModalVisible={setIsEditMemberModalVisible}
setIsAddMemberModalVisible={setIsAddMemberModalVisible}

View file

@ -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(
<TeamMembersComponent
teamData={createMockTeamData()}
canEditTeam={true}
handleMemberDelete={mockHandleMemberDelete}
onMemberSpendReset={mockOnMemberSpendReset}
setSelectedEditMember={mockSetSelectedEditMember}
setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible}
setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible}
/>,
);
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();
});
});
});

View file

@ -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<Member | null>(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 (
<MemberTable
key={teamData.team_id}
members={teamData.team_info.members_with_roles}
canEdit={canEditTeam}
onEdit={(record) => {
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)
}
/>
<>
<MemberTable
key={teamData.team_id}
members={teamData.team_info.members_with_roles}
canEdit={canEditTeam}
onEdit={(record) => {
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)
}
/>
<Dialog open={memberToResetSpend !== null} onOpenChange={(open) => !open && setMemberToResetSpend(null)}>
<DialogContent>
<DialogHeader>
<DialogTitle>Reset Team Member Spend</DialogTitle>
</DialogHeader>
<p>
Reset current cycle spend for{" "}
<strong>{memberToResetSpend?.user_email || memberToResetSpend?.user_id}</strong> in this team to{" "}
<strong>$0</strong>?
</p>
<p className="text-sm text-muted-foreground">
Current cycle spend:{" "}
<strong>${formatNumberWithCommas(getUserCurrentCycleSpend(memberToResetSpend?.user_id ?? null), 4)}</strong>
. This is the value checked against the member&apos;s budget. Total spend and logs are preserved.
</p>
<DialogFooter>
<Button variant="outline" onClick={() => setMemberToResetSpend(null)}>
Cancel
</Button>
<Button variant="destructive" onClick={handleResetSpend} disabled={isResettingSpend}>
Reset
</Button>
</DialogFooter>
</DialogContent>
</Dialog>
</>
);
}