Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/openai_files_endpoints/common_utils.py: 35%

504 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import base64 

2import mimetypes 

3import re 

4from collections.abc import Mapping 

5from dataclasses import dataclass, field 

6from types import MappingProxyType 

7from typing import ( 

8 TYPE_CHECKING, 

9 Final, 

10 Literal, 

11 Optional, 

12 Protocol, 

13 cast, # noqa: TID251 # prisma types Json columns as fields.Json but de-serializes them to plain python on read 

14 get_args, 

15 runtime_checkable, 

16) 

17 

18from litellm.batches.batch_utils import batch_cost_is_final 

19from litellm.constants import MAX_FILE_LIST_LIMIT 

20from litellm.proxy._types import ProxyException 

21from litellm.repositories.table_repositories import ( 

22 ManagedFileRepository, 

23 ManagedObjectRepository, 

24) 

25from litellm.types.llms.openai import OpenAIFilesPurpose 

26from litellm.types.utils import SpecialEnums 

27 

28if TYPE_CHECKING: 28 ↛ 29line 28 didn't jump to line 29 because the condition on line 28 was never true

29 from fastapi import Request 

30 from prisma.models import LiteLLM_ManagedObjectTable 

31 

32 from litellm.proxy._types import UserAPIKeyAuth 

33 from litellm.proxy.utils import PrismaClient 

34 from litellm.router import Router 

35 from litellm.types.utils import LiteLLMBatch 

36 

37 

38FILE_LIST_CONTINUATION_CHUNK_SIZE: Final = 500 

39 

40BATCH_CREATE_HIDDEN_PARAM: Final = "batch_create" 

41LITELLM_EXECUTED_BATCH_ID_PREFIX: Final = "litellm_batch_" 

42 

43 

44def validate_file_list_limit(limit: int | None) -> None: 

45 """Reject a ``limit`` outside the range OpenAI documents for GET /v1/files.""" 

46 if limit is None or 1 <= limit <= MAX_FILE_LIST_LIMIT: 

47 return 

48 bound, expected, openai_code = ( 

49 ("below minimum", ">= 1", "integer_below_min_value") 

50 if limit < 1 

51 else ("above maximum", f"<= {MAX_FILE_LIST_LIMIT}", "integer_above_max_value") 

52 ) 

53 raise ProxyException( 

54 message=f"Invalid 'limit': integer {bound} value. Expected a value {expected}, but got {limit} instead.", 

55 type="invalid_request_error", 

56 param="limit", 

57 code=400, 

58 openai_code=openai_code, 

59 ) 

60 

61 

62def validate_file_list_purpose(purpose: str | None) -> None: 

63 """Reject a ``purpose`` filter no upload to this proxy could have stored. 

64 

65 An unknown purpose matches no file, so filtering on it would report an 

66 empty page for what is really a bad request. Rejecting it keeps a managed 

67 listing consistent with the upload route, which refuses the same values 

68 against this same set. The provider-backed listings do not: they pass 

69 ``purpose`` upstream, so a purpose OpenAI accepts before it is added here 

70 is rejected on the managed path while still working on those. 

71 """ 

72 valid_purposes: Final = get_args(OpenAIFilesPurpose) 

73 if purpose is None or purpose in valid_purposes: 

74 return 

75 raise ProxyException( 

76 message=f"Invalid purpose: {purpose}. Must be one of: {valid_purposes}", 

77 type="invalid_request_error", 

78 param="purpose", 

79 code=400, 

80 ) 

81 

82 

83@runtime_checkable 

84class ManagedResourceAccessChecker(Protocol): 

85 async def can_user_call_unified_file_id( 85 ↛ exitline 85 didn't return from function 'can_user_call_unified_file_id' because

86 self, 

87 unified_file_id: str, 

88 user_api_key_dict: "UserAPIKeyAuth", 

89 ) -> bool: ... 

90 

91 async def can_user_call_unified_object_id( 91 ↛ exitline 91 didn't return from function 'can_user_call_unified_object_id' because

92 self, 

93 unified_object_id: str, 

94 user_api_key_dict: "UserAPIKeyAuth", 

95 ) -> bool: ... 

96 

97 

98def _is_base64_encoded_unified_file_id(b64_uid: str) -> str | Literal[False]: 

99 # Ensure b64_uid is a string and not a mock object 

100 if not isinstance(b64_uid, str): 100 ↛ 101line 100 didn't jump to line 101 because the condition on line 100 was never true

101 return False 

102 # Add padding back if needed 

103 padded: Final = b64_uid + "=" * (-len(b64_uid) % 4) 

104 # Decode from base64 

105 try: 

106 decoded: Final = base64.urlsafe_b64decode(padded).decode() 

107 if decoded.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value): 107 ↛ 108line 107 didn't jump to line 108 because the condition on line 107 was never true

108 return decoded 

109 else: 

110 return False 

111 except Exception: 

112 return False 

113 

114 

115def convert_b64_uid_to_unified_uid(b64_uid: str) -> str: 

116 is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(b64_uid) 

117 if is_base64_unified_file_id: 

118 return is_base64_unified_file_id 

119 else: 

120 return b64_uid 

121 

122 

123def resolve_managed_output_file_model_name( 

124 unified_input_file_id: str | None, fallback_model_name: str | None 

125) -> str | None: 

126 if not unified_input_file_id: 

127 return fallback_model_name 

128 target_model_names: Final = get_models_from_unified_file_id(convert_b64_uid_to_unified_uid(unified_input_file_id)) 

129 if target_model_names: 

130 return ",".join(target_model_names) 

131 return fallback_model_name 

132 

133 

134def get_models_from_unified_file_id(unified_file_id: str) -> list[str]: 

135 """ 

136 Extract model names from unified file ID. 

137 

138 Example: 

139 unified_file_id = "litellm_proxy:application/octet-stream;unified_id,c4843482-b176-4901-8292-7523fd0f2c6e;target_model_names,gpt-4o-mini,gemini-2.0-flash" 

140 returns: ["gpt-4o-mini", "gemini-2.0-flash"] 

141 """ 

142 try: 

143 # Ensure unified_file_id is a string and not a mock object 

144 if not isinstance(unified_file_id, str): 

145 return [] 

146 match: Final = re.search(r"target_model_names,([^;]+)", unified_file_id) 

147 if match: 

148 # Split on comma and strip whitespace from each model name 

149 return [model.strip() for model in match.group(1).split(",")] 

150 return [] 

151 except Exception: 

152 return [] 

153 

154 

155def get_model_id_from_unified_batch_id(file_id: str) -> str | None: 

156 """ 

157 Get the model_id from the file_id 

158 

159 Expected format: litellm_proxy;model_id:{};llm_batch_id:{};llm_output_file_id:{} 

160 """ 

161 ## use regex to get the model_id from the file_id 

162 try: 

163 # Ensure file_id is a string and not a mock object 

164 if not isinstance(file_id, str): 

165 return None 

166 return file_id.split("model_id:")[1].split(";")[0] 

167 except Exception: 

168 return None 

169 

170 

171def get_batch_id_from_unified_batch_id(file_id: str) -> str: 

172 ## use regex to get the batch_id from the file_id 

173 # Ensure file_id is a string and not a mock object 

174 if not isinstance(file_id, str): 

175 return "" 

176 if "llm_batch_id" in file_id: 

177 batch_id = file_id.split("llm_batch_id:", 1)[1] 

178 else: 

179 batch_id = file_id.split("generic_response_id:", 1)[1] 

180 return re.split(r"[;,]", batch_id, maxsplit=1)[0] 

181 

182 

183def is_litellm_executed_batch(decoded_unified_batch_id: str) -> bool: 

184 _, marker, batch_id = decoded_unified_batch_id.partition("llm_batch_id:") 

185 return bool(marker) and batch_id.startswith(LITELLM_EXECUTED_BATCH_ID_PREFIX) 

186 

187 

188def encode_file_id_with_model(file_id: str, model: str, id_type: Literal["file", "batch"] = "file") -> str: 

189 """ 

190 Encode a file/batch ID with model routing information. 

191 

192 Format: <prefix><base64(litellm:<original_id>;model,<model_name>)> 

193 The result preserves the original prefix (file-, batch_, etc.) for OpenAI compliance. 

194 

195 Args: 

196 file_id: Original file/batch ID from the provider (e.g., "file-abc123", "batch_xyz") 

197 model: Model name from model_list (e.g., "gpt-4o-litellm") 

198 id_type: Type of ID being encoded. Used to determine the correct prefix when 

199 the raw ID lacks a recognizable prefix (e.g., Vertex AI numeric IDs). 

200 Defaults to "file" for backward compatibility. 

201 

202 Returns: 

203 Encoded ID starting with appropriate prefix and containing routing information 

204 

205 Examples: 

206 encode_file_id_with_model("file-abc123", "gpt-4o-litellm") 

207 -> "file-bGl0ZWxsbTpmaWxlLWFiYzEyMzttb2RlbCxncHQtNG8taWZvb2Q" 

208 

209 encode_file_id_with_model("batch_abc123", "gpt-4o-test") 

210 -> "batch_bGl0ZWxsbTpiYXRjaF9hYmMxMjM7bW9kZWwsZ3B0LTRvLXRlc3Q" 

211 

212 encode_file_id_with_model("3814889423749775360", "gemini-2.5-pro", id_type="batch") 

213 -> "batch_bGl0ZWxsbTozODE0ODg5NDIzNzQ5Nzc1MzYwO21vZGVsLGdlbWluaS0yLjUtcHJv" 

214 """ 

215 encoded_str: Final = f"litellm:{file_id};model,{model}" 

216 encoded_bytes: Final = base64.urlsafe_b64encode(encoded_str.encode()) 

217 encoded_b64: Final = encoded_bytes.decode().rstrip("=") 

218 

219 # Detect the prefix from the original ID (file-, batch_, etc.) 

220 # For provider-specific IDs without a recognizable prefix (e.g., Vertex AI 

221 # numeric batch IDs), fall back to id_type to determine the correct prefix. 

222 if file_id.startswith("batch_"): 

223 prefix = "batch_" 

224 elif file_id.startswith("file-"): 

225 prefix = "file-" 

226 else: 

227 prefix = "batch_" if id_type == "batch" else "file-" 

228 

229 return f"{prefix}{encoded_b64}" 

230 

231 

232def encode_batch_response_ids(response, model: str) -> None: 

233 """Encode all IDs in a batch response with model routing info (in-place).""" 

234 if not response or not hasattr(response, "id") or not response.id: 

235 return 

236 response.id = encode_file_id_with_model(file_id=response.id, model=model, id_type="batch") 

237 for attr in ("output_file_id", "error_file_id", "input_file_id"): 

238 if hasattr(response, attr) and getattr(response, attr): 

239 setattr( 

240 response, 

241 attr, 

242 encode_file_id_with_model(file_id=getattr(response, attr), model=model), 

243 ) 

244 

245 

246def decode_model_from_file_id(encoded_id: str) -> str | None: 

247 """ 

248 Extract model name from an encoded file/batch ID. 

249 Handles IDs that start with "file-" or "batch_" prefix. 

250 """ 

251 try: 

252 if not isinstance(encoded_id, str): 252 ↛ 253line 252 didn't jump to line 253 because the condition on line 252 was never true

253 return None 

254 

255 # Remove prefix if present (file-, batch_, etc.) 

256 if encoded_id.startswith("file-"): 256 ↛ 257line 256 didn't jump to line 257 because the condition on line 256 was never true

257 b64_part = encoded_id[5:] # Remove "file-" 

258 elif encoded_id.startswith("batch_"): 258 ↛ 259line 258 didn't jump to line 259 because the condition on line 258 was never true

259 b64_part = encoded_id[6:] # Remove "batch_" 

260 else: 

261 b64_part = encoded_id 

262 

263 padded: Final = b64_part + "=" * (-len(b64_part) % 4) 

264 decoded: Final = base64.urlsafe_b64decode(padded).decode() 

265 if decoded.startswith("litellm:") and ";model," in decoded: 265 ↛ 266line 265 didn't jump to line 266 because the condition on line 265 was never true

266 match: Final = re.search(r";model,([^;]+)", decoded) 

267 if match: 

268 return match.group(1).strip() 

269 

270 return None 

271 except Exception: 

272 return None 

273 

274 

275def get_original_file_id(encoded_id: str) -> str: 

276 """ 

277 Extract the original provider file/batch ID from an encoded ID. 

278 Handles IDs that start with "file-" or "batch_" prefix. 

279 """ 

280 try: 

281 if not isinstance(encoded_id, str): 

282 return encoded_id 

283 

284 # Remove prefix if present (file-, batch_, etc.) 

285 if encoded_id.startswith("file-"): 

286 b64_part = encoded_id[5:] # Remove "file-" 

287 elif encoded_id.startswith("batch_"): 

288 b64_part = encoded_id[6:] # Remove "batch_" 

289 else: 

290 b64_part = encoded_id 

291 

292 padded: Final = b64_part + "=" * (-len(b64_part) % 4) 

293 decoded: Final = base64.urlsafe_b64decode(padded).decode() 

294 

295 if decoded.startswith("litellm:") and ";model," in decoded: 

296 match: Final = re.search(r"litellm:([^;]+);model,", decoded) 

297 if match: 

298 return match.group(1) 

299 

300 return encoded_id 

301 except Exception: 

302 return encoded_id 

303 

304 

305def is_model_embedded_id(file_id: str) -> bool: 

306 """ 

307 Check if a file/batch ID has model routing information embedded. 

308 """ 

309 return decode_model_from_file_id(file_id) is not None 

310 

311 

312# ============================================================================ 

313# MODEL-BASED CREDENTIAL ROUTING HELPERS 

314# ============================================================================ 

315 

316 

317def extract_model_from_sources( 

318 file_id: str, 

319 request, # FastAPI Request object 

320 data: dict | None = None, 

321) -> tuple[str | None, str | None]: 

322 """ 

323 Extract model information from multiple sources in priority order: 

324 1. Embedded in file_id (highest priority) 

325 2. Request headers (x-litellm-model) 

326 3. Query parameters (?model=) 

327 4. Request body/data dict 

328 

329 Args: 

330 file_id: File ID that may contain embedded model info 

331 request: FastAPI request object 

332 data: Optional request data dictionary 

333 

334 Returns: 

335 Tuple of (model_from_id, model_from_param) 

336 - model_from_id: Model decoded from file ID (if embedded) 

337 - model_from_param: Model from header/query/body 

338 """ 

339 if data is None: 339 ↛ 340line 339 didn't jump to line 340 because the condition on line 339 was never true

340 data = {} 

341 

342 # Check if file_id has embedded model info 

343 model_from_id: Final = decode_model_from_file_id(file_id) 

344 

345 # Check other sources for model parameter 

346 model_from_param = data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model") 

347 

348 return model_from_id, model_from_param 

349 

350 

351def get_credentials_for_model( 

352 llm_router, # Router instance 

353 model_id: str, 

354 operation_context: str = "file operation", 

355): 

356 """ 

357 Retrieve API credentials for a model from the LLM Router. 

358 

359 Does not check whether the caller may use ``model_id``; use 

360 ``get_authorized_credentials_for_model`` for anything driven by a caller-supplied 

361 model name (request body, header, query param, or a model-encoded resource id). 

362 

363 Args: 

364 llm_router: LiteLLM Router instance 

365 model_id: Model name or deployment ID 

366 operation_context: Description for error messages (e.g., "file upload", "batch creation") 

367 

368 Returns: 

369 Dictionary with credentials (api_key, api_base, custom_llm_provider, etc.) 

370 

371 Raises: 

372 HTTPException: If router not initialized or model not found 

373 """ 

374 from fastapi import HTTPException 

375 

376 from litellm.proxy.route_llm_request import ProxyModelNotFoundError 

377 

378 if llm_router is None: 378 ↛ 379line 378 didn't jump to line 379 because the condition on line 378 was never true

379 raise HTTPException( 

380 status_code=500, 

381 detail={"error": "Router not initialized. Cannot use model-based routing."}, 

382 ) 

383 

384 credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id) 

385 

386 if credentials is None: 386 ↛ 391line 386 didn't jump to line 391 because the condition on line 386 was always true

387 raise ProxyModelNotFoundError( 

388 route=operation_context, model_name=model_id, retryable_with_model_read_through=False 

389 ) 

390 

391 return credentials 

392 

393 

394async def authorize_model_for_key( 

395 model_id: str, 

396 llm_router: Optional["Router"], 

397 user_api_key_dict: "UserAPIKeyAuth", 

398) -> None: 

399 """ 

400 Enforce the caller's model grants on a model name the auth layer never saw. 

401 

402 The files and batches routes carry their model in a header, query param, or a 

403 model-encoded resource id rather than the request body, so ``user_api_key_auth`` 

404 cannot check it. Run the same key, team (incl. team-member and access-group 

405 fallbacks), org and project allowlist checks a chat request would get, so a 

406 restricted key cannot borrow another deployment's server-side credentials. 

407 

408 Raises: 

409 ProxyException (403): the caller is not allowed to use ``model_id`` 

410 """ 

411 from litellm.proxy.auth.auth_checks import can_key_call_resolved_model 

412 

413 await can_key_call_resolved_model( 

414 model=model_id, 

415 llm_model_list=None, 

416 valid_token=user_api_key_dict, 

417 llm_router=llm_router, 

418 ) 

419 

420 

421async def get_authorized_credentials_for_model( 

422 llm_router: Optional["Router"], 

423 model_id: str, 

424 user_api_key_dict: "UserAPIKeyAuth", 

425 operation_context: str = "file operation", 

426) -> dict: # mutable-ok: same contract as get_credentials_for_model, callers merge it into request data 

427 """``get_credentials_for_model`` gated by ``authorize_model_for_key``.""" 

428 await authorize_model_for_key(model_id=model_id, llm_router=llm_router, user_api_key_dict=user_api_key_dict) 

429 return get_credentials_for_model( 

430 llm_router=llm_router, 

431 model_id=model_id, 

432 operation_context=operation_context, 

433 ) 

434 

435 

436def get_team_provider_credentials( 

437 llm_router: Optional["Router"], 

438 user_api_key_dict: "UserAPIKeyAuth", 

439 custom_llm_provider: str, 

440) -> dict | None: 

441 """ 

442 Resolve upstream credentials for a provider-scoped file operation 

443 (e.g. GET /v1/files), which doesn't pin a model. 

444 

445 Priority: 

446 1. The team's own (BYOK) deployment for this provider — a deployment whose 

447 ``model_info.team_id`` matches the caller's team. This keeps team-scoped 

448 listings on the team's own provider account/key instead of a shared 

449 global one. 

450 2. Fallback: any deployment the caller is granted access to for this 

451 provider, expanding wildcard routes and the all-proxy-models sentinel. 

452 

453 Credential lookup is scoped to both the team's allowlist and the key's own 

454 model allowlist (``user_api_key_dict.models``), so neither a team nor a 

455 restricted key within a team can resolve a provider key for a deployment 

456 it isn't authorized to use. A key restricted to an explicit model list 

457 only narrows the team scope; sentinel-bearing keys (all-proxy-models / 

458 all-team-models) defer to the team scope instead of widening past it. 

459 Returns None when the router is unavailable or no authorized deployment 

460 matches, so the caller can fall back to default credential resolution. 

461 """ 

462 if llm_router is None: 462 ↛ 463line 462 didn't jump to line 463 because the condition on line 462 was never true

463 return None 

464 

465 from litellm.proxy._types import SpecialModelNames 

466 from litellm.proxy.auth.model_checks import get_complete_model_list, get_key_models 

467 

468 team_id: Final = user_api_key_dict.team_id 

469 team_models: Final = user_api_key_dict.team_models or [] 

470 

471 proxy_model_list: Final = llm_router.get_model_names(team_id=team_id) 

472 model_access_groups: Final = llm_router.get_model_access_groups() 

473 

474 raw_key_models: Final = user_api_key_dict.models or [] 

475 sentinel_values: Final = { 

476 SpecialModelNames.all_proxy_models.value, 

477 SpecialModelNames.all_team_models.value, 

478 } 

479 key_is_restricted: Final = bool(raw_key_models) and not (set(raw_key_models) & sentinel_values) 

480 key_model_allowlist: Final = ( 

481 tuple( 

482 dict.fromkeys( 

483 get_key_models( 

484 user_api_key_dict=user_api_key_dict, 

485 proxy_model_list=proxy_model_list, 

486 model_access_groups=model_access_groups, 

487 ) 

488 ) 

489 ) 

490 if key_is_restricted 

491 else () 

492 ) 

493 key_model_allowlist_set: Final = frozenset(key_model_allowlist) 

494 

495 def _key_may_use(public_model_name: str | None) -> bool: 

496 if not key_model_allowlist_set: 

497 return True 

498 return public_model_name is not None and public_model_name in key_model_allowlist_set 

499 

500 def _provider_credentials(model_id: str) -> dict | None: 

501 credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id) 

502 if credentials is not None and credentials.get("custom_llm_provider") == custom_llm_provider: 

503 return {key: value for key, value in credentials.items() if key != "model"} 

504 return None 

505 

506 # 1. Prefer the team's own BYOK deployment, matched by model_info.team_id. 

507 if team_id is not None: 507 ↛ 508line 507 didn't jump to line 508 because the condition on line 507 was never true

508 for deployment in llm_router.model_list or []: 

509 model_info = deployment.get("model_info") or {} 

510 if model_info.get("team_id") != team_id: 

511 continue 

512 deployment_id = model_info.get("id") 

513 if deployment_id is None: 

514 continue 

515 if not _key_may_use(model_info.get("team_public_model_name") or deployment.get("model_name")): 

516 continue 

517 credentials = _provider_credentials(deployment_id) 

518 if credentials is not None: 

519 return credentials 

520 

521 # 2. Fall back to deployments the caller is allowed to access. The key's 

522 # effective allowlist (sentinels and access groups already expanded by 

523 # get_key_models) wins when set; otherwise the team's allowlist applies. 

524 # The all-proxy-models sentinel isn't expanded by 

525 # get_complete_model_list, so normalize it to an empty allowlist, which 

526 # defers to the team-scoped proxy model list. A team or key with a 

527 # restricted allowlist (e.g. anthropic only) therefore never resolves 

528 # another provider's key. 

529 grants_all_models: Final = SpecialModelNames.all_proxy_models.value in team_models 

530 effective_team_models: Final = [] if grants_all_models else team_models 

531 

532 models_to_try: Final = list( 

533 dict.fromkeys( 

534 get_complete_model_list( 

535 key_models=list(key_model_allowlist), 

536 team_models=effective_team_models, 

537 proxy_model_list=proxy_model_list, 

538 user_model=None, 

539 infer_model_from_keys=False, 

540 return_wildcard_routes=True, 

541 llm_router=llm_router, 

542 model_access_groups=model_access_groups, 

543 include_model_access_groups=True, 

544 team_id=team_id, 

545 ) 

546 ) 

547 ) 

548 for model_name in models_to_try: 

549 credentials = _provider_credentials(model_name) 

550 if credentials is not None: 

551 return credentials 

552 

553 return None 

554 

555 

556def apply_team_provider_credentials( 

557 data: dict, # mutable-ok: credentials are merged into the request payload in place, same contract as prepare_data_with_credentials 

558 llm_router: Optional["Router"], 

559 user_api_key_dict: "UserAPIKeyAuth", 

560 custom_llm_provider: str, 

561) -> None: 

562 """ 

563 Resolve credentials for a provider-only request (no model pinned) via 

564 ``get_team_provider_credentials`` and merge them into ``data`` in-place. 

565 Leaves ``data`` untouched when no authorized deployment matches, so the 

566 caller falls back to environment-variable credentials exactly as before. 

567 """ 

568 credentials: Final = get_team_provider_credentials( 

569 llm_router=llm_router, 

570 user_api_key_dict=user_api_key_dict, 

571 custom_llm_provider=custom_llm_provider, 

572 ) 

573 if credentials is None: 

574 return 

575 prepare_data_with_credentials(data=data, credentials=credentials) 

576 

577 

578def add_internal_model_credentials( 

579 data: dict, 

580 llm_router: "Router", 

581 model_id: str | None, 

582) -> None: 

583 """ 

584 Attach the deployment's immutable server-side credential snapshot to a router-routed 

585 batch call (in-place). 

586 

587 Cost accounting for a completed batch reads the batch's output file, and the Bedrock 

588 file config resolves its bucket only from this snapshot, never from a request param, 

589 because the bucket is what managed file ids are validated against. Without it that 

590 read fails and the batch's cost is never recorded. 

591 """ 

592 if model_id is None: 

593 return 

594 try: 

595 credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id) 

596 except Exception: # noqa: BLE001 # the snapshot only enables cost accounting; a batch whose deployment no longer resolves must still be retrievable 

597 return 

598 if credentials is None: 

599 return 

600 data["_litellm_internal_model_credentials"] = MappingProxyType(dict(credentials)) 

601 

602 

603def add_deployment_model_info( 

604 data: dict, 

605 llm_router: Optional["Router"], 

606 model_id: str, 

607) -> None: 

608 """ 

609 Stamp the resolved deployment's `model_info` onto a direct (non-router) batch call 

610 (in-place), the way the router does for routed calls, so the completed batch is 

611 priced by its deployment id instead of the published model rate. 

612 """ 

613 deployment: Final = llm_router.get_credential_deployment(model_id=model_id) if llm_router is not None else None 

614 if deployment is None: 

615 return 

616 data["litellm_metadata"] = { 

617 **(data.get("litellm_metadata") or {}), 

618 "model_info": deployment.model_info.model_dump(), 

619 } 

620 

621 

622def prepare_data_with_credentials( 

623 data: dict, 

624 credentials: dict, 

625 file_id: str | None = None, 

626 include_internal_credentials: bool = False, 

627) -> None: 

628 """ 

629 Update data dictionary with model credentials (in-place). 

630 

631 Args: 

632 data: Data dictionary to update 

633 credentials: Credentials from router 

634 file_id: Optional original file_id to set (for decoded file IDs) 

635 include_internal_credentials: Preserve an immutable server-side snapshot 

636 for code paths that must distinguish proxy config from request params. 

637 """ 

638 data.update(credentials) 

639 if include_internal_credentials: 639 ↛ 640line 639 didn't jump to line 640 because the condition on line 639 was never true

640 data["_litellm_internal_model_credentials"] = MappingProxyType(dict(credentials)) 

641 data.pop("custom_llm_provider", None) 

642 

643 if file_id is not None: 643 ↛ 644line 643 didn't jump to line 644 because the condition on line 643 was never true

644 data["file_id"] = file_id 

645 

646 

647async def handle_model_based_routing( 

648 file_id: str, 

649 request, # FastAPI Request object 

650 llm_router, # Router instance 

651 data: dict, 

652 user_api_key_dict: "UserAPIKeyAuth", 

653 check_file_id_encoding: bool = True, 

654) -> tuple[bool, str | None, str | None, dict | None]: 

655 """ 

656 Orchestrate model-based credential routing for file operations. 

657 

658 The model name comes from the caller (embedded in the file id, or a header, query 

659 param or body field), so it is authorized against the caller's key, team, org and 

660 project grants before any deployment credentials are resolved. 

661 

662 Args: 

663 file_id: File ID (may contain embedded model info) 

664 request: FastAPI request object 

665 llm_router: LiteLLM Router instance 

666 data: Request data dictionary 

667 user_api_key_dict: The authenticated caller 

668 check_file_id_encoding: Whether to check for embedded model in file_id 

669 

670 Returns: 

671 Tuple of (should_use_model_routing, model_used, original_file_id, credentials) 

672 - should_use_model_routing: True if model-based routing should be used 

673 - model_used: The model name being used 

674 - original_file_id: Decoded file ID (if it was encoded) 

675 - credentials: Model credentials dict 

676 

677 Raises: 

678 HTTPException: If router unavailable or model not found 

679 ProxyException: If the caller is not allowed to use the model 

680 """ 

681 model_from_id, model_from_param = extract_model_from_sources( 

682 file_id=file_id, 

683 request=request, 

684 data=data, 

685 ) 

686 

687 # Priority 1: Model embedded in file_id 

688 if check_file_id_encoding and model_from_id is not None: 688 ↛ 689line 688 didn't jump to line 689 because the condition on line 688 was never true

689 credentials = await get_authorized_credentials_for_model( 

690 llm_router=llm_router, 

691 model_id=model_from_id, 

692 user_api_key_dict=user_api_key_dict, 

693 operation_context="file operation (file created with model)", 

694 ) 

695 original_file_id: Final = get_original_file_id(file_id) 

696 return True, model_from_id, original_file_id, credentials 

697 

698 # Priority 2: Model from header/query/body 

699 elif model_from_param is not None: 699 ↛ 700line 699 didn't jump to line 700 because the condition on line 699 was never true

700 credentials = await get_authorized_credentials_for_model( 

701 llm_router=llm_router, 

702 model_id=model_from_param, 

703 user_api_key_dict=user_api_key_dict, 

704 operation_context="file operation", 

705 ) 

706 return True, model_from_param, None, credentials 

707 

708 # No model-based routing needed 

709 return False, None, None, None 

710 

711 

712# ============================================================================ 

713# MIME TYPE DETECTION AND NORMALIZATION 

714# ============================================================================ 

715 

716 

717# Gemini-supported image MIME types 

718GEMINI_SUPPORTED_IMAGE_TYPES: Final = { 

719 "image/png", 

720 "image/jpeg", 

721 "image/webp", 

722} 

723 

724# Gemini-supported video MIME types 

725GEMINI_SUPPORTED_VIDEO_TYPES: Final = { 

726 "video/3gpp", 

727 "video/wmv", 

728 "video/webm", 

729 "video/mp4", 

730 "video/mpg", 

731 "video/mpegps", 

732 "video/mpeg", 

733 "video/quicktime", 

734 "video/x-flv", 

735} 

736 

737# Gemini-supported audio MIME types 

738GEMINI_SUPPORTED_AUDIO_TYPES: Final = { 

739 "audio/webm", 

740 "audio/wav", 

741 "audio/pcm", 

742 "audio/opus", 

743 "audio/mp4", 

744 "audio/mpga", 

745 "audio/mpeg", 

746 "audio/m4a", 

747 "audio/mp3", 

748 "audio/flac", 

749 "audio/aac", 

750} 

751 

752# Gemini-supported document MIME types 

753GEMINI_SUPPORTED_DOCUMENT_TYPES: Final = { 

754 "text/plain", 

755 "application/pdf", 

756} 

757 

758# Mapping of common file extensions to MIME types 

759# This extends Python's mimetypes with custom mappings 

760EXTENSION_TO_MIME_TYPE: Final = { 

761 ".jpg": "image/jpeg", # Normalize jpg to jpeg 

762 ".jpeg": "image/jpeg", 

763 ".png": "image/png", 

764 ".webp": "image/webp", 

765 ".pdf": "application/pdf", 

766 ".mp3": "audio/mpeg", 

767 ".wav": "audio/wav", 

768 ".m4a": "audio/mp4", 

769} 

770 

771 

772def detect_content_type_from_filename(filename: str) -> str: 

773 """ 

774 Detect content type from filename using extension. 

775 

776 Uses Python's mimetypes module with custom overrides for common cases. 

777 Normalizes jpg to jpeg for consistency. 

778 """ 

779 if not filename: 

780 return "application/octet-stream" 

781 

782 # Try custom mapping first 

783 filename_lower: Final = filename.lower() 

784 for ext, mime_type in EXTENSION_TO_MIME_TYPE.items(): 

785 if filename_lower.endswith(ext): 

786 return mime_type 

787 

788 # Fall back to Python's mimetypes 

789 mime_type_guess, _ = mimetypes.guess_type(filename) 

790 if mime_type_guess is not None: 

791 return mime_type_guess 

792 

793 return "application/octet-stream" 

794 

795 

796def normalize_mime_type_for_provider(mime_type: str, provider: str | None = None) -> str: 

797 """ 

798 Normalize MIME type for specific provider requirements. 

799 

800 Currently handles: 

801 - Gemini: Normalizes image/jpg to image/jpeg 

802 

803 Args: 

804 mime_type: Original MIME type 

805 provider: Provider name (e.g., "gemini", "vertex_ai") 

806 

807 Returns: 

808 str: Normalized MIME type 

809 """ 

810 normalized = mime_type.lower().strip() 

811 

812 # Gemini/Vertex AI requires image/jpeg, not image/jpg 

813 if provider and ("gemini" in provider.lower() or "vertex_ai" in provider.lower()): 

814 if normalized == "image/jpg": 

815 normalized = "image/jpeg" 

816 

817 # General normalization: always normalize jpg to jpeg 

818 if normalized == "image/jpg": 

819 normalized = "image/jpeg" 

820 

821 return normalized 

822 

823 

824def is_gemini_supported_mime_type(mime_type: str) -> bool: 

825 """ 

826 Check if a MIME type is supported by Gemini multimodal models. 

827 

828 Supported categories: 

829 - Images: image/png, image/jpeg, image/webp 

830 - Video: 3gpp, wmv, webm, mp4, mpg, mpegps, mpeg, quicktime, x-flv 

831 - Audio: webm, wav, pcm, opus, mp4, mpga, mpeg, m4a, mp3, flac, aac 

832 - Documents: text/plain, application/pdf 

833 

834 Args: 

835 mime_type: MIME type to check 

836 

837 Returns: 

838 bool: True if supported, False otherwise 

839 """ 

840 normalized: Final = normalize_mime_type_for_provider(mime_type, provider="gemini") 

841 return normalized in ( 

842 GEMINI_SUPPORTED_IMAGE_TYPES 

843 | GEMINI_SUPPORTED_VIDEO_TYPES 

844 | GEMINI_SUPPORTED_AUDIO_TYPES 

845 | GEMINI_SUPPORTED_DOCUMENT_TYPES 

846 ) 

847 

848 

849def get_content_type_from_file_object(file_object: dict | None) -> str: 

850 """ 

851 Determine content type from file object (from database or API response). 

852 

853 Extracts filename from file object and uses detect_content_type_from_filename. 

854 Falls back to default if file object is invalid or filename not found. 

855 

856 Args: 

857 file_object: File object dictionary (can be None) 

858 

859 Returns: 

860 str: MIME type (defaults to "application/octet-stream" if cannot be determined) 

861 """ 

862 if not file_object: 

863 return "application/octet-stream" 

864 

865 # Handle JSON string 

866 if isinstance(file_object, str): 

867 import json 

868 

869 try: 

870 file_object = json.loads(file_object) 

871 except json.JSONDecodeError: 

872 return "application/octet-stream" 

873 

874 if not isinstance(file_object, dict): 

875 return "application/octet-stream" 

876 

877 # Try to get filename 

878 filename: Final = file_object.get("filename", "") 

879 if filename: 

880 return detect_content_type_from_filename(filename) 

881 

882 return "application/octet-stream" 

883 

884 

885# ============================================================================ 

886# REQUEST PARAMETER EXTRACTION 

887# ============================================================================ 

888 

889 

890@dataclass 

891class FileCreationParams: 

892 """ 

893 Structured parameters extracted from file creation requests. 

894 

895 Attributes: 

896 target_storage: Storage backend name (e.g., "azure_storage", "default") 

897 target_model_names: List of model names for managed files 

898 model: Model parameter for multi-account routing 

899 """ 

900 

901 target_storage: str = "default" 

902 target_model_names: list[str] = field(default_factory=list) 

903 model: str | None = None 

904 

905 def __post_init__(self): 

906 """Normalize and validate parameters after initialization.""" 

907 if self.target_model_names is None: 907 ↛ 908line 907 didn't jump to line 908 because the condition on line 907 was never true

908 self.target_model_names = [] 

909 

910 # Normalize target_storage 

911 if not self.target_storage: 911 ↛ 912line 911 didn't jump to line 912 because the condition on line 911 was never true

912 self.target_storage = "default" 

913 

914 # Strip whitespace from model names 

915 self.target_model_names = [name.strip() for name in self.target_model_names if name.strip()] 

916 

917 

918async def extract_file_creation_params( 

919 request: "Request", 

920 request_body: dict | None = None, 

921 target_model_names_form: str | None = None, 

922 target_storage_form: str | None = None, 

923) -> FileCreationParams: 

924 """ 

925 Extract file creation parameters from request. 

926 

927 Args: 

928 request: FastAPI request object 

929 request_body: Optional pre-parsed request body 

930 target_model_names_form: target_model_names from form field (comma-separated string) 

931 target_storage_form: target_storage from form field (defaults to "default") 

932 

933 Returns: 

934 FileCreationParams: Structured parameters extracted from the request 

935 """ 

936 from litellm.proxy.common_utils.http_parsing_utils import _read_request_body 

937 

938 if request_body is None: 938 ↛ 939line 938 didn't jump to line 939 because the condition on line 938 was never true

939 request_body = await _read_request_body(request=request) or {} 

940 

941 # Extract target_storage (simplified - just use form parameter) 

942 target_storage: Final = _extract_target_storage_simple(target_storage_form) 

943 

944 # Extract target_model_names from the form field, then fall back to the raw form 

945 target_model_names = _extract_target_model_names_simple(target_model_names_form) 

946 if not target_model_names: 

947 target_model_names = await _extract_target_model_names_from_form(request) 

948 

949 # Extract model parameter 

950 model: Final = _extract_model_param(request, request_body) 

951 

952 return FileCreationParams( 

953 target_storage=target_storage, 

954 target_model_names=target_model_names, 

955 model=model, 

956 ) 

957 

958 

959def _extract_target_storage_simple(target_storage_form: str | None = None) -> str: 

960 """ 

961 Extract target_storage parameter from form field. 

962 

963 Args: 

964 target_storage_form: target_storage from form field 

965 

966 Returns: 

967 str: Target storage backend name, or "default" 

968 """ 

969 if target_storage_form: 969 ↛ 971line 969 didn't jump to line 971 because the condition on line 969 was always true

970 return target_storage_form.strip() 

971 return "default" 

972 

973 

974def _extract_target_model_names_simple( 

975 target_model_names_form: str | None = None, 

976) -> list[str]: 

977 """ 

978 Extract target_model_names parameter from form field. 

979 """ 

980 if not target_model_names_form: 

981 return [] 

982 

983 # Parse comma-separated string into list 

984 if isinstance(target_model_names_form, str): 984 ↛ 986line 984 didn't jump to line 986 because the condition on line 984 was always true

985 return [name.strip() for name in target_model_names_form.split(",") if name.strip()] 

986 elif isinstance(target_model_names_form, list): 

987 return [str(name).strip() for name in target_model_names_form if name] 

988 

989 return [] 

990 

991 

992def _is_target_model_names_key(key: str) -> bool: 

993 return key == "target_model_names" or (key.startswith("target_model_names[") and key.endswith("]")) 

994 

995 

996async def _extract_target_model_names_from_form(request: "Request") -> list[str]: 

997 """ 

998 Collect target_model_names from the raw multipart form. 

999 

1000 Reads ``request.form()`` directly instead of the parsed request body, which is 

1001 built via ``dict(form_data)`` and keeps only the last value for repeated keys. 

1002 The OpenAI SDK sends a list ``extra_body`` as repeated ``target_model_names[]`` 

1003 fields, so reading the form preserves every value instead of truncating to one. 

1004 Indexed keys like ``target_model_names[0]`` are handled the same way. 

1005 """ 

1006 form_data: Final = await request.form() 

1007 

1008 names: Final[list[str]] = [] 

1009 for key, value in form_data.multi_items(): 

1010 if _is_target_model_names_key(key) and isinstance(value, str): 

1011 names.extend(_extract_target_model_names_simple(value)) 

1012 

1013 seen: Final = set() 

1014 result: Final[list[str]] = [] 

1015 for name in names: 1015 ↛ 1016line 1015 didn't jump to line 1016 because the loop on line 1015 never started

1016 if name and name not in seen: 

1017 seen.add(name) 

1018 result.append(name) 

1019 return result 

1020 

1021 

1022def validate_managed_files_requirement( 

1023 target_model_names: list[str], 

1024 model: str | None = None, 

1025) -> None: 

1026 """ 

1027 Enforce proxy-level managed files when litellm.require_managed_files is enabled. 

1028 

1029 Raises: 

1030 HTTPException: 400 if the upload would bypass the managed-files flow, i.e. 

1031 target_model_names is missing or a model parameter routes the request 

1032 through the direct provider path instead of the managed-files hook. 

1033 """ 

1034 from fastapi import HTTPException 

1035 

1036 import litellm 

1037 

1038 if litellm.require_managed_files is not True: 1038 ↛ 1041line 1038 didn't jump to line 1041 because the condition on line 1038 was always true

1039 return 

1040 

1041 if not target_model_names: 

1042 raise HTTPException( 

1043 status_code=400, 

1044 detail=( 

1045 "target_model_names is required when require_managed_files is enabled " 

1046 "in litellm_settings. Provide one or more model aliases via the " 

1047 "target_model_names form field (e.g. target_model_names=my-model-alias)." 

1048 ), 

1049 ) 

1050 

1051 if model: 

1052 raise HTTPException( 

1053 status_code=400, 

1054 detail=( 

1055 "model is not allowed when require_managed_files is enabled in " 

1056 "litellm_settings. Uploads must go through managed files using " 

1057 "target_model_names instead of the model parameter." 

1058 ), 

1059 ) 

1060 

1061 

1062async def validate_managed_id_requirement( 

1063 resource_id: str | None, 

1064 resource_kind: Literal["file", "batch", "fine-tuning job"], 

1065 user_api_key_dict: "UserAPIKeyAuth", 

1066 managed_files_obj: object | None, 

1067) -> None: 

1068 """ 

1069 Enforce proxy-level managed resources on every route that accepts a provider-issued id 

1070 when ``litellm.require_managed_files`` is enabled, and authenticate managed ids against 

1071 the caller's stored ownership record. 

1072 

1073 Ownership is only recorded for LiteLLM managed ids, so a raw provider id is forwarded to the 

1074 provider under shared credentials without any tenant check; knowing another tenant's provider 

1075 id would be enough to read, reuse, or destroy the object behind it. 

1076 

1077 Raises: 

1078 HTTPException: 400 for a raw id, 403 for an inaccessible managed id, or 500 when 

1079 ownership validation is unavailable. 

1080 """ 

1081 from fastapi import HTTPException 

1082 

1083 import litellm 

1084 

1085 if litellm.require_managed_files is not True: 1085 ↛ 1088line 1085 didn't jump to line 1088 because the condition on line 1085 was always true

1086 return 

1087 

1088 if not resource_id: 

1089 return 

1090 

1091 if not _is_base64_encoded_unified_file_id(resource_id): 

1092 raise HTTPException( 

1093 status_code=400, 

1094 detail=( 

1095 f"Raw provider {resource_kind} ids cannot be used when require_managed_files is enabled in " 

1096 f"litellm_settings. Use the LiteLLM managed {resource_kind} id returned when the " 

1097 f"{resource_kind} was created." 

1098 ), 

1099 ) 

1100 

1101 if not isinstance(managed_files_obj, ManagedResourceAccessChecker): 

1102 raise HTTPException( 

1103 status_code=500, 

1104 detail="Managed resource ownership validation is unavailable.", 

1105 ) 

1106 

1107 can_access: Final = ( 

1108 await managed_files_obj.can_user_call_unified_file_id(resource_id, user_api_key_dict) 

1109 if resource_kind == "file" 

1110 else await managed_files_obj.can_user_call_unified_object_id(resource_id, user_api_key_dict) 

1111 ) 

1112 if can_access: 

1113 return 

1114 

1115 raise HTTPException( 

1116 status_code=403, 

1117 detail=f"The caller does not have access to this managed {resource_kind} id.", 

1118 ) 

1119 

1120 

1121def _extract_model_param(request: "Request", request_body: dict) -> str | None: 

1122 """ 

1123 Extract model parameter from request. 

1124 

1125 Priority: 

1126 1. request_body.model 

1127 2. Query parameter (?model=) 

1128 3. Header (x-litellm-model) 

1129 """ 

1130 return request_body.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model") 

1131 

1132 

1133# ============================================================================ 

1134# BATCH DATABASE OPERATIONS 

1135# ============================================================================ 

1136 

1137 

1138def _batch_response_model_id_candidates( 

1139 response, 

1140 unified_batch_id: str | Literal[False] | None, 

1141) -> tuple[str, ...]: 

1142 response_id: Final = getattr(response, "id", None) 

1143 decoded_response_id: Final = ( 

1144 _is_base64_encoded_unified_file_id(response_id) if isinstance(response_id, str) else False 

1145 ) 

1146 return tuple( 

1147 candidate 

1148 for candidate in ( 

1149 unified_batch_id if isinstance(unified_batch_id, str) else None, 

1150 decoded_response_id or None, 

1151 response_id 

1152 if isinstance(response_id, str) 

1153 and not decoded_response_id 

1154 and response_id.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value) 

1155 else None, 

1156 ) 

1157 if candidate 

1158 ) 

1159 

1160 

1161def _model_id_for_batch_response( 

1162 response: "LiteLLMBatch", 

1163 unified_batch_id: str | Literal[False] | None, 

1164) -> str | None: 

1165 hidden_params: Final = getattr(response, "_hidden_params", None) or {} 

1166 model_id: Final = hidden_params.get("model_id") 

1167 if model_id: 

1168 return model_id 

1169 return next( 

1170 ( 

1171 candidate_model_id 

1172 for candidate in _batch_response_model_id_candidates(response, unified_batch_id) 

1173 if (candidate_model_id := get_model_id_from_unified_batch_id(candidate)) 

1174 ), 

1175 None, 

1176 ) 

1177 

1178 

1179def _model_name_for_batch_response(response: "LiteLLMBatch") -> str | None: 

1180 hidden_params: Final = getattr(response, "_hidden_params", None) or {} 

1181 unified_file_id: Final = hidden_params.get("unified_file_id") 

1182 return resolve_managed_output_file_model_name( 

1183 unified_input_file_id=unified_file_id 

1184 if isinstance(unified_file_id, str) 

1185 else getattr(response, "input_file_id", None), 

1186 fallback_model_name=hidden_params.get("model_name"), 

1187 ) 

1188 

1189 

1190def _batch_owner_auth_from_db_object(db_batch_object: "LiteLLM_ManagedObjectTable") -> "UserAPIKeyAuth | None": 

1191 from litellm.proxy._types import UserAPIKeyAuth 

1192 

1193 created_by: Final = getattr(db_batch_object, "created_by", None) 

1194 if not isinstance(created_by, str) or not created_by: 

1195 return None 

1196 raw_team_id: Final = getattr(db_batch_object, "team_id", None) 

1197 return UserAPIKeyAuth(user_id=created_by, team_id=raw_team_id if isinstance(raw_team_id, str) else None) 

1198 

1199 

1200async def resolve_input_file_id_to_unified(response, prisma_client) -> None: 

1201 """ 

1202 If the batch response contains a raw provider input_file_id (not already a 

1203 unified ID), look up the corresponding unified file ID from the managed file 

1204 table and replace it in-place. 

1205 """ 

1206 if ( 

1207 hasattr(response, "input_file_id") 

1208 and response.input_file_id 

1209 and not _is_base64_encoded_unified_file_id(response.input_file_id) 

1210 and prisma_client 

1211 ): 

1212 try: 

1213 managed_file: Final = await ManagedFileRepository(prisma_client).table.find_first( 

1214 where={"flat_model_file_ids": {"has": response.input_file_id}} 

1215 ) 

1216 if managed_file: 

1217 response.input_file_id = managed_file.unified_file_id 

1218 except Exception: 

1219 pass 

1220 

1221 

1222async def resolve_output_file_ids_to_unified(response, prisma_client) -> None: 

1223 """ 

1224 If the batch response contains raw provider output_file_id or error_file_id 

1225 (not already unified IDs), look up the corresponding unified file IDs from 

1226 the managed file table and replace them in-place. 

1227 """ 

1228 if not prisma_client: 

1229 return 

1230 for attr in ("output_file_id", "error_file_id"): 

1231 raw_id = getattr(response, attr, None) 

1232 if not raw_id or _is_base64_encoded_unified_file_id(raw_id): 

1233 continue 

1234 try: 

1235 managed_file = await ManagedFileRepository(prisma_client).table.find_first( 

1236 where={"flat_model_file_ids": {"has": raw_id}} 

1237 ) 

1238 if managed_file: 

1239 setattr(response, attr, managed_file.unified_file_id) 

1240 except Exception: 

1241 pass 

1242 

1243 

1244async def map_raw_file_ids_to_unified( 

1245 raw_file_ids: frozenset[str], prisma_client: "PrismaClient | None" 

1246) -> Mapping[str, str]: 

1247 if not raw_file_ids or not prisma_client: 1247 ↛ 1249line 1247 didn't jump to line 1249 because the condition on line 1247 was always true

1248 return MappingProxyType({}) 

1249 managed_files: Final = await ManagedFileRepository(prisma_client).table.find_many( 

1250 where={"flat_model_file_ids": {"hasSome": sorted(raw_file_ids)}} # mutable-ok: prisma where is a plain dict 

1251 ) 

1252 return MappingProxyType( 

1253 { 

1254 raw_id: managed_file.unified_file_id 

1255 for managed_file in managed_files 

1256 for raw_id in managed_file.flat_model_file_ids 

1257 if raw_id in raw_file_ids 

1258 } 

1259 ) 

1260 

1261 

1262def apply_unified_file_ids(response: "LiteLLMBatch", unified_id_by_raw_id: Mapping[str, str]) -> None: 

1263 for file_attr, raw_id in ( 

1264 ("input_file_id", getattr(response, "input_file_id", None)), 

1265 ("output_file_id", getattr(response, "output_file_id", None)), 

1266 ("error_file_id", getattr(response, "error_file_id", None)), 

1267 ): 

1268 if isinstance(raw_id, str) and raw_id in unified_id_by_raw_id: 

1269 setattr(response, file_attr, unified_id_by_raw_id[raw_id]) 

1270 

1271 

1272async def ensure_batch_response_managed_file_ids( 

1273 response, 

1274 managed_files_obj, 

1275 prisma_client, 

1276 verbose_proxy_logger, 

1277 user_api_key_dict=None, 

1278 db_batch_object: "LiteLLM_ManagedObjectTable | None" = None, 

1279 unified_batch_id: str | Literal[False] | None = None, 

1280) -> None: 

1281 """Normalize batch file IDs to managed unified IDs before DB persistence.""" 

1282 await resolve_input_file_id_to_unified(response, prisma_client) 

1283 await resolve_output_file_ids_to_unified(response, prisma_client) 

1284 

1285 if managed_files_obj is None: 

1286 return 

1287 

1288 model_id: Final = _model_id_for_batch_response(response, unified_batch_id) 

1289 if not model_id: 

1290 return 

1291 

1292 model_name: Final = _model_name_for_batch_response(response) 

1293 

1294 owner_auth: Final = _batch_owner_auth_from_db_object(db_batch_object) if db_batch_object is not None else None 

1295 effective_auth: Final = owner_auth if owner_auth is not None else user_api_key_dict 

1296 if effective_auth is None: 

1297 return 

1298 

1299 for file_attr in ("output_file_id", "error_file_id"): 

1300 raw_file_id = getattr(response, file_attr, None) 

1301 if not raw_file_id or _is_base64_encoded_unified_file_id(raw_file_id): 

1302 continue 

1303 try: 

1304 new_unified_file_id = managed_files_obj.get_unified_output_file_id( 

1305 output_file_id=raw_file_id, 

1306 model_id=model_id, 

1307 model_name=model_name, 

1308 ) 

1309 await managed_files_obj.store_unified_file_id( 

1310 file_id=new_unified_file_id, 

1311 file_object=None, 

1312 litellm_parent_otel_span=getattr(effective_auth, "parent_otel_span", None), 

1313 model_mappings={model_id: raw_file_id}, 

1314 user_api_key_dict=effective_auth, 

1315 ) 

1316 setattr(response, file_attr, new_unified_file_id) 

1317 verbose_proxy_logger.debug("Converted batch %s %r to managed ID before DB write", file_attr, raw_file_id) 

1318 except Exception as e: 

1319 verbose_proxy_logger.warning( 

1320 "Failed to convert batch %s=%r to managed ID before DB write: %s", file_attr, raw_file_id, e 

1321 ) 

1322 

1323 

1324async def get_batch_from_database( 

1325 batch_id: str, 

1326 unified_batch_id: str | Literal[False], 

1327 managed_files_obj, 

1328 prisma_client, 

1329 verbose_proxy_logger, 

1330): 

1331 """ 

1332 Try to retrieve batch object from ManagedObjectTable for consistent state. 

1333 

1334 Args: 

1335 batch_id: The batch ID (may be unified/encoded) 

1336 unified_batch_id: Result from _is_base64_encoded_unified_file_id() 

1337 managed_files_obj: The managed_files proxy hook object 

1338 prisma_client: Prisma database client 

1339 verbose_proxy_logger: Logger instance 

1340 

1341 Returns: 

1342 Tuple of (db_batch_object, response_batch) 

1343 - db_batch_object: Raw database object (or None) 

1344 - response_batch: Parsed LiteLLMBatch object (or None) 

1345 """ 

1346 import json 

1347 

1348 from litellm.types.utils import LiteLLMBatch 

1349 

1350 if managed_files_obj is None or not unified_batch_id: 1350 ↛ 1353line 1350 didn't jump to line 1353 because the condition on line 1350 was always true

1351 return None, None 

1352 

1353 try: 

1354 if not prisma_client: 

1355 return None, None 

1356 

1357 db_batch_object: Final = await ManagedObjectRepository(prisma_client).table.find_first( 

1358 where={"unified_object_id": batch_id} 

1359 ) 

1360 

1361 if not db_batch_object or not db_batch_object.file_object: 

1362 return None, None 

1363 

1364 # Parse the batch object from database 

1365 file_object: Final = cast( # cast-ok: prisma types the Json column as str; reads return the decoded value 

1366 "Mapping[str, object] | str", db_batch_object.file_object 

1367 ) 

1368 batch_data: Final = json.loads(file_object) if isinstance(file_object, str) else file_object 

1369 response: Final = LiteLLMBatch.model_validate(batch_data) 

1370 response.id = batch_id 

1371 

1372 # The stored batch object may have raw provider file IDs. Register any missing 

1373 # managed-file rows and normalize output/error IDs before returning. 

1374 await ensure_batch_response_managed_file_ids( 

1375 response=response, 

1376 managed_files_obj=managed_files_obj, 

1377 prisma_client=prisma_client, 

1378 verbose_proxy_logger=verbose_proxy_logger, 

1379 db_batch_object=db_batch_object, 

1380 unified_batch_id=unified_batch_id, 

1381 ) 

1382 

1383 verbose_proxy_logger.debug( 

1384 "Retrieved batch %s from ManagedObjectTable with status=%s", batch_id, response.status 

1385 ) 

1386 

1387 return db_batch_object, response 

1388 

1389 except Exception as e: 

1390 verbose_proxy_logger.warning( 

1391 "Failed to retrieve batch from ManagedObjectTable: %s, falling back to provider", e 

1392 ) 

1393 return None, None 

1394 

1395 

1396def batch_cost_poller_is_active() -> bool: 

1397 """ 

1398 Whether the CheckBatchCost poller will account for a managed batch's cost itself. 

1399 

1400 False whenever the poller cannot be relied on: polling disabled by config, the job 

1401 absent from the scheduler because the enterprise import failed, or the poller not 

1402 yet having confirmed that the batch_processed column exists. That last condition 

1403 matters because the poller needs the column both to find outstanding batches and to 

1404 mark them accounted; without it the poller falls back to a query that excludes 

1405 terminal statuses, so a batch the retrieve path has already marked complete becomes 

1406 invisible to it. Defaulting to False until the poller confirms support keeps the 

1407 retrieve path accounting in exactly the cases the poller would drop the batch. 

1408 """ 

1409 from litellm.constants import PROXY_BATCH_POLLING_ENABLED 

1410 

1411 if not PROXY_BATCH_POLLING_ENABLED: 

1412 return False 

1413 try: 

1414 import litellm.proxy.proxy_server as proxy_server_module 

1415 

1416 scheduler = getattr(proxy_server_module, "scheduler", None) 

1417 if scheduler is None: 

1418 return False 

1419 job = scheduler.get_job("check_batch_cost_job") 

1420 if job is None: 

1421 return False 

1422 poller = getattr(getattr(job, "func", None), "__self__", None) 

1423 return getattr(poller, "batch_processed_support_confirmed", False) is True 

1424 except Exception: # noqa: BLE001 # scheduler backends raise varied types from get_job; an unreadable scheduler means the poller cannot be relied on 

1425 return False 

1426 

1427 

1428def _completed_batch_safe_to_retire(response: "LiteLLMBatch") -> bool: 

1429 """Whether a "completed" batch may be retired from cost recovery. 

1430 

1431 ``batch_processed=True`` is the sole re-pickup gate for CheckBatchCost's 

1432 cost-recovery poller, so setting it retires the batch permanently. A batch can 

1433 reach ``status="completed"`` while ``output_file_id`` is still ``None`` (the 

1434 provider response briefly lags before the output id populates). Retiring in that 

1435 window loses the spend record forever. Retire only once we can prove there is 

1436 nothing left to recover: the output file has actually arrived, or the provider 

1437 reported a positive total with zero successful request lines, proving it 

1438 enumerated the batch and none succeeded. A zero or unknown total means counts 

1439 are unreported, so stay eligible and let the next poller pass revisit it. (#37713) 

1440 """ 

1441 return batch_cost_is_final(response) 

1442 

1443 

1444async def update_batch_in_database( 

1445 batch_id: str, 

1446 unified_batch_id: str | Literal[False], 

1447 response, 

1448 managed_files_obj, 

1449 prisma_client, 

1450 verbose_proxy_logger, 

1451 db_batch_object: "LiteLLM_ManagedObjectTable | None" = None, 

1452 operation: str = "update", 

1453 user_api_key_dict=None, 

1454 poller_owns_accounting: bool | None = None, 

1455): 

1456 """ 

1457 Update batch status and object in ManagedObjectTable. 

1458 

1459 Args: 

1460 batch_id: The batch ID (unified/encoded) 

1461 unified_batch_id: Result from _is_base64_encoded_unified_file_id() 

1462 response: The batch response object with updated state 

1463 managed_files_obj: The managed_files proxy hook object 

1464 prisma_client: Prisma database client 

1465 verbose_proxy_logger: Logger instance 

1466 db_batch_object: Optional existing database object; fetched by unified_object_id when omitted 

1467 operation: Description of operation ("update", "cancel", etc.) 

1468 user_api_key_dict: Optional auth context for creating managed file IDs 

1469 poller_owns_accounting: Whether the caller already decided that the cost poller 

1470 owns this batch's accounting. Callers that suppress their own inline 

1471 accounting must pass the same decision they acted on, because re-deciding 

1472 here can observe a poller that became usable in between and leave the batch 

1473 unmarked after it was already accounted for, billing it twice. Left None by 

1474 callers that record no cost themselves. 

1475 """ 

1476 import litellm.utils 

1477 

1478 if managed_files_obj is None or not unified_batch_id: 

1479 return 

1480 

1481 try: 

1482 if not prisma_client: 

1483 return 

1484 

1485 effective_db_batch_object: Final = ( 

1486 db_batch_object 

1487 if db_batch_object is not None 

1488 else await ManagedObjectRepository(prisma_client).table.find_first(where={"unified_object_id": batch_id}) 

1489 ) 

1490 

1491 # Always normalize the response's file IDs to unified managed IDs 

1492 # (mutates in place) so the caller returns unified IDs to the user 

1493 # even when we skip the DB update below for an unchanged status. 

1494 await ensure_batch_response_managed_file_ids( 

1495 response=response, 

1496 managed_files_obj=managed_files_obj, 

1497 prisma_client=prisma_client, 

1498 verbose_proxy_logger=verbose_proxy_logger, 

1499 user_api_key_dict=user_api_key_dict, 

1500 db_batch_object=effective_db_batch_object, 

1501 unified_batch_id=unified_batch_id, 

1502 ) 

1503 

1504 # Only update if status has changed (when db_batch_object is provided) 

1505 if effective_db_batch_object and response.status == effective_db_batch_object.status: 

1506 return 

1507 

1508 if effective_db_batch_object: 

1509 verbose_proxy_logger.info( 

1510 "Updating batch %s status from %s to %s", batch_id, effective_db_batch_object.status, response.status 

1511 ) 

1512 else: 

1513 verbose_proxy_logger.info("Updating batch %s status to %s after %s", batch_id, response.status, operation) 

1514 

1515 # Normalize status for database storage 

1516 db_status: Final = response.status if response.status != "completed" else "complete" 

1517 

1518 update_data: Final[dict[str, object]] = { 

1519 "status": db_status, 

1520 "file_object": response.model_dump_json(), 

1521 "updated_at": litellm.utils.get_utc_datetime(), 

1522 } 

1523 

1524 poller_owns: Final = batch_cost_poller_is_active() if poller_owns_accounting is None else poller_owns_accounting 

1525 if db_status == "complete" and not poller_owns and _completed_batch_safe_to_retire(response): 

1526 update_data["batch_processed"] = True 

1527 

1528 try: 

1529 await ManagedObjectRepository(prisma_client).table.update( 

1530 where={"unified_object_id": batch_id}, 

1531 data=update_data, 

1532 ) 

1533 except Exception as col_err: 

1534 # If the batch_processed column doesn't exist (old schema), 

1535 # retry without it so the status update still succeeds. 

1536 err_str: Final = str(col_err).lower() 

1537 if "batch_processed" in err_str and update_data.get("batch_processed") is not None: 

1538 verbose_proxy_logger.warning( 

1539 "batch_processed column not found, retrying update without it: %s", col_err 

1540 ) 

1541 update_data.pop("batch_processed", None) 

1542 await ManagedObjectRepository(prisma_client).table.update( 

1543 where={"unified_object_id": batch_id}, 

1544 data=update_data, 

1545 ) 

1546 else: 

1547 raise 

1548 except Exception as e: 

1549 verbose_proxy_logger.error("Failed to update batch status in ManagedObjectTable: %s", e)