mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
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:
parent
d3929287fe
commit
2f33727cc9
2 changed files with 214 additions and 47 deletions
|
|
@ -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),
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue