Merge pull request #41279 from BerriAI/litellm_reset_budget_decrement_spend

fix(proxy): reset budgets by decrementing pre-reset spend instead of zeroing rows
This commit is contained in:
Yassin Kortam 2026-09-15 14:45:03 -07:00 committed by GitHub
commit fa09de9e45
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 264 additions and 111 deletions

View file

@ -7,7 +7,7 @@ from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from enum import Enum
from types import MappingProxyType
from typing import Final, Literal, Protocol, TypeVar
from typing import Final, Generic, Literal, Protocol, TypeVar
from typing_extensions import assert_never
@ -68,6 +68,13 @@ from litellm.types.services import ServiceTypes
_RowT = TypeVar("_RowT")
@dataclass(frozen=True, slots=True)
class _RowReset(Generic[_RowT]):
row: _RowT
spend_decrement: float
_LINKED_KEYS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"budget_duration": None, "spend": {"gt": 0}})
_SPENT_ROWS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"spend": {"gt": 0}})
@ -530,10 +537,9 @@ class ResetBudgetJob:
)
@staticmethod
async def _invalidate_spend_counter(counter_key: str, new_spend: float = 0.0) -> None:
"""Overwrite a spend counter with the post-reset value (0, or the carried
overage when budget rollover is enabled) so a DB-row reset takes effect
immediately.
async def _invalidate_spend_counter(counter_key: str) -> None:
"""Drop a spend counter so the next read reseeds from the committed DB
row, the only value that includes increments that raced the reset.
Call AFTER the DB write commits. Clearing Redis before the DB
commit opens a window where get_current_spend reads 0 from Redis
@ -542,10 +548,10 @@ class ResetBudgetJob:
try:
from litellm.proxy.proxy_server import spend_counter_cache
spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=new_spend, ttl=60)
spend_counter_cache.in_memory_cache.delete_cache(key=counter_key)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=new_spend, ttl=60)
await spend_counter_cache.redis_cache.async_delete_cache(key=counter_key)
except Exception as redis_err:
verbose_proxy_logger.warning(
"Failed to reset spend counter %s in Redis: %s. "
@ -730,8 +736,8 @@ class ResetBudgetJob:
uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at)
async def _invalidate_budget_cascade_caches(self, cascade: _BudgetCascade) -> None:
for counter_key, new_spend in cascade.counter_resets:
await self._invalidate_spend_counter(counter_key, new_spend=new_spend)
for counter_key, _ in cascade.counter_resets:
await self._invalidate_spend_counter(counter_key)
for cache_key in cascade.cache_keys:
await self._invalidate_user_api_key_cache_entry(cache_key)
@ -842,7 +848,7 @@ class ResetBudgetJob:
)
return [LiteLLM_EndUserTable.model_validate(row.model_dump()) for row in rows]
async def _write_key_reset_updates(self, updated_keys: list[LiteLLM_VerificationToken]) -> None:
async def _write_key_reset_updates(self, updated_keys: Sequence[_RowReset[LiteLLM_VerificationToken]]) -> None:
"""
Write per-row {spend, budget_reset_at} updates for keys.
@ -858,18 +864,18 @@ class ResetBudgetJob:
reason="reset_budget_write_keys_failure",
)
async def _write_key_reset_updates_once(self, updated_keys: list[LiteLLM_VerificationToken]) -> None:
async def _write_key_reset_updates_once(self, updated_keys: Sequence[_RowReset[LiteLLM_VerificationToken]]) -> None:
async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow:
for k in updated_keys:
if k.token is None:
if k.row.token is None:
continue
uow.keys.queue_spend_reset(
token=k.token,
budget_reset_at=k.budget_reset_at,
spend_decrement=k.max_budget if (k.spend or 0.0) > 0.0 else None,
token=k.row.token,
budget_reset_at=k.row.budget_reset_at,
spend_decrement=k.spend_decrement,
)
async def _write_user_reset_updates(self, updated_users: list[LiteLLM_UserTable]) -> None:
async def _write_user_reset_updates(self, updated_users: Sequence[_RowReset[LiteLLM_UserTable]]) -> None:
"""
Write per-row {spend, budget_reset_at} updates for users.
@ -882,16 +888,16 @@ class ResetBudgetJob:
reason="reset_budget_write_users_failure",
)
async def _write_user_reset_updates_once(self, updated_users: list[LiteLLM_UserTable]) -> None:
async def _write_user_reset_updates_once(self, updated_users: Sequence[_RowReset[LiteLLM_UserTable]]) -> None:
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,
spend_decrement=u.max_budget if (u.spend or 0.0) > 0.0 else None,
user_id=u.row.user_id,
budget_reset_at=u.row.budget_reset_at,
spend_decrement=u.spend_decrement,
)
async def _write_team_reset_updates(self, updated_teams: list[LiteLLM_TeamTable]) -> None:
async def _write_team_reset_updates(self, updated_teams: Sequence[_RowReset[LiteLLM_TeamTable]]) -> None:
"""
Write per-row {spend, budget_reset_at} updates for teams.
@ -904,13 +910,13 @@ class ResetBudgetJob:
reason="reset_budget_write_teams_failure",
)
async def _write_team_reset_updates_once(self, updated_teams: list[LiteLLM_TeamTable]) -> None:
async def _write_team_reset_updates_once(self, updated_teams: Sequence[_RowReset[LiteLLM_TeamTable]]) -> None:
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,
spend_decrement=t.max_budget if (t.spend or 0.0) > 0.0 else None,
team_id=t.row.team_id,
budget_reset_at=t.row.budget_reset_at,
spend_decrement=t.spend_decrement,
)
def _emit_phase_failure(
@ -962,18 +968,24 @@ class ResetBudgetJob:
reason="reset_budget_read_keys_failure",
)
verbose_proxy_logger.debug("Keys to reset %s", _LazyJson(keys_to_reset))
updated_keys: Final[list[LiteLLM_VerificationToken]] = []
updated_keys: Final[list[_RowReset[LiteLLM_VerificationToken]]] = []
failed_keys: Final = []
if keys_to_reset is not None and len(keys_to_reset) > 0:
for key in keys_to_reset:
try:
pre_reset_spend = float(key.spend or 0.0)
updated_key = await ResetBudgetJob._reset_budget_for_key(
key=key,
current_time=now,
reset_settings=self.reset_settings,
)
if updated_key is not None:
updated_keys.append(updated_key)
updated_keys.append(
_RowReset(
row=updated_key,
spend_decrement=pre_reset_spend - float(updated_key.spend or 0.0),
)
)
else:
failed_keys.append({"key": key, "error": "Returned None without exception"})
except Exception as e:
@ -985,15 +997,15 @@ class ResetBudgetJob:
if updated_keys:
await self._write_key_reset_updates(updated_keys=updated_keys)
for k in updated_keys:
token = getattr(k, "token", None)
token = getattr(k.row, "token", None)
if token:
await self._invalidate_spend_counter(f"spend:key:{token}", new_spend=k.spend or 0.0)
await self._invalidate_spend_counter(f"spend:key:{token}")
end_time = time.time()
outcome: Final = _ChunkOutcome(
fetched=len(keys_to_reset) if keys_to_reset else 0,
advanced=_count_advanced(
(k.budget_reset_at for k in updated_keys),
(k.row.budget_reset_at for k in updated_keys),
cutoff=datetime.now(timezone.utc),
),
)
@ -1063,18 +1075,24 @@ class ResetBudgetJob:
),
reason="reset_budget_read_users_failure",
)
updated_users: Final[list[LiteLLM_UserTable]] = []
updated_users: Final[list[_RowReset[LiteLLM_UserTable]]] = []
failed_users: Final = []
if users_to_reset is not None and len(users_to_reset) > 0:
for user in users_to_reset:
try:
pre_reset_spend = float(user.spend or 0.0)
updated_user = await ResetBudgetJob._reset_budget_for_user(
user=user,
current_time=now,
reset_settings=self.reset_settings,
)
if updated_user is not None:
updated_users.append(updated_user)
updated_users.append(
_RowReset(
row=updated_user,
spend_decrement=pre_reset_spend - float(updated_user.spend or 0.0),
)
)
else:
failed_users.append(
{
@ -1090,9 +1108,9 @@ class ResetBudgetJob:
if updated_users:
await self._write_user_reset_updates(updated_users=updated_users)
for u in updated_users:
user_id = getattr(u, "user_id", None)
user_id = getattr(u.row, "user_id", None)
if user_id:
await self._invalidate_spend_counter(f"spend:user:{user_id}", new_spend=u.spend or 0.0)
await self._invalidate_spend_counter(f"spend:user:{user_id}")
if user_id == LITELLM_PROXY_BUDGET_NAME:
await self._invalidate_global_proxy_spend_cache()
@ -1100,7 +1118,7 @@ class ResetBudgetJob:
outcome: Final = _ChunkOutcome(
fetched=len(users_to_reset) if users_to_reset else 0,
advanced=_count_advanced(
(u.budget_reset_at for u in updated_users),
(u.row.budget_reset_at for u in updated_users),
cutoff=datetime.now(timezone.utc),
),
)
@ -1172,18 +1190,24 @@ class ResetBudgetJob:
),
reason="reset_budget_read_teams_failure",
)
updated_teams: Final[list[LiteLLM_TeamTable]] = []
updated_teams: Final[list[_RowReset[LiteLLM_TeamTable]]] = []
failed_teams: Final = []
if teams_to_reset is not None and len(teams_to_reset) > 0:
for team in teams_to_reset:
try:
pre_reset_spend = float(team.spend or 0.0)
updated_team = await ResetBudgetJob._reset_budget_for_team(
team=team,
current_time=now,
reset_settings=self.reset_settings,
)
if updated_team is not None:
updated_teams.append(updated_team)
updated_teams.append(
_RowReset(
row=updated_team,
spend_decrement=pre_reset_spend - float(updated_team.spend or 0.0),
)
)
else:
failed_teams.append(
{
@ -1199,15 +1223,15 @@ class ResetBudgetJob:
if updated_teams:
await self._write_team_reset_updates(updated_teams=updated_teams)
for t in updated_teams:
team_id = getattr(t, "team_id", None)
team_id = getattr(t.row, "team_id", None)
if team_id:
await self._invalidate_spend_counter(f"spend:team:{team_id}", new_spend=t.spend or 0.0)
await self._invalidate_spend_counter(f"spend:team:{team_id}")
end_time = time.time()
outcome: Final = _ChunkOutcome(
fetched=len(teams_to_reset) if teams_to_reset else 0,
advanced=_count_advanced(
(t.budget_reset_at for t in updated_teams),
(t.row.budget_reset_at for t in updated_teams),
cutoff=datetime.now(timezone.utc),
),
)

View file

@ -24,12 +24,8 @@ from typing import Final
from litellm.repositories.prisma_protocols import BatchTable, PrismaBatch
def _spend_reset_data(budget_reset_at: datetime | None, spend_decrement: float | None) -> Mapping[str, object]:
spend: Final[object] = (
{"decrement": spend_decrement} # mutable-ok: prisma update payload must be a dict
if spend_decrement is not None
else 0
)
def _spend_reset_data(budget_reset_at: datetime | None, spend_decrement: float) -> Mapping[str, object]:
spend: Final[object] = {"decrement": spend_decrement} # mutable-ok: prisma update payload must be a dict
return {"spend": spend, "budget_reset_at": budget_reset_at} # mutable-ok: prisma update payload must be a dict
@ -37,9 +33,7 @@ def _spend_reset_data(budget_reset_at: datetime | None, spend_decrement: float |
class KeySpendResetWrites:
table: BatchTable
def queue_spend_reset(
self, token: str, budget_reset_at: datetime | None, spend_decrement: float | None = None
) -> None:
def queue_spend_reset(self, token: str, budget_reset_at: datetime | None, spend_decrement: float) -> None:
self.table.update(
where={"token": token}, # mutable-ok: prisma where filter must be a dict
data=_spend_reset_data(budget_reset_at, spend_decrement),
@ -50,9 +44,7 @@ class KeySpendResetWrites:
class UserSpendResetWrites:
table: BatchTable
def queue_spend_reset(
self, user_id: str, budget_reset_at: datetime | None, spend_decrement: float | None = None
) -> None:
def queue_spend_reset(self, user_id: str, budget_reset_at: datetime | None, spend_decrement: float) -> None:
self.table.update(
where={"user_id": user_id}, # mutable-ok: prisma where filter must be a dict
data=_spend_reset_data(budget_reset_at, spend_decrement),
@ -63,9 +55,7 @@ class UserSpendResetWrites:
class TeamSpendResetWrites:
table: BatchTable
def queue_spend_reset(
self, team_id: str, budget_reset_at: datetime | None, spend_decrement: float | None = None
) -> None:
def queue_spend_reset(self, team_id: str, budget_reset_at: datetime | None, spend_decrement: float) -> None:
self.table.update(
where={"team_id": team_id}, # mutable-ok: prisma where filter must be a dict
data=_spend_reset_data(budget_reset_at, spend_decrement),

View file

@ -19,7 +19,7 @@ from litellm.constants import (
RESET_BUDGET_JOB_LOCK_TTL_SECONDS,
RESET_BUDGET_JOB_NAME,
)
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob, _RowReset
from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings
@ -243,14 +243,18 @@ def test_write_key_reset_updates_skips_none_token_and_still_writes_the_rest(rese
LiteLLM_VerificationToken(token="tok-ok", budget_reset_at=reset_at),
]
asyncio.run(reset_budget_job._write_key_reset_updates(updated_keys=keys))
asyncio.run(
reset_budget_job._write_key_reset_updates(
updated_keys=[_RowReset(row=k, spend_decrement=(k.spend or 0.0)) for k in keys]
)
)
assert _batch_writes(mock_prisma_client, "key") == [
{
"table": "key",
"op": "update",
"where": {"token": "tok-ok"},
"data": {"spend": 0, "budget_reset_at": reset_at},
"data": {"spend": {"decrement": 0.0}, "budget_reset_at": reset_at},
}
]
@ -282,7 +286,7 @@ def test_reset_budget_for_key(reset_budget_job, mock_prisma_client):
assert len(key_writes) == 1
write = key_writes[0]
assert write["where"] == {"token": "tok-key-1"}
assert write["data"]["spend"] == 0
assert write["data"]["spend"] == {"decrement": 100.0}
assert write["data"]["budget_reset_at"] > now
assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
@ -345,7 +349,7 @@ def test_reset_budget_for_user(reset_budget_job, mock_prisma_client):
assert len(user_writes) == 1
write = user_writes[0]
assert write["where"] == {"user_id": "uid-1"}
assert write["data"]["spend"] == 0
assert write["data"]["spend"] == {"decrement": 200.0}
assert write["data"]["budget_reset_at"] > now
assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
@ -374,7 +378,7 @@ def test_reset_budget_for_team(reset_budget_job, mock_prisma_client):
assert len(team_writes) == 1
write = team_writes[0]
assert write["where"] == {"team_id": "tid-1"}
assert write["data"]["spend"] == 0
assert write["data"]["spend"] == {"decrement": 500.0}
assert write["data"]["budget_reset_at"] > now
assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
@ -488,15 +492,15 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
# key/user/team rows are written via batch_().<table>.update — verify each
# one fired exactly once with the narrow {spend, budget_reset_at} payload.
for table_name, where in [
("key", {"token": "tok-all-1"}),
("user", {"user_id": "uid-all-1"}),
("team", {"team_id": "tid-all-1"}),
for table_name, where, decrement in [
("key", {"token": "tok-all-1"}, 100.0),
("user", {"user_id": "uid-all-1"}, 200.0),
("team", {"team_id": "tid-all-1"}, 500.0),
]:
writes = _batch_writes(mock_prisma_client, table_name, op="update")
assert len(writes) == 1, f"expected 1 {table_name} write, got {len(writes)}"
assert writes[0]["where"] == where
assert writes[0]["data"]["spend"] == 0
assert writes[0]["data"]["spend"] == {"decrement": decrement}
assert set(writes[0]["data"].keys()) == {"spend", "budget_reset_at"}
# The budget tier's cascade rides the same batch machinery.
@ -1226,6 +1230,7 @@ def _make_counter_invalidation_job(monkeypatch):
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = MagicMock()
spend_counter_cache.redis_cache.async_set_cache = AsyncMock()
spend_counter_cache.redis_cache.async_delete_cache = AsyncMock()
user_api_key_cache = MagicMock()
user_api_key_cache.async_delete_cache = AsyncMock()
@ -1260,7 +1265,8 @@ def test_reset_budget_for_keys_invalidates_redis_counter(reset_budget_job, mock_
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-abc", value=0.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:key:sk-abc")
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:key:sk-abc")
def test_reset_budget_for_users_invalidates_redis_counter(reset_budget_job, mock_prisma_client, monkeypatch):
@ -1284,7 +1290,8 @@ def test_reset_budget_for_users_invalidates_redis_counter(reset_budget_job, mock
asyncio.run(reset_budget_job.reset_budget_for_litellm_users())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:user:alice", value=0.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:user:alice")
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:user:alice")
def test_reset_budget_for_proxy_budget_row_invalidates_global_spend_cache(
@ -1368,7 +1375,8 @@ def test_reset_budget_for_teams_invalidates_redis_counter(reset_budget_job, mock
asyncio.run(reset_budget_job.reset_budget_for_litellm_teams())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team:team-x", value=0.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:team:team-x")
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:team:team-x")
def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch):
@ -1428,7 +1436,7 @@ def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch):
# assert_not_called() instead of iterating call_args_list, because the
# latter is vacuously true when the list is empty (would pass even if
# the bypass were re-introduced via a different code path).
counter_cache.in_memory_cache.set_cache.assert_not_called()
counter_cache.in_memory_cache.delete_cache.assert_not_called()
def test_reset_budget_for_keys_writes_only_spend_and_reset_at(reset_budget_job, mock_prisma_client):
@ -1526,8 +1534,8 @@ def test_budget_table_reset_invalidates_counters_and_management_cache(
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
counter_cache.in_memory_cache.set_cache.assert_any_call(key=counter_key, value=0.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key=counter_key, value=0.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key=counter_key)
counter_cache.redis_cache.async_delete_cache.assert_any_await(key=counter_key)
deleted = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert cache_keys <= deleted
@ -1565,8 +1573,8 @@ def test_budget_table_reset_invalidates_enduser_counter_and_cache(reset_budget_j
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:customer-42", value=0.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:customer-42", value=0.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:end_user:customer-42")
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:end_user:customer-42")
deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert "end_user_id:customer-42" in deleted
@ -1627,7 +1635,7 @@ def test_access_groups_are_untouched_when_no_budget_is_due(reset_budget_job, moc
assert mock_prisma_client.db.litellm_modelaccessgroupbudgettable.find_many_calls == []
assert _batch_writes(mock_prisma_client, "model_access_group") == []
counter_cache.in_memory_cache.set_cache.assert_not_called()
counter_cache.in_memory_cache.delete_cache.assert_not_called()
counter_cache.user_api_key_cache.async_delete_cache.assert_not_awaited()
@ -1646,7 +1654,7 @@ def test_budget_table_reset_invalidates_every_access_group_not_just_the_first(
deleted = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert deleted == {"model_access_group:group-a", "model_access_group:group-b", "model_access_group:group-c"}
for name in ("group-a", "group-b", "group-c"):
counter_cache.in_memory_cache.set_cache.assert_any_call(key=f"spend:model_access_group:{name}", value=0.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key=f"spend:model_access_group:{name}")
def test_budget_cascade_carries_access_group_overage_when_rollover_enabled(
@ -1678,7 +1686,7 @@ def test_budget_cascade_carries_access_group_overage_when_rollover_enabled(
} in writes
assert _replay_spend_writes(writes, 15.0) == 5.0
assert _replay_spend_writes(writes, 8.0) == 0
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:model_access_group:gpt-4-group", value=5.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:model_access_group:gpt-4-group")
# ---------------------------------------------------------------------------
@ -1769,7 +1777,7 @@ def test_budget_reset_at_is_not_advanced_when_the_cascade_fails(db_factory, monk
assert prisma_client.db.batch_calls == [], "a failed cascade must not persist any write"
assert prisma_client.db.batchers[0].committed is False
assert prisma_client.updated_data["budget"] == [], "budget_reset_at must not be advanced outside the transaction"
counter_cache.in_memory_cache.set_cache.assert_not_called()
counter_cache.in_memory_cache.delete_cache.assert_not_called()
counter_cache.user_api_key_cache.async_delete_cache.assert_not_awaited()
@ -1806,7 +1814,7 @@ def test_caches_are_invalidated_only_after_the_transaction_commits(monkeypatch):
cap while the DB still holds the over-budget spend."""
events = []
counter_cache = _make_counter_invalidation_job(monkeypatch)
counter_cache.in_memory_cache.set_cache.side_effect = lambda **kwargs: events.append("counter")
counter_cache.in_memory_cache.delete_cache.side_effect = lambda **kwargs: events.append("counter")
job, _ = _job_with_expired_budget(OrderRecordingDB(events))
@ -2839,13 +2847,7 @@ _SPEND_ACCRUED_AFTER_COMMIT = 7.5
class AmbiguousCommitClient(MockPrismaClient):
"""A client whose batch commit lands in the database and only then fails in
transit, so the caller cannot tell whether it committed.
The queued spend-zero is applied to `key_spend`, and fresh usage accrues in
the window between that landed commit and any replay, so a replay is
observable as erased spend rather than merely as an extra commit.
"""
"""A client whose batch commit lands in the database and only then fails in transit."""
def __init__(self, *, error: Exception, spend_accrued_after_commit: float):
super().__init__()
@ -2864,7 +2866,12 @@ class AmbiguousCommitClient(MockPrismaClient):
outer.commit_attempts += 1
result = await batch_commit()
for call in batcher.calls:
if call["table"] == "key" and call["data"].get("spend") == 0:
if call["table"] != "key":
continue
spend_field = call["data"].get("spend")
if isinstance(spend_field, dict):
outer.key_spend -= spend_field["decrement"]
elif spend_field == 0:
outer.key_spend = 0.0
if outer.commit_attempts > 1:
return result
@ -2886,22 +2893,19 @@ class AmbiguousCommitClient(MockPrismaClient):
[
(httpx.ReadError("response lost in transit"), 1, _SPEND_ACCRUED_AFTER_COMMIT, []),
(httpx.ReadTimeout("response lost in transit"), 1, _SPEND_ACCRUED_AFTER_COMMIT, []),
(httpx.ConnectError("never left the client"), 2, 0.0, ["reset_budget_write_keys_failure"]),
(
httpx.ConnectError("never left the client"),
2,
_SPEND_ACCRUED_AFTER_COMMIT - _DUE_ROW_SPEND,
["reset_budget_write_keys_failure"],
),
],
ids=["read_error", "read_timeout", "connect_error_erasure_control"],
)
def test_ambiguous_commit_replay_does_not_erase_newly_accrued_spend(
error, expected_commits, expected_spend, expected_reconnects
):
"""A reset zeroes spend unconditionally, so replaying a commit that already
landed erases every dollar spent since it landed (LIT-5372 review finding).
The `connect_error` case is the control: it is the one error class allowed
to replay, and driving it through this same land-then-fail harness proves
the spend assertion can actually observe an erasure. In production a
ConnectError means the statements never reached the database, so its replay
has nothing to erase.
"""
"""Replaying a commit that already landed erases spend accrued since it landed."""
client = AmbiguousCommitClient(error=error, spend_accrued_after_commit=_SPEND_ACCRUED_AFTER_COMMIT)
client.data["key"] = [_due_row("key", "tok-1")]
job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client)
@ -2999,7 +3003,7 @@ def test_direct_reset_carries_overage_when_rollover_enabled(
assert writes[0]["data"]["spend"] == {"decrement": 100.0}
assert writes[0]["data"]["budget_reset_at"] > now
counter_prefix = {"key": "spend:key", "user": "spend:user", "team": "spend:team"}[table]
counter_cache.in_memory_cache.set_cache.assert_any_call(key=f"{counter_prefix}:{id_value}", value=50.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key=f"{counter_prefix}:{id_value}")
def test_direct_reset_zeroes_under_budget_row_even_with_rollover(
@ -3017,8 +3021,8 @@ def test_direct_reset_zeroes_under_budget_row_even_with_rollover(
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
assert _batch_writes(mock_prisma_client, "key")[0]["data"]["spend"] == 0
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:tok-under", value=0.0, ttl=60)
assert _batch_writes(mock_prisma_client, "key")[0]["data"]["spend"] == {"decrement": 40.0}
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:key:tok-under")
def test_direct_reset_zeroes_row_without_max_budget_even_with_rollover(
@ -3037,7 +3041,7 @@ def test_direct_reset_zeroes_row_without_max_budget_even_with_rollover(
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
assert _batch_writes(mock_prisma_client, "key")[0]["data"]["spend"] == 0
assert _batch_writes(mock_prisma_client, "key")[0]["data"]["spend"] == {"decrement": 150.0}
def test_budget_cascade_carries_overage_per_tier_when_rollover_enabled(
@ -3071,7 +3075,7 @@ def test_budget_cascade_carries_overage_per_tier_when_rollover_enabled(
"where": {"budget_id": "budget-roll", "spend": {"gt": 0, "lte": 10.0}},
"data": {"spend": 0},
} in membership_writes
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team_member:member-1:team-1", value=5.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:team_member:member-1:team-1")
def test_budget_cascade_carries_enduser_overage_when_rollover_enabled(
@ -3131,8 +3135,8 @@ def test_budget_cascade_carries_default_tier_enduser_counter_when_rollover_enabl
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:end_user:enduser-implicit", value=5.0, ttl=60)
counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:end_user:enduser-implicit", value=5.0, ttl=60)
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:end_user:enduser-implicit")
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:end_user:enduser-implicit")
deleted: Final = {call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list}
assert "end_user_id:enduser-implicit" in deleted
@ -3243,3 +3247,138 @@ def test_window_reset_zeroes_counter_when_rollover_disabled(monkeypatch):
spend_counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-off:window:1d", value=0.0)
spend_counter_cache.async_get_cache.assert_not_awaited()
def _apply_spend_payload(db_spend: float, spend_field: dict[str, float]) -> float:
return db_spend - spend_field["decrement"]
_RACE_TABLES = [
(
lambda job: job.reset_budget_for_litellm_keys(),
"key",
"token",
"tok-race",
lambda now: type(
"Key",
(),
{"spend": 5.0, "budget_duration": "1d", "budget_reset_at": now, "token": "tok-race"},
),
),
(
lambda job: job.reset_budget_for_litellm_users(),
"user",
"user_id",
"user-race",
lambda now: type(
"User",
(),
{"spend": 5.0, "budget_duration": "7d", "budget_reset_at": now, "user_id": "user-race"},
),
),
(
lambda job: job.reset_budget_for_litellm_teams(),
"team",
"team_id",
"team-race",
lambda now: type(
"Team",
(),
{"spend": 5.0, "budget_duration": "1mo", "budget_reset_at": now, "team_id": "team-race"},
),
),
]
@pytest.mark.parametrize("run_phase, table, id_field, id_value, row_factory", _RACE_TABLES)
def test_reset_decrement_preserves_spend_landed_after_read(
reset_budget_job, mock_prisma_client, run_phase, table, id_field, id_value, row_factory
):
"""LIT-7814: spend flushed between the read and the commit survives the reset."""
now = datetime.now(timezone.utc)
mock_prisma_client.data[table] = [row_factory(now)]
asyncio.run(run_phase(reset_budget_job))
writes = _batch_writes(mock_prisma_client, table)
assert len(writes) == 1
assert writes[0]["where"] == {id_field: id_value}
assert writes[0]["data"]["spend"] == {"decrement": 5.0}
assert writes[0]["data"]["budget_reset_at"] > now
assert _apply_spend_payload(db_spend=5.4, spend_field=writes[0]["data"]["spend"]) == pytest.approx(0.4)
@pytest.mark.parametrize("run_phase, table, id_field, id_value, row_factory", _RACE_TABLES)
def test_reset_decrement_subsumes_rollover_cap(
rollover_enabled, reset_budget_job, mock_prisma_client, run_phase, table, id_field, id_value, row_factory
):
"""Rollover on, spend over the cap decrements by the cap itself."""
now = datetime.now(timezone.utc)
row = row_factory(now)
row.max_budget = 3.0
mock_prisma_client.data[table] = [row]
asyncio.run(run_phase(reset_budget_job))
writes = _batch_writes(mock_prisma_client, table)
assert len(writes) == 1
assert writes[0]["data"]["spend"] == {"decrement": 3.0}
assert _apply_spend_payload(db_spend=5.4, spend_field=writes[0]["data"]["spend"]) == pytest.approx(2.4)
@pytest.mark.parametrize("run_phase, table, id_field, id_value, row_factory", _RACE_TABLES)
def test_reset_decrement_under_cap_with_rollover(
rollover_enabled, reset_budget_job, mock_prisma_client, run_phase, table, id_field, id_value, row_factory
):
"""Rollover on, spend under the cap decrements by the read-time spend."""
now = datetime.now(timezone.utc)
row = row_factory(now)
row.spend = 2.0
row.max_budget = 3.0
mock_prisma_client.data[table] = [row]
asyncio.run(run_phase(reset_budget_job))
writes = _batch_writes(mock_prisma_client, table)
assert len(writes) == 1
assert writes[0]["data"]["spend"] == {"decrement": 2.0}
assert _apply_spend_payload(db_spend=2.4, spend_field=writes[0]["data"]["spend"]) == pytest.approx(0.4)
@pytest.mark.parametrize("run_phase, table, id_field, id_value, row_factory", _RACE_TABLES)
def test_reset_zero_spend_row_writes_noop_decrement(
reset_budget_job, mock_prisma_client, run_phase, table, id_field, id_value, row_factory
):
"""A spend=0 row gets a no-op decrement, never an absolute spend=0."""
now = datetime.now(timezone.utc)
row = row_factory(now)
row.spend = 0.0
mock_prisma_client.data[table] = [row]
asyncio.run(run_phase(reset_budget_job))
writes = _batch_writes(mock_prisma_client, table)
assert len(writes) == 1
assert writes[0]["data"]["spend"] == {"decrement": 0.0}
assert writes[0]["data"]["budget_reset_at"] > now
assert _apply_spend_payload(db_spend=0.4, spend_field=writes[0]["data"]["spend"]) == pytest.approx(0.4)
def test_reset_deletes_spend_counter_instead_of_seeding(reset_budget_job, mock_prisma_client, monkeypatch):
"""A reset deletes the counter so the next read reseeds from the committed row."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
now = datetime.now(timezone.utc)
mock_prisma_client.data["user"] = [
type(
"User",
(),
{"spend": 5.0, "budget_duration": "7d", "budget_reset_at": now, "id": "user-r", "user_id": "carol"},
)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_users())
counter_cache.in_memory_cache.delete_cache.assert_any_call(key="spend:user:carol")
counter_cache.redis_cache.async_delete_cache.assert_any_await(key="spend:user:carol")
counter_cache.in_memory_cache.set_cache.assert_not_called()
counter_cache.redis_cache.async_set_cache.assert_not_awaited()

View file

@ -46,16 +46,16 @@ async def test_updates_across_tables_share_one_batch_and_commit_once():
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)
uow.keys.queue_spend_reset(token="tok-1", budget_reset_at=reset_at, spend_decrement=1.5)
uow.users.queue_spend_reset(user_id="user-1", budget_reset_at=reset_at, spend_decrement=2.5)
uow.teams.queue_spend_reset(team_id="team-1", budget_reset_at=None, spend_decrement=0.0)
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}),
("litellm_verificationtoken", {"token": "tok-1"}, {"spend": {"decrement": 1.5}, "budget_reset_at": reset_at}),
("litellm_usertable", {"user_id": "user-1"}, {"spend": {"decrement": 2.5}, "budget_reset_at": reset_at}),
("litellm_teamtable", {"team_id": "team-1"}, {"spend": {"decrement": 0.0}, "budget_reset_at": None}),
]
@ -64,7 +64,7 @@ async def test_raising_inside_block_skips_commit():
async def _blow_up_mid_transaction():
async with spend_reset_unit_of_work(lambda: batch) as uow:
uow.keys.queue_spend_reset(token="tok-1", budget_reset_at=None)
uow.keys.queue_spend_reset(token="tok-1", budget_reset_at=None, spend_decrement=0.0)
raise RuntimeError("boom")
with pytest.raises(RuntimeError, match="boom"):