fix(proxy): reset budgets by decrementing pre-reset spend instead of zeroing rows

The budget reset job read a row's spend, reset it in place, then wrote
spend: 0 (or decremented by max_budget under rollover) when committing.
Any spend the batch writer incremented into the row between the read and
the commit was erased while LiteLLM_DailyUserSpend kept it, so the daily
rollup permanently exceeded the counters.

Capture each row's spend before _reset_budget_common mutates it and write
a decrement of pre_spend - post_spend, which equals max_budget in the
rollover-over-cap case it replaces. Rows with no spend still get an
absolute spend: 0.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-15 19:24:33 +00:00
parent d3929287fe
commit 2f33727cc9
2 changed files with 214 additions and 47 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}})
@ -842,7 +849,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 +865,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 if k.spend_decrement > 0.0 else None,
)
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 +889,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 if u.spend_decrement > 0.0 else None,
)
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 +911,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 if t.spend_decrement > 0.0 else None,
)
def _emit_phase_failure(
@ -962,18 +969,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 +998,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}", new_spend=k.row.spend or 0.0)
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 +1076,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 +1109,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}", new_spend=u.row.spend or 0.0)
if user_id == LITELLM_PROXY_BUDGET_NAME:
await self._invalidate_global_proxy_spend_cache()
@ -1100,7 +1119,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 +1191,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 +1224,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}", new_spend=t.row.spend or 0.0)
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

@ -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,7 +243,11 @@ 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") == [
{
@ -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.
@ -2864,7 +2868,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,7 +2895,12 @@ 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"],
)
@ -2898,7 +2912,8 @@ def test_ambiguous_commit_replay_does_not_erase_newly_accrued_spend(
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
the spend assertion can actually observe an erasure (the replayed decrement
both erases the accrued spend and over-decrements the row). In production a
ConnectError means the statements never reached the database, so its replay
has nothing to erase.
"""
@ -3017,7 +3032,7 @@ 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
assert _batch_writes(mock_prisma_client, "key")[0]["data"]["spend"] == {"decrement": 40.0}
counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:tok-under", value=0.0, ttl=60)
@ -3037,7 +3052,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(
@ -3243,3 +3258,130 @@ 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()
# ---------------------------------------------------------------------------
# Reset-vs-flush race (LIT-7814): the reset write must decrement by the spend
# captured at read time, not set spend=0 absolutely, so spend the batch writer
# lands between the job's read and its commit survives the reset.
def _apply_spend_payload(db_spend: float, spend_field: Any) -> float:
if isinstance(spend_field, dict):
return db_spend - spend_field["decrement"]
return spend_field
_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
):
"""Regression for LIT-7814: spend flushed between the read and the commit
must survive the reset. spend=5.0 at read, DB row grows to 5.4 before the
write applies; the decrement leaves 0.4, an absolute spend=0 erases it."""
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=5.0 over a max_budget=3.0 cap: decrement by the cap
leaves the 2.0 carry, matching the old max_budget decrement special case."""
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=2.0 under a max_budget=3.0 cap: decrement by the
read-time spend (2.0), which used to be an absolute spend=0 write."""
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_absolute_zero(
reset_budget_job, mock_prisma_client, run_phase, table, id_field, id_value, row_factory
):
"""A row already at spend=0 still needs its window advanced, with an
absolute spend=0 (a decrement of 0 would be a no-op payload)."""
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"] == 0
assert writes[0]["data"]["budget_reset_at"] > now