refactor(team): model the bulk budget audit payload as frozen types

This commit is contained in:
ryan-crabbe-berri 2026-09-17 15:41:06 -07:00
parent 775b83bcf4
commit f04f0258f7
3 changed files with 48 additions and 38 deletions

View file

@ -585,18 +585,18 @@ async def _upsert_budget_and_membership(
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 {}
source: Final[Mapping[str, Any]] = source_row.model_dump() if source_row is not None else MappingProxyType({})
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))},
**MappingProxyType(
{f: source[f] for f in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS if _is_set_budget_value(source.get(f))}
),
**write_data,
}
# 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.
# Restarting an inherited window on an unrelated edit hands the member a free period.
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

View file

@ -11,6 +11,8 @@ from datetime import datetime, timedelta
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from pydantic import BaseModel, ConfigDict
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import (
LiteLLM_TeamTable,
@ -54,14 +56,6 @@ 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]":
@ -93,33 +87,51 @@ 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 _audit_value(value: object) -> object:
return value.isoformat() if isinstance(value, datetime) else value
class _AuditedMemberBudget(BaseModel):
"""One member's limits as the audit log's before/after values record them."""
model_config = ConfigDict(frozen=True)
user_id: str
budget_id: str | None = None
max_budget: float | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
budget_duration: str | None = None
budget_reset_at: datetime | None = None
allowed_models: tuple[str, ...] | None = None
def _limits_audit_value(
rows: "Sequence[prisma_models.LiteLLM_TeamMembership]",
) -> str:
"""Serialize the members' limits for an audit-log value.
class _AuditedMemberBudgets(BaseModel):
"""The audit-log columns hold a JSON object, so the per-member list is nested under a key."""
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.
"""
model_config = ConfigDict(frozen=True)
team_member_budgets: tuple[_AuditedMemberBudget, ...]
def _audited_member_budget(row: "prisma_models.LiteLLM_TeamMembership") -> _AuditedMemberBudget:
budget: Final = row.litellm_budget_table
if budget is None:
return _AuditedMemberBudget(user_id=row.user_id, budget_id=row.budget_id)
return _AuditedMemberBudget(
user_id=row.user_id,
budget_id=row.budget_id,
max_budget=budget.max_budget,
tpm_limit=budget.tpm_limit,
rpm_limit=budget.rpm_limit,
budget_duration=budget.budget_duration,
budget_reset_at=budget.budget_reset_at,
allowed_models=tuple(budget.allowed_models) if budget.allowed_models is not None else None,
)
def _limits_audit_value(rows: "Sequence[prisma_models.LiteLLM_TeamMembership]") -> str:
"""Serialize the members' limits for an audit-log value, dropping the limits they do not set."""
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: _audit_value(getattr(row.litellm_budget_table, field))
for field in _AUDITED_LIMITS
if row.litellm_budget_table is not None and getattr(row.litellm_budget_table, field) is not None
},
}
for row in sorted(rows, key=lambda row: row.user_id)
)
}
_AuditedMemberBudgets(
team_member_budgets=tuple(_audited_member_budget(row) for row in sorted(rows, key=lambda row: row.user_id))
).model_dump(exclude_none=True, mode="json")
)

View file

@ -246,8 +246,6 @@ 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"]
# 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,