Merge pull request #35748 from BerriAI/litellm_budget_reset_uow

refactor(repositories): add prisma protocol seams and a spend-reset unit of work
This commit is contained in:
Mateo Wang 2026-08-04 18:06:06 -07:00 • committed by GitHub
commit 4e5cd0b9f5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 236 additions and 65 deletions

View file

@ -1,7 +1,7 @@
import asyncio
import json
import time
from collections.abc import Callable, Mapping, Sequence
from collections.abc import Callable, Sequence
from datetime import datetime, timezone
from typing import Final, Literal, Protocol, TypeVar
@ -23,50 +23,20 @@ from litellm.proxy.common_utils.timezone_utils import (
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.prisma_protocols import ReadOnlyTable, SpendLinkedTable
from litellm.repositories.table_repositories import (
EndUserRepository,
TagRepository,
TeamMembershipRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.unit_of_work import spend_reset_unit_of_work
from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
from litellm.types.services import ServiceTypes
_RowT = TypeVar("_RowT")
_RowT_co = TypeVar("_RowT_co", covariant=True)
class _PrismaRecord(Protocol):
def dict(self) -> Mapping[str, object]: ...
class _BatchTable(Protocol):
def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
class _ResetBatcher(Protocol):
@property
def litellm_verificationtoken(self) -> _BatchTable: ...
@property
def litellm_usertable(self) -> _BatchTable: ...
@property
def litellm_teamtable(self) -> _BatchTable: ...
async def commit(self) -> None: ...
class _EndUserTable(Protocol):
async def find_many(self, where: Mapping[str, object]) -> Sequence[_PrismaRecord]: ...
class _SpendLinkedTable(Protocol[_RowT_co]):
async def find_many(self, where: Mapping[str, object]) -> Sequence[_RowT_co]: ...
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
class _TeamMembershipRow(Protocol):
@ -227,7 +197,7 @@ class ResetBudgetJob:
async def _cascade_reset_spend_for_budget_link(
self,
budgets_to_reset: list[LiteLLM_BudgetTableFull],
table: "_SpendLinkedTable[_RowT]",
table: SpendLinkedTable[_RowT],
counter_key_fn: Callable[[_RowT], str],
log_subject: str,
extra_where: dict[str, object] | None = None,
@ -466,7 +436,7 @@ class ResetBudgetJob:
rely on the default budget (litellm.max_end_user_budget_id) applied
in-memory during auth checks.
"""
table: Final[_EndUserTable] = EndUserRepository(self.prisma_client).table
table: Final[ReadOnlyTable] = EndUserRepository(self.prisma_client).table
rows: Final = await table.find_many(
where={
"budget_id": None,
@ -486,16 +456,11 @@ class ResetBudgetJob:
aborts the entire batch — silently leaving spend over the cap and
budget_reset_at unchanged forever.
"""
batcher: Final[_ResetBatcher] = self.prisma_client.db.batch_()
for k in updated_keys:
token = getattr(k, "token", None)
if token is None:
continue
batcher.litellm_verificationtoken.update(
where={"token": token},
data={"spend": 0, "budget_reset_at": k.budget_reset_at},
)
await batcher.commit()
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
for k in updated_keys:
if k.token is None:
continue
uow.keys.queue_spend_reset(token=k.token, budget_reset_at=k.budget_reset_at)
async def _write_user_reset_updates(self, updated_users: list[LiteLLM_UserTable]) -> None:
"""
@ -505,16 +470,9 @@ class ResetBudgetJob:
that trips Prisma's DataError on rows carrying unrecognised fields
(see #27730).
"""
batcher: Final[_ResetBatcher] = self.prisma_client.db.batch_()
for u in updated_users:
user_id = getattr(u, "user_id", None)
if user_id is None:
continue
batcher.litellm_usertable.update(
where={"user_id": user_id},
data={"spend": 0, "budget_reset_at": u.budget_reset_at},
)
await batcher.commit()
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
for u in updated_users:
uow.users.queue_spend_reset(user_id=u.user_id, budget_reset_at=u.budget_reset_at)
async def _write_team_reset_updates(self, updated_teams: list[LiteLLM_TeamTable]) -> None:
"""
@ -524,16 +482,9 @@ class ResetBudgetJob:
that trips Prisma's DataError on rows carrying unrecognised fields
(see #27730).
"""
batcher: Final[_ResetBatcher] = self.prisma_client.db.batch_()
for t in updated_teams:
team_id = getattr(t, "team_id", None)
if team_id is None:
continue
batcher.litellm_teamtable.update(
where={"team_id": team_id},
data={"spend": 0, "budget_reset_at": t.budget_reset_at},
)
await batcher.commit()
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
for t in updated_teams:
uow.teams.queue_spend_reset(team_id=t.team_id, budget_reset_at=t.budget_reset_at)
async def reset_budget_for_litellm_keys(self):
"""

View file

@ -10,6 +10,13 @@ from litellm.repositories.object_permission_repository import (
ObjectPermissionRepository,
)
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.prisma_protocols import (
BatchTable,
PrismaBatch,
PrismaRecord,
ReadOnlyTable,
SpendLinkedTable,
)
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import (
AccessGroupRepository,
@ -62,6 +69,13 @@ from litellm.repositories.table_repositories import (
WorkflowRunRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.unit_of_work import (
KeySpendResetWrites,
SpendResetUnitOfWork,
TeamSpendResetWrites,
UserSpendResetWrites,
spend_reset_unit_of_work,
)
from litellm.repositories.user_repository import UserRepository
from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
@ -73,6 +87,7 @@ __all__ = [
"AdaptiveRouterStateRepository",
"AgentsRepository",
"AuditLogRepository",
"BatchTable",
"BudgetRepository",
"CacheConfigRepository",
"ClaudeCodePluginRepository",
@ -91,6 +106,7 @@ __all__ = [
"HealthCheckRepository",
"InvitationLinkRepository",
"JWTKeyMappingRepository",
"KeySpendResetWrites",
"MCPServerRepository",
"MCPToolsetRepository",
"MCPUserCredentialsRepository",
@ -106,24 +122,32 @@ __all__ = [
"OrganizationRepository",
"PolicyAttachmentRepository",
"PolicyRepository",
"PrismaBatch",
"PrismaRecord",
"PrismaTableRepository",
"ProjectRepository",
"PromptRepository",
"ReadOnlyTable",
"SSOConfigRepository",
"SearchToolsRepository",
"SkillsRepository",
"SpendLinkedTable",
"SpendLogGuardrailIndexRepository",
"SpendLogToolIndexRepository",
"SpendLogsRepository",
"SpendResetUnitOfWork",
"TagRepository",
"TeamMembershipRepository",
"TeamRepository",
"TeamSpendResetWrites",
"ToolRepository",
"UISettingsRepository",
"UserNotificationsRepository",
"UserRepository",
"UserSpendResetWrites",
"VerificationTokenRepository",
"WorkflowEventRepository",
"WorkflowMessageRepository",
"WorkflowRunRepository",
"spend_reset_unit_of_work",
]

View file

@ -0,0 +1,43 @@
"""
Typed Protocol seams over prisma-client-py surfaces.
Modules that reach Prisma through an untyped handle (``prisma_client.db`` or a
repository ``.table``) annotate against these Protocols instead of hand-rolling
private ones per file.
"""
from collections.abc import Mapping, Sequence
from typing import Protocol, TypeVar
RowT_co = TypeVar("RowT_co", covariant=True)
class PrismaRecord(Protocol):
def dict(self) -> Mapping[str, object]: ...
class ReadOnlyTable(Protocol):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[PrismaRecord]: ...
class SpendLinkedTable(Protocol[RowT_co]):
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[RowT_co]: ...
async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
class BatchTable(Protocol):
def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
class PrismaBatch(Protocol):
@property
def litellm_verificationtoken(self) -> BatchTable: ...
@property
def litellm_usertable(self) -> BatchTable: ...
@property
def litellm_teamtable(self) -> BatchTable: ...
async def commit(self) -> None: ...

View file

@ -0,0 +1,61 @@
"""
Unit of work over a single Prisma batch.
``spend_reset_unit_of_work`` opens one ``db.batch_()`` and binds a typed write
repository per table to it, so every update queued through the yielded object
lands in the same transaction. The batch commits when the block exits cleanly
and is abandoned, writing nothing, when the block raises.
Each write repository queues narrow ``{spend, budget_reset_at}`` updates
instead of full-model writes, which trip ``prisma.errors.DataError`` on rows
carrying fields the update input type rejects (see #27730).
"""
from collections.abc import AsyncGenerator, Callable
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
from litellm.repositories.prisma_protocols import BatchTable, PrismaBatch
@dataclass(frozen=True, slots=True)
class KeySpendResetWrites:
table: BatchTable
def queue_spend_reset(self, token: str, budget_reset_at: datetime | None) -> None:
self.table.update(where={"token": token}, data={"spend": 0, "budget_reset_at": budget_reset_at})
@dataclass(frozen=True, slots=True)
class UserSpendResetWrites:
table: BatchTable
def queue_spend_reset(self, user_id: str, budget_reset_at: datetime | None) -> None:
self.table.update(where={"user_id": user_id}, data={"spend": 0, "budget_reset_at": budget_reset_at})
@dataclass(frozen=True, slots=True)
class TeamSpendResetWrites:
table: BatchTable
def queue_spend_reset(self, team_id: str, budget_reset_at: datetime | None) -> None:
self.table.update(where={"team_id": team_id}, data={"spend": 0, "budget_reset_at": budget_reset_at})
@dataclass(frozen=True, slots=True)
class SpendResetUnitOfWork:
keys: KeySpendResetWrites
users: UserSpendResetWrites
teams: TeamSpendResetWrites
@asynccontextmanager
async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> AsyncGenerator[SpendResetUnitOfWork, None]:
batch = new_batch()
yield SpendResetUnitOfWork(
keys=KeySpendResetWrites(table=batch.litellm_verificationtoken),
users=UserSpendResetWrites(table=batch.litellm_usertable),
teams=TeamSpendResetWrites(table=batch.litellm_teamtable),
)
await batch.commit()

View file

@ -14,6 +14,7 @@ import pytest
sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings
from litellm.proxy.utils import ProxyLogging
@ -218,6 +219,31 @@ async def run_async_test(coro):
# Tests
def test_write_key_reset_updates_skips_none_token_and_still_writes_the_rest(reset_budget_job, mock_prisma_client):
"""A key with token=None must be skipped, not queued as where={"token": None}.
Queueing a None token makes the prisma batch commit raise and aborts the
whole batch, silently dropping every key reset that cycle (the #27730
blast radius this write path exists to prevent).
"""
reset_at = datetime.now(timezone.utc)
keys = [
LiteLLM_VerificationToken(token=None, budget_reset_at=reset_at),
LiteLLM_VerificationToken(token="tok-ok", budget_reset_at=reset_at),
]
asyncio.run(reset_budget_job._write_key_reset_updates(updated_keys=keys))
key_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "key"]
assert key_writes == [
{
"table": "key",
"where": {"token": "tok-ok"},
"data": {"spend": 0, "budget_reset_at": reset_at},
}
]
def test_reset_budget_for_key(reset_budget_job, mock_prisma_client):
# Setup test data with timezone-aware datetime
now = datetime.now(timezone.utc)

View file

@ -0,0 +1,66 @@
from datetime import datetime, timezone
from typing import Any, Dict, List, Mapping, Tuple
import pytest
from litellm.repositories.unit_of_work import spend_reset_unit_of_work
class FakeBatchTable:
def __init__(self, table_name: str, calls: List[Tuple[str, Dict[str, Any], Dict[str, Any]]]):
self._table_name = table_name
self._calls = calls
def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> None:
self._calls.append((self._table_name, dict(where), dict(data)))
class FakeBatch:
def __init__(self):
self.calls: List[Tuple[str, Dict[str, Any], Dict[str, Any]]] = []
self.commit_count = 0
self.litellm_verificationtoken = FakeBatchTable("litellm_verificationtoken", self.calls)
self.litellm_usertable = FakeBatchTable("litellm_usertable", self.calls)
self.litellm_teamtable = FakeBatchTable("litellm_teamtable", self.calls)
async def commit(self) -> None:
self.commit_count += 1
async def test_updates_across_tables_share_one_batch_and_commit_once():
batch = FakeBatch()
reset_at = datetime(2026, 8, 3, 12, 0, tzinfo=timezone.utc)
async with spend_reset_unit_of_work(lambda: batch) as uow:
uow.keys.queue_spend_reset(token="tok-1", budget_reset_at=reset_at)
uow.users.queue_spend_reset(user_id="user-1", budget_reset_at=reset_at)
uow.teams.queue_spend_reset(team_id="team-1", budget_reset_at=None)
assert batch.commit_count == 0
assert batch.commit_count == 1
assert batch.calls == [
("litellm_verificationtoken", {"token": "tok-1"}, {"spend": 0, "budget_reset_at": reset_at}),
("litellm_usertable", {"user_id": "user-1"}, {"spend": 0, "budget_reset_at": reset_at}),
("litellm_teamtable", {"team_id": "team-1"}, {"spend": 0, "budget_reset_at": None}),
]
async def test_raising_inside_block_skips_commit():
batch = FakeBatch()
with pytest.raises(RuntimeError, match="boom"):
async with spend_reset_unit_of_work(lambda: batch) as uow:
uow.keys.queue_spend_reset(token="tok-1", budget_reset_at=None)
raise RuntimeError("boom")
assert batch.commit_count == 0
async def test_empty_block_still_commits_the_batch():
batch = FakeBatch()
async with spend_reset_unit_of_work(lambda: batch):
pass
assert batch.commit_count == 1
assert batch.calls == []