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

1""" 

2Writer for LiteLLM_BudgetWindowSpend. 

3 

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. 

9 

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""" 

16 

17from collections.abc import Sequence 

18from dataclasses import dataclass 

19from datetime import datetime, timedelta, timezone 

20from typing import TYPE_CHECKING, Final, Protocol 

21 

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) 

29 

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 

32 

33 

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) 

39 

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) 

54 

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) 

63 

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) 

70 

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) 

77 

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) 

83 

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) 

89 

90_UPSERT_TRANSACTION_TIMEOUT: Final = timedelta(seconds=60) 

91 

92 

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.""" 

97 

98 total: float 

99 before_batch: float 

100 

101 

102class WindowSpendLogsAggregate(Protocol): 

103 """Sums LiteLLM_SpendLogs for one entity since window_start, split at the 

104 batch's earliest request. 

105 

106 Injected so the flush can be exercised without a database and so the 

107 expensive aggregate stays swappable. 

108 """ 

109 

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: ... 

118 

119 

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. 

129 

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. 

135 

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 ) 

165 

166 

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) 

171 

172 

173def _primary_key(transaction: WindowSpendTransaction) -> tuple[str, str, str]: 

174 return ( 

175 transaction["entity_type"], 

176 transaction["entity_id"], 

177 transaction["window_duration"], 

178 ) 

179 

180 

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 ()) 

192 

193 

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. 

201 

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"]) 

219 

220 

221def _seed_base(totals: WindowSeedTotals, batch_spend: float) -> float: 

222 """What the window already held before the increments about to be applied. 

223 

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) 

231 

232 

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) 

238 

239 

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 ) 

259 

260 

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. 

267 

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. 

271 

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 

278 

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 ) 

295 

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 ) 

311 

312 

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. 

321 

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 )