fix(proxy): keep persisted spend another pod has not incremented in the window seed

The batch start told the seed which LiteLLM_SpendLogs rows were its own, but
using it as a hard cutoff also dropped rows another pod had already persisted.
Those rows are only repaid by that pod's own increment, so if it died first the
window row stayed permanently under the recorded spend.

The seed now reads both sums in one scan and takes off this batch's own spend,
flooring at the pre-batch total for the case where its log rows have not landed
yet. Redis payloads keep an empty request_ids so a leader from before the field
was dropped can still merge what it pops during a rolling deploy.

Claude-Session: https://claude.ai/code/session_01QvQzYztinxj8ZuD5YxbVdL
This commit is contained in:
ryan-crabbe-berri 2026-08-31 15:49:03 -07:00
parent abfb6adc2b
commit 5eff708d0f
5 changed files with 197 additions and 87 deletions

View file

@ -7,17 +7,15 @@ instead of aggregating LiteLLM_SpendLogs every time a window counter goes cold
(issue #35766). Raw SQL rather than the Prisma upsert helper because the
conditional roll cannot be expressed through the query builder.
Seeding a row that does not exist yet reads LiteLLM_SpendLogs once, summing
only rows that started before the batch being flushed so neither source counts
the same request twice. Anything at or after that cutoff is owed by an
increment that still reaches the row, on this pod's next flush or another
pod's, so a row lags real spend by at most one flush interval of queued
increments: the same lag the SpendLogs aggregate it replaces (and every other
spend column) already has. A request whose increment is lost before it flushes,
which today means the pod dying, is missed by both sources and stays missing.
Seeding a row that does not exist yet reads LiteLLM_SpendLogs once and takes
off what the increments being flushed will add, so neither source counts the
same request twice. A row therefore lags real spend by at most one flush
interval of increments queued elsewhere: the same lag the SpendLogs aggregate
it replaces (and every other spend column) already has.
"""
from collections.abc import Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Final, Protocol
@ -64,32 +62,45 @@ _ROLL_WINDOW_SPEND_SQL: Final = (
)
_SEED_FROM_SPEND_LOGS_KEY_SQL: Final = (
'SELECT COALESCE(SUM(spend), 0.0) AS total FROM "LiteLLM_SpendLogs" '
"WHERE api_key = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC') "
"AND \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')"
"SELECT COALESCE(SUM(spend), 0.0) AS total, "
"COALESCE(SUM(spend) FILTER (WHERE \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')), 0.0) AS before_batch "
'FROM "LiteLLM_SpendLogs" '
"WHERE api_key = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
)
_SEED_FROM_SPEND_LOGS_TEAM_SQL: Final = (
'SELECT COALESCE(SUM(spend), 0.0) AS total FROM "LiteLLM_SpendLogs" '
"WHERE team_id = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC') "
"AND \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')"
"SELECT COALESCE(SUM(spend), 0.0) AS total, "
"COALESCE(SUM(spend) FILTER (WHERE \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')), 0.0) AS before_batch "
'FROM "LiteLLM_SpendLogs" '
"WHERE team_id = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
)
_SEED_FROM_SPEND_LOGS_KEY_UNBOUNDED_SQL: Final = (
'SELECT COALESCE(SUM(spend), 0.0) AS total FROM "LiteLLM_SpendLogs" '
"SELECT COALESCE(SUM(spend), 0.0) AS total, COALESCE(SUM(spend), 0.0) AS before_batch "
'FROM "LiteLLM_SpendLogs" '
"WHERE api_key = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
)
_SEED_FROM_SPEND_LOGS_TEAM_UNBOUNDED_SQL: Final = (
'SELECT COALESCE(SUM(spend), 0.0) AS total FROM "LiteLLM_SpendLogs" '
"SELECT COALESCE(SUM(spend), 0.0) AS total, COALESCE(SUM(spend), 0.0) AS before_batch "
'FROM "LiteLLM_SpendLogs" '
"WHERE team_id = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
)
_UPSERT_TRANSACTION_TIMEOUT: Final = timedelta(seconds=60)
@dataclass(frozen=True, slots=True)
class WindowSeedTotals:
"""The two sums a seed needs: everything persisted for the window, and the
part of it that predates the batch being flushed."""
total: float
before_batch: float
class WindowSpendLogsAggregate(Protocol):
"""Sums LiteLLM_SpendLogs for one entity between window_start and the
"""Sums LiteLLM_SpendLogs for one entity since window_start, split at the
batch's earliest request.
Injected so the flush can be exercised without a database and so the
@ -103,18 +114,18 @@ class WindowSpendLogsAggregate(Protocol):
entity_id: str,
window_start: datetime,
batch_started_at: datetime | None,
) -> float | None: ...
) -> WindowSeedTotals | None: ...
async def spend_logs_total_before_batch(
async def spend_logs_seed_totals(
prisma_client: "PrismaClient",
entity_type: str,
entity_id: str,
window_start: datetime,
batch_started_at: datetime | None,
) -> float | None:
"""LiteLLM_SpendLogs spend for one entity since window_start, stopping
before the requests the increments being flushed already cover.
) -> WindowSeedTotals | None:
"""LiteLLM_SpendLogs spend for one entity since window_start, both in full
and up to the start of the batch being flushed, in one scan.
The spend log writer drains its own queue on a ~2s poll whenever anything
is queued, while window increments flush on the much slower batch tick, so
@ -122,11 +133,12 @@ async def spend_logs_total_before_batch(
already in the table. Counting them in the seed and again in the increment
is what made a fresh row land at twice the true spend.
Every log row at or after the cutoff belongs to a request whose own
increment still reaches this row, on this pod's next flush or another pod's,
so bounding the sum by time needs nothing from the request itself. Without a
known start the whole window is summed: that can only over-count once, which
enforcement tolerates, whereas under-counting is a budget bypass.
Both halves are needed because neither is safe alone: the full sum
double-counts this batch, and the sum before the batch drops spend another
pod has already persisted but not yet incremented. _seed_base picks between
them. Without a known batch start the two are the same sum, so the seed
counts everything: that can only over-count once, which enforcement
tolerates, whereas under-counting is a budget bypass.
"""
if entity_type == Litellm_EntityType.KEY.value:
bounded_sql, unbounded_sql = _SEED_FROM_SPEND_LOGS_KEY_SQL, _SEED_FROM_SPEND_LOGS_KEY_UNBOUNDED_SQL
@ -145,8 +157,11 @@ async def spend_logs_total_before_batch(
)
)
if not rows:
return 0.0
return float(rows[0].get("total") or 0.0)
return WindowSeedTotals(total=0.0, before_batch=0.0)
return WindowSeedTotals(
total=float(rows[0].get("total") or 0.0),
before_batch=float(rows[0].get("before_batch") or 0.0),
)
def _exclusion_upper_bound(started_at: datetime) -> datetime:
@ -186,19 +201,33 @@ async def _seed_base_for_missing_row(
This is the LiteLLM_SpendLogs aggregate the window counter reseed runs on
every cold counter today, but here it runs once per window lifetime and off
the request path, and it stops before the queued increments so they are
the request path, and it discounts the queued increments so they are
counted once.
"""
if _primary_key(transaction) in existing_primary_keys:
return 0.0
base: Final = await spend_logs_aggregate(
totals: Final = await spend_logs_aggregate(
prisma_client=prisma_client,
entity_type=transaction["entity_type"],
entity_id=transaction["entity_id"],
window_start=datetime.fromisoformat(transaction["window_start"]).replace(tzinfo=timezone.utc),
batch_started_at=_transaction_started_at(transaction),
)
return float(base or 0.0)
if totals is None:
return 0.0
return _seed_base(totals=totals, batch_spend=transaction["spend"])
def _seed_base(totals: WindowSeedTotals, batch_spend: float) -> float:
"""What the window already held before the increments about to be applied.
Subtracting the batch's own spend from the full sum keeps every other
request in the seed, including the ones another pod persisted and has not
incremented yet, which a plain cutoff would drop for good if that pod died.
When this batch's own log rows have not landed yet the subtraction takes
spend that was never counted, so the sum before the batch is the floor.
"""
return max(totals.total - batch_spend, totals.before_batch)
def _transaction_started_at(transaction: WindowSpendTransaction) -> datetime | None:
@ -232,7 +261,7 @@ def _upsert_params(
async def commit_window_spend_updates(
prisma_client: "PrismaClient",
transactions: Sequence[WindowSpendTransaction],
spend_logs_aggregate: WindowSpendLogsAggregate = spend_logs_total_before_batch,
spend_logs_aggregate: WindowSpendLogsAggregate = spend_logs_seed_totals,
) -> None:
"""Apply aggregated window increments to LiteLLM_BudgetWindowSpend.

View file

@ -46,6 +46,7 @@ from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdate
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
WindowSpendTransaction,
WindowSpendUpdateQueue,
to_wire_payload,
)
from litellm.secret_managers.main import str_to_bool
from litellm.types.caching import (
@ -298,7 +299,7 @@ class RedisUpdateBuffer:
ServiceTypes.REDIS_DAILY_AGENT_SPEND_UPDATE_QUEUE,
),
(
window_spend_update_transactions,
tuple(map(to_wire_payload, window_spend_update_transactions)),
REDIS_WINDOW_SPEND_UPDATE_BUFFER_KEY,
ServiceTypes.REDIS_WINDOW_SPEND_UPDATE_QUEUE,
),
@ -484,7 +485,12 @@ class RedisUpdateBuffer:
(daily_end_user_spend_update_transactions, REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY),
(daily_agent_spend_update_transactions, REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY),
(daily_tag_spend_update_transactions, REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY),
(window_spend_update_transactions, REDIS_WINDOW_SPEND_UPDATE_BUFFER_KEY),
(
None
if window_spend_update_transactions is None
else tuple(map(to_wire_payload, window_spend_update_transactions)),
REDIS_WINDOW_SPEND_UPDATE_BUFFER_KEY,
),
)
rpush_list: Final = tuple(

View file

@ -27,10 +27,11 @@ class WindowSpendTransaction(TypedDict):
transaction survives the JSON round trip through the Redis buffer.
started_at is the earliest request start in the batch. The one-time seed for
a window that has no row yet sums only LiteLLM_SpendLogs rows that started
before it, because the spend log writer flushes on its own ~2s poll and will
usually have persisted this batch's rows before the window queue flushes;
without the bound the seed and the increment would each count them.
a window that has no row yet uses it to tell this batch's own
LiteLLM_SpendLogs rows from everything else, because the spend log writer
flushes on its own ~2s poll and will usually have persisted this batch's
rows before the window queue flushes; without that split the seed and the
increment would each count them.
"""
entity_type: ReadOnly[str]
@ -41,6 +42,34 @@ class WindowSpendTransaction(TypedDict):
started_at: ReadOnly[str | None]
class WindowSpendWirePayload(WindowSpendTransaction):
"""How an increment is encoded in the shared Redis buffer.
request_ids is dead weight here: workers built before this field was
dropped index it while merging whatever they pop, and the pop is
destructive, so a leader still running one of those during a rolling deploy
would raise on a payload without the key and lose those increments. It is
always empty, which only makes such a leader seed without exclusions.
TODO: remove once no supported version reads it, i.e. one release after the
field stopped being written.
"""
request_ids: ReadOnly[Sequence[str]]
def to_wire_payload(transaction: WindowSpendTransaction) -> WindowSpendWirePayload:
return WindowSpendWirePayload(
entity_type=transaction["entity_type"],
entity_id=transaction["entity_id"],
window_duration=transaction["window_duration"],
window_start=transaction["window_start"],
spend=transaction["spend"],
started_at=transaction.get("started_at"),
request_ids=(),
)
def to_naive_utc(value: datetime) -> datetime:
"""LiteLLM_BudgetWindowSpend.window_start is TIMESTAMP(3), which holds naive UTC."""
if value.tzinfo is None:

View file

@ -523,10 +523,44 @@ async def test_store_in_memory_spend_updates_pushes_budget_window_spend(redis_up
"window_start": "2026-08-01T00:00:00.000000",
"spend": 1.25,
"started_at": "2026-08-10T12:00:00.000000",
"request_ids": [],
}
]
@pytest.mark.asyncio
async def test_budget_window_payloads_keep_request_ids_for_older_workers(redis_update_buffer, mock_redis_cache):
"""A leader from before the field was dropped indexes request_ids while
merging what it popped, and the pop is destructive, so a payload without
the key would cost a rolling deploy those increments."""
from datetime import datetime, timezone
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
WindowSpendUpdateQueue,
build_window_spend_transaction,
)
mock_redis_cache.async_rpush_pipeline = AsyncMock(return_value=[1])
window_queue = WindowSpendUpdateQueue()
await window_queue.add_update(
build_window_spend_transaction(
entity_type="key",
entity_id="hashed-token",
window_duration="30d",
window_start=datetime(2026, 8, 1, tzinfo=timezone.utc),
spend=1.25,
)
)
await redis_update_buffer.restore_transactions_to_redis(
window_spend_update_transactions=await window_queue.flush_and_get_aggregated_window_spend_transactions(),
)
rpush_list = mock_redis_cache.async_rpush_pipeline.call_args.kwargs["rpush_list"]
restored = json.loads(rpush_list[0]["values"][0])
assert [payload["request_ids"] for payload in restored] == [[]]
@pytest.mark.asyncio
async def test_store_in_memory_spend_updates_restores_budget_window_spend_on_rpush_failure(
redis_update_buffer, mock_redis_cache

View file

@ -6,9 +6,10 @@ from typing import Any
import pytest
from litellm.proxy.db.budget_window_spend_writer import (
WindowSeedTotals,
commit_window_spend_updates,
roll_window_spend_row,
spend_logs_total_before_batch,
spend_logs_seed_totals,
)
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
build_window_spend_transaction,
@ -70,10 +71,15 @@ class _FakePrismaClient:
class _RecordingAggregate:
"""Stands in for the LiteLLM_SpendLogs seed aggregate."""
"""Stands in for the LiteLLM_SpendLogs seed aggregate. before_batch
defaults to the full total, the state where none of this batch's own log
rows have been persisted yet."""
def __init__(self, value: float = 5.0) -> None:
self.value = value
def __init__(self, total: float = 5.0, before_batch: float | None = None) -> None:
self.totals = WindowSeedTotals(
total=total,
before_batch=total if before_batch is None else before_batch,
)
self.calls: list[dict[str, Any]] = []
async def __call__(
@ -83,7 +89,7 @@ class _RecordingAggregate:
entity_id: str,
window_start: datetime,
batch_started_at: datetime | None,
) -> float | None:
) -> WindowSeedTotals | None:
self.calls.append(
{
"entity_type": entity_type,
@ -92,13 +98,13 @@ class _RecordingAggregate:
"batch_started_at": batch_started_at,
}
)
return self.value
return self.totals
class _SpendLogsFake:
"""Sums the LiteLLM_SpendLogs rows (request_id, spend, startTime) it holds,
honouring the cutoff exactly as the real aggregate's
startTime < bound does."""
splitting them at the batch start exactly as the real aggregate's
SUM(...) FILTER (WHERE startTime < bound) does."""
def __init__(self, rows: tuple[tuple[str, float, datetime], ...]) -> None:
self.rows = rows
@ -110,11 +116,14 @@ class _SpendLogsFake:
entity_id: str,
window_start: datetime,
batch_started_at: datetime | None,
) -> float | None:
return math.fsum(
spend
for _request_id, spend, started_at in self.rows
if batch_started_at is None or started_at < batch_started_at
) -> WindowSeedTotals | None:
return WindowSeedTotals(
total=math.fsum(spend for _request_id, spend, _started_at in self.rows),
before_batch=math.fsum(
spend
for _request_id, spend, started_at in self.rows
if batch_started_at is None or started_at < batch_started_at
),
)
@ -151,7 +160,7 @@ async def test_missing_row_is_seeded_from_spend_logs_once():
existed, so a brand new primary key inserts the SpendLogs total plus this
increment."""
db = _FakeDB(existing_rows=[])
aggregate = _RecordingAggregate(value=5.0)
aggregate = _RecordingAggregate(total=5.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -178,7 +187,7 @@ async def test_existing_row_is_never_reseeded():
"""The seed is a full LiteLLM_SpendLogs scan; running it for a row that is
already maintained would both cost a scan and double count."""
db = _FakeDB(existing_rows=[_existing("key", "k1", "30d")])
aggregate = _RecordingAggregate(value=5.0)
aggregate = _RecordingAggregate(total=5.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -195,7 +204,7 @@ async def test_existing_row_is_never_reseeded():
@pytest.mark.asyncio
async def test_seed_runs_only_for_the_primary_keys_that_are_missing():
db = _FakeDB(existing_rows=[_existing("key", "k1", "30d")])
aggregate = _RecordingAggregate(value=5.0)
aggregate = _RecordingAggregate(total=5.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -218,7 +227,7 @@ async def test_insert_spend_and_increment_differ_only_when_a_row_is_seeded():
"""The conflict arm adds the increment alone so two pods that both seed the
same new window cannot add the SpendLogs base twice."""
db = _FakeDB(existing_rows=[])
aggregate = _RecordingAggregate(value=9.0)
aggregate = _RecordingAggregate(total=9.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -255,7 +264,7 @@ async def test_upsert_sql_adds_for_a_current_window_and_replaces_for_a_newer_one
@pytest.mark.asyncio
async def test_upsert_never_interpolates_values_into_the_sql():
db = _FakeDB(existing_rows=[])
aggregate = _RecordingAggregate(value=0.0)
aggregate = _RecordingAggregate(total=0.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -273,7 +282,7 @@ async def test_upserts_are_ordered_by_primary_key_then_window_start():
"""Cross-pod lock ordering, plus an older window must be applied before the
roll that supersedes it or the roll would be undone."""
db = _FakeDB(existing_rows=[])
aggregate = _RecordingAggregate(value=0.0)
aggregate = _RecordingAggregate(total=0.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -301,7 +310,7 @@ async def test_upserts_are_ordered_by_primary_key_then_window_start():
@pytest.mark.asyncio
async def test_existing_row_lookup_sends_every_primary_key_as_array_params():
db = _FakeDB(existing_rows=[])
aggregate = _RecordingAggregate(value=0.0)
aggregate = _RecordingAggregate(total=0.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -339,9 +348,7 @@ async def test_unknown_entity_type_contributes_no_seed():
anything else starts from its increment alone."""
db = _FakeDB(existing_rows=[])
async def no_such_column(
prisma_client, entity_type, entity_id, window_start, batch_started_at
):
async def no_such_column(prisma_client, entity_type, entity_id, window_start, batch_started_at):
return None
await commit_window_spend_updates(
@ -396,7 +403,7 @@ async def test_roll_window_spend_row_is_conditional_on_the_stored_window_being_o
@pytest.mark.asyncio
async def test_seed_receives_the_batch_earliest_start_as_its_cutoff():
db = _FakeDB(existing_rows=[])
aggregate = _RecordingAggregate(value=0.0)
aggregate = _RecordingAggregate(total=0.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -410,7 +417,7 @@ async def test_seed_receives_the_batch_earliest_start_as_its_cutoff():
@pytest.mark.asyncio
async def test_seed_passes_no_start_bound_when_the_batch_has_none():
db = _FakeDB(existing_rows=[])
aggregate = _RecordingAggregate(value=0.0)
aggregate = _RecordingAggregate(total=0.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
@ -463,15 +470,19 @@ async def test_new_row_still_covers_spend_that_predates_the_batch():
@pytest.mark.asyncio
async def test_seed_skips_logs_from_requests_this_batch_never_saw():
"""A concurrent request on another pod can land its spend log before this
pod seeds the row. Its increment is still queued over there, so the cutoff
has to drop it from the seed even though this batch has no way to know its
id; counting it here and again on that pod's flush is the double count the
old id list could not catch."""
async def test_seed_keeps_spend_another_pod_persisted_after_this_batch_started():
"""A concurrent request on another pod can land its spend log after this
batch started but before this pod seeds the row. Dropping it on a plain
time cutoff would lose that spend for the rest of the window if that pod
died before flushing its increment, so the seed takes off only this batch's
own spend and keeps everything else."""
db = _FakeDB(existing_rows=[])
spend_logs = _SpendLogsFake(
rows=(("older", 0.5, BEFORE_BATCH), ("other-pod", 0.25, BATCH_STARTED_AT + timedelta(seconds=1))),
rows=(
("older", 0.5, BEFORE_BATCH),
("mine", 0.000047, BATCH_STARTED_AT),
("other-pod", 0.25, BATCH_STARTED_AT + timedelta(seconds=1)),
),
)
await commit_window_spend_updates(
@ -481,7 +492,7 @@ async def test_seed_skips_logs_from_requests_this_batch_never_saw():
)
((_, params),) = db.batcher.calls
assert params[INSERT_SPEND] == pytest.approx(0.500047)
assert params[INSERT_SPEND] == pytest.approx(0.750047)
@pytest.mark.asyncio
@ -506,10 +517,10 @@ async def test_new_row_is_correct_when_the_batch_logs_have_not_flushed_yet():
"entity_type, expected_column",
[("key", "api_key = $1"), ("team", "team_id = $1")],
)
async def test_seed_aggregate_sql_stops_at_the_batch_start(entity_type, expected_column):
db = _FakeDB(existing_rows=[{"total": 1.25}])
async def test_seed_aggregate_sql_splits_the_window_at_the_batch_start(entity_type, expected_column):
db = _FakeDB(existing_rows=[{"total": 1.25, "before_batch": 0.75}])
total = await spend_logs_total_before_batch(
totals = await spend_logs_seed_totals(
prisma_client=_FakePrismaClient(db),
entity_type=entity_type,
entity_id="e1",
@ -517,11 +528,11 @@ async def test_seed_aggregate_sql_stops_at_the_batch_start(entity_type, expected
batch_started_at=BATCH_STARTED_AT,
)
assert total == pytest.approx(1.25)
assert totals == WindowSeedTotals(total=1.25, before_batch=0.75)
((query, params),) = db.query_raw_calls
normalized = " ".join(query.split())
assert expected_column in normalized
assert "AND \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')" in normalized
assert "FILTER (WHERE \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC'))" in normalized
assert 'FROM "LiteLLM_SpendLogs"' in normalized
# startTime is TIMESTAMP(3): the bound is floored to the second so the
# batch's own earliest row cannot round under it.
@ -532,12 +543,13 @@ async def test_seed_aggregate_sql_stops_at_the_batch_start(entity_type, expected
@pytest.mark.asyncio
async def test_seed_aggregate_sums_the_whole_window_without_a_start_bound():
"""A batch with no known start cannot place the cutoff, so the seed counts
everything; at worst that over-counts one batch, which enforcement
tolerates, where under-counting is a budget bypass."""
db = _FakeDB(existing_rows=[{"total": 1.25}])
"""A batch with no known start cannot place the split, so both halves are
the same sum and the seed counts everything; at worst that over-counts one
batch, which enforcement tolerates, where under-counting is a budget
bypass."""
db = _FakeDB(existing_rows=[{"total": 1.25, "before_batch": 1.25}])
total = await spend_logs_total_before_batch(
totals = await spend_logs_seed_totals(
prisma_client=_FakePrismaClient(db),
entity_type="key",
entity_id="e1",
@ -545,7 +557,7 @@ async def test_seed_aggregate_sums_the_whole_window_without_a_start_bound():
batch_started_at=None,
)
assert total == pytest.approx(1.25)
assert totals == WindowSeedTotals(total=1.25, before_batch=1.25)
((query, params),) = db.query_raw_calls
assert '"startTime" <' not in query
assert params == ("e1", WINDOW_A)
@ -555,7 +567,7 @@ async def test_seed_aggregate_sums_the_whole_window_without_a_start_bound():
async def test_seed_aggregate_returns_none_for_an_entity_type_with_no_spend_logs_column():
db = _FakeDB(existing_rows=[])
total = await spend_logs_total_before_batch(
totals = await spend_logs_seed_totals(
prisma_client=_FakePrismaClient(db),
entity_type="user",
entity_id="u1",
@ -563,7 +575,7 @@ async def test_seed_aggregate_returns_none_for_an_entity_type_with_no_spend_logs
batch_started_at=None,
)
assert total is None
assert totals is None
assert db.query_raw_calls == []
@ -571,7 +583,7 @@ async def test_seed_aggregate_returns_none_for_an_entity_type_with_no_spend_logs
async def test_seed_aggregate_treats_an_entity_with_no_rows_as_zero():
db = _FakeDB(existing_rows=[])
total = await spend_logs_total_before_batch(
totals = await spend_logs_seed_totals(
prisma_client=_FakePrismaClient(db),
entity_type="key",
entity_id="k-unknown",
@ -579,4 +591,4 @@ async def test_seed_aggregate_treats_an_entity_with_no_rows_as_zero():
batch_started_at=None,
)
assert total == 0.0
assert totals == WindowSeedTotals(total=0.0, before_batch=0.0)