Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/db/spend_counter_reseed.py: 15%

241 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2Coalesced reseed of spend counters from the authoritative DB. 

3 

4When a Redis spend counter expires (or is missing on a fresh pod), enforcement 

5must read the current spend from somewhere. The in-process management cache 

6(`user_api_key_cache.team_membership.spend`, etc.) is per-pod and lags DB 

7writes from other pods, so trusting it allows budget bypass in multi-pod 

8deployments. This module reseeds from the authoritative DB instead. 

9 

10A per-counter singleflight lock collapses concurrent reseeds on the same pod 

11to one DB query per cold-cache window. The lock dict is bounded LRU to cap 

12memory in long-lived deployments. 

13""" 

14 

15import asyncio 

16from collections import OrderedDict 

17from collections.abc import Mapping 

18from datetime import datetime, timezone 

19from types import MappingProxyType 

20from typing import TYPE_CHECKING, ClassVar, Final, Optional 

21 

22from litellm._logging import verbose_proxy_logger 

23from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE 

24from litellm.litellm_core_utils.duration_parser import duration_in_seconds 

25from litellm.proxy._types import Litellm_EntityType 

26from litellm.proxy.db.db_lookup_gate import bounded_db_lookup, db_lookup_gate 

27from litellm.proxy.spend_tracking.spend_counter_batch import read_batched_spend_counter, record_spend_counter_value 

28from litellm.repositories.organization_repository import OrganizationRepository 

29from litellm.repositories.project_repository import ProjectRepository 

30from litellm.repositories.table_repositories import ( 

31 BudgetWindowSpendRepository, 

32 EndUserRepository, 

33 SpendLogsRepository, 

34 TeamMembershipRepository, 

35) 

36from litellm.repositories.team_repository import TeamRepository 

37from litellm.repositories.user_repository import UserRepository 

38from litellm.repositories.verification_token_repository import ( 

39 VerificationTokenRepository, 

40) 

41 

42if TYPE_CHECKING: 42 ↛ 43line 42 didn't jump to line 43 because the condition on line 42 was never true

43 from prisma.types import LiteLLM_EndUserTableWhereUniqueInput 

44 

45 from litellm.caching.dual_cache import DualCache 

46 from litellm.proxy.utils import PrismaClient 

47 

48 

49_WINDOW_SPEND_ENTITY_TYPES: Final[Mapping[str, str]] = MappingProxyType( 

50 { 

51 "Key": Litellm_EntityType.KEY.value, 

52 "Team": Litellm_EntityType.TEAM.value, 

53 } 

54) 

55 

56END_USER_COUNTER_PREFIX: Final = "spend:end_user:" 

57 

58_WINDOW_SPEND_LOG_FIELDS: Final[Mapping[str, str]] = MappingProxyType( 

59 { 

60 "Key": "api_key", 

61 "Team": "team_id", 

62 } 

63) 

64 

65 

66def _as_utc(value: datetime) -> datetime: 

67 return value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc) 

68 

69 

70class SpendCounterReseed: 

71 """ 

72 Reseeds spend counters from the authoritative DB and warms the cache, 

73 coalesced via per-counter singleflight locks. 

74 

75 Counter key prefixes map to DB tables: 

76 spend:key:{token} -> LiteLLM_VerificationToken.spend 

77 spend:team:{team_id} -> LiteLLM_TeamTable.spend 

78 spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend 

79 spend:user:{user_id} -> LiteLLM_UserTable.spend 

80 spend:org:{org_id} -> LiteLLM_OrganizationTable.spend 

81 spend:project:{project_id} -> LiteLLM_ProjectTable.spend 

82 

83 End-user and tag spend counters intentionally do not reseed here. Their 

84 auth paths already load the corresponding objects via get_end_user_object() 

85 and get_tag_objects_batch(); callers pass those values as fallback_spend. 

86 end_user_from_db is the one end-user read, used only as the budget floor when 

87 a counter sits below that cached spend: a worker that did not run the budget 

88 reset still caches the pre-reset end-user object, and LiteLLM_EndUserTable 

89 is the row the reset zeroed. 

90 """ 

91 

92 _locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict() 

93 _registry_lock: ClassVar[asyncio.Lock | None] = None 

94 

95 @staticmethod 

96 async def _get_lock(counter_key: str) -> asyncio.Lock: 

97 if SpendCounterReseed._registry_lock is None: 

98 SpendCounterReseed._registry_lock = asyncio.Lock() 

99 async with SpendCounterReseed._registry_lock: 

100 lock = SpendCounterReseed._locks.get(counter_key) 

101 if lock is not None: 

102 SpendCounterReseed._locks.move_to_end(counter_key) 

103 return lock 

104 lock = asyncio.Lock() 

105 SpendCounterReseed._locks[counter_key] = lock 

106 if len(SpendCounterReseed._locks) > SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: 

107 SpendCounterReseed._locks.popitem(last=False) 

108 return lock 

109 

110 @staticmethod 

111 async def increment_in_memory(spend_counter_cache: "DualCache", counter_key: str, increment: float) -> float | None: 

112 """Apply local deltas after an in-flight reseed establishes the spend balance.""" 

113 lock: Final = await SpendCounterReseed._get_lock(counter_key) 

114 async with lock: 

115 return await spend_counter_cache.async_increment_cache( 

116 key=counter_key, value=increment, local_only=True, refresh_ttl=True 

117 ) 

118 

119 @staticmethod 

120 async def from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> float | None: 

121 """ 

122 Read the authoritative spend for a counter from the DB. 

123 

124 Returns the spend value (including 0.0) when the DB is reachable 

125 and the row exists. Returns None when prisma is unavailable, the 

126 row is missing, the key format is unrecognized, or the query 

127 raises. Callers use None to fall back to a caller-supplied source. 

128 """ 

129 if prisma_client is None: 

130 return None 

131 # Per-window key/team counters share prefixes with primary counters 

132 # but don't correspond to a DB row. Do not reject arbitrary entity IDs 

133 # or tag names that merely contain ":window:". 

134 if SpendCounterReseed._is_key_or_team_window_counter(counter_key): 

135 return None 

136 try: 

137 row: Final = await bounded_db_lookup( 

138 SpendCounterReseed._counter_row(prisma_client, counter_key), name="spend_counter" 

139 ) 

140 except Exception: 

141 verbose_proxy_logger.exception("SpendCounterReseed.from_db: failed for %s", counter_key) 

142 return None 

143 if row is None: 

144 return None 

145 return float(getattr(row, "spend", 0.0) or 0.0) 

146 

147 @staticmethod 

148 async def _counter_row(prisma_client: "PrismaClient", counter_key: str) -> object | None: 

149 async with db_lookup_gate.current(): 

150 if counter_key.startswith("spend:key:"): 

151 token: Final = counter_key[len("spend:key:") :] 

152 return await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token}) 

153 if counter_key.startswith("spend:team_member:"): 

154 suffix: Final = counter_key[len("spend:team_member:") :] 

155 if ":" not in suffix: 

156 return None 

157 user_id, team_id = suffix.rsplit(":", 1) 

158 return await TeamMembershipRepository(prisma_client).table.find_unique( 

159 where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}} 

160 ) 

161 if counter_key.startswith("spend:team:"): 

162 return await TeamRepository(prisma_client).table.find_unique( 

163 where={"team_id": counter_key[len("spend:team:") :]} 

164 ) 

165 if counter_key.startswith("spend:user:"): 

166 return await UserRepository(prisma_client).table.find_unique( 

167 where={"user_id": counter_key[len("spend:user:") :]} 

168 ) 

169 if counter_key.startswith("spend:org:"): 

170 return await OrganizationRepository(prisma_client).table.find_unique( 

171 where={"organization_id": counter_key[len("spend:org:") :]} 

172 ) 

173 if counter_key.startswith("spend:project:"): 

174 return await ProjectRepository(prisma_client).table.find_unique( 

175 where={"project_id": counter_key[len("spend:project:") :]} 

176 ) 

177 return None 

178 

179 @staticmethod 

180 async def end_user_from_db(prisma_client: Optional["PrismaClient"], counter_key: str) -> float | None: 

181 if prisma_client is None or not counter_key.startswith(END_USER_COUNTER_PREFIX): 

182 return None 

183 where: Final[LiteLLM_EndUserTableWhereUniqueInput] = {"user_id": counter_key[len(END_USER_COUNTER_PREFIX) :]} 

184 try: 

185 row: Final = await bounded_db_lookup( 

186 EndUserRepository(prisma_client).table.find_unique(where=where), name="end_user_spend" 

187 ) 

188 except Exception: # noqa: BLE001 # a failed floor read falls back to the cached spend, like from_db 

189 verbose_proxy_logger.exception("SpendCounterReseed.end_user_from_db: failed for %s", counter_key) 

190 return None 

191 if row is None: 

192 return None 

193 return float(row.spend or 0.0) 

194 

195 @staticmethod 

196 def _is_key_or_team_window_counter(counter_key: str) -> bool: 

197 for prefix in ("spend:key:", "spend:team:"): 

198 if not counter_key.startswith(prefix): 

199 continue 

200 _, separator, duration = counter_key.rpartition(":window:") 

201 if not separator or not duration: 

202 return False 

203 try: 

204 duration_in_seconds(duration) 

205 except Exception: 

206 return False 

207 return True 

208 return False 

209 

210 @staticmethod 

211 async def _read_active_batch(counter_key: str) -> tuple[float | None, bool] | None: 

212 """The request's MGET answers for this counter; a Redis miss there is authoritative.""" 

213 return await read_batched_spend_counter(counter_key) 

214 

215 @staticmethod 

216 async def coalesced( 

217 prisma_client: Optional["PrismaClient"], 

218 spend_counter_cache: "DualCache", 

219 counter_key: str, 

220 require_cache_warm: bool = False, 

221 ) -> float | None: 

222 """ 

223 Reseed a cold spend counter from the DB and warm the cache, 

224 coalesced via a per-counter lock so concurrent callers (read path 

225 + write path) collapse to one DB query per cold-cache window. 

226 

227 Returns the spend value (including 0.0 from a fresh budget reset) 

228 when the DB read succeeds, or None when the DB is unavailable. 

229 """ 

230 lock: Final = await SpendCounterReseed._get_lock(counter_key) 

231 async with lock: 

232 batched: Final = await SpendCounterReseed._read_active_batch(counter_key) 

233 if batched is not None and batched[0] is not None: 

234 return batched[0] 

235 # Re-check after acquiring the lock. Skip in-memory on a clean 

236 # Redis miss - in-memory is per-pod-stale. 

237 redis_clean_miss = batched is not None 

238 if spend_counter_cache.redis_cache is not None and not redis_clean_miss: 

239 try: 

240 val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key) 

241 if val is not None: 

242 return float(val) 

243 redis_clean_miss = True 

244 except Exception: 

245 pass 

246 if not redis_clean_miss: 

247 val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) 

248 if val is not None: 

249 return float(val) 

250 

251 db_spend: Final = await SpendCounterReseed.from_db(prisma_client, counter_key) 

252 if db_spend is None: 

253 return None 

254 # Warm even when 0 so subsequent reads hit cache, not DB. 

255 # 

256 # Seed via SET NX (cross-pod safe): only one pod initializes the 

257 # Redis key with db_spend; concurrent seeders read the winner's 

258 # value. INCRBYFLOAT-of-db_spend from N pods would multiply the 

259 # counter (N x db_spend) and trigger spurious budget alerts. 

260 current_value: float = float(db_spend) 

261 try: 

262 if spend_counter_cache.redis_cache is not None: 

263 seeded: Final = await spend_counter_cache.redis_cache.async_set_cache( 

264 key=counter_key, 

265 value=db_spend, 

266 nx=True, 

267 ) 

268 if seeded: 

269 current_value = float(db_spend) 

270 else: 

271 cached: Final = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key) 

272 current_value = float(cached) if cached is not None else float(db_spend) 

273 spend_counter_cache.in_memory_cache.set_cache( 

274 key=counter_key, 

275 value=current_value, 

276 ) 

277 record_spend_counter_value(counter_key, current_value) 

278 else: 

279 cached_spend: Final = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) 

280 seeded_spend: Final = max(db_spend, float(cached_spend)) if cached_spend is not None else db_spend 

281 spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=seeded_spend) 

282 return seeded_spend 

283 except Exception: 

284 verbose_proxy_logger.exception( 

285 "SpendCounterReseed.coalesced: failed to warm counter %s", 

286 counter_key, 

287 ) 

288 if require_cache_warm: 

289 raise 

290 return current_value 

291 

292 @staticmethod 

293 async def window_from_table( 

294 prisma_client: Optional["PrismaClient"], 

295 entity_type: str, 

296 entity_id: str, 

297 window_duration: str, 

298 expected_window_start: datetime, 

299 ) -> float | None: 

300 """ 

301 Read the maintained per-window spend row by primary key. 

302 

303 Returns the row's spend only when the row belongs to the window the 

304 caller is enforcing, i.e. ``row.window_start >= expected_window_start``. 

305 A row at or past the expected start was rolled by a pod whose reset_at 

306 was at least as fresh as this caller's, so it is trusted; an older row 

307 means the window boundary was crossed and nothing has rolled the row 

308 yet, so its spend belongs to a previous window. 

309 

310 Returns None for a missing, stale or unreadable row so the caller falls 

311 back to the spend-logs aggregate. ``entity_type`` is the counter-facing 

312 label ("Key"/"Team"); anything else has no row and returns None. 

313 """ 

314 if prisma_client is None: 

315 return None 

316 row_entity_type: Final = _WINDOW_SPEND_ENTITY_TYPES.get(entity_type) 

317 if row_entity_type is None: 

318 return None 

319 

320 try: 

321 row: Final = await BudgetWindowSpendRepository(prisma_client).table.find_unique( 

322 where={ 

323 "entity_type_entity_id_window_duration": { 

324 "entity_type": row_entity_type, 

325 "entity_id": entity_id, 

326 "window_duration": window_duration, 

327 } 

328 } 

329 ) 

330 except Exception: # noqa: BLE001 # any read failure (DB, stale prisma client) must degrade to the aggregate path 

331 verbose_proxy_logger.exception( 

332 "SpendCounterReseed.window_from_table: failed for %s=%s window=%s", 

333 entity_type, 

334 entity_id, 

335 window_duration, 

336 ) 

337 return None 

338 

339 if row is None: 

340 return None 

341 if _as_utc(row.window_start) < _as_utc(expected_window_start): 

342 return None 

343 return float(row.spend or 0.0) 

344 

345 @staticmethod 

346 async def window_from_db( 

347 prisma_client: Optional["PrismaClient"], 

348 entity_type: str, 

349 entity_id: str, 

350 window_duration: str | None, 

351 window_start: datetime, 

352 ) -> float | None: 

353 """ 

354 Authoritative window spend: the maintained row first, falling back to 

355 the spend-logs aggregate only when no current row exists. 

356 

357 The aggregate range-scans an unindexed table, so it must stay a 

358 transitional path (window configured before the row existed) rather 

359 than a steady-state read. 

360 """ 

361 if window_duration is not None: 

362 from_table: Final = await SpendCounterReseed.window_from_table( 

363 prisma_client=prisma_client, 

364 entity_type=entity_type, 

365 entity_id=entity_id, 

366 window_duration=window_duration, 

367 expected_window_start=window_start, 

368 ) 

369 if from_table is not None: 

370 return from_table 

371 return await SpendCounterReseed.window_from_spend_logs( 

372 prisma_client=prisma_client, 

373 entity_type=entity_type, 

374 entity_id=entity_id, 

375 window_start=window_start, 

376 ) 

377 

378 @staticmethod 

379 async def window_from_spend_logs( 

380 prisma_client: Optional["PrismaClient"], 

381 entity_type: str, 

382 entity_id: str, 

383 window_start: datetime, 

384 ) -> float | None: 

385 if prisma_client is None: 

386 return None 

387 

388 group_field: Final = _WINDOW_SPEND_LOG_FIELDS.get(entity_type) 

389 if group_field is None: 

390 return None 

391 where: Final = { 

392 group_field: entity_id, 

393 "startTime": {"gte": window_start}, 

394 } 

395 

396 try: 

397 response: Final = await SpendLogsRepository(prisma_client).table.group_by( 

398 by=[group_field], 

399 where=where, 

400 sum={"spend": True}, 

401 ) 

402 except Exception: 

403 verbose_proxy_logger.exception( 

404 "SpendCounterReseed.window_from_spend_logs: failed for %s=%s", 

405 entity_type, 

406 entity_id, 

407 ) 

408 return None 

409 

410 if not response: 

411 return 0.0 

412 first_row: Final = response[0] 

413 sum_row: Final = first_row.get("_sum") if isinstance(first_row, dict) else getattr(first_row, "_sum", None) 

414 spend: Final = sum_row.get("spend") if isinstance(sum_row, dict) else getattr(sum_row, "spend", None) 

415 return float(spend or 0.0) 

416 

417 @staticmethod 

418 async def coalesced_window( 

419 prisma_client: Optional["PrismaClient"], 

420 spend_counter_cache: "DualCache", 

421 counter_key: str, 

422 entity_type: str, 

423 entity_id: str, 

424 window_duration: str | None, 

425 window_start: datetime, 

426 ) -> float | None: 

427 lock: Final = await SpendCounterReseed._get_lock(counter_key) 

428 async with lock: 

429 batched: Final = await SpendCounterReseed._read_active_batch(counter_key) 

430 if batched is not None and batched[0] is not None: 

431 return batched[0] 

432 redis_clean_miss = batched is not None 

433 if spend_counter_cache.redis_cache is not None and not redis_clean_miss: 

434 try: 

435 val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key) 

436 if val is not None: 

437 return float(val) 

438 redis_clean_miss = True 

439 except Exception: 

440 pass 

441 if not redis_clean_miss: 

442 val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) 

443 if val is not None: 

444 return float(val) 

445 

446 window_spend: Final = await SpendCounterReseed.window_from_db( 

447 prisma_client=prisma_client, 

448 entity_type=entity_type, 

449 entity_id=entity_id, 

450 window_duration=window_duration, 

451 window_start=window_start, 

452 ) 

453 if window_spend is None: 

454 return None 

455 try: 

456 if spend_counter_cache.redis_cache is not None: 

457 seeded: Final = await spend_counter_cache.redis_cache.async_set_cache( 

458 key=counter_key, 

459 value=window_spend, 

460 nx=True, 

461 ) 

462 if seeded: 

463 current_value = window_spend 

464 else: 

465 current_cached_value = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key) 

466 if current_cached_value is None: 

467 current_value = await spend_counter_cache.redis_cache.async_increment( 

468 key=counter_key, 

469 value=window_spend, 

470 ) 

471 else: 

472 current_value = float(current_cached_value) 

473 spend_counter_cache.in_memory_cache.set_cache( 

474 key=counter_key, 

475 value=current_value, 

476 ) 

477 record_spend_counter_value(counter_key, float(current_value)) 

478 else: 

479 cached_spend: Final = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) 

480 seeded_spend: Final = ( 

481 max(window_spend, float(cached_spend)) if cached_spend is not None else window_spend 

482 ) 

483 spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=seeded_spend) 

484 return seeded_spend 

485 except Exception: 

486 verbose_proxy_logger.exception( 

487 "SpendCounterReseed.coalesced_window: failed to warm counter %s", 

488 counter_key, 

489 ) 

490 raise 

491 return current_value