Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/spend_tracking/key_metadata_recovery.py: 48%
174 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
1import asyncio
2from collections.abc import Awaitable, Callable, Mapping, Sequence
3from collections.abc import Set as AbstractSet
4from dataclasses import dataclass
5from datetime import datetime, timedelta
6from types import MappingProxyType
7from typing import Final, TypeVar
9from pydantic import BaseModel, TypeAdapter
10from typing_extensions import ReadOnly, TypedDict
12from litellm._logging import verbose_proxy_logger
13from litellm.caching.in_memory_cache import InMemoryCache
14from litellm.constants import (
15 CLI_SESSION_KEY_PREFIX,
16 SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
17 SPEND_LOG_KEY_METADATA_CACHE_TTL,
18 SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL,
19 SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS,
20)
21from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash
22from litellm.proxy.utils import PrismaClient
23from litellm.repositories.user_repository import UserRepository
25_T = TypeVar("_T")
27_ACTIVE_TOKEN_DIGEST_SQL: Final = """
28SELECT encode(sha256(convert_to(token, 'UTF8')), 'hex') AS digest, key_alias, team_id, user_id
29FROM "LiteLLM_VerificationToken"
30WHERE encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY($1::text[])
31"""
33_DELETED_TOKEN_DIGEST_SQL: Final = """
34SELECT DISTINCT ON (token)
35 encode(sha256(convert_to(token, 'UTF8')), 'hex') AS digest, key_alias, team_id, user_id
36FROM "LiteLLM_DeletedVerificationToken"
37WHERE encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY($1::text[])
38ORDER BY token, deleted_at DESC
39"""
41_SPEND_LOG_ALIAS_SQL: Final = """
42SELECT api_key AS digest,
43 MIN(key_alias) AS first_alias,
44 MAX(key_alias) AS last_alias,
45 MIN(team_id) AS first_team,
46 MAX(team_id) AS last_team,
47 MIN(user_id) AS first_owner,
48 MAX(user_id) AS last_owner
49FROM (
50 SELECT api_key,
51 NULLIF(metadata->>'user_api_key_alias', '') AS key_alias,
52 COALESCE(NULLIF(team_id, ''), NULLIF(metadata->>'user_api_key_team_id', '')) AS team_id,
53 COALESCE(NULLIF("user", ''), NULLIF(metadata->>'user_api_key_user_id', '')) AS user_id
54 FROM "LiteLLM_SpendLogs"
55 WHERE api_key = ANY($1::text[])
56 AND "startTime" >= $2::timestamp
57 AND "startTime" < $3::timestamp
58) named
59WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL
60GROUP BY api_key
61"""
63_SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}"
64_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS)
66_HASHED_JWT_PREFIX: Final = "hashed-jwt-"
67_CLI_SESSION_KEY_PREFIX: Final = f"{CLI_SESSION_KEY_PREFIX}-"
70class KeyMetadataDict(TypedDict, total=False):
71 key_alias: ReadOnly[str | None]
72 team_id: ReadOnly[str | None]
73 user_id: ReadOnly[str | None]
74 user_email: ReadOnly[str | None]
75 key_exists: ReadOnly[bool]
78class _TokenDigestRow(BaseModel):
79 digest: str
80 key_alias: str | None = None
81 team_id: str | None = None
82 user_id: str | None = None
85def _unanimous(first: str | None, last: str | None) -> str | None:
86 return first if first == last else None
89class _SpendLogDigestRow(BaseModel):
90 digest: str
91 first_alias: str | None = None
92 last_alias: str | None = None
93 first_team: str | None = None
94 last_team: str | None = None
95 first_owner: str | None = None
96 last_owner: str | None = None
98 def metadata(self) -> KeyMetadataDict:
99 return KeyMetadataDict(
100 key_alias=_unanimous(self.first_alias, self.last_alias),
101 team_id=_unanimous(self.first_team, self.last_team),
102 user_id=_unanimous(self.first_owner, self.last_owner),
103 )
106_TOKEN_DIGEST_ROWS: Final = TypeAdapter(tuple[_TokenDigestRow, ...])
107_SPEND_LOG_DIGEST_ROWS: Final = TypeAdapter(tuple[_SpendLogDigestRow, ...])
108_CACHED_KEY_METADATA: Final = TypeAdapter(KeyMetadataDict)
109_SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
110 max_size_in_memory=SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
111 default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL,
112)
113_SPEND_LOG_QUERY_LOCK: Final = asyncio.Lock()
114_EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({})
117async def _db_or_empty(
118 load: Callable[[], Awaitable[_T]],
119 warning: str,
120 count: int,
121) -> _T | None:
122 from prisma.errors import PrismaError
124 try:
125 return await load()
126 except PrismaError as e:
127 verbose_proxy_logger.warning(warning, count, e)
128 return None
131async def _reverse_hash_key_metadata(
132 prisma_client: PrismaClient,
133 sql: str,
134 wanted: AbstractSet[str],
135 *,
136 warning: str,
137) -> Mapping[str, KeyMetadataDict]:
138 rows: Final = await _db_or_empty(
139 lambda: prisma_client.db.query_raw(sql, sorted(wanted)),
140 warning,
141 len(wanted),
142 )
143 if rows is None:
144 return _EMPTY_KEY_METADATA
145 return MappingProxyType(
146 {
147 row.digest: KeyMetadataDict(key_alias=row.key_alias, team_id=row.team_id, user_id=row.user_id)
148 for row in _TOKEN_DIGEST_ROWS.validate_python(rows)
149 if row.digest in wanted
150 }
151 )
154@dataclass(frozen=True, slots=True)
155class _UserDetails:
156 email: str | None
157 only_team: str | None
160_EMPTY_USER_DETAILS: Final[Mapping[str, _UserDetails]] = MappingProxyType({})
163async def _details_for_user_ids(
164 prisma_client: PrismaClient,
165 user_ids: AbstractSet[str],
166) -> Mapping[str, _UserDetails]:
167 if not user_ids: 167 ↛ 169line 167 didn't jump to line 169 because the condition on line 167 was always true
168 return _EMPTY_USER_DETAILS
169 users: Final = await _db_or_empty(
170 lambda: UserRepository(prisma_client).table.find_many(
171 where={"user_id": {"in": list(user_ids)}}, # mutable-ok: Prisma find_many where= is a dict
172 ),
173 "Failed user detail recovery for %d user ids: %s",
174 len(user_ids),
175 )
176 if users is None:
177 return _EMPTY_USER_DETAILS
178 return MappingProxyType(
179 {
180 user.user_id: _UserDetails(
181 email=getattr(user, "user_email", None) or None,
182 only_team=_only_team(getattr(user, "teams", None)),
183 )
184 for user in users
185 if getattr(user, "user_id", None)
186 }
187 )
190def _only_team(teams: object) -> str | None:
191 if not isinstance(teams, list) or len(teams) != 1:
192 return None
193 team: Final = teams[0]
194 return team if isinstance(team, str) and team else None
197def _is_cli_session_key(api_key: str) -> bool:
198 return api_key.startswith(_CLI_SESSION_KEY_PREFIX) and len(api_key) > len(_CLI_SESSION_KEY_PREFIX)
201def _meta_with_user_details(
202 api_key: str, meta: KeyMetadataDict, details: Mapping[str, _UserDetails]
203) -> KeyMetadataDict:
204 user_id: Final = meta.get("user_id")
205 if not isinstance(user_id, str) or user_id not in details:
206 return meta
207 user: Final = details[user_id]
208 email: Final = meta.get("user_email") or user.email
209 team_id: Final = meta.get("team_id") or (user.only_team if _is_cli_session_key(api_key) else None)
210 updated: Final[KeyMetadataDict] = {
211 **meta,
212 **({"user_email": email} if email else {}),
213 **({"team_id": team_id} if team_id else {}),
214 }
215 return updated
218async def attach_user_details(
219 prisma_client: PrismaClient,
220 recovered: Mapping[str, KeyMetadataDict],
221) -> Mapping[str, KeyMetadataDict]:
222 needing_details: Final = frozenset(
223 user_id
224 for api_key, meta in recovered.items()
225 for user_id in (meta.get("user_id"),)
226 if isinstance(user_id, str)
227 and user_id
228 and (not meta.get("user_email") or (_is_cli_session_key(api_key) and not meta.get("team_id")))
229 )
230 details: Final = await _details_for_user_ids(prisma_client, needing_details)
231 if not details: 231 ↛ 233line 231 didn't jump to line 233 because the condition on line 231 was always true
232 return recovered
233 return MappingProxyType(
234 {api_key: _meta_with_user_details(api_key, meta, details) for api_key, meta in recovered.items()}
235 )
238async def recover_cli_session_key_metadata(
239 prisma_client: PrismaClient,
240 missing_keys: AbstractSet[str],
241) -> Mapping[str, KeyMetadataDict]:
242 candidates: Final = MappingProxyType(
243 {key: key.removeprefix(_CLI_SESSION_KEY_PREFIX) for key in missing_keys if _is_cli_session_key(key)}
244 )
245 if not candidates: 245 ↛ 247line 245 didn't jump to line 247 because the condition on line 245 was always true
246 return _EMPTY_KEY_METADATA
247 known_users: Final = await _details_for_user_ids(prisma_client, frozenset(candidates.values()))
248 return MappingProxyType(
249 {
250 key: KeyMetadataDict(key_alias=key, user_id=user_id)
251 for key, user_id in candidates.items()
252 if user_id in known_users
253 }
254 )
257async def recover_double_hashed_key_metadata(
258 prisma_client: PrismaClient,
259 missing_keys: AbstractSet[str],
260) -> Mapping[str, KeyMetadataDict]:
261 sha_missing: Final = frozenset(key for key in missing_keys if is_valid_sha256_hash(key))
262 if not sha_missing: 262 ↛ 265line 262 didn't jump to line 265 because the condition on line 262 was always true
263 return _EMPTY_KEY_METADATA
265 from_active: Final = await _reverse_hash_key_metadata(
266 prisma_client,
267 _ACTIVE_TOKEN_DIGEST_SQL,
268 sha_missing,
269 warning="Failed reverse-hash recovery against active keys for %d missing keys: %s",
270 )
271 still_missing: Final = sha_missing - frozenset(from_active)
272 if not still_missing:
273 return from_active
274 from_deleted: Final = await _reverse_hash_key_metadata(
275 prisma_client,
276 _DELETED_TOKEN_DIGEST_SQL,
277 still_missing,
278 warning="Failed reverse-hash recovery against deleted keys for %d missing keys: %s",
279 )
280 return MappingProxyType({**from_active, **from_deleted})
283def _is_spend_log_digest(key: str) -> bool:
284 return is_valid_sha256_hash(key.removeprefix(_HASHED_JWT_PREFIX))
287def _spend_log_cache_key(digest: str, window: tuple[datetime, datetime]) -> str:
288 start, end = window
289 return f"spend_log_key_metadata:{digest}:{start.isoformat()}:{end.isoformat()}"
292def _cached_spend_log_metadata(
293 cache: InMemoryCache,
294 digests: AbstractSet[str],
295 window: tuple[datetime, datetime],
296) -> Mapping[str, KeyMetadataDict]:
297 return MappingProxyType(
298 {
299 digest: _CACHED_KEY_METADATA.validate_python(cached)
300 for digest in digests
301 for cached in (cache.get_cache(_spend_log_cache_key(digest, window)),)
302 if cached is not None
303 }
304 )
307async def _spend_log_rows_within_the_statement_timeout(
308 prisma_client: PrismaClient,
309 digests: AbstractSet[str],
310 window: tuple[datetime, datetime],
311) -> Sequence[Mapping[str, object]]:
312 start, end = window
313 async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction:
314 await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL)
315 return await transaction.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end)
318async def _query_spend_log_metadata(
319 prisma_client: PrismaClient,
320 digests: AbstractSet[str],
321 window: tuple[datetime, datetime],
322) -> Mapping[str, KeyMetadataDict] | None:
323 rows: Final = await _db_or_empty(
324 lambda: _spend_log_rows_within_the_statement_timeout(prisma_client, digests, window),
325 "Failed spend-log alias recovery for %d missing keys: %s",
326 len(digests),
327 )
328 if rows is None:
329 return None
330 return MappingProxyType(
331 {
332 row.digest: meta
333 for row in _SPEND_LOG_DIGEST_ROWS.validate_python(rows)
334 for meta in (row.metadata(),)
335 if row.digest in digests and any(meta.values())
336 }
337 )
340def _remember_spend_log_metadata(
341 cache: InMemoryCache, digest: str, window: tuple[datetime, datetime], meta: KeyMetadataDict | None
342) -> None:
343 key: Final = _spend_log_cache_key(digest, window)
344 if meta is not None:
345 cache.set_cache(key, meta)
346 return
347 missed_before: Final = f"{key}:missed-before"
348 if cache.get_cache(missed_before) is not None:
349 cache.set_cache(key, KeyMetadataDict())
350 return
351 cache.set_cache(key, KeyMetadataDict(), ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL)
352 cache.set_cache(missed_before, True)
355async def _spend_log_metadata_one_query_at_a_time(
356 prisma_client: PrismaClient,
357 cache: InMemoryCache,
358 lock: asyncio.Lock,
359 digests: AbstractSet[str],
360 window: tuple[datetime, datetime],
361) -> Mapping[str, KeyMetadataDict]:
362 async with lock:
363 settled: Final = _cached_spend_log_metadata(cache, digests, window)
364 pending: Final = digests - frozenset(settled)
365 fresh: Final = (
366 await _query_spend_log_metadata(prisma_client, pending, window) if pending else _EMPTY_KEY_METADATA
367 )
368 found: Final = fresh if fresh is not None else _EMPTY_KEY_METADATA
369 for digest in pending:
370 _remember_spend_log_metadata(cache, digest, window, found.get(digest))
371 return MappingProxyType({**settled, **found})
374async def recover_key_metadata_from_spend_logs(
375 prisma_client: PrismaClient,
376 missing_keys: AbstractSet[str],
377 window: tuple[datetime, datetime],
378 cache: InMemoryCache = _SPEND_LOG_METADATA_CACHE,
379 lock: asyncio.Lock = _SPEND_LOG_QUERY_LOCK,
380) -> Mapping[str, KeyMetadataDict]:
381 digests: Final = frozenset(key for key in missing_keys if _is_spend_log_digest(key))
382 if not digests: 382 ↛ 384line 382 didn't jump to line 384 because the condition on line 382 was always true
383 return _EMPTY_KEY_METADATA
384 cached: Final = _cached_spend_log_metadata(cache, digests, window)
385 uncached: Final = digests - frozenset(cached)
386 settled: Final = (
387 await _spend_log_metadata_one_query_at_a_time(prisma_client, cache, lock, uncached, window)
388 if uncached
389 else _EMPTY_KEY_METADATA
390 )
391 return MappingProxyType({digest: meta for digest, meta in (*cached.items(), *settled.items()) if meta})
394def _row_with_recovered_fields(
395 row: Mapping[str, object],
396 recovered: Mapping[str, KeyMetadataDict],
397 *,
398 api_key_field: str,
399 alias_field: str,
400 team_id_field: str,
401 user_email_field: str,
402) -> Mapping[str, object]:
403 api_key: Final = row.get(api_key_field)
404 if not isinstance(api_key, str) or api_key not in recovered:
405 return row
406 meta: Final = recovered[api_key]
407 return MappingProxyType(
408 {
409 **row,
410 alias_field: meta.get("key_alias") or row.get(alias_field),
411 team_id_field: meta.get("team_id") or row.get(team_id_field),
412 user_email_field: row.get(user_email_field) or meta.get("user_email"),
413 }
414 )
417async def fill_missing_api_key_aliases(
418 prisma_client: PrismaClient,
419 rows: Sequence[Mapping[str, object]],
420 *,
421 api_key_field: str = "api_key",
422 alias_field: str = "api_key_alias",
423 team_id_field: str = "team_id",
424 user_email_field: str = "user_email",
425) -> tuple[Mapping[str, object], ...]:
426 """
427 Fill null api_key_alias / team_id / user_email on export rows whose api_key
428 was double-hashed.
430 Used by CloudZero and Focus, which join DailyUserSpend.api_key to
431 VerificationToken.token and otherwise export null aliases for those rows.
432 """
433 missing_keys: Final = frozenset(
434 key
435 for row in rows
436 for key in (row.get(api_key_field),)
437 if isinstance(key, str) and key and row.get(alias_field) in (None, "")
438 )
439 if not missing_keys: 439 ↛ 442line 439 didn't jump to line 442 because the condition on line 439 was always true
440 return tuple(rows)
442 from_session_keys: Final = await recover_cli_session_key_metadata(prisma_client, missing_keys)
443 recovered: Final = await attach_user_details(
444 prisma_client,
445 MappingProxyType(
446 {
447 **from_session_keys,
448 **await recover_double_hashed_key_metadata(prisma_client, missing_keys - frozenset(from_session_keys)),
449 }
450 ),
451 )
452 if not recovered:
453 return tuple(rows)
455 return tuple(
456 _row_with_recovered_fields(
457 row,
458 recovered,
459 api_key_field=api_key_field,
460 alias_field=alias_field,
461 team_id_field=team_id_field,
462 user_email_field=user_email_field,
463 )
464 for row in rows
465 )