diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 450dc12a9ff..39fdc0216a0 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -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): """ diff --git a/litellm/repositories/__init__.py b/litellm/repositories/__init__.py index 1fc3d8dadaf..4f020480f9e 100644 --- a/litellm/repositories/__init__.py +++ b/litellm/repositories/__init__.py @@ -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", ] diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py new file mode 100644 index 00000000000..6aff196ff10 --- /dev/null +++ b/litellm/repositories/prisma_protocols.py @@ -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: ... diff --git a/litellm/repositories/unit_of_work.py b/litellm/repositories/unit_of_work.py new file mode 100644 index 00000000000..682e69d11eb --- /dev/null +++ b/litellm/repositories/unit_of_work.py @@ -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() diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index f04d6f3cf5a..616ad8a0981 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -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) diff --git a/tests/test_litellm/repositories/test_unit_of_work.py b/tests/test_litellm/repositories/test_unit_of_work.py new file mode 100644 index 00000000000..35f102bbb9d --- /dev/null +++ b/tests/test_litellm/repositories/test_unit_of_work.py @@ -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 == []