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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Enqueued-token accounting for batch submissions.
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"""
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
20from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
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
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
30 from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
32 Span = _Span
33 InternalUsageCache = _InternalUsageCache
35BATCH_ENQUEUED_REFUND_STATUSES: Final[frozenset[str]] = frozenset(
36 {"completed", "complete", "failed", "expired", "cancelled", "cancelling"}
37)
39ScopeKey: TypeAlias = Literal["api_key", "team"]
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"""
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"""
62SAVE_RESERVATION_SCRIPT: Final = """
63redis.call('SET', KEYS[1], ARGV[1], 'EX', tonumber(ARGV[2]))
64return 1
65"""
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"""
76@dataclass(frozen=True, slots=True)
77class BatchEnqueuedTokenScope:
78 key: ScopeKey
79 value: str
80 limit: int
83ReservationBackend: TypeAlias = Literal["redis", "memory"]
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)
95@dataclass(frozen=True, slots=True)
96class BatchEnqueuedTokenOverLimit:
97 scope: BatchEnqueuedTokenScope
98 enqueued: int
101BatchEnqueuedTokenOutcome: TypeAlias = BatchEnqueuedTokenReservation | BatchEnqueuedTokenOverLimit
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)
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
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
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)
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 )
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)
162class _BatchResponseView(BaseModel):
163 model_config = ConfigDict(extra="ignore")
165 id: str
166 status: str
167 object: Literal["batch"]
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
177class BatchEnqueuedTokenStore:
178 """Tracks enqueued batch tokens per scope, plus per-batch reservation records for refunds.
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 """
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 )
216 @staticmethod
217 def _counter_key(scope: BatchEnqueuedTokenScope) -> str:
218 return f"batch_enqueued_tokens:{scope.key}:{scope.value}"
220 @staticmethod
221 def _record_key(batch_id: str) -> str:
222 return f"batch_enqueued_token_reservation:{batch_id}"
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)
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 )
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
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 )
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,))
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 )
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)
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 )
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 )
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
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
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
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
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 )