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
« 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.
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.
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"""
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
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)
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
45 from litellm.caching.dual_cache import DualCache
46 from litellm.proxy.utils import PrismaClient
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)
56END_USER_COUNTER_PREFIX: Final = "spend:end_user:"
58_WINDOW_SPEND_LOG_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
59 {
60 "Key": "api_key",
61 "Team": "team_id",
62 }
63)
66def _as_utc(value: datetime) -> datetime:
67 return value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc)
70class SpendCounterReseed:
71 """
72 Reseeds spend counters from the authoritative DB and warms the cache,
73 coalesced via per-counter singleflight locks.
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
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 """
92 _locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict()
93 _registry_lock: ClassVar[asyncio.Lock | None] = None
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
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 )
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.
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)
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
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)
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
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)
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.
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)
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
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.
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.
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
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
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)
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.
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 )
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
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 }
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
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)
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)
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