Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/db/budget_window_spend_writer.py: 41%
70 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Writer for LiteLLM_BudgetWindowSpend.
4The table holds one row per configured budget window whose window_start rolls
5forward in place, so budget enforcement can read a maintained running total
6instead of aggregating LiteLLM_SpendLogs every time a window counter goes cold
7(issue #35766). Raw SQL rather than the Prisma upsert helper because the
8conditional roll cannot be expressed through the query builder.
10Seeding a row that does not exist yet reads LiteLLM_SpendLogs once and takes
11off what the increments being flushed will add, so neither source counts the
12same request twice. A row therefore lags real spend by at most one flush
13interval of increments queued elsewhere: the same lag the SpendLogs aggregate
14it replaces (and every other spend column) already has.
15"""
17from collections.abc import Sequence
18from dataclasses import dataclass
19from datetime import datetime, timedelta, timezone
20from typing import TYPE_CHECKING, Final, Protocol
22from litellm._logging import verbose_proxy_logger
23from litellm.proxy._types import Litellm_EntityType
24from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
25 WindowSpendTransaction,
26 to_naive_utc,
27 window_spend_group_key,
28)
30if TYPE_CHECKING: 30 ↛ 31line 30 didn't jump to line 31 because the condition on line 30 was never true
31 from litellm.proxy.utils import PrismaClient
34_SELECT_EXISTING_ROWS_SQL: Final = (
35 'SELECT entity_type, entity_id, window_duration FROM "LiteLLM_BudgetWindowSpend" '
36 "WHERE (entity_type, entity_id, window_duration) "
37 "IN (SELECT * FROM unnest($1::text[], $2::text[], $3::text[]))"
38)
40_UPSERT_WINDOW_SPEND_SQL: Final = (
41 'INSERT INTO "LiteLLM_BudgetWindowSpend" '
42 "(entity_type, entity_id, window_duration, window_start, spend, created_at, updated_at) "
43 "VALUES ($1, $2, $3, ($4::timestamptz AT TIME ZONE 'UTC'), $5, "
44 "($7::timestamptz AT TIME ZONE 'UTC'), ($7::timestamptz AT TIME ZONE 'UTC')) "
45 "ON CONFLICT (entity_type, entity_id, window_duration) DO UPDATE SET "
46 "spend = CASE "
47 'WHEN "LiteLLM_BudgetWindowSpend".window_start >= EXCLUDED.window_start '
48 'THEN "LiteLLM_BudgetWindowSpend".spend + $6 '
49 "ELSE EXCLUDED.spend "
50 "END, "
51 'window_start = GREATEST("LiteLLM_BudgetWindowSpend".window_start, EXCLUDED.window_start), '
52 "updated_at = ($7::timestamptz AT TIME ZONE 'UTC')"
53)
55_ROLL_WINDOW_SPEND_SQL: Final = (
56 'UPDATE "LiteLLM_BudgetWindowSpend" SET '
57 "window_start = ($4::timestamptz AT TIME ZONE 'UTC'), "
58 "spend = 0, "
59 "updated_at = ($5::timestamptz AT TIME ZONE 'UTC') "
60 "WHERE entity_type = $1 AND entity_id = $2 AND window_duration = $3 "
61 "AND window_start < ($4::timestamptz AT TIME ZONE 'UTC')"
62)
64_SEED_FROM_SPEND_LOGS_KEY_SQL: Final = (
65 "SELECT COALESCE(SUM(spend), 0.0) AS total, "
66 "COALESCE(SUM(spend) FILTER (WHERE \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')), 0.0) AS before_batch "
67 'FROM "LiteLLM_SpendLogs" '
68 "WHERE api_key = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
69)
71_SEED_FROM_SPEND_LOGS_TEAM_SQL: Final = (
72 "SELECT COALESCE(SUM(spend), 0.0) AS total, "
73 "COALESCE(SUM(spend) FILTER (WHERE \"startTime\" < ($3::timestamptz AT TIME ZONE 'UTC')), 0.0) AS before_batch "
74 'FROM "LiteLLM_SpendLogs" '
75 "WHERE team_id = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
76)
78_SEED_FROM_SPEND_LOGS_KEY_UNBOUNDED_SQL: Final = (
79 "SELECT COALESCE(SUM(spend), 0.0) AS total, COALESCE(SUM(spend), 0.0) AS before_batch "
80 'FROM "LiteLLM_SpendLogs" '
81 "WHERE api_key = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
82)
84_SEED_FROM_SPEND_LOGS_TEAM_UNBOUNDED_SQL: Final = (
85 "SELECT COALESCE(SUM(spend), 0.0) AS total, COALESCE(SUM(spend), 0.0) AS before_batch "
86 'FROM "LiteLLM_SpendLogs" '
87 "WHERE team_id = $1 AND \"startTime\" >= ($2::timestamptz AT TIME ZONE 'UTC')"
88)
90_UPSERT_TRANSACTION_TIMEOUT: Final = timedelta(seconds=60)
93@dataclass(frozen=True, slots=True)
94class WindowSeedTotals:
95 """The two sums a seed needs: everything persisted for the window, and the
96 part of it that predates the batch being flushed."""
98 total: float
99 before_batch: float
102class WindowSpendLogsAggregate(Protocol):
103 """Sums LiteLLM_SpendLogs for one entity since window_start, split at the
104 batch's earliest request.
106 Injected so the flush can be exercised without a database and so the
107 expensive aggregate stays swappable.
108 """
110 async def __call__( 110 ↛ exitline 110 didn't return from function '__call__' because
111 self,
112 prisma_client: "PrismaClient",
113 entity_type: str,
114 entity_id: str,
115 window_start: datetime,
116 batch_started_at: datetime | None,
117 ) -> WindowSeedTotals | None: ...
120async def spend_logs_seed_totals(
121 prisma_client: "PrismaClient",
122 entity_type: str,
123 entity_id: str,
124 window_start: datetime,
125 batch_started_at: datetime | None,
126) -> WindowSeedTotals | None:
127 """LiteLLM_SpendLogs spend for one entity since window_start, both in full
128 and up to the start of the batch being flushed, in one scan.
130 The spend log writer drains its own queue on a ~2s poll whenever anything
131 is queued, while window increments flush on the much slower batch tick, so
132 by the time a window row is seeded its batch's log rows are normally
133 already in the table. Counting them in the seed and again in the increment
134 is what made a fresh row land at twice the true spend.
136 Both halves are needed because neither is safe alone: the full sum
137 double-counts this batch, and the sum before the batch drops spend another
138 pod has already persisted but not yet incremented. _seed_base picks between
139 them. Without a known batch start the two are the same sum, so the seed
140 counts everything: that can only over-count once, which enforcement
141 tolerates, whereas under-counting is a budget bypass.
142 """
143 if entity_type == Litellm_EntityType.KEY.value:
144 bounded_sql, unbounded_sql = _SEED_FROM_SPEND_LOGS_KEY_SQL, _SEED_FROM_SPEND_LOGS_KEY_UNBOUNDED_SQL
145 elif entity_type == Litellm_EntityType.TEAM.value:
146 bounded_sql, unbounded_sql = _SEED_FROM_SPEND_LOGS_TEAM_SQL, _SEED_FROM_SPEND_LOGS_TEAM_UNBOUNDED_SQL
147 else:
148 return None
149 rows: Final = (
150 await prisma_client.db.query_raw(unbounded_sql, entity_id, window_start)
151 if batch_started_at is None
152 else await prisma_client.db.query_raw(
153 bounded_sql,
154 entity_id,
155 window_start,
156 _exclusion_upper_bound(batch_started_at),
157 )
158 )
159 if not rows:
160 return WindowSeedTotals(total=0.0, before_batch=0.0)
161 return WindowSeedTotals(
162 total=float(rows[0].get("total") or 0.0),
163 before_batch=float(rows[0].get("before_batch") or 0.0),
164 )
167def _exclusion_upper_bound(started_at: datetime) -> datetime:
168 """LiteLLM_SpendLogs.startTime is TIMESTAMP(3); floor to the second so a
169 millisecond rounding of the batch's own earliest row cannot slip under it."""
170 return to_naive_utc(started_at).replace(microsecond=0)
173def _primary_key(transaction: WindowSpendTransaction) -> tuple[str, str, str]:
174 return (
175 transaction["entity_type"],
176 transaction["entity_id"],
177 transaction["window_duration"],
178 )
181async def _existing_primary_keys(
182 prisma_client: "PrismaClient",
183 transactions: tuple[WindowSpendTransaction, ...],
184) -> frozenset[tuple[str, str, str]]:
185 rows: Final = await prisma_client.db.query_raw(
186 _SELECT_EXISTING_ROWS_SQL,
187 tuple(transaction["entity_type"] for transaction in transactions),
188 tuple(transaction["entity_id"] for transaction in transactions),
189 tuple(transaction["window_duration"] for transaction in transactions),
190 )
191 return frozenset((row["entity_type"], row["entity_id"], row["window_duration"]) for row in rows or ())
194async def _seed_base_for_missing_row(
195 prisma_client: "PrismaClient",
196 transaction: WindowSpendTransaction,
197 existing_primary_keys: frozenset[tuple[str, str, str]],
198 spend_logs_aggregate: WindowSpendLogsAggregate,
199) -> float:
200 """Spend already recorded for a window that has no row yet.
202 This is the LiteLLM_SpendLogs aggregate the window counter reseed runs on
203 every cold counter today, but here it runs once per window lifetime and off
204 the request path, and it discounts the queued increments so they are
205 counted once.
206 """
207 if _primary_key(transaction) in existing_primary_keys:
208 return 0.0
209 totals: Final = await spend_logs_aggregate(
210 prisma_client=prisma_client,
211 entity_type=transaction["entity_type"],
212 entity_id=transaction["entity_id"],
213 window_start=datetime.fromisoformat(transaction["window_start"]).replace(tzinfo=timezone.utc),
214 batch_started_at=_transaction_started_at(transaction),
215 )
216 if totals is None:
217 return 0.0
218 return _seed_base(totals=totals, batch_spend=transaction["spend"])
221def _seed_base(totals: WindowSeedTotals, batch_spend: float) -> float:
222 """What the window already held before the increments about to be applied.
224 Subtracting the batch's own spend from the full sum keeps every other
225 request in the seed, including the ones another pod persisted and has not
226 incremented yet, which a plain cutoff would drop for good if that pod died.
227 When this batch's own log rows have not landed yet the subtraction takes
228 spend that was never counted, so the sum before the batch is the floor.
229 """
230 return max(totals.total - batch_spend, totals.before_batch)
233def _transaction_started_at(transaction: WindowSpendTransaction) -> datetime | None:
234 started_at: Final = transaction.get("started_at")
235 if started_at is None:
236 return None
237 return datetime.fromisoformat(started_at).replace(tzinfo=timezone.utc)
240def _upsert_params(
241 transaction: WindowSpendTransaction,
242 seed_base: float,
243 now: datetime,
244) -> tuple[str, str, str, datetime, float, float, datetime]:
245 """$5 is what a brand new row starts at (pre-existing spend plus this
246 increment); $6 is the increment alone, which is all an already-current row
247 may add. They are equal for every row that already existed, so a row is
248 never seeded twice when two pods flush the same new window."""
249 increment: Final = float(transaction["spend"])
250 return (
251 transaction["entity_type"],
252 transaction["entity_id"],
253 transaction["window_duration"],
254 datetime.fromisoformat(transaction["window_start"]),
255 seed_base + increment,
256 increment,
257 now,
258 )
261async def commit_window_spend_updates(
262 prisma_client: "PrismaClient",
263 transactions: Sequence[WindowSpendTransaction],
264 spend_logs_aggregate: WindowSpendLogsAggregate = spend_logs_seed_totals,
265) -> None:
266 """Apply aggregated window increments to LiteLLM_BudgetWindowSpend.
268 An increment at or behind the row's window_start adds into the row (this is
269 how in-flight requests that raced a reset carry into the new window); an
270 increment ahead of it rolls the window and starts from that increment.
272 Statements are ordered by primary key so concurrent pods take row locks in
273 the same order, with window_start breaking ties so an older window is
274 applied before the roll that supersedes it.
275 """
276 if not transactions: 276 ↛ 279line 276 didn't jump to line 279 because the condition on line 276 was always true
277 return
279 ordered: Final = tuple(sorted(transactions, key=window_spend_group_key))
280 existing_primary_keys: Final = await _existing_primary_keys(
281 prisma_client=prisma_client,
282 transactions=ordered,
283 )
284 seed_bases: Final = tuple(
285 [
286 await _seed_base_for_missing_row(
287 prisma_client=prisma_client,
288 transaction=transaction,
289 existing_primary_keys=existing_primary_keys,
290 spend_logs_aggregate=spend_logs_aggregate,
291 )
292 for transaction in ordered
293 ]
294 )
296 now: Final = to_naive_utc(datetime.now(timezone.utc))
297 verbose_proxy_logger.debug(
298 "Spend tracking - committing %d budget window spend upserts over %d existing rows",
299 len(ordered),
300 len(existing_primary_keys),
301 )
302 async with (
303 prisma_client.db.tx(timeout=_UPSERT_TRANSACTION_TIMEOUT) as db_transaction,
304 db_transaction.batch_() as batcher,
305 ):
306 for transaction, seed_base in zip(ordered, seed_bases):
307 batcher.execute_raw(
308 _UPSERT_WINDOW_SPEND_SQL,
309 *_upsert_params(transaction=transaction, seed_base=seed_base, now=now),
310 )
313async def roll_window_spend_row(
314 prisma_client: "PrismaClient",
315 entity_type: str,
316 entity_id: str,
317 window_duration: str,
318 new_window_start: datetime,
319) -> None:
320 """Move a row onto the window that just started and zero its spend.
322 Conditional on the stored window_start still being behind the new one so a
323 pod that already rolled the row (or increments that arrived under the new
324 window) are not clobbered.
325 """
326 await prisma_client.db.execute_raw(
327 _ROLL_WINDOW_SPEND_SQL,
328 entity_type,
329 entity_id,
330 window_duration,
331 to_naive_utc(new_window_start),
332 to_naive_utc(datetime.now(timezone.utc)),
333 )