From 7a6a677b72cd3822e3fe0ae29329fa8b0e17d577 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 15:09:05 -0700 Subject: [PATCH] feat(proxy): enqueued-token rate limiting for batches with refund on completion and cancellation --- litellm/constants.py | 6 + litellm/proxy/hooks/batch_enqueued_tokens.py | 357 ++++++++++++++++++ litellm/proxy/hooks/batch_rate_limiter.py | 80 +++- .../hooks/parallel_request_limiter_v3.py | 42 +++ tests/e2e/batches/test_batches_e2e.py | 146 ++++++- .../coverage_registry/quota_management.yaml | 3 + tests/e2e/models.py | 1 + .../proxy/hooks/test_batch_enqueued_tokens.py | 190 ++++++++++ .../proxy/hooks/test_batch_file_validation.py | 192 ++++++++++ .../hooks/test_parallel_request_limiter_v3.py | 120 ++++++ 10 files changed, 1133 insertions(+), 4 deletions(-) create mode 100644 litellm/proxy/hooks/batch_enqueued_tokens.py create mode 100644 tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py diff --git a/litellm/constants.py b/litellm/constants.py index 39a49e55f0d..d7ddacf5fac 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1766,6 +1766,12 @@ PTU_LAPSED_ALERT_LIMIT: Final[int] = 10 # one is seconds old, so a few minutes separates them. PTU_PRUNE_SKEW_GRACE_SECONDS: Final[int] = 300 +# How long enqueued-token reservations for batches live without a refund. Providers +# complete or expire batches within their completion window (24h for OpenAI), so a +# reservation still unrefunded after 8 days belongs to a batch whose terminal state +# was never observed (e.g. proxy restart); expiry returns the tokens to the caller. +BATCH_ENQUEUED_TOKEN_TTL_SECONDS: Final[int] = 8 * 24 * 60 * 60 + # Shared read-only empty mapping, for defaulting optional Mapping parameters without # constructing a fresh mutable dict at each call site. EMPTY_MAPPING: Final = MappingProxyType({}) diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py new file mode 100644 index 00000000000..040c8b687a7 --- /dev/null +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -0,0 +1,357 @@ +""" +Enqueued-token accounting for batch submissions. + +Opt-in via ``batch_enqueued_token_limit`` in key or team metadata: batch +submissions reserve their estimated token count against a long-lived +enqueued-token allowance instead of the per-minute rate-limit windows, and +the reservation is refunded when the batch reaches a terminal state +(completed, failed, expired, or cancelled). +""" + +import asyncio +from collections.abc import Awaitable, Mapping, Sequence +from dataclasses import dataclass +from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias + +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError + +from litellm._logging import verbose_proxy_logger +from litellm.constants import BATCH_ENQUEUED_TOKEN_TTL_SECONDS +from litellm.proxy._types import UserAPIKeyAuth + +if TYPE_CHECKING: + from opentelemetry.trace import Span as _Span + + from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache + + Span = _Span + InternalUsageCache = _InternalUsageCache + +BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit" + +BATCH_ENQUEUED_REFUND_STATUSES: Final[frozenset[str]] = frozenset( + {"completed", "complete", "failed", "expired", "cancelled", "cancelling"} +) + +ScopeKey: TypeAlias = Literal["api_key", "team"] + +RESERVE_ENQUEUED_TOKENS_SCRIPT: Final = """ +local amount = tonumber(ARGV[1]) +local ttl = tonumber(ARGV[2]) +for i = 1, #KEYS do + local limit = tonumber(ARGV[2 + i]) + local current = tonumber(redis.call('GET', KEYS[i]) or '0') + if current + amount > limit then + return {0, i - 1, current} + end +end +for i = 1, #KEYS do + redis.call('INCRBY', KEYS[i], amount) + redis.call('EXPIRE', KEYS[i], ttl) +end +return {1, -1, 0} +""" + +REFUND_ENQUEUED_TOKENS_SCRIPT: Final = """ +local amount = tonumber(ARGV[1]) +for i = 1, #KEYS do + local updated = redis.call('DECRBY', KEYS[i], amount) + if updated <= 0 then + redis.call('DEL', KEYS[i]) + end +end +return 1 +""" + +SAVE_RESERVATION_SCRIPT: Final = """ +redis.call('SET', KEYS[1], ARGV[1], 'EX', tonumber(ARGV[2])) +return 1 +""" + +POP_RESERVATION_SCRIPT: Final = """ +local value = redis.call('GET', KEYS[1]) +if value then + redis.call('DEL', KEYS[1]) +end +return value +""" + + +@dataclass(frozen=True, slots=True) +class BatchEnqueuedTokenScope: + key: ScopeKey + value: str + limit: int + + +@dataclass(frozen=True, slots=True) +class BatchEnqueuedTokenReservation: + tokens: int + scopes: tuple[BatchEnqueuedTokenScope, ...] + + +@dataclass(frozen=True, slots=True) +class BatchEnqueuedTokenOverLimit: + scope: BatchEnqueuedTokenScope + enqueued: int + + +BatchEnqueuedTokenOutcome: TypeAlias = BatchEnqueuedTokenReservation | BatchEnqueuedTokenOverLimit + +_LIMIT_ADAPTER: Final = TypeAdapter(Annotated[int, Field(gt=0)]) +_RESERVE_RESULT_ADAPTER: Final = TypeAdapter(tuple[int, int, int]) +_POPPED_VALUE_ADAPTER: Final = TypeAdapter(str | bytes | None) +_STORED_COUNTER_ADAPTER: Final = TypeAdapter(int | None) +_RESERVATION_ADAPTER: Final = TypeAdapter(BatchEnqueuedTokenReservation) + + +class _ScriptRunner(Protocol): + def __call__(self, keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> Awaitable[object]: ... + + +def _read_metadata_limit(metadata: Mapping[str, object] | None) -> int | None: + if not metadata: + return None + raw: Final = metadata.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY) + if raw is None: + return None + try: + return _LIMIT_ADAPTER.validate_python(raw) + except ValidationError: + verbose_proxy_logger.warning( + "Ignoring invalid %s value %r; expected a positive integer", + BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, + raw, + ) + return None + + +def resolve_batch_enqueued_token_scopes( + user_api_key_dict: UserAPIKeyAuth, +) -> tuple[BatchEnqueuedTokenScope, ...]: + key_limit: Final = _read_metadata_limit(user_api_key_dict.metadata) + team_limit: Final = _read_metadata_limit(user_api_key_dict.team_metadata) + candidates: Final = ( + BatchEnqueuedTokenScope(key="api_key", value=user_api_key_dict.api_key, limit=key_limit) + if key_limit is not None and user_api_key_dict.api_key + else None, + BatchEnqueuedTokenScope(key="team", value=user_api_key_dict.team_id, limit=team_limit) + if team_limit is not None and user_api_key_dict.team_id + else None, + ) + return tuple(scope for scope in candidates if scope is not None) + + +def canonical_provider_batch_id(batch_id: str) -> str: + from litellm.proxy.openai_files_endpoints.common_utils import ( + _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper + get_batch_id_from_unified_batch_id, + get_original_file_id, + ) + + decoded: Final = _is_base64_encoded_unified_file_id(batch_id) + if isinstance(decoded, str): + if "llm_batch_id" in decoded or "generic_response_id" in decoded: + return get_batch_id_from_unified_batch_id(decoded) + return decoded + return get_original_file_id(batch_id) + + +class _BatchResponseView(BaseModel): + model_config = ConfigDict(extra="ignore") + + id: str + status: str + object: Literal["batch"] + + +def batch_response_view(response: object) -> _BatchResponseView | None: + try: + return _BatchResponseView.model_validate(response, from_attributes=True) + except ValidationError: + return None + + +class BatchEnqueuedTokenStore: + """Tracks enqueued batch tokens per scope, plus per-batch reservation records for refunds. + + Counters and records live in Redis (via atomic Lua scripts) when Redis is + configured; otherwise a single-process in-memory fallback guarded by one + asyncio lock is used. Everything expires after + ``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the + terminal-state refund can never leak tokens forever. + """ + + def __init__(self, internal_usage_cache: "InternalUsageCache") -> None: + self.internal_usage_cache = internal_usage_cache + self._lock = asyncio.Lock() + redis_cache = internal_usage_cache.dual_cache.redis_cache + self._reserve_script: _ScriptRunner | None = ( + redis_cache.async_register_script(RESERVE_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None + ) + self._refund_script: _ScriptRunner | None = ( + redis_cache.async_register_script(REFUND_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None + ) + self._save_script: _ScriptRunner | None = ( + redis_cache.async_register_script(SAVE_RESERVATION_SCRIPT) if redis_cache is not None else None + ) + self._pop_script: _ScriptRunner | None = ( + redis_cache.async_register_script(POP_RESERVATION_SCRIPT) if redis_cache is not None else None + ) + + @staticmethod + def _counter_key(scope: BatchEnqueuedTokenScope) -> str: + return f"batch_enqueued_tokens:{scope.key}:{scope.value}" + + @staticmethod + def _record_key(batch_id: str) -> str: + return f"batch_enqueued_token_reservation:{batch_id}" + + async def reserve( + self, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + litellm_parent_otel_span: "Span | None" = None, + ) -> BatchEnqueuedTokenOutcome: + if tokens <= 0 or not scopes: + return BatchEnqueuedTokenReservation(tokens=max(tokens, 0), scopes=scopes) + if self._reserve_script is not None: + try: + raw_result = await self._reserve_script( + tuple(self._counter_key(scope) for scope in scopes), + (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, *(scope.limit for scope in scopes)), + ) + result = _RESERVE_RESULT_ADAPTER.validate_python(raw_result) + if result[0] == 1: + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) + return BatchEnqueuedTokenOverLimit(scope=scopes[result[1]], enqueued=result[2]) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters + verbose_proxy_logger.warning( + "Redis enqueued-token reserve failed, falling back to in-memory: %s", str(e) + ) + return await self._reserve_in_memory(tokens=tokens, scopes=scopes, span=litellm_parent_otel_span) + + async def _reserve_in_memory( + self, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + span: "Span | None", + ) -> BatchEnqueuedTokenOutcome: + async with self._lock: + currents: Final = tuple([await self._get_local_counter(scope, span) for scope in scopes]) + for scope, current in zip(scopes, currents): + if current + tokens > scope.limit: + return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=current) + for scope, current in zip(scopes, currents): + await self._set_local_counter(scope, current + tokens, span) + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) + + async def refund( + self, + reservation: BatchEnqueuedTokenReservation, + litellm_parent_otel_span: "Span | None" = None, + ) -> None: + if reservation.tokens <= 0 or not reservation.scopes: + return + if self._refund_script is not None: + try: + await self._refund_script( + tuple(self._counter_key(scope) for scope in reservation.scopes), + (reservation.tokens,), + ) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters + verbose_proxy_logger.warning( + "Redis enqueued-token refund failed, falling back to in-memory: %s", str(e) + ) + else: + return + async with self._lock: + for scope in reservation.scopes: + current = await self._get_local_counter(scope, litellm_parent_otel_span) + remaining = current - reservation.tokens + if remaining <= 0: + self.internal_usage_cache.dual_cache.in_memory_cache.delete_cache(key=self._counter_key(scope)) + else: + await self._set_local_counter(scope, remaining, litellm_parent_otel_span) + + async def save_reservation( + self, + batch_id: str, + reservation: BatchEnqueuedTokenReservation, + litellm_parent_otel_span: "Span | None" = None, + ) -> None: + serialized: Final = _RESERVATION_ADAPTER.dump_json(reservation).decode("utf-8") + if self._save_script is not None: + try: + await self._save_script( + (self._record_key(batch_id),), + (serialized, BATCH_ENQUEUED_TOKEN_TTL_SECONDS), + ) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record + verbose_proxy_logger.warning( + "Redis enqueued-token reservation save failed, falling back to in-memory: %s", str(e) + ) + else: + return + await self.internal_usage_cache.async_set_cache( + key=self._record_key(batch_id), + value=serialized, + ttl=BATCH_ENQUEUED_TOKEN_TTL_SECONDS, + litellm_parent_otel_span=litellm_parent_otel_span, + local_only=True, + ) + + async def pop_reservation( + self, + batch_id: str, + litellm_parent_otel_span: "Span | None" = None, + ) -> BatchEnqueuedTokenReservation | None: + raw: object = None + if self._pop_script is not None: + try: + raw = _POPPED_VALUE_ADAPTER.validate_python(await self._pop_script((self._record_key(batch_id),), ())) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record + verbose_proxy_logger.warning( + "Redis enqueued-token reservation pop failed, falling back to in-memory: %s", str(e) + ) + raw = await self._pop_local_record(batch_id, litellm_parent_otel_span) + else: + raw = await self._pop_local_record(batch_id, litellm_parent_otel_span) + if raw is None: + return None + try: + if isinstance(raw, (str, bytes)): + return _RESERVATION_ADAPTER.validate_json(raw) + return _RESERVATION_ADAPTER.validate_python(raw) + except ValidationError: + verbose_proxy_logger.warning("Discarding malformed enqueued-token reservation record for %s", batch_id) + return None + + async def _pop_local_record(self, batch_id: str, span: "Span | None") -> object: + async with self._lock: + stored = await self.internal_usage_cache.async_get_cache( + key=self._record_key(batch_id), + litellm_parent_otel_span=span, + local_only=True, + ) + if stored is None: + return None + self.internal_usage_cache.dual_cache.in_memory_cache.delete_cache(key=self._record_key(batch_id)) + return stored + + async def _get_local_counter(self, scope: BatchEnqueuedTokenScope, span: "Span | None") -> int: + stored = await self.internal_usage_cache.async_get_cache( + key=self._counter_key(scope), + litellm_parent_otel_span=span, + local_only=True, + ) + return _STORED_COUNTER_ADAPTER.validate_python(stored) or 0 + + async def _set_local_counter(self, scope: BatchEnqueuedTokenScope, value: int, span: "Span | None") -> None: + await self.internal_usage_cache.async_set_cache( + key=self._counter_key(scope), + value=value, + ttl=BATCH_ENQUEUED_TOKEN_TTL_SECONDS, + litellm_parent_otel_span=span, + local_only=True, + ) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index d6229fb80a6..5b814ad28fd 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -46,9 +46,16 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, map_v3_rate_limit_type, ) +from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenOverLimit, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + resolve_batch_enqueued_token_scopes, +) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY, + get_or_create_request_stash, ) from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit @@ -291,6 +298,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): self, data: dict, user_api_key_dict: UserAPIKeyAuth, + has_enqueued_scopes: bool = False, ) -> tuple[bool, list["RateLimitDescriptor"] | None]: """ Skip downloading batch input files when the operator disabled batch @@ -343,8 +351,10 @@ class _PROXY_BatchRateLimiter(CustomLogger): user_api_key_dict=user_api_key_dict, data=data, ) - if not self._has_applicable_batch_rate_limits(descriptors) and not self._project_has_any_io_token_limits( - user_api_key_dict + if ( + not has_enqueued_scopes + and not self._has_applicable_batch_rate_limits(descriptors) + and not self._project_has_any_io_token_limits(user_api_key_dict) ): verbose_proxy_logger.debug("Skipping batch input file processing: no rate limits configured") return True, None @@ -511,6 +521,59 @@ class _PROXY_BatchRateLimiter(CustomLogger): return file_id, fetch_kwargs + async def _reserve_batch_enqueued_tokens( + self, + user_api_key_dict: UserAPIKeyAuth, + data: Mapping[str, object], + batch_usage: BatchFileUsage, + scopes: tuple[BatchEnqueuedTokenScope, ...], + ) -> None: + """Reserve the batch's estimated tokens against the caller's enqueued-token allowance. + + Runs instead of the per-minute counter charge when the key or team + opted in via ``batch_enqueued_token_limit`` metadata. The reservation + is stashed on the request so the v3 limiter's post-call hooks can + persist it (keyed by the provider batch id) and refund it when the + batch reaches a terminal state. + """ + outcome: Final = await self.parallel_request_limiter.batch_enqueued_token_store.reserve( + tokens=batch_usage.total_tokens, + scopes=scopes, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + match outcome: + case BatchEnqueuedTokenOverLimit(): + self._raise_enqueued_limit_error(over_limit=outcome, data=data, batch_usage=batch_usage) + case BatchEnqueuedTokenReservation(): + get_or_create_request_stash().batch_enqueued_reservation = outcome + + def _raise_enqueued_limit_error( + self, + over_limit: BatchEnqueuedTokenOverLimit, + data: Mapping[str, object], + batch_usage: BatchFileUsage, + ) -> NoReturn: + scope: Final = over_limit.scope + remaining: Final = max(0, scope.limit - over_limit.enqueued) + detail: Final = ( + f"Batch enqueued token limit exceeded for {scope.key}: {scope.value}. " + f"Batch requires {batch_usage.total_tokens} tokens but only {remaining} enqueued tokens remaining " + f"out of {scope.limit} enqueued token limit. " + f"Tokens free up as running batches complete or are cancelled." + ) + raw_model: Final = data.get("model") + resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( + raw_model if isinstance(raw_model, str) else None + ) + raise ProxyRateLimitError( + detail=detail, + headers=MappingProxyType({"rate_limit_type": "tokens"}), + category=RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT, + rate_limit_type=map_v3_rate_limit_type("tokens"), + model=resolved_model, + llm_provider=llm_provider, + ) + def _raise_rate_limit_error( self, status: "RateLimitStatus", @@ -1039,8 +1102,9 @@ class _PROXY_BatchRateLimiter(CustomLogger): verbose_proxy_logger.debug("No input_file_id in batch request, skipping rate limiting") return data + enqueued_scopes: Final = resolve_batch_enqueued_token_scopes(user_api_key_dict) should_skip, batch_rate_limit_descriptors = self._should_skip_batch_input_file_processing( - data=data, user_api_key_dict=user_api_key_dict + data=data, user_api_key_dict=user_api_key_dict, has_enqueued_scopes=bool(enqueued_scopes) ) if should_skip: return data @@ -1066,6 +1130,16 @@ class _PROXY_BatchRateLimiter(CustomLogger): data["_batch_token_count"] = batch_usage.total_tokens data["_batch_request_count"] = batch_usage.request_count + if enqueued_scopes: + await self._reserve_batch_enqueued_tokens( + user_api_key_dict=user_api_key_dict, + data=data, + batch_usage=batch_usage, + scopes=enqueued_scopes, + ) + verbose_proxy_logger.debug("Batch enqueued-token reservation succeeded") + return data + # Directly increment counters by batch amounts (check happens atomically) # This will raise HTTPException if limits are exceeded await self._check_and_increment_batch_counters( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 2d799ded752..b295da69255 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -44,6 +44,13 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, map_v3_rate_limit_type, ) +from litellm.proxy.hooks.batch_enqueued_tokens import ( + BATCH_ENQUEUED_REFUND_STATUSES, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenStore, + batch_response_view, + canonical_provider_batch_id, +) from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage @@ -515,6 +522,7 @@ class RequestRateLimiterStash: otpm_reserved_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]] = field( default_factory=frozenset ) + batch_enqueued_reservation: BatchEnqueuedTokenReservation | None = None reservation_released: bool = False @@ -619,6 +627,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Batch rate limiter (lazy loaded) self._batch_rate_limiter: CallTypeRateLimiter | None = None + self.batch_enqueued_token_store = BatchEnqueuedTokenStore(internal_usage_cache=internal_usage_cache) # Serializes multi-phase check+increment sequences (batch + dynamic # limiters) within this process to close the TOCTOU window between @@ -4673,6 +4682,32 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): except Exception as e: verbose_proxy_logger.exception("Error in rate limit post-call hook: %s", e) + try: + await self._handle_batch_enqueued_post_call(user_api_key_dict=user_api_key_dict, response=response) + except Exception as e: # noqa: BLE001 # post-call batch accounting must never fail the response + verbose_proxy_logger.exception("Error in batch enqueued-token post-call hook: %s", e) + + async def _handle_batch_enqueued_post_call(self, user_api_key_dict: UserAPIKeyAuth, response: object) -> None: + view: Final = batch_response_view(response) + if view is None: + return + span: Final = user_api_key_dict.parent_otel_span + stash: Final = get_request_stash() + if stash is not None and stash.batch_enqueued_reservation is not None: + await self.batch_enqueued_token_store.save_reservation( + batch_id=canonical_provider_batch_id(view.id), + reservation=stash.batch_enqueued_reservation, + litellm_parent_otel_span=span, + ) + stash.batch_enqueued_reservation = None + if view.status in BATCH_ENQUEUED_REFUND_STATUSES: + popped: Final = await self.batch_enqueued_token_store.pop_reservation( + batch_id=canonical_provider_batch_id(view.id), + litellm_parent_otel_span=span, + ) + if popped is not None: + await self.batch_enqueued_token_store.refund(reservation=popped, litellm_parent_otel_span=span) + async def async_post_call_failure_hook( self, request_data: dict, @@ -4706,6 +4741,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) stash.parallel_slot = None + if stash.batch_enqueued_reservation is not None: + await self.batch_enqueued_token_store.refund( + reservation=stash.batch_enqueued_reservation, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + stash.batch_enqueued_reservation = None + if stash.reservation_released: return reserved_tokens: Final = stash.reserved_tokens diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py index 1376bdbed38..53bf9739983 100644 --- a/tests/e2e/batches/test_batches_e2e.py +++ b/tests/e2e/batches/test_batches_e2e.py @@ -17,6 +17,7 @@ from __future__ import annotations import json import os +import re import time from datetime import datetime, timedelta, timezone from typing import Callable @@ -57,7 +58,7 @@ from e2e_http import ( unwrap, ) from lifecycle import ResourceManager -from models import KeyGenerateBody, LiteLLMParamsBody, SpendLogRow +from models import KeyGenerateBody, KeyMetadata, LiteLLMParamsBody, SpendLogRow pytestmark = pytest.mark.e2e @@ -685,6 +686,149 @@ class TestBatchRateLimitErrorMapping: ) +BATCH_ENQUEUED_HEADROOM_TOKENS = 100_000 +_BATCH_REQUIRES_TOKENS = re.compile(r"Batch requires (\d+) tokens") + + +class TestBatchEnqueuedTokenLimit: + """Opt-in enqueued-token allowance governs batch submission instead of RPM/TPM. + + A key whose metadata carries batch_enqueued_token_limit reserves the batch's + token estimate against that allowance at create time: per-minute limits no + longer gate batch submission, exhausting the allowance rejects the create + before it reaches the provider, and cancelling a running batch refunds its + reservation so blocked submissions go through again (LIT-5273). + """ + + def _upload_batch_file( + self, client: BatchClient, resources: ResourceManager, key: str + ) -> FileObject: + file = unwrap( + client.upload_file( + content=_multi_request_jsonl("gpt-4o-mini", BATCH_RL_REQUEST_LINES), + form=FileUploadForm(purpose="batch"), + model=OPENAI_BATCH_MODEL, + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + return file + + def _generate_enqueued_key( + self, + client: BatchClient, + resources: ResourceManager, + *, + limit: int, + marker: str, + rpm_limit: int | None = None, + ) -> str: + key = client.proxy.generate_key( + KeyGenerateBody( + models=[], + rpm_limit=rpm_limit, + user_id=f"e2e-batch-enq-{marker}-{unique_marker()}", + metadata=KeyMetadata(batch_enqueued_token_limit=limit), + ) + ) + resources.defer(lambda: client.proxy.delete_key(key)) + return key + + @pytest.mark.covers( + "quota_management.ratelimit.batch_enqueued_tokens.accepts_over_rpm", + exercised_on=["batches"], + ) + def test_enqueued_allowance_accepts_batch_over_key_rpm( + self, client: BatchClient, resources: ResourceManager, batch_deployments: None + ) -> None: + key = self._generate_enqueued_key( + client, + resources, + limit=BATCH_ENQUEUED_HEADROOM_TOKENS, + marker="rpm", + rpm_limit=BATCH_RL_RPM_LIMIT, + ) + file = self._upload_batch_file(client, resources, key) + + created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + + assert created.status_code != 429, ( + f"enqueued-token allowance must govern batch submission instead of the " + f"key RPM ({BATCH_RL_RPM_LIMIT} < {BATCH_RL_REQUEST_LINES} rows); " + f"got 429: {created.body[:400]}" + ) + require_successful_call(created) + batch = BatchObject.model_validate_json(created.body) + resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key))) + + @pytest.mark.covers( + "quota_management.ratelimit.batch_enqueued_tokens.blocks_when_exhausted", + exercised_on=["batches"], + ) + @pytest.mark.covers( + "quota_management.ratelimit.batch_enqueued_tokens.refunds_on_cancel", + exercised_on=["batches"], + ) + def test_exhausted_allowance_blocks_until_cancel_refunds( + self, client: BatchClient, resources: ResourceManager, batch_deployments: None + ) -> None: + sizing_key = self._generate_enqueued_key( + client, resources, limit=1, marker="size" + ) + sizing_file = self._upload_batch_file(client, resources, sizing_key) + sized = client.create_batch( + body=BatchCreateBody(input_file_id=sizing_file.id), key=sizing_key + ) + assert sized.status_code == 429, ( + f"a 1-token allowance must reject any batch before it reaches the " + f"provider, got {sized.status_code}: {sized.body[:400]}" + ) + assert "batch enqueued token limit exceeded" in sized.body.lower(), ( + f"429 body must name the enqueued token limit, got: {sized.body[:400]}" + ) + requires = _BATCH_REQUIRES_TOKENS.search(sized.body) + assert requires is not None, ( + f"429 body must report the batch token requirement so callers can size " + f"allowances, got: {sized.body[:400]}" + ) + batch_tokens = int(requires.group(1)) + assert batch_tokens > 1 + + key = self._generate_enqueued_key( + client, resources, limit=batch_tokens + batch_tokens // 2, marker="refund" + ) + file = self._upload_batch_file(client, resources, key) + + first = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + require_successful_call(first) + first_batch = BatchObject.model_validate_json(first.body) + resources.defer(quietly(lambda: client.cancel_batch(first_batch.id, key=key))) + + blocked = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + assert blocked.status_code == 429, ( + f"second batch must not fit the remaining allowance while the first is " + f"enqueued, got {blocked.status_code}: {blocked.body[:400]}" + ) + assert "batch enqueued token limit exceeded" in blocked.body.lower(), ( + f"429 body must name the enqueued token limit, got: {blocked.body[:400]}" + ) + + cancelled = cancel_batch(client, first_batch.id, key=key, provider=None) + assert cancelled.status in {"cancelling", "cancelled"}, ( + f"cancel must reach a cancel state for the refund to fire, " + f"got {cancelled.status}" + ) + + retried = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + assert retried.status_code != 429, ( + f"cancelling the first batch must refund its reservation so the retry " + f"fits the allowance, got 429: {retried.body[:400]}" + ) + require_successful_call(retried) + retry_batch = BatchObject.model_validate_json(retried.body) + resources.defer(quietly(lambda: client.cancel_batch(retry_batch.id, key=key))) + + ASSUME_ROLE_RAW_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" diff --git a/tests/e2e/coverage_registry/quota_management.yaml b/tests/e2e/coverage_registry/quota_management.yaml index 2dfa7adddea..4b8aa1da002 100644 --- a/tests/e2e/coverage_registry/quota_management.yaml +++ b/tests/e2e/coverage_registry/quota_management.yaml @@ -2,6 +2,9 @@ # litellm/proxy/hooks/ + litellm/proxy/auth/auth_checks.py + litellm/proxy/spend_tracking/. - {id: quota_management.ratelimit.rpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: rpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "parallel_request_limiter_v3.py", rationale: "v3 limiter enforces RPM per key/team/model; 429 on breach"} - {id: quota_management.ratelimit.batch_rpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_rpm, assertions: [blocks_over_limit], exercised_on: [batches], source: "batch_rate_limiter.py", rationale: "Batch create that exceeds key RPM returns mapped 429 with retry-after"} +- {id: quota_management.ratelimit.batch_enqueued_tokens.accepts_over_rpm, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [accepts_over_rpm], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Key with an enqueued-token allowance submits a batch whose row count exceeds its RPM and the create is accepted"} +- {id: quota_management.ratelimit.batch_enqueued_tokens.blocks_when_exhausted, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [blocks_when_exhausted], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Batch create is rejected with a 429 naming the enqueued token limit once the allowance cannot fit the file"} +- {id: quota_management.ratelimit.batch_enqueued_tokens.refunds_on_cancel, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [refunds_on_cancel], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Cancelling a running batch returns its reserved tokens so a previously blocked submission succeeds"} - {id: quota_management.ratelimit.tpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: tpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "parallel_request_limiter_v3.py", rationale: "v3 limiter enforces TPM per key/team/model; 429 on breach"} - {id: quota_management.ratelimit.tpm.excludes_cached_tokens, module: quota_management, tier: P0, behavior: ratelimit, variant: tpm, assertions: [excludes_cached_tokens], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py:_get_total_tokens_from_usage", rationale: "Cached prompt tokens must not count toward TPM (LIT-1930)"} - {id: quota_management.ratelimit.redis_backed.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: redis_backed, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py", rationale: "With Redis configured, RPM still enforces 429 across the shared limiter path customers run multi-replica"} diff --git a/tests/e2e/models.py b/tests/e2e/models.py index df5cb841fad..150be8966ee 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -46,6 +46,7 @@ class KeyLoggingCallback(BaseModel): class KeyMetadata(BaseModel): logging: list[KeyLoggingCallback] | None = None priority: str | None = None + batch_enqueued_token_limit: int | None = None class ObjectPermission(BaseModel): diff --git a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py new file mode 100644 index 00000000000..a917206f33a --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -0,0 +1,190 @@ +""" +LIT-5273: enqueued-token accounting for batch submissions. + +Covers the ``BatchEnqueuedTokenStore`` (reserve / refund / reservation +records), the metadata-driven scope resolution, and the batch-id and +response-shape helpers the v3 limiter's post-call hooks rely on. +""" + +import base64 +import socket +import uuid +from types import SimpleNamespace + +import pytest + +from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenOverLimit, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + BatchEnqueuedTokenStore, + batch_response_view, + canonical_provider_batch_id, + resolve_batch_enqueued_token_scopes, +) +from litellm.proxy.utils import InternalUsageCache + + +def _in_memory_store() -> BatchEnqueuedTokenStore: + return BatchEnqueuedTokenStore(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) + + +def _scope(limit: int, key: str = "api_key") -> BatchEnqueuedTokenScope: + return BatchEnqueuedTokenScope(key=key, value=f"{key}-{uuid.uuid4().hex}", limit=limit) + + +def test_scope_resolution_reads_key_and_team_metadata(): + user = UserAPIKeyAuth( + api_key="hashed-key", + metadata={"batch_enqueued_token_limit": 100}, + team_id="team-1", + team_metadata={"batch_enqueued_token_limit": "150"}, + ) + scopes = resolve_batch_enqueued_token_scopes(user) + assert scopes == ( + BatchEnqueuedTokenScope(key="api_key", value="hashed-key", limit=100), + BatchEnqueuedTokenScope(key="team", value="team-1", limit=150), + ) + + +def test_scope_resolution_returns_empty_without_opt_in(): + assert resolve_batch_enqueued_token_scopes(UserAPIKeyAuth(api_key="k")) == () + assert resolve_batch_enqueued_token_scopes(UserAPIKeyAuth(api_key="k", metadata={}, team_metadata=None)) == () + + +@pytest.mark.parametrize("bad_value", ["not-a-number", 0, -5, None, [1000]]) +def test_scope_resolution_ignores_invalid_limits(bad_value): + user = UserAPIKeyAuth(api_key="k", metadata={"batch_enqueued_token_limit": bad_value}) + assert resolve_batch_enqueued_token_scopes(user) == () + + +def test_scope_resolution_skips_team_scope_without_team_id(): + user = UserAPIKeyAuth(api_key="k", team_metadata={"batch_enqueued_token_limit": 100}) + assert resolve_batch_enqueued_token_scopes(user) == () + + +@pytest.mark.asyncio +async def test_reserve_rejects_once_allowance_is_exhausted(): + store = _in_memory_store() + scope = _scope(limit=100) + first = await store.reserve(tokens=80, scopes=(scope,)) + assert isinstance(first, BatchEnqueuedTokenReservation) + second = await store.reserve(tokens=30, scopes=(scope,)) + assert second == BatchEnqueuedTokenOverLimit(scope=scope, enqueued=80) + third = await store.reserve(tokens=20, scopes=(scope,)) + assert isinstance(third, BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_reserve_is_all_or_nothing_across_scopes(): + store = _in_memory_store() + key_scope = _scope(limit=100, key="api_key") + team_scope = _scope(limit=50, key="team") + over = await store.reserve(tokens=60, scopes=(key_scope, team_scope)) + assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0) + exact_fit = await store.reserve(tokens=50, scopes=(key_scope, team_scope)) + assert isinstance(exact_fit, BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_refund_restores_allowance_and_never_goes_negative(): + store = _in_memory_store() + scope = _scope(limit=100) + reservation = await store.reserve(tokens=30, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + await store.refund(reservation) + await store.refund(reservation) + refill = await store.reserve(tokens=100, scopes=(scope,)) + assert isinstance(refill, BatchEnqueuedTokenReservation) + assert isinstance(await store.reserve(tokens=1, scopes=(scope,)), BatchEnqueuedTokenOverLimit) + + +@pytest.mark.asyncio +async def test_reservation_record_roundtrip_pops_exactly_once(): + store = _in_memory_store() + scope = _scope(limit=100) + reservation = await store.reserve(tokens=40, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + await store.save_reservation("batch_abc", reservation) + assert await store.pop_reservation("batch_abc") == reservation + assert await store.pop_reservation("batch_abc") is None + assert await store.pop_reservation("batch_never_saved") is None + + +@pytest.mark.asyncio +async def test_zero_token_reserve_charges_nothing(): + store = _in_memory_store() + scope = _scope(limit=100) + empty = await store.reserve(tokens=0, scopes=(scope,)) + assert empty == BatchEnqueuedTokenReservation(tokens=0, scopes=(scope,)) + full = await store.reserve(tokens=100, scopes=(scope,)) + assert isinstance(full, BatchEnqueuedTokenReservation) + + +def test_canonical_provider_batch_id_passes_raw_ids_through(): + assert canonical_provider_batch_id("batch_abc123") == "batch_abc123" + + +def test_canonical_provider_batch_id_decodes_unified_batch_ids(): + unified = "litellm_proxy;model_id:m-1;llm_batch_id:batch_prov_9;llm_output_file_id:file-9" + encoded = base64.urlsafe_b64encode(unified.encode()).decode().rstrip("=") + assert canonical_provider_batch_id(encoded) == "batch_prov_9" + + +def test_canonical_provider_batch_id_decodes_model_embedded_ids(): + from litellm.proxy.openai_files_endpoints.common_utils import encode_file_id_with_model + + encoded = encode_file_id_with_model(file_id="batch_prov_7", model="my-alias", id_type="batch") + assert canonical_provider_batch_id(encoded) == "batch_prov_7" + + +def test_batch_response_view_accepts_batch_objects_only(): + batch = SimpleNamespace(id="batch_1", status="completed", object="batch") + view = batch_response_view(batch) + assert view is not None and view.id == "batch_1" and view.status == "completed" + assert batch_response_view({"id": "chatcmpl-1", "object": "chat.completion"}) is None + assert batch_response_view(None) is None + assert batch_response_view("batch_1") is None + + +def _local_redis_port() -> int | None: + for port in (6379,): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.settimeout(0.2) + if sock.connect_ex(("127.0.0.1", port)) == 0: + return port + return None + + +@pytest.mark.asyncio +@pytest.mark.skipif(_local_redis_port() is None, reason="requires a local Redis on 6379 for the Lua script path") +async def test_redis_lua_path_full_lifecycle(): + from litellm.caching.redis_cache import RedisCache + + port = _local_redis_port() + redis_cache = RedisCache(host="127.0.0.1", port=port) + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=redis_cache, default_in_memory_ttl=60)) + ) + key_scope = _scope(limit=100, key="api_key") + team_scope = _scope(limit=50, key="team") + + over = await store.reserve(tokens=60, scopes=(key_scope, team_scope)) + assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0) + + reservation = await store.reserve(tokens=50, scopes=(key_scope, team_scope)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + assert isinstance(await store.reserve(tokens=1, scopes=(key_scope, team_scope)), BatchEnqueuedTokenOverLimit) + + batch_id = f"batch_{uuid.uuid4().hex}" + await store.save_reservation(batch_id, reservation) + popped = await store.pop_reservation(batch_id) + assert popped == reservation + assert await store.pop_reservation(batch_id) is None + + await store.refund(popped) + refill = await store.reserve(tokens=50, scopes=(key_scope, team_scope)) + assert isinstance(refill, BatchEnqueuedTokenReservation) + await store.refund(refill) diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index 1ce1a2f3e51..8698ee8ba12 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -2094,3 +2094,195 @@ def test_estimate_entry_output_tokens_multiplies_candidate_count(body_extra, exp } assert rate_limiter._estimate_entry_output_tokens(entry, None) == expected + + +# --------------------------------------------------------------------------- +# LIT-5273: enqueued-token limits govern batch submission when opted in +# --------------------------------------------------------------------------- + + +def _enqueued_rate_limiter(): + from litellm import DualCache + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + ) + from litellm.proxy.utils import InternalUsageCache + + local_cache = DualCache(default_in_memory_ttl=60) + internal_usage_cache = InternalUsageCache(local_cache) + parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache) + rate_limiter = _PROXY_BatchRateLimiter( + internal_usage_cache=internal_usage_cache, + parallel_request_limiter=parallel_request_limiter, + ) + return rate_limiter, local_cache + + +_ENQUEUED_BATCH_FILE_CONTENT = ( + b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n' + b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n' + b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n' +) + + +def _enqueued_batch_patches(): + mock_content = MagicMock() + mock_content.content = _ENQUEUED_BATCH_FILE_CONTENT + afile_content_mock = AsyncMock(return_value=mock_content) + return afile_content_mock, ( + patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model", + return_value={"custom_llm_provider": "openai"}, + ), + ) + + +@pytest.mark.asyncio +async def test_enqueued_limit_accepts_batch_over_per_minute_limits(): + """The headline LIT-5273 behavior: a key that opted into an enqueued-token + allowance submits a batch whose row count and token count both exceed its + per-minute RPM/TPM limits, and the batch is accepted (repeatedly) because + only the enqueued allowance governs. Without the opt-in the same key is + rejected on RPM before the batch reaches the provider.""" + from litellm.proxy.hooks.parallel_request_limiter_v3 import get_request_stash + + rate_limiter, local_cache = _enqueued_rate_limiter() + afile_content_mock, patches = _enqueued_batch_patches() + + legacy_user = UserAPIKeyAuth(api_key="sk-legacy-rpm", models=["*"], rpm_limit=1, tpm_limit=10) + opted_in_user = UserAPIKeyAuth( + api_key="sk-enqueued-rpm", + models=["*"], + rpm_limit=1, + tpm_limit=10, + metadata={"batch_enqueued_token_limit": 100000}, + ) + + with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock): + with pytest.raises(HTTPException) as legacy_exc: + await rate_limiter.async_pre_call_hook( + user_api_key_dict=legacy_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + assert legacy_exc.value.status_code == 429 + + first_data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"} + result = await rate_limiter.async_pre_call_hook( + user_api_key_dict=opted_in_user, + cache=local_cache, + data=first_data, + call_type="acreate_batch", + ) + assert result is first_data + stash = get_request_stash() + assert stash is not None and stash.batch_enqueued_reservation is not None + assert stash.batch_enqueued_reservation.tokens == first_data["_batch_token_count"] > 0 + + second = await rate_limiter.async_pre_call_hook( + user_api_key_dict=opted_in_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + assert second is not None + + +@pytest.mark.asyncio +async def test_enqueued_limit_rejects_when_allowance_is_exhausted(): + """Submissions are rejected pre-provider once the enqueued allowance can't + fit the batch, even for a key with no per-minute limits at all (which + previously skipped batch rate limiting entirely).""" + rate_limiter, local_cache = _enqueued_rate_limiter() + afile_content_mock, patches = _enqueued_batch_patches() + + sizing_user = UserAPIKeyAuth( + api_key="sk-enqueued-sizing", models=["*"], metadata={"batch_enqueued_token_limit": 1000000} + ) + with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock): + sizing_data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"} + await rate_limiter.async_pre_call_hook( + user_api_key_dict=sizing_user, + cache=local_cache, + data=sizing_data, + call_type="acreate_batch", + ) + batch_tokens = sizing_data["_batch_token_count"] + assert batch_tokens > 0 + + capped_user = UserAPIKeyAuth( + api_key="sk-enqueued-capped", + models=["*"], + metadata={"batch_enqueued_token_limit": batch_tokens + batch_tokens // 2}, + ) + await rate_limiter.async_pre_call_hook( + user_api_key_dict=capped_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + with pytest.raises(HTTPException) as exc: + await rate_limiter.async_pre_call_hook( + user_api_key_dict=capped_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + + assert exc.value.status_code == 429 + assert "Batch enqueued token limit exceeded for api_key" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_enqueued_team_limit_applies_to_batch_submission(): + rate_limiter, local_cache = _enqueued_rate_limiter() + afile_content_mock, patches = _enqueued_batch_patches() + + team_user = UserAPIKeyAuth( + api_key="sk-enqueued-team-key", + models=["*"], + team_id="team-enqueued-batch", + team_metadata={"batch_enqueued_token_limit": 10}, + ) + with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock): + with pytest.raises(HTTPException) as exc: + await rate_limiter.async_pre_call_hook( + user_api_key_dict=team_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + + assert exc.value.status_code == 429 + assert "Batch enqueued token limit exceeded for team: team-enqueued-batch" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_disable_flag_still_skips_batch_processing_with_enqueued_limits(): + rate_limiter, local_cache = _enqueued_rate_limiter() + afile_content_mock, _ = _enqueued_batch_patches() + + opted_in_user = UserAPIKeyAuth( + api_key="sk-enqueued-disabled", + models=["*"], + metadata={"batch_enqueued_token_limit": 10}, + ) + with ( + patch("litellm.proxy.proxy_server.general_settings", {"disable_batch_input_file_rate_limiting": True}), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.afile_content", new=afile_content_mock), + ): + data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"} + result = await rate_limiter.async_pre_call_hook( + user_api_key_dict=opted_in_user, + cache=local_cache, + data=data, + call_type="acreate_batch", + ) + + assert result is data + afile_content_mock.assert_not_awaited() diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 54366226dfb..68c219fd72d 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -5858,3 +5858,123 @@ async def test_conflicting_token_limits_cannot_bypass_tpm_reservation(): ) assert exc_info.value.status_code == 429 + + +# --------------------------------------------------------------------------- +# LIT-5273: batch enqueued-token reservations in the post-call hooks +# --------------------------------------------------------------------------- + + +def _enqueued_test_handler() -> _PROXY_MaxParallelRequestsHandler: + return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) + + +def _batch_response(batch_id: str, status: str): + from types import SimpleNamespace + + return SimpleNamespace(id=batch_id, status=status, object="batch") + + +@pytest.mark.asyncio +async def test_success_hook_persists_batch_enqueued_reservation_and_refunds_on_completion(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenOverLimit, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-enqueued-key", limit=100) + user = UserAPIKeyAuth(api_key="hashed-enqueued-key") + + reservation = await store.reserve(tokens=60, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "validating") + ) + assert get_request_stash().batch_enqueued_reservation is None + assert isinstance(await store.reserve(tokens=50, scopes=(scope,)), BatchEnqueuedTokenOverLimit) + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "completed") + ) + refill = await store.reserve(tokens=40, scopes=(scope,)) + assert isinstance(refill, BatchEnqueuedTokenReservation) + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "completed") + ) + assert isinstance(await store.reserve(tokens=70, scopes=(scope,)), BatchEnqueuedTokenOverLimit) + + +@pytest.mark.asyncio +async def test_success_hook_refunds_batch_enqueued_reservation_on_cancellation(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="team", value="team-enqueued", limit=100) + user = UserAPIKeyAuth(api_key="hashed-enqueued-key", team_id="team-enqueued") + + reservation = await store.reserve(tokens=90, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_2", "validating") + ) + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_2", "cancelling") + ) + assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_failure_hook_refunds_stashed_batch_enqueued_reservation(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-failing-key", limit=100) + user = UserAPIKeyAuth(api_key="hashed-failing-key") + + reservation = await store.reserve(tokens=80, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + + await handler.async_post_call_failure_hook( + request_data={}, original_exception=Exception("guardrail rejected"), user_api_key_dict=user + ) + assert get_request_stash().batch_enqueued_reservation is None + assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_success_hook_leaves_stash_untouched_for_non_batch_responses(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-chat-key", limit=100) + user = UserAPIKeyAuth(api_key="hashed-chat-key") + + reservation = await store.reserve(tokens=10, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=ModelResponse(usage=Usage(total_tokens=5)) + ) + assert get_request_stash().batch_enqueued_reservation == reservation