fix(team): keep a forked member budget's reset window and audit bulk member budget writes

Forking a shared budget row rebuilt budget_reset_at from the duration, so editing
an unrelated limit restarted the member's window while their spend carried over:
a tpm bump quietly handed them a fresh period. The fork now inherits the source
row's deadline, and only recomputes when the patch actually sets budget_duration.

The bulk member budget route now writes one audit entry per call, a team-scoped
'updated' row carrying every written member's limits before and after, matching
what /team/member_add already records for membership changes. It honors the
litellm-changed-by header like the other audited team routes.
This commit is contained in:
ryan-crabbe-berri 2026-09-17 15:17:00 -07:00
parent 5396810bb6
commit 909a30d6a1
6 changed files with 190 additions and 37 deletions

View file

@ -582,25 +582,25 @@ async def _upsert_budget_and_membership(
)
return
create_data: Final[dict[str, Any]] = {
source_row: Final = (
await tx.litellm_budgettable.find_unique(where={"budget_id": existing_budget_id}) if is_shared_default else None
)
source: Final[Mapping[str, Any]] = source_row.model_dump() if source_row is not None else {}
create_data: Final[dict[str, Any]] = { # mutable-ok: Prisma create payloads are dict-shaped
"created_by": user_api_key_dict.user_id or "",
"updated_by": user_api_key_dict.user_id or "",
**{f: source[f] for f in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS if _is_set_budget_value(source.get(f))},
**write_data,
}
if is_shared_default:
default_budget_row: Final = await tx.litellm_budgettable.find_unique(where={"budget_id": existing_budget_id})
if default_budget_row is not None:
default_budget_dict: Final = default_budget_row.model_dump()
for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS:
value = default_budget_dict.get(field)
if _is_set_budget_value(value):
create_data[field] = value
create_data.update(write_data)
if create_data.get("budget_duration") is not None:
create_data["budget_reset_at"] = get_budget_reset_time(budget_duration=create_data["budget_duration"])
else:
# A patch that leaves the reset cadence alone must not move the deadline: the clone
# inherits the source row's window instead of restarting it from now, which would
# silently grant a member a fresh period whenever any unrelated limit is edited.
carried: Final = source.get("budget_reset_at") if "budget_duration" not in budget_patch else None
if carried is not None:
create_data["budget_reset_at"] = carried
if create_data.get("budget_reset_at") is None:
create_data.pop("budget_reset_at", None)
if not _has_meaningful_budget_limit(create_data):

View file

@ -2,7 +2,7 @@
from typing import Annotated, Final
from fastapi import APIRouter, Depends
from fastapi import APIRouter, Depends, Header
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
@ -108,6 +108,12 @@ async def bulk_update_team_member_budgets_action(
team_id: str,
data: BulkTeamMemberBudgetUpdateRequest,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
litellm_changed_by: Annotated[
str | None,
Header(
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
),
] = None,
) -> BulkTeamMemberBudgetUpdateResponse:
"""
Set per-member limits for up to 500 members of one team in one call. Same
@ -135,7 +141,7 @@ async def bulk_update_team_member_budgets_action(
```
"""
try:
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client, user_api_key_cache
if prisma_client is None:
raise ManagementProblem(
@ -153,6 +159,8 @@ async def bulk_update_team_member_budgets_action(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_admin_name=litellm_proxy_admin_name,
litellm_changed_by=litellm_changed_by,
)
return BulkTeamMemberBudgetUpdateResponse(data=results)

View file

@ -11,7 +11,14 @@ from datetime import timedelta
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, Member, UserAPIKeyAuth
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import (
LiteLLM_TeamTable,
LitellmTableNames,
LitellmUserRoles,
Member,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import invalidate_team_member_spend_state
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
@ -21,6 +28,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage] # the single-member write, shared so the two surfaces cannot drift
member_budget_patch,
)
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
from litellm.proxy.management_helpers.bulk_user_deletion import (
_duplicate_member_indexes, # pyright: ignore[reportPrivateUsage] # same duplicate rule as members/bulk_delete
_eq_filter, # pyright: ignore[reportPrivateUsage] # same prisma filter shape as members/bulk_delete
@ -46,6 +54,14 @@ if TYPE_CHECKING:
_BATCH_TX_TIMEOUT: Final = timedelta(seconds=60)
_NO_METADATA: Final = MappingProxyType({})
_WITH_BUDGET: Final = MappingProxyType({"litellm_budget_table": True})
_AUDITED_LIMITS: Final = (
"max_budget",
"tpm_limit",
"rpm_limit",
"budget_duration",
"budget_reset_at",
"allowed_models",
)
def _membership_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamMembership]":
@ -77,6 +93,32 @@ async def _shared_budget_ids(tx: "Prisma", budget_ids: frozenset[str]) -> frozen
return frozenset(budget_id for budget_id in budget_ids if sum(1 for row in rows if row.budget_id == budget_id) > 1)
def _limits_audit_value(
rows: "Sequence[prisma_models.LiteLLM_TeamMembership]",
) -> str:
"""Serialize the members' limits for an audit-log value.
The audit-log columns hold a JSON object, so the per-member list is nested under a
key rather than serialized as a top-level array.
"""
return safe_dumps(
{ # mutable-ok: the audit-log JSON column rejects a top-level array, so this value must be an object
"team_member_budgets": tuple(
{
"user_id": row.user_id,
"budget_id": row.budget_id,
**{
field: getattr(row.litellm_budget_table, field)
for field in _AUDITED_LIMITS
if row.litellm_budget_table is not None
},
}
for row in sorted(rows, key=lambda row: row.user_id)
)
}
)
def _result(
member: TeamMemberBudgetPatch,
user_id: str | None,
@ -114,6 +156,8 @@ async def bulk_update_team_member_budgets(
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
litellm_proxy_admin_name: str,
litellm_changed_by: str | None = None,
) -> tuple[TeamMemberBudgetUpdateResult, ...]:
"""Apply one merge patch of per-member limits per requested member, in one transaction."""
team: Final = await TeamRepository(WriterPinnedClient(prisma_client.db)).find_by_id(team_id)
@ -151,7 +195,7 @@ async def bulk_update_team_member_budgets(
team_members_filter: Final = _team_users_filter(team_id, user_ids)
async with prisma_client.tx(timeout=_BATCH_TX_TIMEOUT) as tx:
memberships: Final = await _membership_tx_db(tx).find_many(where=team_members_filter)
memberships: Final = await _membership_tx_db(tx).find_many(where=team_members_filter, include=_WITH_BUDGET)
budget_id_of: Final = MappingProxyType({m.user_id: m.budget_id for m in memberships})
shared: Final = await _shared_budget_ids(
tx, frozenset(budget_id for budget_id in budget_id_of.values() if budget_id is not None)
@ -179,6 +223,17 @@ async def bulk_update_team_member_budgets(
user_id=user_id, team_id=team_id, user_api_key_cache=user_api_key_cache
)
await create_object_audit_log(
object_id=team_id,
action="updated",
litellm_changed_by=litellm_changed_by,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
table_name=LitellmTableNames.TEAM_TABLE_NAME,
before_value=_limits_audit_value(memberships),
after_value=_limits_audit_value(written),
)
budget_of: Final = MappingProxyType({m.user_id: m.litellm_budget_table for m in written})
return tuple(
_result(

View file

@ -1,6 +1,6 @@
# tests/litellm/proxy/common_utils/test_upsert_budget_membership.py
import types
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -27,9 +27,7 @@ def mock_tx():
budget = MagicMock()
budget.update = AsyncMock()
budget.find_unique = AsyncMock(return_value=None)
budget.create = AsyncMock(
return_value=types.SimpleNamespace(budget_id="new-budget-123")
)
budget.create = AsyncMock(return_value=types.SimpleNamespace(budget_id="new-budget-123"))
tx = MagicMock()
tx.litellm_teammembership = membership
@ -83,9 +81,7 @@ async def test_empty_patch_is_noop(mock_tx, fake_user):
# member falls back to the team default instead of keeping an empty private row.
@pytest.mark.asyncio
async def test_clearing_all_limits_disconnects(mock_tx, fake_user):
mock_tx.litellm_budgettable.find_unique = AsyncMock(
return_value=budget_row(max_budget=100.0)
)
mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=100.0))
await _upsert_budget_and_membership(
mock_tx,
@ -136,9 +132,7 @@ async def test_clear_one_field_keeps_others(mock_tx, fake_user):
# budget_reset_at, so the budget rolls over without waiting for the reset cron.
@pytest.mark.asyncio
async def test_update_in_place_seeds_reset_at(mock_tx, fake_user):
mock_tx.litellm_budgettable.find_unique = AsyncMock(
return_value=budget_row(max_budget=20.0)
)
mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=20.0))
await _upsert_budget_and_membership(
mock_tx,
@ -163,9 +157,7 @@ async def test_update_in_place_seeds_reset_at(mock_tx, fake_user):
# budget_duration must not get a (re)computed reset time.
@pytest.mark.asyncio
async def test_update_in_place_single_field_leaves_reset_at_alone(mock_tx, fake_user):
mock_tx.litellm_budgettable.find_unique = AsyncMock(
return_value=budget_row(max_budget=50.0)
)
mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=50.0))
await _upsert_budget_and_membership(
mock_tx,
@ -225,6 +217,7 @@ async def test_create_seeds_reset_at_and_links(mock_tx, fake_user):
@pytest.mark.asyncio
async def test_clone_on_write_from_shared_default(mock_tx, fake_user):
shared_default_id = "team-default-budget-1"
shared_reset_at = datetime.now(timezone.utc) + timedelta(hours=3)
mock_tx.litellm_budgettable.find_unique = AsyncMock(
return_value=budget_row(
budget_id=shared_default_id,
@ -235,6 +228,7 @@ async def test_clone_on_write_from_shared_default(mock_tx, fake_user):
rpm_limit=None,
model_max_budget=None,
budget_duration="1d",
budget_reset_at=shared_reset_at,
allowed_models=[],
)
)
@ -252,7 +246,9 @@ async def test_clone_on_write_from_shared_default(mock_tx, fake_user):
mock_tx.litellm_budgettable.update.assert_not_called()
mock_tx.litellm_budgettable.create.assert_awaited_once()
create_data = mock_tx.litellm_budgettable.create.await_args.kwargs["data"]
assert_future_reset_time(create_data.pop("budget_reset_at"))
# The patch never touched budget_duration, so the fork keeps the window it
# inherited: restarting it here would hand the member a fresh period for free.
assert create_data.pop("budget_reset_at") == shared_reset_at
assert create_data == {
"created_by": fake_user.user_id,
"updated_by": fake_user.user_id,
@ -318,9 +314,7 @@ async def test_clone_on_write_clears_duration(mock_tx, fake_user):
# team default), we update it in place rather than forking another row.
@pytest.mark.asyncio
async def test_private_budget_updates_in_place(mock_tx, fake_user):
mock_tx.litellm_budgettable.find_unique = AsyncMock(
return_value=budget_row(max_budget=10.0)
)
mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=10.0))
await _upsert_budget_and_membership(
mock_tx,

View file

@ -7,6 +7,7 @@ table and the membership/budget relation the bulk budget writer needs.
"""
import copy
import json
from collections.abc import Mapping, Sequence
from contextlib import asynccontextmanager
from datetime import datetime, timedelta, timezone
@ -244,6 +245,7 @@ def _budget(
tpm_limit: int | None = None,
rpm_limit: int | None = None,
budget_duration: str | None = None,
budget_reset_at: datetime | None = None,
) -> _BudgetRow:
return _BudgetRow(
budget_id=budget_id,
@ -251,6 +253,7 @@ def _budget(
tpm_limit=tpm_limit,
rpm_limit=rpm_limit,
budget_duration=budget_duration,
budget_reset_at=budget_reset_at,
)
@ -267,6 +270,7 @@ async def _bulk_update(
user_api_key_dict=caller,
prisma_client=prisma, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient
user_api_key_cache=cache or UserApiKeyCache(),
litellm_proxy_admin_name="default_user_id",
)
@ -631,6 +635,95 @@ async def test_the_roster_authz_read_runs_on_the_writer_so_a_lagging_replica_can
assert writer.db.litellm_budgettable.rows["priv-m1"].max_budget == 1.0
@pytest.mark.asyncio
async def test_the_batch_writes_one_audit_entry_carrying_every_written_members_limits_before_and_after(monkeypatch):
import litellm
from litellm.proxy._types import LitellmTableNames
monkeypatch.setattr(litellm, "store_audit_logs", True)
captured: list[object] = []
async def capture(request_data):
captured.append(request_data)
monkeypatch.setattr("litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update", capture)
prisma = _FakePrisma(
teams=[_team("m1", "m2")],
memberships=[_membership("m1", "priv-m1"), _membership("m2", "priv-m2")],
budgets=[_budget("priv-m1", max_budget=1.0), _budget("priv-m2", max_budget=2.0)],
)
await _bulk_update(prisma, [{"user_id": "m1", "max_budget_in_team": 10}])
assert len(captured) == 1
entry = captured[0]
assert (entry.object_id, entry.action, entry.table_name) == (
TEAM_ID,
"updated",
LitellmTableNames.TEAM_TABLE_NAME,
)
before = {row["user_id"]: row for row in json.loads(entry.before_value)["team_member_budgets"]}
after = {row["user_id"]: row for row in json.loads(entry.updated_values)["team_member_budgets"]}
assert (before["m1"]["max_budget"], after["m1"]["max_budget"]) == (1.0, 10.0)
assert "m2" not in before and "m2" not in after
@pytest.mark.asyncio
async def test_no_audit_entry_is_written_when_audit_logging_is_off(monkeypatch):
import litellm
monkeypatch.setattr(litellm, "store_audit_logs", False)
captured: list[object] = []
async def capture(request_data):
captured.append(request_data)
monkeypatch.setattr("litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update", capture)
prisma = _FakePrisma(
teams=[_team("m1")],
memberships=[_membership("m1", "priv-m1")],
budgets=[_budget("priv-m1", max_budget=1.0)],
)
await _bulk_update(prisma, [{"user_id": "m1", "max_budget_in_team": 10}])
assert captured == []
assert _budget_of(prisma, "m1").max_budget == 10.0
@pytest.mark.asyncio
async def test_forking_a_shared_row_keeps_its_reset_window_so_an_unrelated_limit_edit_grants_no_free_period():
shared_reset_at = datetime.now(timezone.utc) + timedelta(days=3)
prisma = _FakePrisma(
teams=[_team("m1", "m2")],
memberships=[_membership("m1", "shared-b"), _membership("m2", "shared-b")],
budgets=[_budget("shared-b", max_budget=100.0, budget_duration="30d", budget_reset_at=shared_reset_at)],
)
results = await _bulk_update(prisma, [{"user_id": "m1", "tpm_limit": 9}])
assert [(r.success, r.budget_duration) for r in results] == [(True, "30d")]
assert _budget_id_of(prisma, "m1") not in (None, "shared-b")
assert _budget_of(prisma, "m1").budget_reset_at == shared_reset_at
assert prisma.db.litellm_budgettable.rows["shared-b"].budget_reset_at == shared_reset_at
@pytest.mark.asyncio
async def test_forking_a_shared_row_does_restart_the_window_when_the_patch_sets_a_new_duration():
shared_reset_at = datetime.now(timezone.utc) + timedelta(days=3)
prisma = _FakePrisma(
teams=[_team("m1", "m2")],
memberships=[_membership("m1", "shared-b"), _membership("m2", "shared-b")],
budgets=[_budget("shared-b", max_budget=100.0, budget_duration="30d", budget_reset_at=shared_reset_at)],
)
await _bulk_update(prisma, [{"user_id": "m1", "budget_duration": "1d"}])
forked = _budget_of(prisma, "m1").budget_reset_at
assert forked is not None and forked != shared_reset_at
assert forked <= datetime.now(timezone.utc) + timedelta(days=1)
app = FastAPI()

View file

@ -52186,7 +52186,10 @@ export interface operations {
bulk_update_team_member_budgets_action_management_v1_teams__team_id__members_bulk_update_post: {
parameters: {
query?: never;
header?: never;
header?: {
/** @description The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability */
"litellm-changed-by"?: string | null;
};
path: {
team_id: string;
};