fix(proxy): bound the window spend seed exclusion to the batch's own start time

request_id can be chosen by the client through x-litellm-call-id, so an
unbounded NOT (request_id = ANY(batch)) let a replayed old id drop that id's
historical LiteLLM_SpendLogs row from the one-time seed while its increment
still landed. The increment now carries the request start, the batch keeps
the earliest one, and the seed only excludes ids whose startTime is at or
after it.
This commit is contained in:
ryan-crabbe-berri 2026-08-29 11:05:57 -07:00
parent 83cedfd2b2
commit 041cae8280
9 changed files with 245 additions and 66 deletions

View file

@ -65,13 +65,23 @@ _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 NOT (request_id = ANY($3::text[]))"
"AND NOT (request_id = ANY($3::text[]) AND \"startTime\" >= ($4::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 NOT (request_id = ANY($3::text[]))"
"AND NOT (request_id = ANY($3::text[]) AND \"startTime\" >= ($4::timestamptz AT TIME ZONE 'UTC'))"
)
_SEED_FROM_SPEND_LOGS_KEY_UNBOUNDED_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')"
)
_SEED_FROM_SPEND_LOGS_TEAM_UNBOUNDED_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')"
)
_UPSERT_TRANSACTION_TIMEOUT: Final = timedelta(seconds=60)
@ -92,6 +102,7 @@ class WindowSpendLogsAggregate(Protocol):
entity_id: str,
window_start: datetime,
exclude_request_ids: Sequence[str],
exclude_started_at: datetime | None,
) -> float | None: ...
@ -101,6 +112,7 @@ async def spend_logs_total_excluding(
entity_id: str,
window_start: datetime,
exclude_request_ids: Sequence[str],
exclude_started_at: datetime | None,
) -> float | None:
"""LiteLLM_SpendLogs spend for one entity since window_start, minus the
requests already accounted for by the increments being flushed.
@ -110,22 +122,43 @@ async def spend_logs_total_excluding(
by the time a window row is seeded its batch's log rows are normally
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.
The exclusion is bounded to rows that started at or after the batch's
earliest request. request_id can be chosen by the client
(x-litellm-call-id), so an unbounded exclusion would let a replayed old id
erase a historical row from the seed while its increment still lands.
Without a known start the batch's ids are not excluded at all: that can
only over-count once, which enforcement tolerates, whereas under-counting
is a budget bypass.
"""
if entity_type == Litellm_EntityType.KEY.value:
rows = await prisma_client.db.query_raw(
_SEED_FROM_SPEND_LOGS_KEY_SQL, entity_id, window_start, tuple(exclude_request_ids)
)
bounded_sql, unbounded_sql = _SEED_FROM_SPEND_LOGS_KEY_SQL, _SEED_FROM_SPEND_LOGS_KEY_UNBOUNDED_SQL
elif entity_type == Litellm_EntityType.TEAM.value:
rows = await prisma_client.db.query_raw(
_SEED_FROM_SPEND_LOGS_TEAM_SQL, entity_id, window_start, tuple(exclude_request_ids)
)
bounded_sql, unbounded_sql = _SEED_FROM_SPEND_LOGS_TEAM_SQL, _SEED_FROM_SPEND_LOGS_TEAM_UNBOUNDED_SQL
else:
return None
rows: Final = (
await prisma_client.db.query_raw(unbounded_sql, entity_id, window_start)
if exclude_started_at is None or not exclude_request_ids
else await prisma_client.db.query_raw(
bounded_sql,
entity_id,
window_start,
tuple(exclude_request_ids),
_exclusion_lower_bound(exclude_started_at),
)
)
if not rows:
return 0.0
return float(rows[0].get("total") or 0.0)
def _exclusion_lower_bound(started_at: datetime) -> datetime:
"""LiteLLM_SpendLogs.startTime is TIMESTAMP(3); floor to the second so a
millisecond rounding of the batch's own earliest row cannot slip under it."""
return to_naive_utc(started_at).replace(microsecond=0)
def _primary_key(transaction: WindowSpendTransaction) -> tuple[str, str, str]:
return (
transaction["entity_type"],
@ -168,10 +201,18 @@ async def _seed_base_for_missing_row(
entity_id=transaction["entity_id"],
window_start=datetime.fromisoformat(transaction["window_start"]).replace(tzinfo=timezone.utc),
exclude_request_ids=transaction["request_ids"],
exclude_started_at=_transaction_started_at(transaction),
)
return float(base or 0.0)
def _transaction_started_at(transaction: WindowSpendTransaction) -> datetime | None:
started_at: Final = transaction.get("started_at")
if started_at is None:
return None
return datetime.fromisoformat(started_at).replace(tzinfo=timezone.utc)
def _upsert_params(
transaction: WindowSpendTransaction,
seed_base: float,

View file

@ -30,6 +30,11 @@ class WindowSpendTransaction(TypedDict):
own ~2s poll and will usually have persisted these rows before the window
queue flushes; without the exclusion the seed and the increment would each
count them.
started_at is the earliest request start in the batch. The seed only
subtracts a request_id whose LiteLLM_SpendLogs.startTime is at or after it,
so a client that replays an old id through x-litellm-call-id cannot make the
seed drop the historical row that id already paid for.
"""
entity_type: str
@ -38,6 +43,7 @@ class WindowSpendTransaction(TypedDict):
window_start: str
spend: float
request_ids: Sequence[str]
started_at: str | None
def to_naive_utc(value: datetime) -> datetime:
@ -65,6 +71,7 @@ def build_window_spend_transaction(
window_start: datetime,
spend: float,
request_id: str | None = None,
started_at: datetime | None = None,
) -> WindowSpendTransaction:
return WindowSpendTransaction(
entity_type=entity_type,
@ -73,6 +80,9 @@ def build_window_spend_transaction(
window_start=to_naive_utc(window_start).isoformat(timespec="microseconds"),
spend=spend,
request_ids=() if request_id is None else (request_id,),
started_at=None
if started_at is None
else to_naive_utc(started_at.astimezone(timezone.utc)).isoformat(timespec="microseconds"),
)
@ -80,6 +90,9 @@ def _merge_window_spend_transactions(
payloads: tuple[WindowSpendTransaction, ...],
) -> WindowSpendTransaction:
first: Final = payloads[0]
started_ats: Final = tuple(
started_at for payload in payloads if (started_at := payload.get("started_at")) is not None
)
return WindowSpendTransaction(
entity_type=first["entity_type"],
entity_id=first["entity_id"],
@ -87,6 +100,7 @@ def _merge_window_spend_transactions(
window_start=first["window_start"],
spend=math.fsum(payload["spend"] for payload in payloads),
request_ids=tuple(sorted(frozenset(chain.from_iterable(payload["request_ids"] for payload in payloads)))),
started_at=min(started_ats) if started_ats else None,
)

View file

@ -535,6 +535,7 @@ async def _update_database_and_spend_counters(
end_user_id=end_user_id,
tags=request_tags,
request_id=spend_log_request_id,
request_started_at=start_time,
)
except Exception:
if budget_reservation is not None:

View file

@ -2396,6 +2396,7 @@ async def increment_spend_counters(
end_user_id: str | None = None,
tags: list[str] | None = None,
request_id: str | None = None,
request_started_at: datetime | None = None,
):
"""
Atomically increment spend counters for budget enforcement.
@ -2468,6 +2469,7 @@ async def increment_spend_counters(
window_start=key_window_start,
increment=cost,
request_id=request_id,
request_started_at=request_started_at,
)
async def _team_scope(scope_team_id: str) -> None:
@ -2510,6 +2512,7 @@ async def increment_spend_counters(
window_start=team_window_start,
increment=cost,
request_id=request_id,
request_started_at=request_started_at,
)
async def _team_member_scope(scope_user_id: str, scope_team_id: str) -> None:
@ -2703,14 +2706,15 @@ async def _enqueue_window_spend_row_update(
window_start: datetime | None,
increment: float,
request_id: str | None,
request_started_at: datetime | None,
) -> None:
"""Queue this request's cost against the LiteLLM_BudgetWindowSpend row for
the window, so enforcement can read a maintained total instead of
aggregating LiteLLM_SpendLogs.
request_id is the LiteLLM_SpendLogs id this cost was recorded under; the
flush uses it to keep the one-time seed from counting a request that its
increment already covers.
request_id is the LiteLLM_SpendLogs id this cost was recorded under and
request_started_at its startTime; the flush uses them to keep the one-time
seed from counting a request that its increment already covers.
Enqueued even when the cache increment was skipped for a reserved counter:
the reservation only pre-charged the counter, and the row still owes the
@ -2732,6 +2736,7 @@ async def _enqueue_window_spend_row_update(
window_start=window_start,
spend=increment,
request_id=request_id,
started_at=request_started_at,
)
)
except Exception as e: # noqa: BLE001 # spend tracking must never fail the cost callback

View file

@ -442,6 +442,7 @@ async def test_store_in_memory_spend_updates_pushes_budget_window_spend(
window_start=datetime(2026, 8, 1, tzinfo=timezone.utc),
spend=1.25,
request_id="req-1",
started_at=datetime(2026, 8, 10, 12, 0, 0, tzinfo=timezone.utc),
)
)
@ -466,6 +467,7 @@ async def test_store_in_memory_spend_updates_pushes_budget_window_spend(
"window_start": "2026-08-01T00:00:00.000000",
"spend": 1.25,
"request_ids": ["req-1"],
"started_at": "2026-08-10T12:00:00.000000",
}]

View file

@ -24,6 +24,7 @@ def _txn(
duration: str = "30d",
entity_type: str = "key",
request_id: str | None = None,
started_at: datetime | None = None,
):
return build_window_spend_transaction(
entity_type=entity_type,
@ -32,6 +33,7 @@ def _txn(
window_start=window_start,
spend=spend,
request_id=request_id,
started_at=started_at,
)
@ -47,9 +49,35 @@ def test_build_window_spend_transaction_stores_naive_utc_iso():
"window_start": "2026-08-02T00:00:00.000000",
"spend": 1.0,
"request_ids": ("req-1",),
"started_at": None,
}
def test_build_window_spend_transaction_stores_started_at_as_naive_utc_iso():
"""started_at is compared against LiteLLM_SpendLogs.startTime, which the
spend log writer stores after converting the request start to UTC."""
non_utc = datetime(2026, 8, 10, 8, 30, 15, 123456, tzinfo=timezone(timedelta(hours=-4)))
assert _txn("k1", WINDOW_A, 1.0, started_at=non_utc)["started_at"] == "2026-08-10T12:30:15.123456"
@pytest.mark.asyncio
async def test_aggregation_keeps_the_earliest_started_at_of_the_batch():
"""The seed bounds its request-id exclusion at the batch's earliest start,
so a later start must never win the merge."""
queue = WindowSpendUpdateQueue()
earliest = datetime(2026, 8, 10, 12, 0, 0, tzinfo=timezone.utc)
await queue.add_update(_txn("k1", WINDOW_A, 1.0, request_id="req-2", started_at=earliest + timedelta(seconds=5)))
await queue.add_update(_txn("k1", WINDOW_A, 1.0, request_id="req-1", started_at=earliest))
await queue.add_update(_txn("k1", WINDOW_A, 1.0, request_id="req-3"))
aggregated = await queue.flush_and_get_aggregated_window_spend_transactions()
assert len(aggregated) == 1
assert aggregated[0]["started_at"] == "2026-08-10T12:00:00.000000"
assert aggregated[0]["request_ids"] == ("req-1", "req-2", "req-3")
def test_to_naive_utc_leaves_naive_values_alone():
naive = datetime(2026, 8, 1, 12, 0)
assert to_naive_utc(naive) == naive

View file

@ -2,7 +2,7 @@ import math
import os
import sys
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from typing import Any
sys.path.insert(0, os.path.abspath("../../../.."))
@ -20,6 +20,8 @@ from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
WINDOW_A = datetime(2026, 8, 1, tzinfo=timezone.utc)
WINDOW_B = datetime(2026, 8, 31, tzinfo=timezone.utc)
BATCH_STARTED_AT = datetime(2026, 8, 10, 12, 0, 0, 250_000, tzinfo=timezone.utc)
BEFORE_BATCH = BATCH_STARTED_AT - timedelta(hours=1)
ENTITY_TYPE, ENTITY_ID, WINDOW_DURATION, WINDOW_START, INSERT_SPEND, INCREMENT, NOW = range(7)
@ -85,6 +87,7 @@ class _RecordingAggregate:
entity_id: str,
window_start: datetime,
exclude_request_ids: Any,
exclude_started_at: datetime | None,
) -> float | None:
self.calls.append(
{
@ -92,16 +95,18 @@ class _RecordingAggregate:
"entity_id": entity_id,
"window_start": window_start,
"exclude_request_ids": tuple(exclude_request_ids),
"exclude_started_at": exclude_started_at,
}
)
return self.value
class _SpendLogsFake:
"""Sums the LiteLLM_SpendLogs rows it holds, honouring the request-id
exclusion exactly as the real aggregate's NOT (request_id = ANY(...)) does."""
"""Sums the LiteLLM_SpendLogs rows (request_id, spend, startTime) it holds,
honouring the exclusion exactly as the real aggregate's
NOT (request_id = ANY(...) AND startTime >= bound) does."""
def __init__(self, rows: tuple[tuple[str, float], ...]) -> None:
def __init__(self, rows: tuple[tuple[str, float, datetime], ...]) -> None:
self.rows = rows
async def __call__(
@ -111,9 +116,26 @@ class _SpendLogsFake:
entity_id: str,
window_start: datetime,
exclude_request_ids: Any,
exclude_started_at: datetime | None,
) -> float | None:
excluded = frozenset(exclude_request_ids)
return math.fsum(spend for request_id, spend in self.rows if request_id not in excluded)
excluded = frozenset(exclude_request_ids) if exclude_started_at is not None else frozenset()
return math.fsum(
spend
for request_id, spend, started_at in self.rows
if not (request_id in excluded and started_at >= exclude_started_at)
)
def _batch(request_ids: tuple[str, ...], spend: float, started_at: datetime | None = BATCH_STARTED_AT) -> dict:
return {
"entity_type": "key",
"entity_id": "k1",
"window_duration": "30d",
"window_start": "2026-08-01T00:00:00.000000",
"spend": spend,
"request_ids": request_ids,
"started_at": None if started_at is None else started_at.replace(tzinfo=None).isoformat(timespec="microseconds"),
}
def _existing(entity_type: str, entity_id: str, window_duration: str) -> dict[str, str]:
@ -321,7 +343,9 @@ 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, exclude_request_ids):
async def no_such_column(
prisma_client, entity_type, entity_id, window_start, exclude_request_ids, exclude_started_at
):
return None
await commit_window_spend_updates(
@ -338,7 +362,7 @@ async def test_unknown_entity_type_contributes_no_seed():
async def test_unavailable_spend_logs_aggregate_seeds_zero_rather_than_failing():
db = _FakeDB(existing_rows=[])
async def unavailable(prisma_client, entity_type, entity_id, window_start, exclude_request_ids):
async def unavailable(prisma_client, entity_type, entity_id, window_start, exclude_request_ids, exclude_started_at):
return None
await commit_window_spend_updates(
@ -374,26 +398,32 @@ 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_request_ids_to_exclude():
async def test_seed_receives_the_batch_request_ids_and_earliest_start_to_exclude():
db = _FakeDB(existing_rows=[])
aggregate = _RecordingAggregate(value=0.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(
{
"entity_type": "key",
"entity_id": "k1",
"window_duration": "30d",
"window_start": "2026-08-01T00:00:00.000000",
"spend": 3.0,
"request_ids": ("req-1", "req-2", "req-3"),
},
),
transactions=(_batch(("req-1", "req-2", "req-3"), 3.0),),
spend_logs_aggregate=aggregate,
)
assert aggregate.calls[0]["exclude_request_ids"] == ("req-1", "req-2", "req-3")
assert aggregate.calls[0]["exclude_started_at"] == BATCH_STARTED_AT
@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)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(_batch(("req-1",), 1.0, started_at=None),),
spend_logs_aggregate=aggregate,
)
assert aggregate.calls[0]["exclude_started_at"] is None
@pytest.mark.asyncio
@ -404,21 +434,16 @@ async def test_new_row_is_not_double_counted_when_the_batch_logs_already_flushed
a fresh row land at exactly twice the true spend."""
db = _FakeDB(existing_rows=[])
already_flushed = _SpendLogsFake(
rows=(("req-1", 0.000047), ("req-2", 0.000047), ("req-3", 0.000047)),
rows=(
("req-1", 0.000047, BATCH_STARTED_AT),
("req-2", 0.000047, BATCH_STARTED_AT + timedelta(seconds=1)),
("req-3", 0.000047, BATCH_STARTED_AT + timedelta(seconds=2)),
),
)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(
{
"entity_type": "key",
"entity_id": "k1",
"window_duration": "30d",
"window_start": "2026-08-01T00:00:00.000000",
"spend": 0.000141,
"request_ids": ("req-1", "req-2", "req-3"),
},
),
transactions=(_batch(("req-1", "req-2", "req-3"), 0.000141),),
spend_logs_aggregate=already_flushed,
)
@ -430,20 +455,30 @@ async def test_new_row_is_not_double_counted_when_the_batch_logs_already_flushed
async def test_new_row_still_covers_spend_that_predates_the_batch():
"""The exclusion must not throw away the pre-existing spend the seed is for."""
db = _FakeDB(existing_rows=[])
spend_logs = _SpendLogsFake(rows=(("older", 0.5), ("req-1", 0.000047)))
spend_logs = _SpendLogsFake(rows=(("older", 0.5, BEFORE_BATCH), ("req-1", 0.000047, BATCH_STARTED_AT)))
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(
{
"entity_type": "key",
"entity_id": "k1",
"window_duration": "30d",
"window_start": "2026-08-01T00:00:00.000000",
"spend": 0.000047,
"request_ids": ("req-1",),
},
),
transactions=(_batch(("req-1",), 0.000047),),
spend_logs_aggregate=spend_logs,
)
(_, params), = db.batcher.calls
assert params[INSERT_SPEND] == pytest.approx(0.500047)
@pytest.mark.asyncio
async def test_replayed_request_id_cannot_erase_historical_spend_from_the_seed():
"""request_id can be chosen by the client via x-litellm-call-id. A request
that replays an id from before this batch writes no new LiteLLM_SpendLogs
row (the insert skips duplicates), so the seed must keep counting the
historical row that id belongs to; only its increment is new."""
db = _FakeDB(existing_rows=[])
spend_logs = _SpendLogsFake(rows=(("replayed", 0.5, BEFORE_BATCH),))
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(_batch(("replayed",), 0.000047),),
spend_logs_aggregate=spend_logs,
)
@ -460,16 +495,7 @@ async def test_new_row_is_correct_when_the_batch_logs_have_not_flushed_yet():
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(
{
"entity_type": "key",
"entity_id": "k1",
"window_duration": "30d",
"window_start": "2026-08-01T00:00:00.000000",
"spend": 0.000141,
"request_ids": ("req-1", "req-2", "req-3"),
},
),
transactions=(_batch(("req-1", "req-2", "req-3"), 0.000141),),
spend_logs_aggregate=nothing_flushed,
)
@ -482,7 +508,9 @@ 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_excludes_the_request_ids_by_parameter(entity_type, expected_column):
async def test_seed_aggregate_sql_excludes_the_request_ids_only_within_the_batch_start_bound(
entity_type, expected_column
):
db = _FakeDB(existing_rows=[{"total": 1.25}])
total = await spend_logs_total_excluding(
@ -491,19 +519,49 @@ async def test_seed_aggregate_sql_excludes_the_request_ids_by_parameter(entity_t
entity_id="e1",
window_start=WINDOW_A,
exclude_request_ids=("req-1", "req-2"),
exclude_started_at=BATCH_STARTED_AT,
)
assert total == pytest.approx(1.25)
(query, params), = db.query_raw_calls
normalized = " ".join(query.split())
assert expected_column in normalized
assert "NOT (request_id = ANY($3::text[]))" in normalized
assert "NOT (request_id = ANY($3::text[]) AND \"startTime\" >= ($4::timestamptz AT TIME ZONE 'UTC'))" in normalized
assert 'FROM "LiteLLM_SpendLogs"' in normalized
assert params == ("e1", WINDOW_A, ("req-1", "req-2"))
# startTime is TIMESTAMP(3): the bound is floored to the second so the
# batch's own earliest row cannot round under it.
assert params == ("e1", WINDOW_A, ("req-1", "req-2"), datetime(2026, 8, 10, 12, 0, 0))
# The ids are bound, never spliced into the statement.
assert "req-1" not in query
@pytest.mark.asyncio
@pytest.mark.parametrize(
"exclude_request_ids, exclude_started_at",
[(("req-1",), None), ((), BATCH_STARTED_AT)],
)
async def test_seed_aggregate_excludes_nothing_without_both_ids_and_a_start_bound(
exclude_request_ids, exclude_started_at
):
"""Ids without a start bound would reopen the replayed-id hole, so the
seed counts everything instead; at worst that over-counts one batch."""
db = _FakeDB(existing_rows=[{"total": 1.25}])
total = await spend_logs_total_excluding(
prisma_client=_FakePrismaClient(db),
entity_type="key",
entity_id="e1",
window_start=WINDOW_A,
exclude_request_ids=exclude_request_ids,
exclude_started_at=exclude_started_at,
)
assert total == pytest.approx(1.25)
(query, params), = db.query_raw_calls
assert "request_id" not in query
assert params == ("e1", WINDOW_A)
@pytest.mark.asyncio
async def test_seed_aggregate_returns_none_for_an_entity_type_with_no_spend_logs_column():
db = _FakeDB(existing_rows=[])
@ -514,6 +572,7 @@ async def test_seed_aggregate_returns_none_for_an_entity_type_with_no_spend_logs
entity_id="u1",
window_start=WINDOW_A,
exclude_request_ids=(),
exclude_started_at=None,
)
assert total is None
@ -530,6 +589,7 @@ async def test_seed_aggregate_treats_an_entity_with_no_rows_as_zero():
entity_id="k-unknown",
window_start=WINDOW_A,
exclude_request_ids=(),
exclude_started_at=None,
)
assert total == 0.0

View file

@ -445,6 +445,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
)
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
start_time = datetime.now()
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
@ -456,7 +457,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
start_time=start_time,
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
@ -474,6 +475,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
end_user_id="test_end_user_id",
tags=["tag-a"],
request_id="chatcmpl-abc123",
request_started_at=start_time,
)

View file

@ -11234,6 +11234,32 @@ async def test_window_spend_row_carries_the_spend_log_request_id():
assert enqueued[0]["request_ids"] == ("chatcmpl-abc123",)
@pytest.mark.asyncio
async def test_window_spend_row_carries_the_request_start_time():
"""The seed only excludes a batch id whose LiteLLM_SpendLogs.startTime is at
or after this, so it must be the same start the spend log was written with."""
from litellm.proxy.proxy_server import increment_spend_counters
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
key_obj = MagicMock()
key_obj.budget_limits = [
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
]
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
await increment_spend_counters(
token="hashed-token",
team_id=None,
user_id=None,
response_cost=0.25,
request_id="chatcmpl-abc123",
request_started_at=datetime(2026, 8, 10, 12, 0, 0, 500_000, tzinfo=timezone.utc),
)
enqueued = await _drain(queue)
assert enqueued[0]["started_at"] == "2026-08-10T12:00:00.500000"
@pytest.mark.asyncio
async def test_window_spend_row_without_a_request_id_excludes_nothing():
from litellm.proxy.proxy_server import increment_spend_counters