mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
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:
parent
abfb6adc2b
commit
5eff708d0f
5 changed files with 197 additions and 87 deletions
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue