mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
4e5cd0b9f5
6 changed files with 236 additions and 65 deletions
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
43
litellm/repositories/prisma_protocols.py
Normal file
43
litellm/repositories/prisma_protocols.py
Normal 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: ...
|
||||
61
litellm/repositories/unit_of_work.py
Normal file
61
litellm/repositories/unit_of_work.py
Normal 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()
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
66
tests/test_litellm/repositories/test_unit_of_work.py
Normal file
66
tests/test_litellm/repositories/test_unit_of_work.py
Normal 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 == []
|
||||
Loading…
Add table
Reference in a new issue