mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-05 08:07:05 +00:00
Merge branch 'litellm_window_spend_writer' into litellm_window_spend_reader
This commit is contained in:
commit
7745fe887f
9 changed files with 245 additions and 66 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -2401,6 +2401,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.
|
||||
|
|
@ -2474,6 +2475,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:
|
||||
|
|
@ -2517,6 +2519,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:
|
||||
|
|
@ -2710,14 +2713,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
|
||||
|
|
@ -2739,6 +2743,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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -11242,6 +11242,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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue