Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/batch_rate_limiter.py: 14%
400 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"""
2Batch Rate Limiter Hook
4This hook implements rate limiting for batch API requests by:
51. Reading batch input files to count requests and estimate tokens at submission
62. Validating actual usage from output files when batches complete
73. Integrating with the existing parallel request limiter infrastructure
9## Integration & Calling
10This hook is automatically registered and called by the proxy system.
11See BATCH_RATE_LIMITER_INTEGRATION.md for complete integration details.
13Quick summary:
14- Add to PROXY_HOOKS in litellm/proxy/hooks/__init__.py
15- Gets auto-instantiated on proxy startup via _add_proxy_hooks()
16- async_pre_call_hook() fires on POST /v1/batches (batch submission)
17- async_log_success_event() fires on GET /v1/batches/{id} (batch completion)
18"""
20import json
21from collections.abc import Callable, Iterable, Mapping, Sequence
22from datetime import datetime, timezone
23from types import MappingProxyType
24from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, TypeAlias
26from fastapi import HTTPException
27from pydantic import BaseModel, Field, TypeAdapter, ValidationError
29import litellm
30from litellm._logging import verbose_proxy_logger
31from litellm.batches.batch_utils import (
32 _count_entry_tokens,
33 _estimate_batch_entry_tokens,
34 _extract_file_access_credentials,
35 _iter_batch_input_lines,
36)
37from litellm.constants import BATCH_TPD_DESCRIPTOR_SUFFIX, BATCH_TPD_WINDOW_SECONDS
38from litellm.exceptions import RateLimitErrorCategory
39from litellm.integrations.custom_logger import CustomLogger
40from litellm.proxy._types import (
41 ProxyErrorTypes,
42 ProxyException,
43 SpecialModelNames,
44 UserAPIKeyAuth,
45)
46from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata
47from litellm.proxy.common_utils.proxy_rate_limit_error import (
48 ProxyRateLimitError,
49 map_v3_rate_limit_type,
50)
51from litellm.proxy.hooks.batch_enqueued_tokens import (
52 BatchEnqueuedTokenOverLimit,
53 BatchEnqueuedTokenReservation,
54 BatchEnqueuedTokenScope,
55 resolve_batch_enqueued_token_scopes,
56)
57from litellm.proxy.hooks.parallel_request_limiter_v3 import (
58 PROJECT_ITPM_DESCRIPTOR_KEY,
59 PROJECT_OTPM_DESCRIPTOR_KEY,
60 ReservationAwareIncrementOperation,
61 get_or_create_request_stash,
62)
63from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
65if TYPE_CHECKING: 65 ↛ 66line 65 didn't jump to line 66 because the condition on line 65 was never true
66 from opentelemetry.trace import Span as _Span
68 from litellm.caching.caching import DualCache
69 from litellm.proxy.hooks.parallel_request_limiter_v3 import (
70 RateLimitDescriptor as _RateLimitDescriptor,
71 )
72 from litellm.proxy.hooks.parallel_request_limiter_v3 import (
73 RateLimitStatus as _RateLimitStatus,
74 )
75 from litellm.proxy.hooks.parallel_request_limiter_v3 import (
76 _PROXY_MaxParallelRequestsHandler_v3 as _ParallelRequestLimiter,
77 )
78 from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
79 from litellm.router import Router as _Router
80 from litellm.types.llms.openai import HttpxBinaryResponseContent
82 Span = _Span
83 InternalUsageCache = _InternalUsageCache
84 Router = _Router
85 ParallelRequestLimiter = _ParallelRequestLimiter
86 RateLimitStatus = _RateLimitStatus
87 RateLimitDescriptor = _RateLimitDescriptor
88else:
89 Span = Any
90 InternalUsageCache = Any
91 Router = Any
92 ParallelRequestLimiter = Any
93 RateLimitStatus = dict[str, Any]
94 RateLimitDescriptor = dict[str, Any]
97_BATCH_BODY_ADAPTER: Final = TypeAdapter(dict[str, object])
98_WINDOW_START_ADAPTER: Final[TypeAdapter[int | float | str | None]] = TypeAdapter(int | float | str | None)
100IncrementAmounts: TypeAlias = dict[Literal["requests", "tokens"], int]
103class BatchFileUsage(BaseModel):
104 """
105 Internal model for batch file usage tracking, used for batch rate limiting
106 """
108 total_tokens: int
109 request_count: int
110 output_tokens: int = 0
111 # Keyed by each row's own `body.model`, distinct from `total_tokens`/
112 # `output_tokens` (the whole-file totals charged to the file-bound/
113 # top-level routing model's key/team/model limits). A batch's rows can
114 # each target a different model, so the project's per-model ITPM/OTPM
115 # quota for a row's actual model must be charged with that row's own
116 # tokens -- see `_create_project_io_descriptors_for_models`.
117 per_model_usage: dict[str, dict[str, int]] = Field(default_factory=dict)
120class _PROXY_BatchRateLimiter(CustomLogger):
121 """
122 Rate limiter for batch API requests.
124 Handles rate limiting at two points:
125 1. Batch submission - reads input file and reserves capacity
126 2. Batch completion - reads output file and adjusts for actual usage
127 """
129 def __init__(
130 self,
131 internal_usage_cache: InternalUsageCache,
132 parallel_request_limiter: ParallelRequestLimiter,
133 time_provider: Callable[[], datetime] | None = None,
134 ):
135 """
136 Initialize the batch rate limiter.
138 Note: These dependencies are automatically injected by ProxyLogging._add_proxy_hooks()
139 when this hook is registered in PROXY_HOOKS. See BATCH_RATE_LIMITER_INTEGRATION.md.
141 Args:
142 internal_usage_cache: Cache for storing rate limit data (auto-injected)
143 parallel_request_limiter: Existing rate limiter to integrate with (needs custom injection)
144 time_provider: Clock used for rate limit reset times (defaults to ``datetime.now``)
145 """
146 self.internal_usage_cache = internal_usage_cache
147 self.parallel_request_limiter = parallel_request_limiter
148 self._time_provider: Final = time_provider or datetime.now
149 self._warned_unsupported_model_skip = False
151 def _get_file_bound_batch_model(self, data: dict) -> str | None:
152 """Resolve the model bound to the batch input file ID.
154 ``create_batch`` routes a file-bound id (model-embedded ``file-...`` or
155 unified managed file) on that bound model and ignores the top-level
156 ``model``, so this is the authoritative routing model whenever the file
157 binds one. The provider is then read from that deployment's trusted
158 credentials for the provider-level skip decision.
159 """
160 input_file_id: Final = data.get("input_file_id")
161 if not isinstance(input_file_id, str) or not input_file_id:
162 return None
164 from litellm.proxy.openai_files_endpoints.common_utils import (
165 _is_base64_encoded_unified_file_id,
166 decode_model_from_file_id,
167 get_models_from_unified_file_id,
168 )
170 model_from_file_id: Final = decode_model_from_file_id(input_file_id)
171 if model_from_file_id:
172 return model_from_file_id
174 unified_file_id: Final = _is_base64_encoded_unified_file_id(input_file_id)
175 if unified_file_id:
176 target_model_names: Final = get_models_from_unified_file_id(unified_file_id)
177 if target_model_names:
178 return target_model_names[0]
180 return None
182 def _get_batch_routing_model(self, data: dict) -> str | None:
183 """Resolve the deployment/model used for this batch from request data.
185 Mirrors ``create_batch`` routing precedence: a model bound to the input
186 file id wins over the top-level ``model``, because the batch endpoint
187 ignores the top-level model for file-bound ids. Resolving the provider
188 skip from the top-level model first would let a caller point ``model``
189 at a skip-listed provider while the file routes a rate-limited one.
190 """
191 file_bound_model: Final = self._get_file_bound_batch_model(data)
192 if file_bound_model:
193 return file_bound_model
195 model: Final = data.get("model")
196 if isinstance(model, str) and model:
197 return model
199 return None
201 def _resolve_batch_provider(self, batch_model: str | None) -> str | None:
202 """Resolve the provider from the deployment that serves ``batch_model``.
204 The provider is read from trusted router credentials rather than the
205 user-supplied ``custom_llm_provider`` request field, so a caller cannot
206 spoof a skip-listed provider to bypass batch rate limiting.
207 """
208 if not batch_model:
209 return None
211 from litellm.proxy.openai_files_endpoints.common_utils import (
212 get_credentials_for_model,
213 )
214 from litellm.proxy.proxy_server import llm_router
216 if llm_router is None:
217 return None
219 try:
220 credentials: Final = get_credentials_for_model(
221 llm_router=llm_router,
222 model_id=batch_model,
223 operation_context="batch input file read (rate limiting)",
224 )
225 except HTTPException:
226 return None
228 provider: Final = credentials.get("custom_llm_provider")
229 return provider if isinstance(provider, str) and provider else None
231 def _create_batch_rate_limit_descriptors(
232 self,
233 user_api_key_dict: UserAPIKeyAuth,
234 data: dict,
235 ) -> list["RateLimitDescriptor"]:
236 """Build the standard key/user/team/model descriptor list a batch is charged against.
238 Deliberately excludes the project-scoped ITPM/OTPM descriptors: those
239 are charged per the JSONL row's own `body.model` once the file is
240 parsed (`_create_project_io_descriptors_for_models`), not the
241 file-bound/top-level routing model this function resolves. Charging
242 project quotas here would let a caller bind the file to a model
243 without a quota while rows execute against a quota-limited model.
245 Scopes with a ``tpd_limit`` (key, team, end user) are charged against a
246 daily token descriptor instead of their per-minute RPM/TPM descriptor,
247 because a batch's rows are scheduled by the provider and never share a
248 minute with the submission. The daily descriptor uses its own key so
249 its 24h window never collides with the online limiter's counters.
250 """
251 descriptors: Final = self.parallel_request_limiter._create_rate_limit_descriptors(
252 user_api_key_dict=user_api_key_dict,
253 data=data,
254 rpm_limit_type=None,
255 tpm_limit_type=None,
256 model_has_failures=False,
257 )
258 tpd_limits: Final[Mapping[str, tuple[str, int]]] = MappingProxyType(
259 {
260 key: (value, limit)
261 for key, value, limit in (
262 ("api_key", user_api_key_dict.api_key, user_api_key_dict.tpd_limit),
263 ("team", user_api_key_dict.team_id, user_api_key_dict.team_tpd_limit),
264 ("end_user", user_api_key_dict.end_user_id, user_api_key_dict.end_user_tpd_limit),
265 )
266 if value and limit is not None
267 }
268 )
269 if not tpd_limits:
270 return descriptors
271 return [
272 *(d for d in descriptors if d["key"] not in tpd_limits),
273 *(
274 RateLimitDescriptor(
275 key=f"{key}{BATCH_TPD_DESCRIPTOR_SUFFIX}",
276 value=value,
277 rate_limit={
278 "requests_per_unit": None,
279 "tokens_per_unit": limit,
280 "window_size": BATCH_TPD_WINDOW_SECONDS,
281 },
282 )
283 for key, (value, limit) in tpd_limits.items()
284 ),
285 ]
287 @staticmethod
288 def _project_has_any_io_token_limits(user_api_key_dict: UserAPIKeyAuth) -> bool:
289 """True when the project has any per-model ITPM/OTPM quota configured.
291 Used to stop the "skip batch input file processing" fast path from
292 bypassing a project quota configured for a model other than the
293 batch's file-bound/top-level routing model: the row models that
294 actually drive execution and billing aren't known until the JSONL
295 is parsed, so the file must be read whenever *any* model could be
296 quota-limited, not only when the routing model itself is.
297 """
298 if user_api_key_dict.project_id is None:
299 return False
300 return bool(
301 get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_itpm_limit")
302 ) or bool(get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_otpm_limit"))
304 def _create_project_io_descriptors_for_models(
305 self,
306 user_api_key_dict: UserAPIKeyAuth,
307 per_model_usage: Mapping[str, Mapping[str, int]],
308 ) -> tuple[list["RateLimitDescriptor"], list[IncrementAmounts]]: # mutable-ok: see below
309 """Build project ITPM/OTPM descriptors charged against each row's own model.
311 One descriptor pair per distinct `body.model` found in the JSONL,
312 each incremented only by that model's own counted usage -- never the
313 whole-batch total -- so a quota-limited model can't hide behind an
314 unlimited routing model, and an unrelated model's rows can't inflate
315 a different model's counter.
316 """
317 extra_descriptors: Final[list[RateLimitDescriptor]] = [] # mutable-ok: see above
318 extra_increments: Final[list[IncrementAmounts]] = [] # mutable-ok: see above
319 for model, usage in per_model_usage.items():
320 model_descriptors: list[RateLimitDescriptor] = [] # mutable-ok: reset per loop iteration, not module state
321 self.parallel_request_limiter.add_project_io_token_rate_limit_descriptors_from_metadata(
322 user_api_key_dict=user_api_key_dict,
323 requested_model=model,
324 descriptors=model_descriptors,
325 )
326 for descriptor in model_descriptors:
327 extra_descriptors.append(descriptor)
328 extra_increments.append(
329 { # mutable-ok: atomic limiter API requires mutable increment records
330 "requests": 0,
331 "tokens": usage.get("output_tokens", 0)
332 if descriptor["key"] == PROJECT_OTPM_DESCRIPTOR_KEY
333 else usage.get("total_tokens", 0),
334 }
335 )
336 return extra_descriptors, extra_increments
338 def _should_skip_batch_input_file_processing(
339 self,
340 data: dict,
341 user_api_key_dict: UserAPIKeyAuth,
342 has_enqueued_scopes: bool = False,
343 ) -> tuple[bool, list["RateLimitDescriptor"] | None]:
344 """
345 Skip downloading batch input files when the operator disabled batch
346 input-file rate limiting, when the batch runs entirely on a skip-listed
347 provider, or when there is nothing to enforce (no applicable rate
348 limits).
350 A skip is only honored for keys with unrestricted model access. When
351 the key has a model allowlist, the JSONL must still be downloaded so
352 ``_enforce_batch_file_model_access`` can validate every ``body.model``
353 entry, otherwise a restricted key could smuggle unauthorized models
354 into the file via an admin-configured skip.
356 The skip is never keyed on a specific model name. The models a batch
357 actually runs are its JSONL ``body.model`` entries, and any model
358 identifier the caller can influence (the top-level ``model`` or the
359 unsigned model embedded in a ``file-...`` id) can be pointed at a
360 skip-listed deployment while the file routes a different, rate-limited
361 model. The provider skip is safe because the provider is read from the
362 routing deployment's trusted credentials and the batch is constrained
363 to run on that provider.
365 The no-limits check also treats any project-configured ITPM/OTPM
366 quota as an applicable limit, even when it isn't scoped to the
367 routing model: a row can target a different, quota-limited model,
368 and that isn't knowable without parsing the JSONL.
370 Returns ``(should_skip, descriptors)`` where ``descriptors`` is the
371 rate-limit descriptor list computed for the no-limits check, so the
372 caller can reuse it for counter enforcement without recomputing.
373 """
374 from litellm.proxy.proxy_server import general_settings
376 self._warn_if_unsupported_model_skip_configured(general_settings)
378 if self._key_requires_batch_model_access_check(user_api_key_dict):
379 return False, None
381 if general_settings.get("disable_batch_input_file_rate_limiting") is True:
382 return True, None
384 skip_providers: Final = general_settings.get("skip_batch_input_file_rate_limiting_for_providers") or []
385 if skip_providers:
386 batch_provider: Final = self._resolve_batch_provider(self._get_batch_routing_model(data))
387 if batch_provider and batch_provider in skip_providers:
388 verbose_proxy_logger.debug("Skipping batch input file processing for provider=%s", batch_provider)
389 return True, None
391 descriptors: Final = self._create_batch_rate_limit_descriptors(
392 user_api_key_dict=user_api_key_dict,
393 data=data,
394 )
395 if (
396 not has_enqueued_scopes
397 and not self._has_applicable_batch_rate_limits(descriptors)
398 and not self._project_has_any_io_token_limits(user_api_key_dict)
399 ):
400 verbose_proxy_logger.debug("Skipping batch input file processing: no rate limits configured")
401 return True, None
403 return False, descriptors
405 def _warn_if_unsupported_model_skip_configured(self, general_settings: dict) -> None:
406 """Warn once that ``skip_batch_input_file_rate_limiting_for_models`` is a no-op.
408 A per-model skip is intentionally not honored because the model a batch
409 runs on is caller-influenced and can be pointed at a skip-listed
410 deployment while the JSONL routes a different, rate-limited model.
411 """
412 if self._warned_unsupported_model_skip:
413 return
414 if general_settings.get("skip_batch_input_file_rate_limiting_for_models"):
415 self._warned_unsupported_model_skip = True
416 verbose_proxy_logger.warning(
417 "general_settings.skip_batch_input_file_rate_limiting_for_models is not "
418 "supported and has no effect. Use "
419 "skip_batch_input_file_rate_limiting_for_providers or "
420 "disable_batch_input_file_rate_limiting instead."
421 )
423 @staticmethod
424 def _key_requires_batch_model_access_check(
425 user_api_key_dict: UserAPIKeyAuth,
426 ) -> bool:
427 """True when the key may only call a subset of models (JSONL must be checked)."""
428 models: Final = user_api_key_dict.models or []
429 if "*" in models:
430 return False
431 if SpecialModelNames.all_proxy_models.value in models:
432 return False
433 if user_api_key_dict.access_group_ids:
434 return True
435 if not models:
436 return False
437 return True
439 def _estimate_entry_output_tokens(
440 self,
441 entry: Mapping[str, object],
442 min_configured_otpm_limit: int | None,
443 ) -> int:
444 """Conservative per-row output-token estimate for the project OTPM reservation.
446 Batch completion never reconciles actual usage back into the rate
447 limiter, so this pre-call estimate is the only OTPM enforcement a
448 batch gets. Mirrors the real-time no-``max_tokens`` floor so a row
449 that omits an output cap can't be used to bypass OTPM the way an
450 unbounded streaming request could.
452 Embeddings rows are identified by the row's own ``url`` (the OpenAI
453 batch schema puts the target route there, e.g. ``/v1/embeddings``),
454 never by body shape: a `/v1/responses` row also carries `body.input`
455 with no `messages`/`prompt`, so guessing from body shape alone would
456 misclassify a token-generating Responses row as a zero-output
457 embeddings row and let it skip the OTPM reservation entirely.
458 """
459 url: Final = entry.get("url")
460 if isinstance(url, str) and "embeddings" in url:
461 return 0 # embeddings: no output tokens
462 raw_body: Final = entry.get("body")
463 body: Final[Mapping[str, object]] = (
464 MappingProxyType(_BATCH_BODY_ADAPTER.validate_python(raw_body))
465 if isinstance(raw_body, Mapping)
466 else MappingProxyType({})
467 )
468 # `max_tokens`/`max_completion_tokens` cap chat completions; `/v1/responses`
469 # rows cap output with `max_output_tokens` instead -- omitting it here
470 # would fall through to the floor estimate for every capped Responses row.
471 explicit_cap: Final = next(
472 (
473 v
474 for v in (
475 body.get("max_tokens"),
476 body.get("max_completion_tokens"),
477 body.get("max_output_tokens"),
478 )
479 if v is not None
480 ),
481 None,
482 )
483 candidate_count: Final = self.parallel_request_limiter.get_output_candidate_count(body)
484 if explicit_cap is not None:
485 try:
486 return max(0, int(explicit_cap)) * candidate_count
487 except (TypeError, ValueError, OverflowError):
488 pass
489 return self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) * candidate_count
491 @staticmethod
492 def _has_applicable_batch_rate_limits(
493 descriptors: list["RateLimitDescriptor"],
494 ) -> bool:
495 for descriptor in descriptors:
496 rate_limit = descriptor.get("rate_limit") or {}
497 if (
498 rate_limit.get("requests_per_unit") is not None
499 or rate_limit.get("tokens_per_unit") is not None
500 or rate_limit.get("max_parallel_requests") is not None
501 ):
502 return True
503 return False
505 def _resolve_batch_input_file_fetch_params(
506 self,
507 file_id: str,
508 custom_llm_provider: str,
509 data: dict,
510 ) -> tuple[str, dict[str, Any]]:
511 """
512 Map proxy-facing file IDs to provider file IDs and credentials.
514 Model-embedded IDs (``file-<base64>``) are not unified managed-file IDs;
515 without decoding them, ``afile_content`` is called with the encoded ID
516 and the upstream provider returns 404.
517 """
518 from litellm.proxy.openai_files_endpoints.common_utils import (
519 decode_model_from_file_id,
520 get_credentials_for_model,
521 get_original_file_id,
522 )
523 from litellm.proxy.proxy_server import llm_router
525 fetch_kwargs: Final[dict[str, Any]] = {
526 "custom_llm_provider": custom_llm_provider,
527 }
529 model_from_file_id: Final = decode_model_from_file_id(file_id)
530 if model_from_file_id:
531 if llm_router is not None:
532 try:
533 credentials = get_credentials_for_model(
534 llm_router=llm_router,
535 model_id=model_from_file_id,
536 operation_context="batch input file read (rate limiting)",
537 )
538 fetch_kwargs.update(_extract_file_access_credentials(credentials))
539 fetch_kwargs["model"] = model_from_file_id
540 provider = credentials.get("custom_llm_provider")
541 if provider:
542 fetch_kwargs["custom_llm_provider"] = provider
543 except HTTPException:
544 pass
545 return get_original_file_id(file_id), fetch_kwargs
547 request_model: Final = data.get("model")
548 if isinstance(request_model, str) and request_model and llm_router is not None:
549 try:
550 credentials = get_credentials_for_model(
551 llm_router=llm_router,
552 model_id=request_model,
553 operation_context="batch input file read (rate limiting)",
554 )
555 fetch_kwargs.update(_extract_file_access_credentials(credentials))
556 fetch_kwargs["model"] = request_model
557 provider = credentials.get("custom_llm_provider")
558 if provider:
559 fetch_kwargs["custom_llm_provider"] = provider
560 except HTTPException:
561 pass
563 return file_id, fetch_kwargs
565 async def _reserve_batch_enqueued_tokens(
566 self,
567 user_api_key_dict: UserAPIKeyAuth,
568 data: Mapping[str, object],
569 batch_usage: BatchFileUsage,
570 scopes: tuple[BatchEnqueuedTokenScope, ...],
571 ) -> None:
572 """Reserve the batch's estimated tokens against the caller's enqueued-token allowance.
574 Runs instead of the per-minute counter charge when the key or team
575 opted in via ``batch_enqueued_token_limit`` metadata. The reservation
576 is stashed on the request so the v3 limiter's post-call hooks can
577 persist it (keyed by the provider batch id) and refund it when the
578 batch reaches a terminal state.
579 """
580 outcome: Final = await self.parallel_request_limiter.batch_enqueued_token_store.reserve(
581 tokens=batch_usage.total_tokens,
582 scopes=scopes,
583 litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
584 )
585 match outcome:
586 case BatchEnqueuedTokenOverLimit():
587 self._raise_enqueued_limit_error(over_limit=outcome, data=data, batch_usage=batch_usage)
588 case BatchEnqueuedTokenReservation():
589 get_or_create_request_stash().batch_enqueued_reservation = outcome
591 def _raise_enqueued_limit_error(
592 self,
593 over_limit: BatchEnqueuedTokenOverLimit,
594 data: Mapping[str, object],
595 batch_usage: BatchFileUsage,
596 ) -> NoReturn:
597 scope: Final = over_limit.scope
598 remaining: Final = max(0, scope.limit - over_limit.enqueued)
599 detail: Final = (
600 f"Batch enqueued token limit exceeded for {scope.key}: {scope.value}. "
601 f"Batch requires {batch_usage.total_tokens} tokens but only {remaining} enqueued tokens remaining "
602 f"out of {scope.limit} enqueued token limit. "
603 f"Tokens free up as running batches complete or are cancelled."
604 )
605 raw_model: Final = data.get("model")
606 resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
607 raw_model if isinstance(raw_model, str) else None
608 )
609 raise ProxyRateLimitError(
610 detail=detail,
611 headers=MappingProxyType({"rate_limit_type": "tokens"}),
612 category=RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT,
613 rate_limit_type=map_v3_rate_limit_type("tokens"),
614 model=resolved_model,
615 llm_provider=llm_provider,
616 )
618 def _raise_rate_limit_error(
619 self,
620 status: "RateLimitStatus",
621 descriptors: list["RateLimitDescriptor"],
622 batch_usage: BatchFileUsage,
623 limit_type: str,
624 requested_model: str | None = None,
625 window_start: int | None = None,
626 ) -> NoReturn:
627 """Raise :class:`ProxyRateLimitError` (a 429) for batch rate limit exceeded.
629 ``window_start`` is the active counter window's start (unix seconds) when
630 known, so the reset time reflects that window's actual end rather than a
631 full window from now.
632 """
634 # Find the descriptor for this status. Matching on (key, value) is
635 # required, not key alone: a batch can carry several project ITPM/OTPM
636 # descriptors sharing one key (e.g. `model_per_project_otpm`) but
637 # scoped to different models via `value`
638 # ("{project_id}:{model}") -- key-only matching would always resolve
639 # to the first same-keyed descriptor regardless of which one was
640 # actually over its limit. Falls back to key-only matching for
641 # statuses that predate `descriptor_value` (e.g. from should_rate_limit).
642 status_descriptor_value: Final = status.get("descriptor_value")
643 descriptor_index: Final = next(
644 (
645 i
646 for i, d in enumerate(descriptors)
647 if d.get("key") == status.get("descriptor_key")
648 and (status_descriptor_value is None or d.get("value") == status_descriptor_value)
649 ),
650 0,
651 )
652 descriptor: Final[RateLimitDescriptor] = (
653 descriptors[descriptor_index] if descriptors else {"key": "", "value": "", "rate_limit": None}
654 )
656 now: Final = self._time_provider().timestamp()
657 window_size: Final = (descriptor.get("rate_limit") or {}).get(
658 "window_size"
659 ) or self.parallel_request_limiter.window_size
660 reset_time: Final = now + window_size if window_start is None else window_start + window_size
661 retry_after: Final = max(0, int(reset_time - now))
662 reset_time_formatted: Final = datetime.fromtimestamp(reset_time, tz=timezone.utc).strftime(
663 "%Y-%m-%d %H:%M:%S UTC"
664 )
666 remaining_display: Final = max(0, status["limit_remaining"])
667 current_limit: Final = status["current_limit"]
669 if limit_type == "requests":
670 detail = (
671 f"Batch rate limit exceeded for {descriptor.get('key', 'unknown')}: {descriptor.get('value', 'unknown')}. "
672 f"Batch contains {batch_usage.request_count} requests but only {remaining_display} requests remaining "
673 f"out of {current_limit} RPM limit. "
674 f"Limit resets at: {reset_time_formatted}"
675 )
676 else: # tokens
677 # Project ITPM/OTPM descriptors are keyed "{project_id}:{model}" and
678 # charged with that model's own rows (see
679 # `_create_project_io_descriptors_for_models`), not the whole
680 # batch's totals -- report the matching per-model figure when one
681 # is available so the error reflects what was actually charged.
682 descriptor_model: Final = (
683 descriptor.get("value", "").split(":", 1)[-1]
684 if descriptor.get("key") in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY)
685 else None
686 )
687 model_usage: Final = batch_usage.per_model_usage.get(descriptor_model) if descriptor_model else None
688 batch_token_count: Final = (
689 (model_usage or {}).get("output_tokens", batch_usage.output_tokens)
690 if descriptor.get("key") == PROJECT_OTPM_DESCRIPTOR_KEY
691 else (model_usage or {}).get("total_tokens", batch_usage.total_tokens)
692 if descriptor.get("key") == PROJECT_ITPM_DESCRIPTOR_KEY
693 else batch_usage.total_tokens
694 )
695 token_limit_label: Final = (
696 "TPD" if descriptor.get("key", "").endswith(BATCH_TPD_DESCRIPTOR_SUFFIX) else "TPM"
697 )
698 detail = (
699 f"Batch rate limit exceeded for {descriptor.get('key', 'unknown')}: {descriptor.get('value', 'unknown')}. "
700 f"Batch contains {batch_token_count} tokens but only {remaining_display} tokens remaining "
701 f"out of {current_limit} {token_limit_label} limit. "
702 f"Limit resets at: {reset_time_formatted}"
703 )
705 resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(requested_model)
706 raise ProxyRateLimitError(
707 detail=detail,
708 headers={
709 "retry-after": str(retry_after),
710 "rate_limit_type": limit_type,
711 "reset_at": reset_time_formatted,
712 },
713 category=RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT,
714 rate_limit_type=map_v3_rate_limit_type(limit_type),
715 model=resolved_model,
716 llm_provider=llm_provider,
717 )
719 async def _check_and_increment_batch_counters(
720 self,
721 user_api_key_dict: UserAPIKeyAuth,
722 data: dict,
723 batch_usage: BatchFileUsage,
724 descriptors: list["RateLimitDescriptor"] | None = None,
725 ) -> None:
726 """
727 Atomically check + increment rate-limit counters by the batch amounts.
729 Raises HTTPException if any descriptor would exceed its limit; in that
730 case no counter is modified. Backed by `atomic_check_and_increment_by_n`
731 which uses a Redis Lua script when available (multi-process atomic) and
732 falls back to a per-process asyncio.Lock + in-memory operation.
734 ``descriptors`` may be passed in by the pre-call hook to reuse the list
735 already computed when deciding whether to skip file processing. It
736 never contains project ITPM/OTPM descriptors (those are model-specific
737 and only knowable once ``batch_usage.per_model_usage`` is populated by
738 parsing the JSONL), so this always builds and appends them here.
739 """
740 if descriptors is None:
741 descriptors = self._create_batch_rate_limit_descriptors(
742 user_api_key_dict=user_api_key_dict,
743 data=data,
744 )
746 increments: list[IncrementAmounts] = [ # mutable-ok: reassigned below to append project IO increments
747 { # mutable-ok: atomic limiter API requires mutable increment records
748 "requests": batch_usage.request_count,
749 "tokens": batch_usage.total_tokens,
750 }
751 for _d in descriptors
752 ]
754 project_io_descriptors, project_io_increments = self._create_project_io_descriptors_for_models(
755 user_api_key_dict=user_api_key_dict,
756 per_model_usage=batch_usage.per_model_usage,
757 )
758 descriptors = [*descriptors, *project_io_descriptors]
759 increments = [*increments, *project_io_increments]
761 rate_limit_response: Final = await self.parallel_request_limiter.atomic_check_and_increment_by_n(
762 descriptors=descriptors,
763 increments=increments,
764 parent_otel_span=user_api_key_dict.parent_otel_span,
765 )
767 stash: Final = get_or_create_request_stash()
768 stash.batch_tpd_refund_ops = ()
769 if rate_limit_response["overall_code"] == "OVER_LIMIT":
770 requested_model: Final = data.get("model") if data else None
771 for status in rate_limit_response["statuses"]:
772 if status["code"] == "OVER_LIMIT":
773 self._raise_rate_limit_error(
774 status,
775 descriptors,
776 batch_usage,
777 status["rate_limit_type"],
778 requested_model=requested_model,
779 window_start=await self._read_tpd_window_start(
780 status=status, parent_otel_span=user_api_key_dict.parent_otel_span
781 ),
782 )
784 stash.batch_tpd_refund_ops = self._build_tpd_refund_ops(
785 descriptors=descriptors,
786 tokens=batch_usage.total_tokens,
787 reservation_windows=rate_limit_response.get("reservation_windows", frozenset()),
788 )
790 async def _read_tpd_window_start(self, status: "RateLimitStatus", parent_otel_span: "Span | None") -> int | None:
791 descriptor_key: Final = status.get("descriptor_key") or ""
792 if not descriptor_key.endswith(BATCH_TPD_DESCRIPTOR_SUFFIX):
793 return None
794 try:
795 window_start: Final = _WINDOW_START_ADAPTER.validate_python(
796 await self.parallel_request_limiter.internal_usage_cache.async_get_cache(
797 key=f"{{{descriptor_key}:{status.get('descriptor_value') or ''}}}:window",
798 litellm_parent_otel_span=parent_otel_span,
799 ),
800 strict=True,
801 )
802 return None if window_start is None else int(float(window_start))
803 except (ValidationError, ValueError):
804 return None
806 def _build_tpd_refund_ops(
807 self,
808 descriptors: Sequence["RateLimitDescriptor"],
809 tokens: int,
810 reservation_windows: frozenset[tuple[str, str, Literal["redis", "local"]]],
811 ) -> tuple[ReservationAwareIncrementOperation, ...]:
812 """Refund operations for the daily token counters this batch charged.
814 The v3 limiter's failure hook applies them when the submission fails
815 after the counters were incremented. Each operation carries the window
816 identity the charge landed in, so the refund is skipped once that
817 window has rolled over.
818 """
819 if tokens <= 0 or not reservation_windows:
820 return ()
821 tpd_descriptors_by_counter: Final[Mapping[str, RateLimitDescriptor]] = MappingProxyType(
822 {
823 self.parallel_request_limiter.create_rate_limit_keys(
824 descriptor["key"], descriptor["value"], "tokens"
825 ): descriptor
826 for descriptor in descriptors
827 if descriptor["key"].endswith(BATCH_TPD_DESCRIPTOR_SUFFIX)
828 }
829 )
830 return tuple(
831 ReservationAwareIncrementOperation(
832 key=counter_key,
833 increment_value=-tokens,
834 ttl=BATCH_TPD_WINDOW_SECONDS,
835 window_key=f"{{{descriptor['key']}:{descriptor['value']}}}:window",
836 expected_window_start=window_start,
837 reservation_backend=backend,
838 )
839 for counter_key, window_start, backend in sorted(reservation_windows)
840 if (descriptor := tpd_descriptors_by_counter.get(counter_key)) is not None
841 )
843 async def count_input_file_usage(
844 self,
845 file_id: str,
846 custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
847 user_api_key_dict: UserAPIKeyAuth | None = None,
848 data: dict | None = None,
849 descriptors: Sequence["RateLimitDescriptor"] | None = None,
850 ) -> BatchFileUsage:
851 """
852 Count number of requests and tokens in a batch input file.
854 Args:
855 file_id: The file ID to read
856 custom_llm_provider: The custom LLM provider to use for token encoding
857 user_api_key_dict: User authentication information for file access (required for managed files)
858 descriptors: Rate limit descriptors already computed for this batch, so the
859 configured project OTPM limit can scale the no-``max_tokens`` output floor
861 Returns:
862 BatchFileUsage with total_tokens, output_tokens, request_count, and
863 per_model_usage (each row's own totals, keyed by its `body.model`)
864 """
865 descriptor_otpm_limits: Final = tuple(
866 int(v)
867 for d in (descriptors or ())
868 if d.get("key") == PROJECT_OTPM_DESCRIPTOR_KEY
869 for rate_limit in (d.get("rate_limit"),)
870 for v in (rate_limit.get("tokens_per_unit") if rate_limit is not None else None,)
871 if v is not None
872 )
873 # `descriptors` only ever carries the routing model's own OTPM limit
874 # (see `_create_batch_rate_limit_descriptors`), but a row can target
875 # any project-configured model. Folding in every configured model's
876 # OTPM limit keeps the no-`max_tokens` floor from drifting wide just
877 # because a row's specific model isn't known until parsed below.
878 project_otpm_limits: Final = (
879 tuple(int(v) for v in project_otpm_limit_map.values())
880 if user_api_key_dict is not None
881 and (
882 project_otpm_limit_map := get_model_rate_limit_from_metadata(
883 user_api_key_dict, "project_metadata", "model_otpm_limit"
884 )
885 )
886 else ()
887 )
888 min_configured_otpm_limit: Final = min((*descriptor_otpm_limits, *project_otpm_limits), default=None)
889 try:
890 # Check if this is a managed file (base64 encoded unified file ID)
891 from litellm.proxy.openai_files_endpoints.common_utils import (
892 _is_base64_encoded_unified_file_id,
893 get_models_from_unified_file_id,
894 )
896 # Managed files require bypassing the HTTP endpoint (which runs access-check hooks)
897 # and calling the managed files hook directly with the user's credentials.
898 is_managed_file: Final = _is_base64_encoded_unified_file_id(file_id)
899 # For managed files the unified file id encodes the proxy model
900 # alias(es) the file was uploaded for; auth validates against those.
901 target_model_names: Final = get_models_from_unified_file_id(is_managed_file) if is_managed_file else []
902 if is_managed_file and user_api_key_dict is not None:
903 file_content = await self._fetch_managed_file_content(
904 file_id=file_id,
905 user_api_key_dict=user_api_key_dict,
906 )
907 else:
908 provider_file_id, fetch_kwargs = self._resolve_batch_input_file_fetch_params(
909 file_id=file_id,
910 custom_llm_provider=custom_llm_provider,
911 data=data or {},
912 )
913 # For non-managed files, use the standard litellm.afile_content
914 file_content = await litellm.afile_content(
915 file_id=provider_file_id,
916 user_api_key_dict=user_api_key_dict,
917 **fetch_kwargs,
918 )
920 file_content_bytes: Final = getattr(file_content, "content", None)
921 if not isinstance(file_content_bytes, bytes):
922 raise ValueError(
923 f"Expected bytes content from file retrieval for {file_id}, got {type(file_content_bytes)}"
924 )
926 # Single streaming pass over the JSONL lines, accounting each row
927 # independently. One bad row can never abort the pass: a malformed
928 # line is skipped (its request can't run upstream anyway) and a row
929 # the token counter can't measure falls back to a conservative
930 # size-based estimate. This guarantees two things a restricted caller
931 # must not be able to break by crafting a row that raises:
932 # 1. The allowlist check below always sees every parseable
933 # ``body.model`` (the loop never stops early), so models can't be
934 # smuggled in after a bad row.
935 # 2. The token total is never silently zeroed, so the TPM limit
936 # can't be evaded by sending uncountable rows.
937 # Counting stays best-effort, so a legitimate (e.g. multimodal) row
938 # the counter can't measure is estimated, not hard-rejected.
939 models: Final[set] = set()
940 # Keyed by each row's own `body.model`, so the project ITPM/OTPM
941 # quota for that model is charged with only its own rows' tokens,
942 # never the whole batch's -- see `_create_project_io_descriptors_for_models`.
943 per_model_usage: Final[dict[str, dict[str, int]]] = {}
944 total_tokens = 0
945 output_tokens = 0 # rebind-ok: accumulated per JSONL row in the loop below
946 request_count = 0
947 for raw_line in _iter_batch_input_lines(file_content_bytes):
948 request_count += 1
949 try:
950 entry = json.loads(raw_line)
951 except Exception:
952 entry_total_tokens = _estimate_batch_entry_tokens(raw_line)
953 entry_output_tokens = self.parallel_request_limiter.no_max_tokens_output_floor(
954 min_configured_otpm_limit
955 )
956 total_tokens += entry_total_tokens
957 output_tokens += entry_output_tokens
958 continue
960 model: str | None = (entry.get("body") or {}).get("model") if isinstance(entry, dict) else None
961 if model:
962 models.add(model)
964 if isinstance(entry, dict):
965 entry_output_tokens = self._estimate_entry_output_tokens(entry, min_configured_otpm_limit)
966 else:
967 entry_output_tokens = self.parallel_request_limiter.no_max_tokens_output_floor(
968 min_configured_otpm_limit
969 )
970 output_tokens += entry_output_tokens
972 try:
973 entry_total_tokens = _count_entry_tokens(entry)
974 except Exception:
975 entry_total_tokens = _estimate_batch_entry_tokens(raw_line)
976 total_tokens += entry_total_tokens
978 if model:
979 model_usage = per_model_usage.setdefault(
980 model, {"total_tokens": 0, "output_tokens": 0, "request_count": 0}
981 )
982 model_usage["total_tokens"] += entry_total_tokens
983 model_usage["output_tokens"] += entry_output_tokens
984 model_usage["request_count"] += 1
986 # Validate every model named in the batch JSONL against the
987 # caller's per-key model allowlist. Without this, a caller
988 # could smuggle restricted/expensive models inside the file
989 # and the upstream provider would execute the batch under
990 # the proxy's shared API key.
991 if user_api_key_dict is not None:
992 await self._enforce_batch_file_model_access(
993 user_api_key_dict=user_api_key_dict,
994 models=models,
995 target_model_names=target_model_names or None,
996 )
998 return BatchFileUsage(
999 total_tokens=total_tokens,
1000 request_count=request_count,
1001 output_tokens=output_tokens,
1002 per_model_usage=per_model_usage,
1003 )
1005 except HTTPException as e:
1006 # Distinguish intentional 403s from `_enforce_batch_file_model_access`
1007 # from genuine I/O failures so security-relevant rejections show up
1008 # in the access log instead of getting buried in error noise.
1009 if e.status_code == 403:
1010 verbose_proxy_logger.warning(
1011 "Batch rejected: caller not authorized for a model named in %s: %s", file_id, e.detail
1012 )
1013 else:
1014 verbose_proxy_logger.error(
1015 "Batch input file rejected for %s: status=%s detail=%s", file_id, e.status_code, e.detail
1016 )
1017 raise
1018 except Exception as e:
1019 verbose_proxy_logger.error("Error counting input file usage for %s: %s", file_id, e)
1020 raise
1022 async def _enforce_batch_file_model_access(
1023 self,
1024 user_api_key_dict: UserAPIKeyAuth,
1025 models: Iterable[str] | None = None,
1026 target_model_names: list[str] | None = None,
1027 ) -> None:
1028 """Reject the batch if the caller is not authorized for the upload target.
1030 For managed files, ``target_model_names`` (from the unified file id) is
1031 the proxy alias the file was uploaded for and is checked directly.
1032 Otherwise the ``body.model`` values collected from the JSONL (``models``)
1033 are checked.
1035 Reuses standard auth helpers so the same model access rules the proxy
1036 enforces on `/chat/completions` apply here.
1037 """
1038 from litellm.proxy.auth.auth_checks import (
1039 _check_team_member_model_access,
1040 _key_access_group_grants_model,
1041 can_key_call_model,
1042 can_team_access_model,
1043 get_team_object,
1044 )
1045 from litellm.proxy.proxy_server import llm_router, prisma_client, proxy_logging_obj, user_api_key_cache
1047 if target_model_names:
1048 models = target_model_names
1050 if not models:
1051 return
1053 team_object = None
1054 if (
1055 SpecialModelNames.all_team_models.value in (user_api_key_dict.models or [])
1056 and user_api_key_dict.team_id is not None
1057 and prisma_client is not None
1058 ):
1059 try:
1060 team_object = await get_team_object(
1061 team_id=user_api_key_dict.team_id,
1062 prisma_client=prisma_client,
1063 user_api_key_cache=user_api_key_cache,
1064 parent_otel_span=user_api_key_dict.parent_otel_span,
1065 proxy_logging_obj=proxy_logging_obj,
1066 )
1067 except HTTPException:
1068 raise
1069 except Exception as e:
1070 raise HTTPException(
1071 status_code=403,
1072 detail={
1073 "error": ("Batch input file model access could not be validated against the current team.")
1074 },
1075 ) from e
1077 llm_model_list: Final = llm_router.model_list if llm_router is not None else None
1078 for model in models:
1079 model_to_check = model
1080 try:
1081 if team_object is not None:
1082 try:
1083 await can_team_access_model(
1084 model=model_to_check,
1085 team_object=team_object,
1086 llm_router=llm_router,
1087 team_model_aliases=user_api_key_dict.team_model_aliases,
1088 )
1089 except ProxyException as team_denial:
1090 if team_denial.type != ProxyErrorTypes.team_model_access_denied:
1091 raise
1092 if not await _key_access_group_grants_model(
1093 model=model_to_check,
1094 valid_token=user_api_key_dict,
1095 team_object=team_object,
1096 llm_router=llm_router,
1097 ):
1098 raise
1099 await _check_team_member_model_access(
1100 model=model_to_check,
1101 team_object=team_object,
1102 valid_token=user_api_key_dict,
1103 llm_router=llm_router,
1104 prisma_client=prisma_client,
1105 user_api_key_cache=user_api_key_cache,
1106 proxy_logging_obj=proxy_logging_obj,
1107 )
1108 else:
1109 await can_key_call_model(
1110 model=model_to_check,
1111 llm_model_list=llm_model_list,
1112 valid_token=user_api_key_dict,
1113 llm_router=llm_router,
1114 )
1115 except HTTPException:
1116 raise
1117 except Exception as e:
1118 raise HTTPException(
1119 status_code=403,
1120 detail={
1121 "error": (
1122 "Batch input file references a model the caller is "
1123 f"not authorized to use: model={model_to_check}, reason={e}"
1124 )
1125 },
1126 )
1128 async def _fetch_managed_file_content(
1129 self,
1130 file_id: str,
1131 user_api_key_dict: UserAPIKeyAuth,
1132 ) -> "HttpxBinaryResponseContent":
1133 """
1134 Fetch file content from managed files hook.
1136 This is needed for managed files because they require proper user context
1137 to verify file ownership and access permissions.
1139 Args:
1140 file_id: The managed file ID (base64 encoded)
1141 user_api_key_dict: User authentication information
1143 Returns:
1144 HttpxBinaryResponseContent with the file content
1145 """
1146 from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
1148 # Import proxy_server dependencies at runtime to avoid circular imports
1149 try:
1150 from litellm.proxy.proxy_server import llm_router, proxy_logging_obj
1151 except ImportError as e:
1152 raise ValueError(
1153 f"Cannot import proxy_server dependencies: {e}. Managed files require proxy_server to be initialized."
1154 )
1156 # Get the managed files hook
1157 if proxy_logging_obj is None:
1158 raise ValueError("proxy_logging_obj not available. Cannot access managed files hook.")
1160 managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files")
1161 if managed_files_obj is None:
1162 raise ValueError("Managed files hook not found. Cannot access managed file.")
1164 if not isinstance(managed_files_obj, BaseFileEndpoints):
1165 raise ValueError("Managed files hook is not a BaseFileEndpoints instance.")
1167 if llm_router is None:
1168 raise ValueError("llm_router not available. Cannot access managed files.")
1170 # Use the managed files hook to get file content
1171 # This properly handles user permissions and file ownership
1172 file_content: Final = await managed_files_obj.afile_content(
1173 file_id=file_id,
1174 litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
1175 llm_router=llm_router,
1176 )
1178 return file_content
1180 async def async_pre_call_hook(
1181 self,
1182 user_api_key_dict: UserAPIKeyAuth,
1183 cache: "DualCache",
1184 data: dict,
1185 call_type: str,
1186 ) -> Exception | str | dict | None:
1187 """
1188 Pre-call hook for batch operations.
1190 Only handles batch creation (acreate_batch):
1191 - Reads input file
1192 - Counts tokens and requests
1193 - Reserves rate limit capacity via parallel_request_limiter
1195 Args:
1196 user_api_key_dict: User authentication information
1197 cache: Cache instance (not used directly)
1198 data: Request data
1199 call_type: Type of call being made
1201 Returns:
1202 Modified data dict or None
1204 Raises:
1205 HTTPException: 429 if rate limit would be exceeded
1206 """
1207 # Only handle batch creation
1208 if call_type != "acreate_batch": 1208 ↛ 1209line 1208 didn't jump to line 1209 because the condition on line 1208 was never true
1209 verbose_proxy_logger.debug(
1210 "Batch rate limiter: Not handling batch creation rate limiting for call type: %s", call_type
1211 )
1212 return data
1214 verbose_proxy_logger.debug("Batch rate limiter: Handling batch creation rate limiting")
1216 try:
1217 # Extract input_file_id from data
1218 input_file_id: Final = data.get("input_file_id")
1219 if not input_file_id: 1219 ↛ 1223line 1219 didn't jump to line 1223 because the condition on line 1219 was always true
1220 verbose_proxy_logger.debug("No input_file_id in batch request, skipping rate limiting")
1221 return data
1223 enqueued_scopes: Final = resolve_batch_enqueued_token_scopes(user_api_key_dict)
1224 should_skip, batch_rate_limit_descriptors = self._should_skip_batch_input_file_processing(
1225 data=data, user_api_key_dict=user_api_key_dict, has_enqueued_scopes=bool(enqueued_scopes)
1226 )
1227 if should_skip:
1228 return data
1230 # Get custom_llm_provider for token counting
1231 custom_llm_provider: Final = data.get("custom_llm_provider", "openai")
1233 # Count tokens and requests from input file
1234 verbose_proxy_logger.debug("Counting tokens from batch input file: %s", input_file_id)
1235 batch_usage: Final = await self.count_input_file_usage(
1236 file_id=input_file_id,
1237 custom_llm_provider=custom_llm_provider,
1238 user_api_key_dict=user_api_key_dict,
1239 data=data,
1240 descriptors=batch_rate_limit_descriptors,
1241 )
1243 verbose_proxy_logger.debug(
1244 "Batch input file usage - Tokens: %s, Requests: %s", batch_usage.total_tokens, batch_usage.request_count
1245 )
1247 # Store batch usage in data for later reference
1248 data["_batch_token_count"] = batch_usage.total_tokens
1249 data["_batch_request_count"] = batch_usage.request_count
1251 if enqueued_scopes:
1252 await self._reserve_batch_enqueued_tokens(
1253 user_api_key_dict=user_api_key_dict,
1254 data=data,
1255 batch_usage=batch_usage,
1256 scopes=enqueued_scopes,
1257 )
1258 verbose_proxy_logger.debug("Batch enqueued-token reservation succeeded")
1259 return data
1261 # Directly increment counters by batch amounts (check happens atomically)
1262 # This will raise HTTPException if limits are exceeded
1263 await self._check_and_increment_batch_counters(
1264 user_api_key_dict=user_api_key_dict,
1265 data=data,
1266 batch_usage=batch_usage,
1267 descriptors=batch_rate_limit_descriptors,
1268 )
1270 verbose_proxy_logger.debug("Batch rate limit check passed, counters incremented")
1271 return data
1273 except HTTPException:
1274 # Re-raise HTTP exceptions (rate limit exceeded)
1275 raise
1276 except Exception as e:
1277 verbose_proxy_logger.error("Error in batch rate limiting: %s", e, exc_info=True)
1278 # Don't block the request if rate limiting fails
1279 return data