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>
This commit is contained in:
ryan 2026-09-16 02:53:47 +00:00 • committed by ryan-crabbe-berri
parent 7c068a4cf7
commit efef6ab684
5 changed files with 87 additions and 164 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
@ -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::<value>::user_id::<value>"; 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::<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]
_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:

View file

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

View file

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

View file

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

View file

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