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
« 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)
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
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
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
38FILE_LIST_CONTINUATION_CHUNK_SIZE: Final = 500
40BATCH_CREATE_HIDDEN_PARAM: Final = "batch_create"
41LITELLM_EXECUTED_BATCH_ID_PREFIX: Final = "litellm_batch_"
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 )
62def validate_file_list_purpose(purpose: str | None) -> None:
63 """Reject a ``purpose`` filter no upload to this proxy could have stored.
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 )
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: ...
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: ...
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
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
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
134def get_models_from_unified_file_id(unified_file_id: str) -> list[str]:
135 """
136 Extract model names from unified file ID.
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 []
155def get_model_id_from_unified_batch_id(file_id: str) -> str | None:
156 """
157 Get the model_id from the file_id
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
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]
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)
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.
192 Format: <prefix><base64(litellm:<original_id>;model,<model_name>)>
193 The result preserves the original prefix (file-, batch_, etc.) for OpenAI compliance.
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.
202 Returns:
203 Encoded ID starting with appropriate prefix and containing routing information
205 Examples:
206 encode_file_id_with_model("file-abc123", "gpt-4o-litellm")
207 -> "file-bGl0ZWxsbTpmaWxlLWFiYzEyMzttb2RlbCxncHQtNG8taWZvb2Q"
209 encode_file_id_with_model("batch_abc123", "gpt-4o-test")
210 -> "batch_bGl0ZWxsbTpiYXRjaF9hYmMxMjM7bW9kZWwsZ3B0LTRvLXRlc3Q"
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("=")
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-"
229 return f"{prefix}{encoded_b64}"
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 )
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
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
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()
270 return None
271 except Exception:
272 return None
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
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
292 padded: Final = b64_part + "=" * (-len(b64_part) % 4)
293 decoded: Final = base64.urlsafe_b64decode(padded).decode()
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)
300 return encoded_id
301 except Exception:
302 return encoded_id
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
312# ============================================================================
313# MODEL-BASED CREDENTIAL ROUTING HELPERS
314# ============================================================================
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
329 Args:
330 file_id: File ID that may contain embedded model info
331 request: FastAPI request object
332 data: Optional request data dictionary
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 = {}
342 # Check if file_id has embedded model info
343 model_from_id: Final = decode_model_from_file_id(file_id)
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")
348 return model_from_id, model_from_param
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.
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).
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")
368 Returns:
369 Dictionary with credentials (api_key, api_base, custom_llm_provider, etc.)
371 Raises:
372 HTTPException: If router not initialized or model not found
373 """
374 from fastapi import HTTPException
376 from litellm.proxy.route_llm_request import ProxyModelNotFoundError
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 )
384 credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
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 )
391 return credentials
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.
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.
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
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 )
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 )
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.
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.
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
465 from litellm.proxy._types import SpecialModelNames
466 from litellm.proxy.auth.model_checks import get_complete_model_list, get_key_models
468 team_id: Final = user_api_key_dict.team_id
469 team_models: Final = user_api_key_dict.team_models or []
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()
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)
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
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
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
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
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
553 return None
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)
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).
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))
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 }
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).
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)
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
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.
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.
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
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
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 )
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
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
708 # No model-based routing needed
709 return False, None, None, None
712# ============================================================================
713# MIME TYPE DETECTION AND NORMALIZATION
714# ============================================================================
717# Gemini-supported image MIME types
718GEMINI_SUPPORTED_IMAGE_TYPES: Final = {
719 "image/png",
720 "image/jpeg",
721 "image/webp",
722}
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}
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}
752# Gemini-supported document MIME types
753GEMINI_SUPPORTED_DOCUMENT_TYPES: Final = {
754 "text/plain",
755 "application/pdf",
756}
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}
772def detect_content_type_from_filename(filename: str) -> str:
773 """
774 Detect content type from filename using extension.
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"
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
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
793 return "application/octet-stream"
796def normalize_mime_type_for_provider(mime_type: str, provider: str | None = None) -> str:
797 """
798 Normalize MIME type for specific provider requirements.
800 Currently handles:
801 - Gemini: Normalizes image/jpg to image/jpeg
803 Args:
804 mime_type: Original MIME type
805 provider: Provider name (e.g., "gemini", "vertex_ai")
807 Returns:
808 str: Normalized MIME type
809 """
810 normalized = mime_type.lower().strip()
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"
817 # General normalization: always normalize jpg to jpeg
818 if normalized == "image/jpg":
819 normalized = "image/jpeg"
821 return normalized
824def is_gemini_supported_mime_type(mime_type: str) -> bool:
825 """
826 Check if a MIME type is supported by Gemini multimodal models.
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
834 Args:
835 mime_type: MIME type to check
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 )
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).
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.
856 Args:
857 file_object: File object dictionary (can be None)
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"
865 # Handle JSON string
866 if isinstance(file_object, str):
867 import json
869 try:
870 file_object = json.loads(file_object)
871 except json.JSONDecodeError:
872 return "application/octet-stream"
874 if not isinstance(file_object, dict):
875 return "application/octet-stream"
877 # Try to get filename
878 filename: Final = file_object.get("filename", "")
879 if filename:
880 return detect_content_type_from_filename(filename)
882 return "application/octet-stream"
885# ============================================================================
886# REQUEST PARAMETER EXTRACTION
887# ============================================================================
890@dataclass
891class FileCreationParams:
892 """
893 Structured parameters extracted from file creation requests.
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 """
901 target_storage: str = "default"
902 target_model_names: list[str] = field(default_factory=list)
903 model: str | None = None
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 = []
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"
914 # Strip whitespace from model names
915 self.target_model_names = [name.strip() for name in self.target_model_names if name.strip()]
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.
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")
933 Returns:
934 FileCreationParams: Structured parameters extracted from the request
935 """
936 from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
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 {}
941 # Extract target_storage (simplified - just use form parameter)
942 target_storage: Final = _extract_target_storage_simple(target_storage_form)
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)
949 # Extract model parameter
950 model: Final = _extract_model_param(request, request_body)
952 return FileCreationParams(
953 target_storage=target_storage,
954 target_model_names=target_model_names,
955 model=model,
956 )
959def _extract_target_storage_simple(target_storage_form: str | None = None) -> str:
960 """
961 Extract target_storage parameter from form field.
963 Args:
964 target_storage_form: target_storage from form field
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"
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 []
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]
989 return []
992def _is_target_model_names_key(key: str) -> bool:
993 return key == "target_model_names" or (key.startswith("target_model_names[") and key.endswith("]"))
996async def _extract_target_model_names_from_form(request: "Request") -> list[str]:
997 """
998 Collect target_model_names from the raw multipart form.
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()
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))
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
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.
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
1036 import litellm
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
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 )
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 )
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.
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.
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
1083 import litellm
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
1088 if not resource_id:
1089 return
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 )
1101 if not isinstance(managed_files_obj, ManagedResourceAccessChecker):
1102 raise HTTPException(
1103 status_code=500,
1104 detail="Managed resource ownership validation is unavailable.",
1105 )
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
1115 raise HTTPException(
1116 status_code=403,
1117 detail=f"The caller does not have access to this managed {resource_kind} id.",
1118 )
1121def _extract_model_param(request: "Request", request_body: dict) -> str | None:
1122 """
1123 Extract model parameter from request.
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")
1133# ============================================================================
1134# BATCH DATABASE OPERATIONS
1135# ============================================================================
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 )
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 )
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 )
1190def _batch_owner_auth_from_db_object(db_batch_object: "LiteLLM_ManagedObjectTable") -> "UserAPIKeyAuth | None":
1191 from litellm.proxy._types import UserAPIKeyAuth
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)
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
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
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 )
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])
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)
1285 if managed_files_obj is None:
1286 return
1288 model_id: Final = _model_id_for_batch_response(response, unified_batch_id)
1289 if not model_id:
1290 return
1292 model_name: Final = _model_name_for_batch_response(response)
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
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 )
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.
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
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
1348 from litellm.types.utils import LiteLLMBatch
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
1353 try:
1354 if not prisma_client:
1355 return None, None
1357 db_batch_object: Final = await ManagedObjectRepository(prisma_client).table.find_first(
1358 where={"unified_object_id": batch_id}
1359 )
1361 if not db_batch_object or not db_batch_object.file_object:
1362 return None, None
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
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 )
1383 verbose_proxy_logger.debug(
1384 "Retrieved batch %s from ManagedObjectTable with status=%s", batch_id, response.status
1385 )
1387 return db_batch_object, response
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
1396def batch_cost_poller_is_active() -> bool:
1397 """
1398 Whether the CheckBatchCost poller will account for a managed batch's cost itself.
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
1411 if not PROXY_BATCH_POLLING_ENABLED:
1412 return False
1413 try:
1414 import litellm.proxy.proxy_server as proxy_server_module
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
1428def _completed_batch_safe_to_retire(response: "LiteLLMBatch") -> bool:
1429 """Whether a "completed" batch may be retired from cost recovery.
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)
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.
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
1478 if managed_files_obj is None or not unified_batch_id:
1479 return
1481 try:
1482 if not prisma_client:
1483 return
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 )
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 )
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
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)
1515 # Normalize status for database storage
1516 db_status: Final = response.status if response.status != "completed" else "complete"
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 }
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
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)