Merge pull request #38851 from BerriAI/litellm_window_spend_seed_by_time

refactor(proxy): bound the budget window seed by time instead of request ids
This commit is contained in:
ryan-crabbe-berri 2026-08-31 17:21:45 -07:00 committed by GitHub
commit 50eed7efa2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 256 additions and 421 deletions

View file

@ -7,20 +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, excluding
the requests whose increments are in the same batch so neither source counts
them twice. One gap survives that exclusion: without the Redis transaction
buffer every pod flushes its own increments, so a row seeded by one pod can
include spend logs whose increments are still queued on another pod, and those
increments are added again when that pod flushes. That is bounded by a single
flush interval, happens at most once per window row, and only ever over-counts:
the seed never omits spend, because every increment not yet in the row still
reaches it on its own pod's next flush. A row therefore 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.
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
@ -67,33 +62,46 @@ _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 \"startTime\" >= ($4::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 NOT (request_id = ANY($3::text[]) AND \"startTime\" >= ($4::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 since window_start, ignoring the
requests whose ids are handed in.
"""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
expensive aggregate stays swappable.
@ -105,21 +113,19 @@ class WindowSpendLogsAggregate(Protocol):
entity_type: str,
entity_id: str,
window_start: datetime,
exclude_request_ids: Sequence[str],
exclude_started_at: datetime | None,
) -> float | None: ...
batch_started_at: datetime | None,
) -> WindowSeedTotals | None: ...
async def spend_logs_total_excluding(
async def spend_logs_seed_totals(
prisma_client: "PrismaClient",
entity_type: str,
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.
batch_started_at: datetime | None,
) -> 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
@ -127,13 +133,12 @@ async def spend_logs_total_excluding(
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.
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
@ -143,21 +148,23 @@ async def spend_logs_total_excluding(
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
if batch_started_at is None
else await prisma_client.db.query_raw(
bounded_sql,
entity_id,
window_start,
tuple(exclude_request_ids),
_exclusion_lower_bound(exclude_started_at),
_exclusion_upper_bound(batch_started_at),
)
)
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_lower_bound(started_at: datetime) -> datetime:
def _exclusion_upper_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)
@ -194,20 +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 excludes this batch's own requests so they are
counted by their increments alone.
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),
exclude_request_ids=transaction["request_ids"],
exclude_started_at=_transaction_started_at(transaction),
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:
@ -241,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_excluding,
spend_logs_aggregate: WindowSpendLogsAggregate = spend_logs_seed_totals,
) -> None:
"""Apply aggregated window increments to LiteLLM_BudgetWindowSpend.

View file

@ -215,11 +215,7 @@ class DBSpendUpdateWriter:
start_time: datetime | None,
end_time: datetime | None,
response_cost: float | None,
) -> str | None:
"""Returns the LiteLLM_SpendLogs request_id this call was recorded
under, so the caller can tell the budget-window writer which log rows
its increments already cover. None when the payload could not be built.
"""
) -> None:
from litellm.proxy.proxy_server import (
disable_spend_logs,
litellm_proxy_budget_name,
@ -236,7 +232,7 @@ class DBSpendUpdateWriter:
team_id,
)
if ProxyUpdateSpend.disable_spend_updates() is True:
return None
return
if token is not None and isinstance(token, str) and token.startswith("sk-"):
hashed_token = hash_token(token=token)
else:
@ -310,7 +306,6 @@ class DBSpendUpdateWriter:
)
verbose_proxy_logger.debug("Runs spend update on all tables")
return payload.get("request_id")
except Exception:
spend_log_error(
"Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue "
@ -323,7 +318,7 @@ class DBSpendUpdateWriter:
org_id,
end_user_id,
)
return None
return
async def _enqueue_tool_usage_transaction(
self,

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

@ -26,17 +26,12 @@ class WindowSpendTransaction(TypedDict):
window_start is an ISO-8601 string rather than a datetime so the
transaction survives the JSON round trip through the Redis buffer.
request_ids carries the LiteLLM_SpendLogs ids this spend came from. The
one-time seed for a window that has no row yet subtracts them from its
LiteLLM_SpendLogs aggregate, because the spend log writer flushes on its
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.
started_at is the earliest request start in the batch. The one-time seed for
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]
@ -44,10 +39,37 @@ class WindowSpendTransaction(TypedDict):
window_duration: ReadOnly[str]
window_start: ReadOnly[str]
spend: ReadOnly[float]
request_ids: ReadOnly[Sequence[str]]
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:
@ -72,7 +94,6 @@ def build_window_spend_transaction(
window_duration: str,
window_start: datetime,
spend: float,
request_id: str | None = None,
started_at: datetime | None = None,
) -> WindowSpendTransaction:
return WindowSpendTransaction(
@ -81,7 +102,6 @@ def build_window_spend_transaction(
window_duration=window_duration,
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"),
@ -101,7 +121,6 @@ def _merge_window_spend_transactions(
window_duration=first["window_duration"],
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

@ -587,7 +587,7 @@ async def _update_database_and_spend_counters(
model_access_groups: Sequence[str] | None = None,
) -> None:
try:
spend_log_request_id = await proxy_logging_obj.db_spend_update_writer.update_database(
await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
@ -623,7 +623,6 @@ async def _update_database_and_spend_counters(
budget_reservation=budget_reservation,
end_user_id=end_user_id,
tags=request_tags,
request_id=spend_log_request_id,
request_started_at=start_time,
model_access_groups=model_access_groups,
)

View file

@ -2658,7 +2658,6 @@ async def increment_spend_counters(
budget_reservation: dict | None = None,
end_user_id: str | None = None,
tags: list[str] | None = None,
request_id: str | None = None,
request_started_at: datetime | None = None,
model_access_groups: Sequence[str] | None = None,
):
@ -2733,7 +2732,6 @@ async def increment_spend_counters(
window_duration=duration,
window_start=key_window_start,
increment=cost,
request_id=request_id,
request_started_at=request_started_at,
)
@ -2777,7 +2775,6 @@ async def increment_spend_counters(
window_duration=duration,
window_start=team_window_start,
increment=cost,
request_id=request_id,
request_started_at=request_started_at,
)
@ -3005,16 +3002,15 @@ async def _enqueue_window_spend_row_update(
window_duration: str,
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 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.
request_started_at is this request's LiteLLM_SpendLogs startTime; the flush
stops the one-time seed there so a request its increment already covers is
not counted twice.
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
@ -3035,7 +3031,6 @@ async def _enqueue_window_spend_row_update(
window_duration=window_duration,
window_start=window_start,
spend=increment,
request_id=request_id,
started_at=request_started_at,
)
)

View file

@ -203,7 +203,7 @@ async def test_get_all_transactions_from_redis_buffer_pipeline(redis_update_buff
"window_duration": "30d",
"window_start": "2026-08-01T00:00:00.000000",
"spend": 3.0,
"request_ids": ["req-1"],
"started_at": None,
}
]
)
@ -233,13 +233,11 @@ async def test_get_all_transactions_from_redis_buffer_pipeline(redis_update_buff
window_spend,
) = result
# Budget window spend from two pods is summed per window, not overwritten,
# and both pods' request ids reach the seed exclusion.
# Budget window spend from two pods is summed per window, not overwritten.
assert window_spend is not None
assert len(window_spend) == 1
assert window_spend[0]["spend"] == 6.0
assert window_spend[0]["entity_id"] == "hashed-token"
assert window_spend[0]["request_ids"] == ("req-1",)
# Verify db spend was parsed correctly
assert db_spend is not None
@ -326,7 +324,6 @@ async def test_restored_window_spend_transactions_drain_back_unchanged(redis_upd
window_duration="30d",
window_start=datetime(2026, 8, 1, tzinfo=timezone.utc),
spend=3.0,
request_id="req-1",
started_at=datetime(2026, 8, 10, 12, 0, tzinfo=timezone.utc),
),
)
@ -500,7 +497,6 @@ async def test_store_in_memory_spend_updates_pushes_budget_window_spend(redis_up
window_duration="30d",
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),
)
)
@ -526,12 +522,45 @@ async def test_store_in_memory_spend_updates_pushes_budget_window_spend(redis_up
"window_duration": "30d",
"window_start": "2026-08-01T00:00:00.000000",
"spend": 1.25,
"request_ids": ["req-1"],
"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

@ -19,7 +19,6 @@ def _txn(
spend: float,
duration: str = "30d",
entity_type: str = "key",
request_id: str | None = None,
started_at: datetime | None = None,
):
return build_window_spend_transaction(
@ -28,7 +27,6 @@ def _txn(
window_duration=duration,
window_start=window_start,
spend=spend,
request_id=request_id,
started_at=started_at,
)
@ -38,13 +36,12 @@ def test_build_window_spend_transaction_stores_naive_utc_iso():
TIMESTAMP(3) column, so a non-UTC input must be converted, not truncated."""
non_utc = datetime(2026, 8, 1, 20, 0, tzinfo=timezone(timedelta(hours=-4)))
assert _txn("k1", non_utc, 1.0, request_id="req-1") == {
assert _txn("k1", non_utc, 1.0) == {
"entity_type": "key",
"entity_id": "k1",
"window_duration": "30d",
"window_start": "2026-08-02T00:00:00.000000",
"spend": 1.0,
"request_ids": ("req-1",),
"started_at": None,
}
@ -59,19 +56,19 @@ def test_build_window_spend_transaction_stores_started_at_as_naive_utc_iso():
@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."""
"""The seed stops at the batch's earliest start, so a later start must never
win the merge: it would push the cutoff forward and count a request the
increments already cover."""
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"))
await queue.add_update(_txn("k1", WINDOW_A, 1.0, started_at=earliest + timedelta(seconds=5)))
await queue.add_update(_txn("k1", WINDOW_A, 1.0, started_at=earliest))
await queue.add_update(_txn("k1", WINDOW_A, 1.0))
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():
@ -208,62 +205,12 @@ def test_aggregation_survives_the_redis_json_round_trip():
assert reloaded == aggregated
@pytest.mark.asyncio
async def test_aggregation_unions_the_request_ids_of_merged_increments():
"""The seed excludes exactly the requests its batch already covers, so every
merged increment's id has to survive aggregation."""
queue = WindowSpendUpdateQueue()
await queue.add_update(_txn("k1", WINDOW_A, 1.0, request_id="req-1"))
await queue.add_update(_txn("k1", WINDOW_A, 2.0, request_id="req-2"))
aggregated = await queue.flush_and_get_aggregated_window_spend_transactions()
assert len(aggregated) == 1
assert aggregated[0]["request_ids"] == ("req-1", "req-2")
@pytest.mark.asyncio
async def test_request_ids_stay_with_their_own_window():
queue = WindowSpendUpdateQueue()
await queue.add_update(_txn("k1", WINDOW_A, 1.0, request_id="req-a"))
await queue.add_update(_txn("k1", WINDOW_B, 2.0, request_id="req-b"))
aggregated = await queue.flush_and_get_aggregated_window_spend_transactions()
assert {payload["window_start"]: payload["request_ids"] for payload in aggregated} == {
"2026-08-01T00:00:00.000000": ("req-a",),
"2026-08-31T00:00:00.000000": ("req-b",),
}
@pytest.mark.asyncio
async def test_request_ids_are_deduplicated_and_ordered():
queue = WindowSpendUpdateQueue()
await queue.add_update(_txn("k1", WINDOW_A, 1.0, request_id="req-b"))
await queue.add_update(_txn("k1", WINDOW_A, 1.0, request_id="req-a"))
await queue.add_update(_txn("k1", WINDOW_A, 1.0, request_id="req-a"))
aggregated = await queue.flush_and_get_aggregated_window_spend_transactions()
assert aggregated[0]["request_ids"] == ("req-a", "req-b")
@pytest.mark.asyncio
async def test_increment_without_a_request_id_carries_no_exclusion():
queue = WindowSpendUpdateQueue()
await queue.add_update(_txn("k1", WINDOW_A, 1.0))
aggregated = await queue.flush_and_get_aggregated_window_spend_transactions()
assert aggregated[0]["request_ids"] == ()
def test_request_ids_survive_the_redis_json_round_trip():
def test_started_at_survives_the_redis_json_round_trip():
aggregated = WindowSpendUpdateQueue.get_aggregated_window_spend_transactions(
[(_txn("k1", WINDOW_A, 1.0, request_id="req-1"),)]
[(_txn("k1", WINDOW_A, 1.0, started_at=datetime(2026, 8, 10, 12, 0, tzinfo=timezone.utc)),)]
)
reloaded = WindowSpendUpdateQueue.get_aggregated_window_spend_transactions([json.loads(json.dumps(aggregated))])
assert reloaded[0]["request_ids"] == ("req-1",)
assert reloaded[0]["started_at"] == "2026-08-10T12:00:00.000000"
assert reloaded[0]["spend"] == 1.0

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_excluding,
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__(
@ -82,25 +88,23 @@ class _RecordingAggregate:
entity_type: str,
entity_id: str,
window_start: datetime,
exclude_request_ids: Any,
exclude_started_at: datetime | None,
) -> float | None:
batch_started_at: datetime | None,
) -> WindowSeedTotals | None:
self.calls.append(
{
"entity_type": entity_type,
"entity_id": entity_id,
"window_start": window_start,
"exclude_request_ids": tuple(exclude_request_ids),
"exclude_started_at": exclude_started_at,
"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 exclusion exactly as the real aggregate's
NOT (request_id = ANY(...) AND 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
@ -111,25 +115,25 @@ class _SpendLogsFake:
entity_type: str,
entity_id: str,
window_start: datetime,
exclude_request_ids: Any,
exclude_started_at: datetime | None,
) -> float | None:
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)
batch_started_at: datetime | None,
) -> 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
),
)
def _batch(request_ids: tuple[str, ...], spend: float, started_at: datetime | None = BATCH_STARTED_AT) -> dict:
def _batch(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"),
@ -156,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),
@ -183,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),
@ -200,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),
@ -223,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),
@ -260,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),
@ -278,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),
@ -306,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),
@ -344,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, exclude_request_ids, exclude_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(
@ -363,7 +365,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, exclude_started_at):
async def unavailable(prisma_client, entity_type, entity_id, window_start, batch_started_at):
return None
await commit_window_spend_updates(
@ -399,32 +401,31 @@ 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_and_earliest_start_to_exclude():
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),
transactions=(_batch(("req-1", "req-2", "req-3"), 3.0),),
transactions=(_batch(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
assert aggregate.calls[0]["batch_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)
aggregate = _RecordingAggregate(total=0.0)
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(_batch(("req-1",), 1.0, started_at=None),),
transactions=(_batch(1.0, started_at=None),),
spend_logs_aggregate=aggregate,
)
assert aggregate.calls[0]["exclude_started_at"] is None
assert aggregate.calls[0]["batch_started_at"] is None
@pytest.mark.asyncio
@ -444,7 +445,7 @@ async def test_new_row_is_not_double_counted_when_the_batch_logs_already_flushed
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(_batch(("req-1", "req-2", "req-3"), 0.000141),),
transactions=(_batch(0.000141),),
spend_logs_aggregate=already_flushed,
)
@ -460,7 +461,7 @@ async def test_new_row_still_covers_spend_that_predates_the_batch():
await commit_window_spend_updates(
prisma_client=_FakePrismaClient(db),
transactions=(_batch(("req-1",), 0.000047),),
transactions=(_batch(0.000047),),
spend_logs_aggregate=spend_logs,
)
@ -469,22 +470,29 @@ async def test_new_row_still_covers_spend_that_predates_the_batch():
@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."""
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=(("replayed", 0.5, BEFORE_BATCH),))
spend_logs = _SpendLogsFake(
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(
prisma_client=_FakePrismaClient(db),
transactions=(_batch(("replayed",), 0.000047),),
transactions=(_batch(0.000047),),
spend_logs_aggregate=spend_logs,
)
((_, params),) = db.batcher.calls
assert params[INSERT_SPEND] == pytest.approx(0.500047)
assert params[INSERT_SPEND] == pytest.approx(0.750047)
@pytest.mark.asyncio
@ -496,7 +504,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=(_batch(("req-1", "req-2", "req-3"), 0.000141),),
transactions=(_batch(0.000141),),
spend_logs_aggregate=nothing_flushed,
)
@ -509,57 +517,49 @@ 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_only_within_the_batch_start_bound(
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_excluding(
totals = await spend_logs_seed_totals(
prisma_client=_FakePrismaClient(db),
entity_type=entity_type,
entity_id="e1",
window_start=WINDOW_A,
exclude_request_ids=("req-1", "req-2"),
exclude_started_at=BATCH_STARTED_AT,
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 "NOT (request_id = ANY($3::text[]) AND \"startTime\" >= ($4::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.
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
assert params == ("e1", WINDOW_A, datetime(2026, 8, 10, 12, 0, 0))
# Nothing the caller supplied reaches the statement text.
assert "e1" 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}])
async def test_seed_aggregate_sums_the_whole_window_without_a_start_bound():
"""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_excluding(
totals = await spend_logs_seed_totals(
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,
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 "request_id" not in query
assert '"startTime" <' not in query
assert params == ("e1", WINDOW_A)
@ -567,16 +567,15 @@ async def test_seed_aggregate_excludes_nothing_without_both_ids_and_a_start_boun
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_excluding(
totals = await spend_logs_seed_totals(
prisma_client=_FakePrismaClient(db),
entity_type="user",
entity_id="u1",
window_start=WINDOW_A,
exclude_request_ids=(),
exclude_started_at=None,
batch_started_at=None,
)
assert total is None
assert totals is None
assert db.query_raw_calls == []
@ -584,13 +583,12 @@ 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_excluding(
totals = await spend_logs_seed_totals(
prisma_client=_FakePrismaClient(db),
entity_type="key",
entity_id="k-unknown",
window_start=WINDOW_A,
exclude_request_ids=(),
exclude_started_at=None,
batch_started_at=None,
)
assert total == 0.0
assert totals == WindowSeedTotals(total=0.0, before_batch=0.0)

View file

@ -2581,7 +2581,6 @@ async def test_failed_window_spend_commit_requeues_the_increments_and_continues_
window_duration="30d",
window_start=datetime(2026, 8, 1, tzinfo=timezone.utc),
spend=0.5,
request_id="req-1",
)
await db_writer.window_spend_update_queue.add_update(transaction)
db = _WindowSpendFakeDB()
@ -2611,7 +2610,6 @@ async def test_failed_window_spend_commit_from_redis_is_restored_to_redis():
window_duration="7d",
window_start=datetime(2026, 8, 1, tzinfo=timezone.utc),
spend=2.0,
request_id="req-1",
),
)
mock_redis_update_buffer = AsyncMock()
@ -2638,74 +2636,6 @@ async def test_failed_window_spend_commit_from_redis_is_restored_to_redis():
db_writer.pod_lock_manager.release_lock.assert_awaited_once()
@pytest.mark.asyncio
async def test_update_database_returns_the_spend_log_request_id():
"""The budget-window seed excludes the log rows its increments already
cover, so the caller needs the id this call was recorded under. It cannot
be re-derived: cache hits append time.time() to the id."""
db_writer = DBSpendUpdateWriter()
db_writer._insert_spend_log_to_db = AsyncMock()
db_writer._enqueue_tool_usage_transaction = AsyncMock()
with (
patch.multiple( # test-quality-ok: update_database lazily imports these proxy_server globals; no injection seam
"litellm.proxy.proxy_server",
disable_spend_logs=False,
prisma_client=MagicMock(),
litellm_proxy_budget_name="test-budget",
)
):
request_id = await db_writer.update_database(
token="test-token",
user_id="test-user",
end_user_id=None,
team_id="test-team",
org_id=None,
kwargs={"model": "gpt-4", "custom_llm_provider": "openai", "litellm_call_id": "call-xyz"},
completion_response=MagicMock(),
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.1,
)
await asyncio.sleep(0)
assert request_id is not None
# Same id the spend log row was queued under.
assert request_id == db_writer._insert_spend_log_to_db.call_args[1]["payload"]["request_id"]
@pytest.mark.asyncio
async def test_update_database_returns_none_when_the_payload_cannot_be_built():
db_writer = DBSpendUpdateWriter()
with (
patch.multiple( # test-quality-ok: update_database lazily imports these proxy_server globals; no injection seam
"litellm.proxy.proxy_server",
disable_spend_logs=False,
prisma_client=MagicMock(),
litellm_proxy_budget_name="test-budget",
),
patch( # test-quality-ok: the payload builder is called by name inside update_database; no injection seam
"litellm.proxy.spend_tracking.spend_tracking_utils.get_logging_payload",
side_effect=Exception("payload boom"),
),
):
request_id = await db_writer.update_database(
token="test-token",
user_id="test-user",
end_user_id=None,
team_id="test-team",
org_id=None,
kwargs={"model": "gpt-4"},
completion_response=MagicMock(),
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.1,
)
assert request_id is None
@pytest.mark.asyncio
async def test_commit_spend_updates_to_db_does_not_stamp_key_settings_updated_at():
"""Spend flushes must leave settings_updated_at alone, or it decays into

View file

@ -567,9 +567,7 @@ async def test_update_database_and_spend_counters_preserves_db_exception_when_re
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_updates_counters_after_db_update():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(
return_value="chatcmpl-abc123"
)
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock()
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
start_time = datetime.now()
@ -602,7 +600,6 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
budget_reservation=budget_reservation,
end_user_id="test_end_user_id",
tags=["tag-a"],
request_id="chatcmpl-abc123",
request_started_at=start_time,
model_access_groups=("premium",),
)
@ -1884,61 +1881,6 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request(
)
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_forwards_the_spend_log_request_id():
"""The budget-window flush excludes the log rows its increments already
cover. That only works if the id update_database recorded the row under is
handed to the counter update, so this seam is load-bearing."""
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(
return_value="chatcmpl-abc123"
)
increment_spend_counters = AsyncMock()
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=None,
)
assert increment_spend_counters.await_args.kwargs["request_id"] == "chatcmpl-abc123"
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_forwards_a_missing_request_id_as_none():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(return_value=None)
increment_spend_counters = AsyncMock()
await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=None,
)
assert increment_spend_counters.await_args.kwargs["request_id"] is None
class _FakeDeploymentLookup:
"""Deployment lookup returning the access groups each deployment declares."""

View file

@ -11399,35 +11399,10 @@ async def test_no_window_spend_row_enqueued_without_budget_limits():
assert enqueued == []
@pytest.mark.asyncio
async def test_window_spend_row_carries_the_spend_log_request_id():
"""The flush excludes these ids from its one-time seed, so the id threaded
here has to be the same one the LiteLLM_SpendLogs row was written under."""
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",
)
enqueued = await _drain(queue)
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."""
"""The seed sums LiteLLM_SpendLogs only up to this point, so it must be the
same start the spend log row was written with."""
from litellm.proxy.proxy_server import increment_spend_counters
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
@ -11442,7 +11417,6 @@ async def test_window_spend_row_carries_the_request_start_time():
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)
@ -11451,26 +11425,7 @@ async def test_window_spend_row_carries_the_request_start_time():
@pytest.mark.asyncio
async def test_window_spend_row_without_a_request_id_excludes_nothing():
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
)
enqueued = await _drain(queue)
assert enqueued[0]["request_ids"] == ()
@pytest.mark.asyncio
async def test_team_window_spend_row_carries_the_request_id():
async def test_team_window_spend_row_carries_the_request_start_time():
from litellm.proxy.proxy_server import increment_spend_counters
reset_at = datetime.now(timezone.utc) + timedelta(days=3)
@ -11485,11 +11440,11 @@ async def test_team_window_spend_row_carries_the_request_id():
team_id="team-1",
user_id=None,
response_cost=1.5,
request_id="chatcmpl-team",
request_started_at=datetime(2026, 8, 10, 12, 0, 0, 500_000, tzinfo=timezone.utc),
)
enqueued = await _drain(queue)
assert enqueued[0]["request_ids"] == ("chatcmpl-team",)
assert enqueued[0]["started_at"] == "2026-08-10T12:00:00.500000"
def _mock_startup_prisma_client(health_check_error=None, connect_error=None):