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

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 

10 

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 

17 

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 

37 

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 

40 

41 from litellm.router import Router 

42 

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({}) 

75 

76 

77class _ErrorDetail(TypedDict): 

78 message: ReadOnly[str] 

79 type: ReadOnly[str] 

80 param: ReadOnly[None] 

81 code: ReadOnly[None] 

82 

83 

84class _ErrorBody(TypedDict): 

85 error: ReadOnly[_ErrorDetail] 

86 

87 

88class _ResultResponse(TypedDict): 

89 status_code: ReadOnly[int] 

90 request_id: ReadOnly[str] 

91 body: ReadOnly[Mapping[str, object]] 

92 

93 

94class _LineError(TypedDict): 

95 code: ReadOnly[str] 

96 message: ReadOnly[str] 

97 

98 

99class _ResultLine(TypedDict): 

100 id: ReadOnly[str] 

101 custom_id: ReadOnly[str] 

102 response: ReadOnly[_ResultResponse | None] 

103 error: ReadOnly[_LineError | None] 

104 

105 

106class BatchInputLine(BaseModel): 

107 model_config = ConfigDict(extra="forbid", frozen=True) 

108 

109 custom_id: str 

110 method: Literal["POST"] 

111 url: str 

112 body: Mapping[str, object] 

113 

114 

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

116class InvalidBatchInput: 

117 line_number: int | None 

118 reason: str 

119 

120 def describe(self) -> str: 

121 return f"line {self.line_number}: {self.reason}" if self.line_number is not None else self.reason 

122 

123 

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

125class RowOutcome: 

126 custom_id: str 

127 status_code: int 

128 body: Mapping[str, object] 

129 succeeded: bool 

130 

131 

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

133class ExpiredRow: 

134 custom_id: str 

135 

136 

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 

147 

148 

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

152 

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: ... 

156 

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: ... 

169 

170 

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

173 

174 

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]: ... 

186 

187 

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

191 

192 

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 

199 

200 

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: ... 

205 

206 

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

209 

210 

211class BodyRejection(Protocol): 

212 def __call__(self, body: Mapping[str, object], /) -> str | None: ... 212 ↛ exitline 212 didn't return from function '__call__' because

213 

214 

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 

230 

231 

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) 

248 

249 

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 

260 

261 

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) 

270 

271 

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 

279 

280 

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 ) 

286 

287 

288def _accept_every_body(_body: Mapping[str, object]) -> str | None: 

289 return None 

290 

291 

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 

307 

308 

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 

325 

326 

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) 

330 

331 

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") 

337 

338 

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 

342 

343 

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 

349 

350 

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 

360 

361 

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 

367 

368 

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 

377 

378 

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__}") 

383 

384 

385def _resolve_transition(current_status: str, requested: BatchStatus) -> BatchStatus: 

386 if current_status != "cancelling": 

387 return requested 

388 return _CANCELLING_TRANSITIONS.get(requested, requested) 

389 

390 

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 

395 

396 

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 

403 

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 

413 

414 

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 

441 

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 

500 

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) 

516 

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 

533 

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 

546 

547 return reject 

548 

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)) 

564 

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() 

581 

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) 

589 

590 async def _touch(self, run: _BatchRun) -> None: 

591 await self.batches.touch(run.unified_batch_id, run.user_api_key_dict.user_id) 

592 

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 ) 

620 

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) 

634 

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) 

643 

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)) 

649 

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 

655 

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 } 

664 

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 

687 

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) 

704 

705 

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 )