mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
5396810bb6
commit
909a30d6a1
6 changed files with 190 additions and 37 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue