Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/batch_enqueued_tokens.py: 33%

216 statements  

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

1""" 

2Enqueued-token accounting for batch submissions. 

3 

4Opt-in via admin-set ``batch_enqueued_token_limit`` in key or team metadata: batch 

5submissions reserve their estimated token count against a long-lived 

6enqueued-token allowance instead of the per-minute rate-limit windows, and 

7the reservation is refunded when the batch reaches a terminal state 

8(completed, failed, expired, or cancelled). 

9""" 

10 

11import asyncio 

12import logging 

13import math 

14import time 

15import uuid 

16from collections.abc import Awaitable, Callable, Mapping, Sequence 

17from dataclasses import dataclass, field 

18from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias 

19 

20from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError 

21 

22from litellm._logging import verbose_proxy_logger 

23from litellm.caching.redis_cache import log_redis_failure 

24from litellm.constants import BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, BATCH_ENQUEUED_TOKEN_TTL_SECONDS 

25from litellm.proxy._types import UserAPIKeyAuth 

26 

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

28 from opentelemetry.trace import Span as _Span 

29 

30 from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache 

31 

32 Span = _Span 

33 InternalUsageCache = _InternalUsageCache 

34 

35BATCH_ENQUEUED_REFUND_STATUSES: Final[frozenset[str]] = frozenset( 

36 {"completed", "complete", "failed", "expired", "cancelled", "cancelling"} 

37) 

38 

39ScopeKey: TypeAlias = Literal["api_key", "team"] 

40 

41RESERVE_ENQUEUED_TOKENS_SCRIPT: Final = """ 

42local amount = tonumber(ARGV[1]) 

43local ttl = tonumber(ARGV[2]) 

44local limit = tonumber(ARGV[3]) 

45local current = tonumber(redis.call('GET', KEYS[1]) or '0') 

46if current + amount > limit then 

47 return {0, current} 

48end 

49local updated = redis.call('INCRBY', KEYS[1], amount) 

50redis.call('EXPIRE', KEYS[1], ttl) 

51return {1, updated} 

52""" 

53 

54REFUND_ENQUEUED_TOKENS_SCRIPT: Final = """ 

55local updated = redis.call('DECRBY', KEYS[1], tonumber(ARGV[1])) 

56if updated <= 0 then 

57 redis.call('DEL', KEYS[1]) 

58end 

59return 1 

60""" 

61 

62SAVE_RESERVATION_SCRIPT: Final = """ 

63redis.call('SET', KEYS[1], ARGV[1], 'EX', tonumber(ARGV[2])) 

64return 1 

65""" 

66 

67POP_RESERVATION_SCRIPT: Final = """ 

68local value = redis.call('GET', KEYS[1]) 

69if value and value ~= '' then 

70 redis.call('SET', KEYS[1], '', 'EX', tonumber(ARGV[1])) 

71end 

72return value 

73""" 

74 

75 

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

77class BatchEnqueuedTokenScope: 

78 key: ScopeKey 

79 value: str 

80 limit: int 

81 

82 

83ReservationBackend: TypeAlias = Literal["redis", "memory"] 

84 

85 

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

87class BatchEnqueuedTokenReservation: 

88 tokens: int 

89 scopes: tuple[BatchEnqueuedTokenScope, ...] 

90 backend: ReservationBackend = "redis" 

91 owner: str = "" 

92 reserved_at_monotonic: float = field(default_factory=time.monotonic, compare=False) 

93 

94 

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

96class BatchEnqueuedTokenOverLimit: 

97 scope: BatchEnqueuedTokenScope 

98 enqueued: int 

99 

100 

101BatchEnqueuedTokenOutcome: TypeAlias = BatchEnqueuedTokenReservation | BatchEnqueuedTokenOverLimit 

102 

103_LIMIT_ADAPTER: Final = TypeAdapter(Annotated[int, Field(gt=0)]) 

104_RESERVE_RESULT_ADAPTER: Final = TypeAdapter(tuple[int, int]) 

105_POPPED_VALUE_ADAPTER: Final = TypeAdapter(str | bytes | None) 

106_STORED_COUNTER_ADAPTER: Final = TypeAdapter(int | None) 

107_RESERVATION_ADAPTER: Final = TypeAdapter(BatchEnqueuedTokenReservation) 

108 

109 

110class _ScriptRunner(Protocol): 

111 def __call__(self, keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> Awaitable[object]: ... 111 ↛ exitline 111 didn't return from function '__call__' because

112 

113 

114def _read_metadata_limit(metadata: Mapping[str, object] | None) -> int | None: 

115 if not metadata: 

116 return None 

117 raw: Final = metadata.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY) 

118 if raw is None: 

119 return None 

120 try: 

121 return _LIMIT_ADAPTER.validate_python(raw) 

122 except ValidationError: 

123 verbose_proxy_logger.warning( 

124 "Ignoring invalid %s value %r; expected a positive integer", 

125 BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, 

126 raw, 

127 ) 

128 return None 

129 

130 

131def resolve_batch_enqueued_token_scopes( 

132 user_api_key_dict: UserAPIKeyAuth, 

133) -> tuple[BatchEnqueuedTokenScope, ...]: 

134 key_limit: Final = _read_metadata_limit(user_api_key_dict.metadata) 

135 team_limit: Final = _read_metadata_limit(user_api_key_dict.team_metadata) 

136 candidates: Final = ( 

137 BatchEnqueuedTokenScope(key="api_key", value=user_api_key_dict.api_key, limit=key_limit) 

138 if key_limit is not None and user_api_key_dict.api_key 

139 else None, 

140 BatchEnqueuedTokenScope(key="team", value=user_api_key_dict.team_id, limit=team_limit) 

141 if team_limit is not None and user_api_key_dict.team_id 

142 else None, 

143 ) 

144 return tuple(scope for scope in candidates if scope is not None) 

145 

146 

147def canonical_provider_batch_id(batch_id: str) -> str: 

148 from litellm.proxy.openai_files_endpoints.common_utils import ( 

149 _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper 

150 get_batch_id_from_unified_batch_id, 

151 get_original_file_id, 

152 ) 

153 

154 decoded: Final = _is_base64_encoded_unified_file_id(batch_id) 

155 if isinstance(decoded, str): 

156 if "llm_batch_id" in decoded or "generic_response_id" in decoded: 

157 return get_batch_id_from_unified_batch_id(decoded) 

158 return decoded 

159 return get_original_file_id(batch_id) 

160 

161 

162class _BatchResponseView(BaseModel): 

163 model_config = ConfigDict(extra="ignore") 

164 

165 id: str 

166 status: str 

167 object: Literal["batch"] 

168 

169 

170def batch_response_view(response: object) -> _BatchResponseView | None: 

171 try: 

172 return _BatchResponseView.model_validate(response, from_attributes=True) 

173 except ValidationError: 

174 return None 

175 

176 

177class BatchEnqueuedTokenStore: 

178 """Tracks enqueued batch tokens per scope, plus per-batch reservation records for refunds. 

179 

180 Counters and records live in Redis when Redis is configured, through 

181 single-key Lua scripts issued one scope at a time (Redis Cluster safe: no 

182 cross-slot commands), with an over-limit or failing scope rolling back the 

183 scopes reserved before it; otherwise a single-process in-memory fallback 

184 guarded by one asyncio lock is used. Reservations remember which backend 

185 granted them, and in-memory grants also remember the granting worker, so a 

186 refund never debits counters the grant did not charge. Everything expires after 

187 ``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the 

188 terminal-state refund can never leak tokens forever, and reservation records 

189 expire no later than the counters they would refund, so a stale record can 

190 never debit an allowance re-granted after its counters expired. 

191 """ 

192 

193 def __init__( 

194 self, 

195 internal_usage_cache: "InternalUsageCache", 

196 monotonic: Callable[[], float] = time.monotonic, 

197 ) -> None: 

198 self.internal_usage_cache = internal_usage_cache 

199 self._monotonic: Final = monotonic 

200 self._lock = asyncio.Lock() 

201 self._owner_token = uuid.uuid4().hex 

202 redis_cache = internal_usage_cache.dual_cache.redis_cache 

203 self._reserve_script: _ScriptRunner | None = ( 

204 redis_cache.async_register_script(RESERVE_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None 

205 ) 

206 self._refund_script: _ScriptRunner | None = ( 

207 redis_cache.async_register_script(REFUND_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None 

208 ) 

209 self._save_script: _ScriptRunner | None = ( 

210 redis_cache.async_register_script(SAVE_RESERVATION_SCRIPT) if redis_cache is not None else None 

211 ) 

212 self._pop_script: _ScriptRunner | None = ( 

213 redis_cache.async_register_script(POP_RESERVATION_SCRIPT) if redis_cache is not None else None 

214 ) 

215 

216 @staticmethod 

217 def _counter_key(scope: BatchEnqueuedTokenScope) -> str: 

218 return f"batch_enqueued_tokens:{scope.key}:{scope.value}" 

219 

220 @staticmethod 

221 def _record_key(batch_id: str) -> str: 

222 return f"batch_enqueued_token_reservation:{batch_id}" 

223 

224 async def reserve( 

225 self, 

226 tokens: int, 

227 scopes: tuple[BatchEnqueuedTokenScope, ...], 

228 litellm_parent_otel_span: "Span | None" = None, 

229 ) -> BatchEnqueuedTokenOutcome: 

230 if tokens <= 0 or not scopes: 

231 return BatchEnqueuedTokenReservation(tokens=max(tokens, 0), scopes=scopes) 

232 reserve_script: Final = self._reserve_script 

233 refund_script: Final = self._refund_script 

234 if reserve_script is not None and refund_script is not None: 

235 try: 

236 return await self._reserve_via_redis(reserve_script, refund_script, tokens=tokens, scopes=scopes) 

237 except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters 

238 log_redis_failure( 

239 verbose_proxy_logger, 

240 logging.WARNING, 

241 "Redis enqueued-token reserve failed, falling back to in-memory", 

242 e, 

243 ) 

244 return await self._reserve_in_memory(tokens=tokens, scopes=scopes, span=litellm_parent_otel_span) 

245 

246 async def _reserve_via_redis( 

247 self, 

248 reserve_script: _ScriptRunner, 

249 refund_script: _ScriptRunner, 

250 tokens: int, 

251 scopes: tuple[BatchEnqueuedTokenScope, ...], 

252 ) -> BatchEnqueuedTokenOutcome: 

253 started: Final = self._monotonic() 

254 for index, scope in enumerate(scopes): 

255 result = await self._run_reserve_script( 

256 reserve_script, 

257 refund_script, 

258 tokens=tokens, 

259 scope=scope, 

260 already_reserved=scopes[:index], 

261 ) 

262 if result[0] != 1: 

263 await self._rollback_partial_reserve(refund_script, tokens=tokens, scopes=scopes[:index]) 

264 return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=result[1]) 

265 return BatchEnqueuedTokenReservation( 

266 tokens=tokens, scopes=scopes, backend="redis", reserved_at_monotonic=started 

267 ) 

268 

269 async def _run_reserve_script( 

270 self, 

271 reserve_script: _ScriptRunner, 

272 refund_script: _ScriptRunner, 

273 tokens: int, 

274 scope: BatchEnqueuedTokenScope, 

275 already_reserved: tuple[BatchEnqueuedTokenScope, ...], 

276 ) -> tuple[int, int]: 

277 try: 

278 raw_result: Final = await reserve_script( 

279 (self._counter_key(scope),), 

280 (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, scope.limit), 

281 ) 

282 return _RESERVE_RESULT_ADAPTER.validate_python(raw_result) 

283 except Exception: 

284 await self._rollback_partial_reserve(refund_script, tokens=tokens, scopes=already_reserved) 

285 raise 

286 

287 async def _rollback_partial_reserve( 

288 self, 

289 refund_script: _ScriptRunner, 

290 tokens: int, 

291 scopes: tuple[BatchEnqueuedTokenScope, ...], 

292 ) -> None: 

293 try: 

294 await self._refund_via_redis(refund_script, tokens=tokens, scopes=scopes) 

295 except Exception as e: # noqa: BLE001 # best-effort rollback: the leak is TTL-bounded and only tightens the allowance 

296 verbose_proxy_logger.warning( 

297 "Rollback of partially reserved enqueued tokens failed; leaked increments expire with the TTL: %s", 

298 str(e), 

299 ) 

300 

301 async def _refund_via_redis( 

302 self, 

303 refund_script: _ScriptRunner, 

304 tokens: int, 

305 scopes: tuple[BatchEnqueuedTokenScope, ...], 

306 ) -> None: 

307 for scope in scopes: 

308 await refund_script((self._counter_key(scope),), (tokens,)) 

309 

310 async def _reserve_in_memory( 

311 self, 

312 tokens: int, 

313 scopes: tuple[BatchEnqueuedTokenScope, ...], 

314 span: "Span | None", 

315 ) -> BatchEnqueuedTokenOutcome: 

316 started: Final = self._monotonic() 

317 async with self._lock: 

318 currents: Final = tuple([await self._get_local_counter(scope, span) for scope in scopes]) 

319 for scope, current in zip(scopes, currents): 

320 if current + tokens > scope.limit: 

321 return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=current) 

322 for scope, current in zip(scopes, currents): 

323 await self._set_local_counter(scope, current + tokens, span) 

324 return BatchEnqueuedTokenReservation( 

325 tokens=tokens, scopes=scopes, backend="memory", owner=self._owner_token, reserved_at_monotonic=started 

326 ) 

327 

328 async def refund( 

329 self, 

330 reservation: BatchEnqueuedTokenReservation, 

331 litellm_parent_otel_span: "Span | None" = None, 

332 ) -> None: 

333 if reservation.tokens <= 0 or not reservation.scopes: 

334 return 

335 if reservation.backend == "redis": 

336 await self._refund_redis_reservation(reservation) 

337 return 

338 if reservation.owner != self._owner_token: 

339 verbose_proxy_logger.warning( 

340 "Skipping enqueued-token refund granted in another worker's memory; its counters expire with the TTL" 

341 ) 

342 return 

343 async with self._lock: 

344 for scope in reservation.scopes: 

345 current = await self._get_local_counter(scope, litellm_parent_otel_span) 

346 remaining = current - reservation.tokens 

347 if remaining <= 0: 

348 self.internal_usage_cache.dual_cache.in_memory_cache.delete_cache(key=self._counter_key(scope)) 

349 else: 

350 await self._set_local_counter(scope, remaining, litellm_parent_otel_span) 

351 

352 async def _refund_redis_reservation(self, reservation: BatchEnqueuedTokenReservation) -> None: 

353 refund_script: Final = self._refund_script 

354 if refund_script is None: 

355 verbose_proxy_logger.warning( 

356 "No Redis client for a Redis-granted enqueued-token refund; leaked increments expire with the TTL" 

357 ) 

358 return 

359 try: 

360 await self._refund_via_redis(refund_script, tokens=reservation.tokens, scopes=reservation.scopes) 

361 except Exception as e: # noqa: BLE001 # best-effort refund: the leak is TTL-bounded and only tightens the allowance 

362 verbose_proxy_logger.warning( 

363 "Redis enqueued-token refund failed; leaked increments expire with the TTL: %s", str(e) 

364 ) 

365 

366 async def save_reservation( 

367 self, 

368 batch_id: str, 

369 reservation: BatchEnqueuedTokenReservation, 

370 litellm_parent_otel_span: "Span | None" = None, 

371 ) -> None: 

372 serialized: Final = _RESERVATION_ADAPTER.dump_json(reservation).decode("utf-8") 

373 elapsed: Final = self._monotonic() - reservation.reserved_at_monotonic 

374 ttl: Final = max(1, BATCH_ENQUEUED_TOKEN_TTL_SECONDS - math.ceil(elapsed)) 

375 if self._save_script is not None: 

376 try: 

377 await self._save_script( 

378 (self._record_key(batch_id),), 

379 (serialized, ttl), 

380 ) 

381 except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record 

382 log_redis_failure( 

383 verbose_proxy_logger, 

384 logging.WARNING, 

385 "Redis enqueued-token reservation save failed, falling back to in-memory", 

386 e, 

387 ) 

388 else: 

389 return 

390 await self.internal_usage_cache.async_set_cache( 

391 key=self._record_key(batch_id), 

392 value=serialized, 

393 ttl=ttl, 

394 litellm_parent_otel_span=litellm_parent_otel_span, 

395 local_only=True, 

396 ) 

397 

398 async def pop_reservation( 

399 self, 

400 batch_id: str, 

401 litellm_parent_otel_span: "Span | None" = None, 

402 ) -> BatchEnqueuedTokenReservation | None: 

403 redis_raw: Final = await self._pop_redis_record(batch_id) 

404 if redis_raw is not None and not redis_raw: 

405 # The Redis pop tombstones popped records in place, so a hit on the empty 

406 # tombstone means the batch was already refunded elsewhere; a local copy 

407 # left behind by a save that raised after landing must not refund again. 

408 await self._pop_local_record(batch_id, litellm_parent_otel_span) 

409 return None 

410 raw: Final = ( 

411 redis_raw if redis_raw is not None else await self._pop_local_record(batch_id, litellm_parent_otel_span) 

412 ) 

413 if raw is None: 

414 return None 

415 try: 

416 if isinstance(raw, (str, bytes)): 

417 return _RESERVATION_ADAPTER.validate_json(raw) 

418 return _RESERVATION_ADAPTER.validate_python(raw) 

419 except ValidationError: 

420 verbose_proxy_logger.warning("Discarding malformed enqueued-token reservation record for %s", batch_id) 

421 return None 

422 

423 async def _pop_redis_record(self, batch_id: str) -> str | bytes | None: 

424 pop_script: Final = self._pop_script 

425 if pop_script is None: 

426 return None 

427 try: 

428 return _POPPED_VALUE_ADAPTER.validate_python( 

429 await pop_script((self._record_key(batch_id),), (BATCH_ENQUEUED_TOKEN_TTL_SECONDS,)) 

430 ) 

431 except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record 

432 log_redis_failure( 

433 verbose_proxy_logger, 

434 logging.WARNING, 

435 "Redis enqueued-token reservation pop failed, falling back to in-memory", 

436 e, 

437 ) 

438 return None 

439 

440 async def _pop_local_record(self, batch_id: str, span: "Span | None") -> object: 

441 async with self._lock: 

442 stored = await self.internal_usage_cache.async_get_cache( 

443 key=self._record_key(batch_id), 

444 litellm_parent_otel_span=span, 

445 local_only=True, 

446 ) 

447 if stored is None: 

448 return None 

449 self.internal_usage_cache.dual_cache.in_memory_cache.delete_cache(key=self._record_key(batch_id)) 

450 return stored 

451 

452 async def _get_local_counter(self, scope: BatchEnqueuedTokenScope, span: "Span | None") -> int: 

453 stored = await self.internal_usage_cache.async_get_cache( 

454 key=self._counter_key(scope), 

455 litellm_parent_otel_span=span, 

456 local_only=True, 

457 ) 

458 return _STORED_COUNTER_ADAPTER.validate_python(stored) or 0 

459 

460 async def _set_local_counter(self, scope: BatchEnqueuedTokenScope, value: int, span: "Span | None") -> None: 

461 await self.internal_usage_cache.async_set_cache( 

462 key=self._counter_key(scope), 

463 value=value, 

464 ttl=BATCH_ENQUEUED_TOKEN_TTL_SECONDS, 

465 litellm_parent_otel_span=span, 

466 local_only=True, 

467 )