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

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 

8 

9from pydantic import BaseModel, TypeAdapter 

10from typing_extensions import ReadOnly, TypedDict 

11 

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 

24 

25_T = TypeVar("_T") 

26 

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

32 

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

40 

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

62 

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) 

65 

66_HASHED_JWT_PREFIX: Final = "hashed-jwt-" 

67_CLI_SESSION_KEY_PREFIX: Final = f"{CLI_SESSION_KEY_PREFIX}-" 

68 

69 

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] 

76 

77 

78class _TokenDigestRow(BaseModel): 

79 digest: str 

80 key_alias: str | None = None 

81 team_id: str | None = None 

82 user_id: str | None = None 

83 

84 

85def _unanimous(first: str | None, last: str | None) -> str | None: 

86 return first if first == last else None 

87 

88 

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 

97 

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 ) 

104 

105 

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

115 

116 

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 

123 

124 try: 

125 return await load() 

126 except PrismaError as e: 

127 verbose_proxy_logger.warning(warning, count, e) 

128 return None 

129 

130 

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 ) 

152 

153 

154@dataclass(frozen=True, slots=True) 

155class _UserDetails: 

156 email: str | None 

157 only_team: str | None 

158 

159 

160_EMPTY_USER_DETAILS: Final[Mapping[str, _UserDetails]] = MappingProxyType({}) 

161 

162 

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 ) 

188 

189 

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 

195 

196 

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) 

199 

200 

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 

216 

217 

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 ) 

236 

237 

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 ) 

255 

256 

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 

264 

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

281 

282 

283def _is_spend_log_digest(key: str) -> bool: 

284 return is_valid_sha256_hash(key.removeprefix(_HASHED_JWT_PREFIX)) 

285 

286 

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

290 

291 

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 ) 

305 

306 

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) 

316 

317 

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 ) 

338 

339 

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) 

353 

354 

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

372 

373 

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

392 

393 

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 ) 

415 

416 

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. 

429 

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) 

441 

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) 

454 

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 )