Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/batches_endpoints/litellm_executed_batches.py: 35%
407 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
1import asyncio
2import json
3import time
4from collections.abc import Awaitable, Callable, Mapping, Sequence
5from dataclasses import dataclass
6from datetime import datetime, timedelta, timezone
7from itertools import pairwise
8from types import MappingProxyType
9from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable
11import httpx
12from openai.types.batch import Errors
13from openai.types.batch_error import BatchError
14from openai.types.batch_request_counts import BatchRequestCounts
15from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
16from typing_extensions import ReadOnly, TypedDict
18import litellm
19from litellm._logging import verbose_proxy_logger
20from litellm._uuid import uuid as uuid_module
21from litellm.constants import LITELLM_EXECUTED_BATCH_CONCURRENCY
22from litellm.integrations.prometheus import PrometheusLogger
23from litellm.llms.base_llm.files.litellm_db_storage_backend import LITELLM_DB_STORAGE_BACKEND_NAME
24from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend
25from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
26from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
27from litellm.models.managed_files import LiteLLM_ManagedFileTable
28from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth
29from litellm.proxy.auth.auth_utils import is_request_body_safe
30from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
31from litellm.proxy.openai_files_endpoints.common_utils import LITELLM_EXECUTED_BATCH_ID_PREFIX
32from litellm.proxy.openai_files_endpoints.storage_backend_service import StorageBackendFileService
33from litellm.proxy.utils import PrismaClient, ProxyLogging
34from litellm.repositories.managed_batch_repository import ManagedBatchRepository
35from litellm.types.llms.openai import LiteLLMBatchCreateRequest, OpenAIFileObject, OpenAIFilesPurpose
36from litellm.types.utils import LITELLM_EXECUTED_BATCH_PROVIDERS, ExtractedFileData, LiteLLMBatch, LlmProviders
38if TYPE_CHECKING: 38 ↛ 39line 38 didn't jump to line 39 because the condition on line 38 was never true
39 from prisma import types as prisma_types
41 from litellm.router import Router
43BatchEndpoint: TypeAlias = Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"]
44BatchStatus: TypeAlias = Literal[
45 "in_progress", "finalizing", "completed", "failed", "cancelling", "cancelled", "expired"
46]
47TERMINAL_BATCH_STATUSES: Final[frozenset[str]] = frozenset({"completed", "failed", "cancelled", "expired"})
48_STOP_STATUSES: Final[frozenset[str]] = TERMINAL_BATCH_STATUSES | frozenset({"cancelling"})
49_BATCH_ENDPOINT_ADAPTER: Final[TypeAdapter[BatchEndpoint]] = TypeAdapter(BatchEndpoint)
50_CANCEL_POLL_SECONDS: Final = 1.0
51_HEARTBEAT_SECONDS: Final = 30.0
52_STALE_AFTER_SECONDS: Final = 180.0
53_FILES_API_PROBE_TIMEOUT_SECONDS: Final = 5.0
54_COMPLETION_WINDOW_SECONDS: Final = 24 * 60 * 60
55_RUNNER_LOST_MESSAGE: Final = "the proxy replica running this batch stopped before it finished; resubmit the batch"
56_EXPIRED_MESSAGE: Final = "This request could not be executed before the completion window expired."
57_ROUTER_METHODS: Final[Mapping[BatchEndpoint, str]] = MappingProxyType(
58 {
59 "/v1/chat/completions": "acompletion",
60 "/v1/completions": "atext_completion",
61 "/v1/embeddings": "aembedding",
62 "/v1/responses": "aresponses",
63 }
64)
65_CANCELLING_TRANSITIONS: Final[Mapping[BatchStatus, BatchStatus]] = MappingProxyType(
66 {"completed": "cancelled", "expired": "cancelled", "in_progress": "cancelling", "finalizing": "cancelling"}
67)
68LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE: Final = (
69 "upload it through POST /v1/files with purpose=batch and either the x-litellm-model header or the "
70 "target_model_names form field naming the model, so LiteLLM keeps the file and runs the batch itself"
71)
72_RUNNING_BATCHES: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong references keep running batch tasks alive
73_NO_FIELDS: Final[Mapping[str, object]] = MappingProxyType({})
74_NO_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
77class _ErrorDetail(TypedDict):
78 message: ReadOnly[str]
79 type: ReadOnly[str]
80 param: ReadOnly[None]
81 code: ReadOnly[None]
84class _ErrorBody(TypedDict):
85 error: ReadOnly[_ErrorDetail]
88class _ResultResponse(TypedDict):
89 status_code: ReadOnly[int]
90 request_id: ReadOnly[str]
91 body: ReadOnly[Mapping[str, object]]
94class _LineError(TypedDict):
95 code: ReadOnly[str]
96 message: ReadOnly[str]
99class _ResultLine(TypedDict):
100 id: ReadOnly[str]
101 custom_id: ReadOnly[str]
102 response: ReadOnly[_ResultResponse | None]
103 error: ReadOnly[_LineError | None]
106class BatchInputLine(BaseModel):
107 model_config = ConfigDict(extra="forbid", frozen=True)
109 custom_id: str
110 method: Literal["POST"]
111 url: str
112 body: Mapping[str, object]
115@dataclass(frozen=True, slots=True)
116class InvalidBatchInput:
117 line_number: int | None
118 reason: str
120 def describe(self) -> str:
121 return f"line {self.line_number}: {self.reason}" if self.line_number is not None else self.reason
124@dataclass(frozen=True, slots=True)
125class RowOutcome:
126 custom_id: str
127 status_code: int
128 body: Mapping[str, object]
129 succeeded: bool
132@dataclass(frozen=True, slots=True)
133class ExpiredRow:
134 custom_id: str
137@dataclass(frozen=True, slots=True)
138class _BatchRun:
139 unified_batch_id: str
140 llm_batch_id: str
141 model: str
142 endpoint: BatchEndpoint
143 lines: tuple[BatchInputLine, ...]
144 user_api_key_dict: UserAPIKeyAuth
145 request_tags: tuple[str, ...]
146 deadline: float
149@runtime_checkable
150class ManagedBatchStore(Protocol):
151 def get_unified_batch_id(self, batch_id: str, model_id: str) -> str: ... 151 ↛ exitline 151 didn't return from function 'get_unified_batch_id' because
153 async def get_unified_file_id( 153 ↛ exitline 153 didn't return from function 'get_unified_file_id' because
154 self, file_id: str, litellm_parent_otel_span: object | None = None
155 ) -> LiteLLM_ManagedFileTable | None: ...
157 async def store_unified_object_id( 157 ↛ exitline 157 didn't return from function 'store_unified_object_id' because
158 self,
159 unified_object_id: str,
160 file_object: LiteLLMBatch,
161 litellm_parent_otel_span: object | None,
162 model_object_id: str,
163 file_purpose: Literal["batch", "fine-tune", "response"],
164 user_api_key_dict: UserAPIKeyAuth,
165 request_tags: Sequence[str] | None = None,
166 persist_attribution: bool = False,
167 batch_processed: bool = False,
168 ) -> None: ...
171class _StorageBackendFactory(Protocol):
172 def __call__(self, backend_type: str, prisma_client: PrismaClient | None = None) -> BaseFileStorageBackend: ... 172 ↛ exitline 172 didn't return from function '__call__' because
175class _ResultFileUploader(Protocol):
176 def __call__( 176 ↛ exitline 176 didn't return from function '__call__' because
177 self,
178 file_data: Mapping[str, object],
179 target_storage: str,
180 target_model_names: Sequence[str],
181 purpose: OpenAIFilesPurpose,
182 proxy_logging_obj: ProxyLogging,
183 user_api_key_dict: UserAPIKeyAuth,
184 prisma_client: PrismaClient | None = None,
185 ) -> Awaitable[OpenAIFileObject]: ...
188@runtime_checkable
189class _RouterCall(Protocol):
190 def __call__(self, **params: object) -> Awaitable[object]: ... # kwargs-ok: the request body is passed as keywords 190 ↛ exitline 190 didn't return from function '__call__' because
193def litellm_executed_provider_of(credentials: Mapping[str, object]) -> str | None:
194 explicit_provider: Final = credentials.get("custom_llm_provider")
195 provider: Final = (
196 explicit_provider if isinstance(explicit_provider, str) else _provider_of(credentials.get("model"))
197 )
198 return provider if provider in LITELLM_EXECUTED_BATCH_PROVIDERS else None
201class _HttpGetter(Protocol):
202 async def get( 202 ↛ exitline 202 didn't return from function 'get' because
203 self, url: str, *, headers: dict[str, str] | None = None, timeout: float | httpx.Timeout | None = None
204 ) -> httpx.Response: ...
207class FilesApiProbe(Protocol):
208 async def __call__(self, api_base: str, api_key: str | None) -> bool: ... 208 ↛ exitline 208 didn't return from function '__call__' because
211class BodyRejection(Protocol):
212 def __call__(self, body: Mapping[str, object], /) -> str | None: ... 212 ↛ exitline 212 didn't return from function '__call__' because
215async def upstream_lacks_files_api(api_base: str, api_key: str | None, http_client: _HttpGetter | None = None) -> bool:
216 client: Final = http_client or get_async_httpx_client(llm_provider=LlmProviders.HOSTED_VLLM)
217 try:
218 response: Final = await client.get(
219 f"{api_base.rstrip('/')}/files",
220 headers=(
221 {"Authorization": f"Bearer {api_key}"} # mutable-ok: AsyncHTTPHandler.get wants a plain dict
222 if api_key
223 else None
224 ),
225 timeout=_FILES_API_PROBE_TIMEOUT_SECONDS,
226 )
227 except httpx.HTTPError:
228 return False
229 return response.status_code == httpx.codes.NOT_FOUND
232def _upstream_of(credentials: Mapping[str, object], provider: str) -> tuple[str, str | None] | None:
233 model: Final = credentials.get("model")
234 api_base: Final = credentials.get("api_base")
235 api_key: Final = credentials.get("api_key")
236 if not isinstance(model, str):
237 return None
238 try:
239 _, _, resolved_api_key, resolved_api_base = litellm.get_llm_provider(
240 model=model,
241 custom_llm_provider=provider,
242 api_base=api_base if isinstance(api_base, str) else None,
243 api_key=api_key if isinstance(api_key, str) else None,
244 )
245 except Exception: # noqa: BLE001 # get_llm_provider raises on a model it cannot map, which means nothing to probe
246 return None
247 return None if resolved_api_base is None else (resolved_api_base, resolved_api_key)
250async def litellm_executed_provider_for(
251 credentials: Mapping[str, object], lacks_files_api: FilesApiProbe = upstream_lacks_files_api
252) -> str | None:
253 provider: Final = litellm_executed_provider_of(credentials)
254 if provider is None:
255 return None
256 upstream: Final = _upstream_of(credentials, provider)
257 if upstream is None:
258 return None
259 return provider if await lacks_files_api(*upstream) else None
262async def resolve_litellm_executed_provider(
263 llm_router: "Router",
264 model: str,
265 team_id: str | None,
266 lacks_files_api: FilesApiProbe = upstream_lacks_files_api,
267) -> str | None:
268 credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model, team_id=team_id)
269 return None if credentials is None else await litellm_executed_provider_for(credentials, lacks_files_api)
272def _provider_of(model: object) -> str | None:
273 if not isinstance(model, str):
274 return None
275 try:
276 return litellm.get_llm_provider(model=model)[1]
277 except Exception: # noqa: BLE001 # get_llm_provider raises on an unknown model, which means no provider
278 return None
281def _validation_reason(error: ValidationError) -> str:
282 return "; ".join(
283 f"{'.'.join(str(part) for part in item['loc'])}: {item['msg']}" if item["loc"] else item["msg"]
284 for item in error.errors()
285 )
288def _accept_every_body(_body: Mapping[str, object]) -> str | None:
289 return None
292def _parse_line(
293 line_number: int, raw: bytes, endpoint: BatchEndpoint, reject_body: BodyRejection
294) -> BatchInputLine | InvalidBatchInput:
295 try:
296 line: Final = BatchInputLine.model_validate_json(raw)
297 except ValidationError as e:
298 return InvalidBatchInput(line_number, _validation_reason(e))
299 if line.url != endpoint:
300 return InvalidBatchInput(line_number, f"url {line.url!r} does not match the batch endpoint {endpoint!r}")
301 if line.body.get("stream"):
302 return InvalidBatchInput(line_number, "streaming requests are not supported in a batch")
303 rejection: Final = reject_body(line.body)
304 if rejection is not None:
305 return InvalidBatchInput(line_number, rejection)
306 return line
309def parse_batch_input(
310 content: bytes, endpoint: BatchEndpoint, reject_body: BodyRejection = _accept_every_body
311) -> tuple[BatchInputLine, ...] | InvalidBatchInput:
312 raw_lines: Final = tuple((number, raw) for number, raw in enumerate(content.splitlines(), start=1) if raw.strip())
313 if not raw_lines:
314 return InvalidBatchInput(None, "the input file has no requests")
315 parsed: Final = tuple(_parse_line(number, raw, endpoint, reject_body) for number, raw in raw_lines)
316 first_invalid: Final = next((item for item in parsed if isinstance(item, InvalidBatchInput)), None)
317 if first_invalid is not None:
318 return first_invalid
319 lines: Final = tuple(item for item in parsed if isinstance(item, BatchInputLine))
320 custom_ids: Final = sorted(line.custom_id for line in lines)
321 duplicate: Final = next((first for first, second in pairwise(custom_ids) if first == second), None)
322 if duplicate is not None:
323 return InvalidBatchInput(None, f"custom_id {duplicate!r} is used more than once")
324 return lines
327def batch_error(status_code: int, message: str) -> ProxyException:
328 error_type: Final = "invalid_request_error" if status_code < 500 else ProxyErrorTypes.internal_server_error.value
329 return ProxyException(message=message, type=error_type, param=None, code=status_code)
332def _validate_endpoint(endpoint: object) -> BatchEndpoint:
333 try:
334 return _BATCH_ENDPOINT_ADAPTER.validate_python(endpoint)
335 except ValidationError:
336 raise batch_error(400, f"endpoint {endpoint!r} is not supported for a LiteLLM-executed batch")
339def _status_code_of(error: Exception) -> int:
340 status_code: Final[object] = getattr(error, "status_code", None)
341 return status_code if isinstance(status_code, int) else 500
344def _error_body(error: Exception) -> _ErrorBody:
345 body: Final[_ErrorBody] = {
346 "error": {"message": str(error), "type": type(error).__name__, "param": None, "code": None}
347 }
348 return body
351def _line_response(outcome: RowOutcome | ExpiredRow) -> _ResultResponse | None:
352 if isinstance(outcome, ExpiredRow):
353 return None
354 response: Final[_ResultResponse] = {
355 "status_code": outcome.status_code,
356 "request_id": f"req_{uuid_module.uuid4().hex[:24]}",
357 "body": outcome.body,
358 }
359 return response
362def _line_error(outcome: RowOutcome | ExpiredRow) -> _LineError | None:
363 if isinstance(outcome, RowOutcome):
364 return None
365 error: Final[_LineError] = {"code": "batch_expired", "message": _EXPIRED_MESSAGE}
366 return error
369def _result_line(outcome: RowOutcome | ExpiredRow) -> _ResultLine:
370 line: Final[_ResultLine] = {
371 "id": f"batch_req_{uuid_module.uuid4().hex[:24]}",
372 "custom_id": outcome.custom_id,
373 "response": _line_response(outcome),
374 "error": _line_error(outcome),
375 }
376 return line
379def _dump(response: object) -> Mapping[str, object]:
380 if isinstance(response, BaseModel):
381 return response.model_dump(mode="json")
382 raise TypeError(f"Batch rows must return a single response object, got {type(response).__name__}")
385def _resolve_transition(current_status: str, requested: BatchStatus) -> BatchStatus:
386 if current_status != "cancelling":
387 return requested
388 return _CANCELLING_TRANSITIONS.get(requested, requested)
391def executed_batch_runner_lost(status: str, updated_at: datetime) -> bool:
392 if status in TERMINAL_BATCH_STATUSES:
393 return False
394 return (datetime.now(timezone.utc) - updated_at).total_seconds() > _STALE_AFTER_SECONDS
397class _StopWatch:
398 def __init__(self, load_status: Callable[[], Awaitable[str | None]], interval_seconds: float) -> None:
399 self._load_status = load_status
400 self._interval_seconds = interval_seconds
401 self._checked_at = float("-inf")
402 self._stopped = False
404 async def stopped(self) -> bool:
405 if self._stopped:
406 return True
407 now: Final = time.monotonic()
408 if now - self._checked_at < self._interval_seconds:
409 return False
410 self._checked_at = now
411 self._stopped = await self._load_status() in _STOP_STATUSES
412 return self._stopped
415class LiteLLMExecutedBatchRunner:
416 def __init__(
417 self,
418 llm_router: "Router",
419 prisma_client: PrismaClient,
420 managed_files: ManagedBatchStore,
421 batches: ManagedBatchRepository,
422 proxy_logging_obj: ProxyLogging,
423 general_settings: Mapping[str, object],
424 concurrency: int = LITELLM_EXECUTED_BATCH_CONCURRENCY,
425 heartbeat_seconds: float = _HEARTBEAT_SECONDS,
426 completion_window_seconds: float = _COMPLETION_WINDOW_SECONDS,
427 storage_backend_factory: _StorageBackendFactory = get_storage_backend,
428 upload_result_file: _ResultFileUploader = StorageBackendFileService.upload_file_to_storage_backend,
429 ) -> None:
430 self.llm_router = llm_router
431 self.prisma_client = prisma_client
432 self.managed_files = managed_files
433 self.batches = batches
434 self.proxy_logging_obj = proxy_logging_obj
435 self.general_settings = general_settings
436 self.concurrency = concurrency
437 self.heartbeat_seconds = heartbeat_seconds
438 self.completion_window_seconds = completion_window_seconds
439 self.storage_backend_factory = storage_backend_factory
440 self.upload_result_file = upload_result_file
442 async def create(
443 self,
444 create_request: LiteLLMBatchCreateRequest,
445 unified_input_file_id: str,
446 model: str,
447 provider: str,
448 user_api_key_dict: UserAPIKeyAuth,
449 request_tags: Sequence[str] | None,
450 ) -> LiteLLMBatch:
451 endpoint: Final = _validate_endpoint(create_request.get("endpoint"))
452 content: Final = await self._download_input(unified_input_file_id, user_api_key_dict)
453 parsed: Final = parse_batch_input(content, endpoint, self._body_rejection(model))
454 if isinstance(parsed, InvalidBatchInput):
455 raise batch_error(400, f"Invalid batch input file: {parsed.describe()}")
456 llm_batch_id: Final = f"{LITELLM_EXECUTED_BATCH_ID_PREFIX}{uuid_module.uuid4().hex}"
457 model_id: Final = next(iter(self.llm_router.get_model_ids(model_name=model)), model)
458 unified_batch_id: Final = self.managed_files.get_unified_batch_id(batch_id=llm_batch_id, model_id=model_id)
459 now: Final = time.time()
460 created_at: Final = int(now)
461 batch: Final = LiteLLMBatch(
462 id=unified_batch_id,
463 object="batch",
464 endpoint=endpoint,
465 input_file_id=unified_input_file_id,
466 completion_window="24h",
467 status="validating",
468 created_at=created_at,
469 expires_at=created_at + int(self.completion_window_seconds),
470 metadata=create_request.get("metadata"),
471 model=model,
472 request_counts=BatchRequestCounts(completed=0, failed=0, total=len(parsed)),
473 )
474 await self.managed_files.store_unified_object_id(
475 unified_object_id=unified_batch_id,
476 file_object=batch,
477 litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
478 model_object_id=llm_batch_id,
479 file_purpose="batch",
480 user_api_key_dict=user_api_key_dict,
481 request_tags=request_tags,
482 persist_attribution=True,
483 batch_processed=True,
484 )
485 _record_batch_created(model, provider, user_api_key_dict)
486 run: Final = _BatchRun(
487 unified_batch_id=unified_batch_id,
488 llm_batch_id=llm_batch_id,
489 model=model,
490 endpoint=endpoint,
491 lines=parsed,
492 user_api_key_dict=user_api_key_dict,
493 request_tags=tuple(request_tags or ()),
494 deadline=now + self.completion_window_seconds,
495 )
496 task: Final = asyncio.create_task(self._run(run))
497 _RUNNING_BATCHES.add(task)
498 task.add_done_callback(_RUNNING_BATCHES.discard)
499 return batch
501 async def cancel(self, unified_batch_id: str, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch:
502 current: Final = await self.batches.load_batch(unified_batch_id)
503 if current is None:
504 raise batch_error(404, f"Batch {unified_batch_id} not found")
505 if current.status in TERMINAL_BATCH_STATUSES:
506 raise batch_error(400, f"Cannot cancel a batch with status '{current.status}'")
507 if current.status == "cancelling":
508 return current
509 cancelling: Final = current.model_copy(
510 update=MappingProxyType({"status": "cancelling", "cancelling_at": int(time.time())})
511 )
512 unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status}
513 if await self.batches.compare_and_set(cancelling, unchanged, user_api_key_dict.user_id):
514 return cancelling
515 return await self.cancel(unified_batch_id, user_api_key_dict)
517 async def fail_abandoned(self, batch: LiteLLMBatch, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch:
518 error: Final = BatchError(message=_RUNNER_LOST_MESSAGE, code="runner_lost")
519 errors: Final = Errors(data=[error], object="list") # mutable-ok: Errors.data is typed as a list
520 failed: Final = batch.model_copy(
521 update=MappingProxyType({"status": "failed", "failed_at": int(time.time()), "errors": errors})
522 )
523 untouched: Final[prisma_types.DateTimeFilter] = {
524 "lt": datetime.now(timezone.utc) - timedelta(seconds=_STALE_AFTER_SECONDS)
525 }
526 still_abandoned: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {
527 "status": batch.status,
528 "updated_at": untouched,
529 }
530 if await self.batches.compare_and_set(failed, still_abandoned, user_api_key_dict.user_id):
531 return failed
532 return await self.batches.load_batch(batch.id) or batch
534 def _body_rejection(self, model: str) -> BodyRejection:
535 def reject(body: Mapping[str, object]) -> str | None:
536 try:
537 is_request_body_safe(
538 request_body=dict(body), # mutable-ok: is_request_body_safe takes a dict
539 general_settings=dict(self.general_settings), # mutable-ok: is_request_body_safe takes a dict
540 llm_router=self.llm_router,
541 model=model,
542 )
543 except ValueError as e:
544 return str(e)
545 return None
547 return reject
549 async def _download_input(self, unified_input_file_id: str, user_api_key_dict: UserAPIKeyAuth) -> bytes:
550 stored: Final = await self.managed_files.get_unified_file_id(
551 unified_input_file_id, litellm_parent_otel_span=user_api_key_dict.parent_otel_span
552 )
553 if stored is None or not stored.storage_backend or not stored.storage_url:
554 raise batch_error(
555 400,
556 f"LiteLLM does not hold the content of input file {unified_input_file_id}: "
557 f"{LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE}",
558 )
559 try:
560 backend: Final = self.storage_backend_factory(stored.storage_backend, prisma_client=self.prisma_client)
561 return await backend.download_file(stored.storage_url)
562 except ValueError as e:
563 raise batch_error(400, str(e))
565 async def _run(self, run: _BatchRun) -> None:
566 heartbeat: Final = asyncio.create_task(self._heartbeat(run))
567 try:
568 await self._execute(run)
569 except Exception as e: # noqa: BLE001 # whatever fails, the batch must end up marked failed
570 verbose_proxy_logger.exception("LiteLLM-executed batch %s failed: %s", run.unified_batch_id, e)
571 error: Final = BatchError(message=str(e), code="internal_error")
572 errors: Final = Errors(data=[error], object="list") # mutable-ok: Errors.data is typed as a list
573 try:
574 await self._advance(run, "failed", MappingProxyType({"errors": errors}))
575 except Exception as advance_error: # noqa: BLE001 # a failed status write is logged, never raised
576 verbose_proxy_logger.exception(
577 "LiteLLM-executed batch %s could not be marked failed: %s", run.unified_batch_id, advance_error
578 )
579 finally:
580 heartbeat.cancel()
582 async def _heartbeat(self, run: _BatchRun) -> None:
583 while True:
584 await asyncio.sleep(self.heartbeat_seconds)
585 try:
586 await self._touch(run)
587 except Exception as e: # noqa: BLE001 # a missed beat is logged and the next one retries
588 verbose_proxy_logger.warning("LiteLLM-executed batch %s heartbeat failed: %s", run.unified_batch_id, e)
590 async def _touch(self, run: _BatchRun) -> None:
591 await self.batches.touch(run.unified_batch_id, run.user_api_key_dict.user_id)
593 async def _execute(self, run: _BatchRun) -> None:
594 await self._advance(run, "in_progress")
595 watch: Final = _StopWatch(lambda: self.batches.load_status(run.unified_batch_id), _CANCEL_POLL_SECONDS)
596 semaphore: Final = asyncio.Semaphore(self.concurrency)
597 results: Final = await asyncio.gather(*(self._run_row(run, line, watch, semaphore) for line in run.lines))
598 outcomes: Final = tuple(outcome for outcome in results if outcome is not None)
599 if await self._advance(run, "finalizing") is None:
600 return
601 succeeded: Final = tuple(
602 outcome for outcome in outcomes if isinstance(outcome, RowOutcome) and outcome.succeeded
603 )
604 failed: Final = tuple(
605 outcome for outcome in outcomes if isinstance(outcome, ExpiredRow) or not outcome.succeeded
606 )
607 output_file_id: Final = await self._upload_results(run, "output", succeeded)
608 error_file_id: Final = await self._upload_results(run, "error", failed)
609 request_counts: Final = BatchRequestCounts(completed=len(succeeded), failed=len(failed), total=len(run.lines))
610 final_status: Final[BatchStatus] = (
611 "expired" if any(isinstance(outcome, ExpiredRow) for outcome in outcomes) else "completed"
612 )
613 await self._advance(
614 run,
615 final_status,
616 MappingProxyType(
617 {"output_file_id": output_file_id, "error_file_id": error_file_id, "request_counts": request_counts}
618 ),
619 )
621 async def _run_row(
622 self, run: _BatchRun, line: BatchInputLine, watch: _StopWatch, semaphore: asyncio.Semaphore
623 ) -> RowOutcome | ExpiredRow | None:
624 async with semaphore:
625 if await watch.stopped():
626 return None
627 remaining: Final = run.deadline - time.time()
628 if remaining <= 0:
629 return ExpiredRow(custom_id=line.custom_id)
630 try:
631 return await asyncio.wait_for(self._row_outcome(run, line), timeout=remaining)
632 except asyncio.TimeoutError:
633 return ExpiredRow(custom_id=line.custom_id)
635 async def _row_outcome(self, run: _BatchRun, line: BatchInputLine) -> RowOutcome:
636 try:
637 body: Final = await self._dispatch(run, line)
638 except Exception as e: # noqa: BLE001 # a provider error becomes the row's error line, never a crashed batch
639 return RowOutcome(
640 custom_id=line.custom_id, status_code=_status_code_of(e), body=_error_body(e), succeeded=False
641 )
642 return RowOutcome(custom_id=line.custom_id, status_code=200, body=body, succeeded=True)
644 async def _dispatch(self, run: _BatchRun, line: BatchInputLine) -> Mapping[str, object]:
645 params: Final = MappingProxyType(
646 {**line.body, "model": run.model, "metadata": self._row_metadata(run), "disable_fallbacks": True}
647 )
648 return _dump(await self._router_call(run.endpoint)(**params))
650 def _router_call(self, endpoint: BatchEndpoint) -> _RouterCall:
651 method: Final[object] = getattr(self.llm_router, _ROUTER_METHODS[endpoint], None)
652 if not isinstance(method, _RouterCall):
653 raise TypeError(f"the router has no callable for {endpoint}")
654 return method
656 def _row_metadata(self, run: _BatchRun) -> dict[str, object]: # mutable-ok: router updates metadata in place
657 return { # mutable-ok: the router updates request metadata in place
658 **LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(run.user_api_key_dict),
659 "user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(run.user_api_key_dict),
660 "user_api_end_user_max_budget": run.user_api_key_dict.end_user_max_budget,
661 "tags": list(run.request_tags), # mutable-ok: litellm types request tags as a list
662 "batch_id": run.unified_batch_id,
663 }
665 async def _upload_results(
666 self, run: _BatchRun, kind: Literal["output", "error"], outcomes: Sequence[RowOutcome | ExpiredRow]
667 ) -> str | None:
668 if not outcomes:
669 return None
670 content: Final = "".join(f"{json.dumps(_result_line(outcome))}\n" for outcome in outcomes).encode()
671 file_data: Final[ExtractedFileData] = {
672 "filename": f"{run.llm_batch_id}_{kind}.jsonl",
673 "content": content,
674 "content_type": "application/jsonl",
675 "headers": _NO_HEADERS,
676 }
677 file_object: Final = await self.upload_result_file(
678 file_data=file_data,
679 target_storage=LITELLM_DB_STORAGE_BACKEND_NAME,
680 target_model_names=(run.model,),
681 purpose="batch_output",
682 proxy_logging_obj=self.proxy_logging_obj,
683 user_api_key_dict=run.user_api_key_dict,
684 prisma_client=self.prisma_client,
685 )
686 return file_object.id
688 async def _advance(
689 self, run: _BatchRun, requested: BatchStatus, fields: Mapping[str, object] = _NO_FIELDS
690 ) -> BatchStatus | None:
691 current: Final = await self.batches.load_batch(run.unified_batch_id)
692 if current is None:
693 raise RuntimeError(f"Batch {run.unified_batch_id} is no longer stored")
694 if current.status in TERMINAL_BATCH_STATUSES:
695 return None
696 status: Final = _resolve_transition(current.status, requested)
697 updated: Final = current.model_copy(
698 update=MappingProxyType({**fields, "status": status, f"{status}_at": int(time.time())})
699 )
700 unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status}
701 if await self.batches.compare_and_set(updated, unchanged, run.user_api_key_dict.user_id):
702 return status
703 return await self._advance(run, requested, fields)
706def _record_batch_created(model: str, provider: str, user_api_key_dict: UserAPIKeyAuth) -> None:
707 prometheus_logger: Final = PrometheusLogger.get_instance()
708 if prometheus_logger is None:
709 return
710 prometheus_logger.record_managed_batch_created(
711 model=model,
712 api_provider=provider,
713 user=user_api_key_dict.user_id or "",
714 user_email=user_api_key_dict.user_email or "",
715 api_key_alias=user_api_key_dict.key_alias or "",
716 )