diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
index 84e319c8c37..2b2d0a82431 100644
--- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
+++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
@@ -732,6 +732,7 @@ class CheckBatchCost:
another pod claimed it. Raises on results-fetch or cost-computation
failures so the caller can leave the job unprocessed and retry it on a
later poll.
+
"""
from litellm.batches.batch_utils import (
count_error_file_failed_requests,
@@ -815,15 +816,14 @@ class CheckBatchCost:
custom_llm_provider=custom_llm_provider,
)
- # CheckBatchCost bypasses async_post_call_success_hook, so convert raw
- # output/error file IDs to managed base64 IDs before the DB write here.
- managed_files_hook = self.proxy_logging_obj.get_proxy_hook("managed_files")
- if managed_files_hook is not None:
+ from litellm.proxy.openai_files_endpoints.common_utils import ManagedBatchOutputFileWriter
+
+ managed_files_hook: Final = self.proxy_logging_obj.get_proxy_hook("managed_files")
+ if isinstance(managed_files_hook, ManagedBatchOutputFileWriter):
+ managed_file_writer: Final = managed_files_hook
from litellm.proxy._types import UserAPIKeyAuth
- managed_file_model_name = self._get_managed_file_model_name(
- job=job, deployment_info=deployment_info
- )
+ managed_file_model_name = self._get_managed_file_model_name(job=job, deployment_info=deployment_info)
_minimal_auth = UserAPIKeyAuth(
user_id=job.created_by or "default-user-id",
team_id=getattr(job, "team_id", None),
@@ -832,17 +832,19 @@ class CheckBatchCost:
_raw_file_id = cast(str | None, getattr(response, _file_attr, None))
if _raw_file_id and not _is_base64_encoded_unified_file_id(_raw_file_id):
try:
- _unified_file_id = managed_files_hook.get_unified_output_file_id(
+ _unified_file_id = managed_file_writer.get_unified_output_file_id(
output_file_id=_raw_file_id,
model_id=model_id,
model_name=managed_file_model_name,
)
- await managed_files_hook.store_unified_file_id(
- file_id=_unified_file_id,
- file_object=None,
+ await managed_file_writer.store_batch_output_file(
+ unified_file_id=_unified_file_id,
+ provider_file_id=_raw_file_id,
+ model_id=model_id,
+ model_name=managed_file_model_name,
+ owner=_minimal_auth,
litellm_parent_otel_span=None,
- model_mappings={model_id: _raw_file_id},
- user_api_key_dict=_minimal_auth,
+ size_bytes=len(content_bytes) if _file_attr == "output_file_id" else None,
)
setattr(response, _file_attr, _unified_file_id)
verbose_proxy_logger.info(
diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
index 28c4cf5e66f..4c06f0b9802 100644
--- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
+++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py
@@ -1,9 +1,11 @@
# What is this?
## This hook is used to check for LiteLLM managed files in the request body, and replace them with model-specific file id
+import asyncio
import base64
import json
-from collections.abc import Iterator, Mapping, Sequence
+import time
+from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
@@ -23,21 +25,25 @@ from uuid import NAMESPACE_URL, uuid5
import httpx
from fastapi import HTTPException
from pydantic import ValidationError
-from typing_extensions import ReadOnly
+from typing_extensions import ReadOnly, Unpack
import litellm
from litellm import Router, verbose_logger
from litellm._internal_context import with_service_target
from litellm._uuid import uuid
from litellm.caching.caching import DualCache
-from litellm.constants import MAX_FILE_LIST_LIMIT
+from litellm.constants import (
+ BATCH_OUTPUT_FILE_FALLBACK_RETRY_AFTER_SECONDS,
+ BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS,
+ MAX_FILE_LIST_LIMIT,
+)
from litellm.files.types import FileRetrieveCallOptions, FileRetrieveProvider
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.hidden_params import HIDDEN_PARAMS_ATTR
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_metadata,
)
-from openai import AsyncOpenAI
+from openai import APIConnectionError, AsyncOpenAI
from openai.types.file_deleted import FileDeleted
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
@@ -55,7 +61,10 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
-from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
+from litellm.proxy.litellm_pre_call_utils import (
+ LiteLLMProxyRequestSetup,
+ sanitize_for_log,
+)
from litellm.proxy.openai_files_endpoints.common_utils import (
BATCH_CREATE_HIDDEN_PARAM,
FILE_LIST_CONTINUATION_CHUNK_SIZE,
@@ -77,6 +86,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import (
request_tags_from_metadata,
)
+from litellm.repositories.managed_file_repository import ManagedFileRepository
from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue]
AllMessageValues,
AsyncCursorPage,
@@ -96,6 +106,17 @@ from litellm.types.utils import (
SpecialEnums,
)
+
+class _ManagedFileRetrieve(Protocol):
+ async def __call__(
+ self,
+ *,
+ file_id: str,
+ _litellm_internal_model_credentials: Mapping[str, object] | None = None,
+ **kwargs: Unpack[FileRetrieveCallOptions],
+ ) -> OpenAIFileObject: ...
+
+
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from prisma.models import (
@@ -146,6 +167,91 @@ def _parse_managed_file_object(raw_file_object: object, unified_file_id: str) ->
return None
+def _batch_output_file_object(
+ unified_file_id: str, raw_file_id: str, size_bytes: int, *, fallback: bool
+) -> OpenAIFileObject:
+ filename: Final = raw_file_id.rsplit("/", 1)[-1] or raw_file_id
+ return OpenAIFileObject(
+ id=unified_file_id,
+ object="file",
+ purpose="batch_output",
+ filename=filename,
+ created_at=int(time.time()),
+ bytes=size_bytes,
+ status="processed",
+ litellm_details_fallback=True if fallback else None,
+ )
+
+
+def _public_file_object(file_object: OpenAIFileObject, unified_file_id: str) -> OpenAIFileObject:
+ return file_object.model_copy(update={"id": unified_file_id, "litellm_details_fallback": None})
+
+
+def _is_transient_file_retrieve_error(error: Exception) -> bool:
+ status_code: Final = getattr(error, "status_code", None)
+ if isinstance(status_code, int):
+ return status_code in {408, 429} or status_code >= 500
+ return isinstance(error, (httpx.TransportError, APIConnectionError, asyncio.TimeoutError))
+
+
+def _proxy_llm_router() -> Router | None:
+ import litellm.proxy.proxy_server as proxy_server_module
+
+ return cast(Router | None, getattr(proxy_server_module, "llm_router", None))
+
+
+def _provider_file_retrieve_credentials(
+ *,
+ llm_router: Router | None,
+ model_id: str | None,
+) -> Mapping[str, object] | None:
+ if llm_router is None or model_id is None:
+ return None
+ try:
+ credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id)
+ except Exception as error:
+ verbose_logger.warning(
+ "Failed to retrieve credentials for provider file "
+ f"model_id={sanitize_for_log(model_id)}: {sanitize_for_log(error)}"
+ )
+ return None
+ return cast(Mapping[str, object], credentials) if credentials else None
+
+
+_PROVIDER_FILE_RETRIEVE_PROVIDERS: Final[frozenset[str]] = frozenset(
+ {
+ "openai",
+ "azure",
+ "gemini",
+ "vertex_ai",
+ "bedrock",
+ "hosted_vllm",
+ "litellm_proxy",
+ "manus",
+ "anthropic",
+ "mistral",
+ "xai",
+ }
+)
+
+
+def _model_name_file_retrieve_provider(model_name: str | None) -> FileRetrieveProvider | None:
+ if model_name is None:
+ return None
+ provider, separator, _ = model_name.partition("/")
+ if not separator or provider not in _PROVIDER_FILE_RETRIEVE_PROVIDERS:
+ return None
+ return cast(FileRetrieveProvider, provider)
+
+
+def _has_provider_file_retrieve_route(
+ *,
+ router_credentials: Mapping[str, object] | None,
+ model_name: str | None,
+) -> bool:
+ return bool(router_credentials) or _model_name_file_retrieve_provider(model_name) is not None
+
+
class _ManagedFileRow(Protocol):
unified_file_id: str
file_object: OpenAIFileObject
@@ -161,6 +267,8 @@ class _ManagedFileRow(Protocol):
class _ManagedFileTableActions(Protocol):
async def find_first(self, where: Mapping[str, object]) -> Optional[_ManagedFileRow]: ...
+ async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
+
async def find_many(
self,
where: Mapping[str, object],
@@ -253,13 +361,21 @@ def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, s
_MANAGED_FILES_TARGET: Final = "managed_files"
+_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS: Final = (0.5, 1.0, 2.0)
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Class variables or attributes
- def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient):
+ def __init__(
+ self,
+ internal_usage_cache: InternalUsageCache,
+ prisma_client: PrismaClient,
+ *,
+ sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
+ ):
self.internal_usage_cache = internal_usage_cache
self.prisma_client = prisma_client
+ self._sleep = sleep
@staticmethod
def _get_prometheus_logger():
@@ -595,14 +711,231 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
for file_id in provider_file_ids:
model_name = decode_model_from_file_id(file_id)
raw_file_id = get_original_file_id(file_id)
- await self.store_unified_file_id(
- file_id=file_id,
- file_object=None,
+ await self.store_batch_output_file(
+ unified_file_id=file_id,
+ provider_file_id=raw_file_id,
+ model_id=model_name or None,
+ model_name=model_name,
litellm_parent_otel_span=litellm_parent_otel_span,
- model_mappings={model_name: raw_file_id} if model_name else {},
- user_api_key_dict=owner_identity,
+ owner=owner_identity,
)
+ async def _afile_retrieve_with_retries(
+ self,
+ *,
+ provider_file_id: str,
+ call_options: Mapping[str, object],
+ internal_model_credentials: Mapping[str, object] | None = None,
+ ) -> OpenAIFileObject:
+ """Retrieve provider file details with transient retries and SDK retries disabled."""
+ retrieve_file: Final = cast(_ManagedFileRetrieve, litellm.afile_retrieve)
+ retrieve_options: Final = cast(
+ FileRetrieveCallOptions,
+ {**call_options, "max_retries": 0},
+ )
+ for attempt in range(len(_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS) + 1):
+ if attempt > 0:
+ await self._sleep(_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS[attempt - 1])
+ try:
+ if internal_model_credentials is None:
+ return await retrieve_file(
+ file_id=provider_file_id,
+ **retrieve_options,
+ )
+ return await retrieve_file(
+ file_id=provider_file_id,
+ _litellm_internal_model_credentials=MappingProxyType(dict(internal_model_credentials)),
+ **retrieve_options,
+ )
+ except Exception as error:
+ if not _is_transient_file_retrieve_error(error) or attempt == len(
+ _PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS
+ ):
+ raise
+ raise RuntimeError("Provider file retrieve retry loop ended without a result")
+
+ async def _fetch_provider_file_object(
+ self,
+ *,
+ unified_file_id: str,
+ provider_file_id: str,
+ model_id: str | None,
+ model_name: str | None,
+ llm_router: Router | None = None,
+ raise_on_failure: bool = False,
+ allow_default_provider: bool = False,
+ ) -> tuple[OpenAIFileObject | None, bool]:
+ """Fetch provider file details through the configured route under a total timeout."""
+ route_llm_router: Final = llm_router if llm_router is not None else _proxy_llm_router()
+ router_credentials: Final = _provider_file_retrieve_credentials(
+ llm_router=route_llm_router,
+ model_id=model_id,
+ )
+ model_name_provider: Final = _model_name_file_retrieve_provider(model_name)
+ fetch_route_available: Final = _has_provider_file_retrieve_route(
+ router_credentials=router_credentials,
+ model_name=model_name,
+ )
+ default_provider_route: Final = allow_default_provider and route_llm_router is not None
+ if not fetch_route_available and not default_provider_route:
+ return None, False
+
+ try:
+ if router_credentials is not None:
+ provider_file_object: Final = await asyncio.wait_for(
+ self._afile_retrieve_with_retries(
+ provider_file_id=provider_file_id,
+ call_options=router_credentials,
+ internal_model_credentials=router_credentials,
+ ),
+ timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS,
+ )
+ return provider_file_object.model_copy(update={"id": unified_file_id}), True
+ if model_name_provider is not None:
+ provider_file_object_by_model_name: Final = await asyncio.wait_for(
+ self._afile_retrieve_with_retries(
+ provider_file_id=provider_file_id,
+ call_options={"custom_llm_provider": model_name_provider},
+ ),
+ timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS,
+ )
+ return (
+ provider_file_object_by_model_name.model_copy(update={"id": unified_file_id}),
+ True,
+ )
+ if default_provider_route:
+ provider_file_object_by_default_route: Final = await asyncio.wait_for(
+ self._afile_retrieve_with_retries(
+ provider_file_id=provider_file_id,
+ call_options=router_credentials or {},
+ ),
+ timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS,
+ )
+ return (
+ provider_file_object_by_default_route.model_copy(update={"id": unified_file_id}),
+ True,
+ )
+ return None, False
+ except Exception as error:
+ verbose_logger.warning(
+ "Failed to retrieve batch file object for "
+ f"provider_file_id={sanitize_for_log(provider_file_id)}: "
+ f"{type(error).__name__} {sanitize_for_log(error)}"
+ )
+ if raise_on_failure:
+ if isinstance(error, TimeoutError) and not str(error):
+ raise TimeoutError(
+ "Provider file retrieve timed out "
+ f"after {BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS} seconds"
+ ) from error
+ raise
+ return None, True
+
+ async def _save_refreshed_file_object(
+ self,
+ stored: LiteLLM_ManagedFileTable,
+ file_object: OpenAIFileObject,
+ ) -> None:
+ if not await ManagedFileRepository(self.prisma_client).update_file_object(stored.unified_file_id, file_object):
+ return
+ refreshed_row: Final = stored.model_copy(update={"file_object": file_object})
+ await self.internal_usage_cache.async_set_cache(
+ key=stored.unified_file_id,
+ value=refreshed_row.model_dump(),
+ litellm_parent_otel_span=None,
+ )
+
+ async def store_batch_output_file(
+ self,
+ *,
+ unified_file_id: str,
+ provider_file_id: str,
+ model_id: str | None,
+ model_name: str | None = None,
+ owner: UserAPIKeyAuth,
+ litellm_parent_otel_span: Span | None,
+ size_bytes: int | None = None,
+ fetch_provider_details: bool = True,
+ ) -> None:
+ """Register batch output or error file metadata, optionally fetching provider details."""
+ stored_file: Final = await self.get_unified_file_id(unified_file_id, litellm_parent_otel_span)
+ stored_object: Final = stored_file.file_object if stored_file is not None else None
+ if not fetch_provider_details:
+ if stored_file is not None:
+ return
+ router_credentials: Final = _provider_file_retrieve_credentials(
+ llm_router=_proxy_llm_router(),
+ model_id=model_id,
+ )
+ file_object_without_provider_details: Final = _batch_output_file_object(
+ unified_file_id,
+ provider_file_id,
+ size_bytes or 0,
+ fallback=_has_provider_file_retrieve_route(
+ router_credentials=router_credentials,
+ model_name=model_name,
+ ),
+ )
+ await self.store_unified_file_id(
+ file_id=unified_file_id,
+ file_object=file_object_without_provider_details,
+ litellm_parent_otel_span=litellm_parent_otel_span,
+ model_mappings={model_id: provider_file_id} if model_id else {},
+ user_api_key_dict=owner,
+ )
+ return
+ if stored_object is not None and not stored_object.litellm_details_fallback:
+ return
+
+ fallback_written_recently: Final = (
+ stored_object is not None
+ and time.time() - stored_object.created_at < BATCH_OUTPUT_FILE_FALLBACK_RETRY_AFTER_SECONDS
+ )
+ provider_fetch_result: Final = (
+ (None, True)
+ if fallback_written_recently
+ else await self._fetch_provider_file_object(
+ unified_file_id=unified_file_id,
+ provider_file_id=provider_file_id,
+ model_id=model_id,
+ model_name=model_name,
+ )
+ )
+ provider_object, fetch_route_available = provider_fetch_result
+ if (
+ stored_file is not None
+ and provider_object is None
+ and (size_bytes is None or (stored_object is not None and stored_object.bytes == size_bytes))
+ ):
+ return
+
+ file_object: Final = (
+ provider_object
+ if provider_object is not None
+ else (
+ stored_object.model_copy(update={"bytes": size_bytes})
+ if stored_object is not None and size_bytes is not None
+ else _batch_output_file_object(
+ unified_file_id,
+ provider_file_id,
+ size_bytes or 0,
+ fallback=fetch_route_available,
+ )
+ )
+ )
+
+ if stored_file is not None:
+ await self._save_refreshed_file_object(stored_file, file_object)
+ return
+
+ await self.store_unified_file_id(
+ file_id=unified_file_id,
+ file_object=file_object,
+ litellm_parent_otel_span=litellm_parent_otel_span,
+ model_mappings={model_id: provider_file_id} if model_id else {},
+ user_api_key_dict=owner,
+ )
+
async def list_user_batches(
self,
user_api_key_dict: UserAPIKeyAuth,
@@ -745,6 +1078,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
user_api_key_dict=user_api_key_dict,
db_batch_object=row,
unified_batch_id=_is_base64_encoded_unified_file_id(row.unified_object_id),
+ fetch_provider_details=False,
)
except Exception as e:
verbose_logger.warning(f"Failed to resolve managed file ids for batch {row.unified_object_id}: {e}")
@@ -806,7 +1140,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
}
)
return [
- parsed_file_object.model_copy(update={"id": row.unified_file_id})
+ _public_file_object(parsed_file_object, row.unified_file_id)
for row in file_ids
if (parsed_file_object := _parse_managed_file_object(row.file_object, row.unified_file_id)) is not None
]
@@ -1474,42 +1808,13 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
setattr(response, file_attr, unified_file_id)
- # Use llm_router credentials when available. Without credentials,
- # Azure and other auth-required providers return 500/401.
- file_object = None
- try:
- # Import module and use getattr for better testability with mocks
- import litellm.proxy.proxy_server as proxy_server_module
-
- _llm_router = getattr(proxy_server_module, "llm_router", None)
- if _llm_router is not None and model_id:
- _creds = _llm_router.get_deployment_credentials_with_provider(model_id) or {}
- file_object = await litellm.afile_retrieve(
- file_id=provider_file_id,
- **cast(FileRetrieveCallOptions, _creds),
- )
- else:
- file_object = await litellm.afile_retrieve(
- custom_llm_provider=cast(
- FileRetrieveProvider,
- model_name.split("/")[0] if model_name and "/" in model_name else "openai",
- ),
- file_id=provider_file_id,
- )
- verbose_logger.debug(
- f"Successfully retrieved file object for {file_attr}={provider_file_id}"
- )
- except Exception as e:
- verbose_logger.warning(
- f"Failed to retrieve file object for {file_attr}={provider_file_id}: {str(e)}. Storing with None and will fetch on-demand."
- )
-
- await self.store_unified_file_id(
- file_id=unified_file_id,
- file_object=file_object,
+ await self.store_batch_output_file(
+ unified_file_id=unified_file_id,
+ provider_file_id=provider_file_id,
+ model_id=model_id,
+ model_name=resolved_model_name,
+ owner=user_api_key_dict,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
- model_mappings={model_id: provider_file_id},
- user_api_key_dict=user_api_key_dict,
)
request_metadata: Final = data.get("litellm_metadata")
await self.store_unified_object_id(
@@ -1612,6 +1917,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
async def afile_retrieve(
self, file_id: str, litellm_parent_otel_span: Optional[Span], llm_router: Optional[Router] = None
) -> OpenAIFileObject:
+ """Return public details for a managed file ID, refreshing a basic entry when possible."""
stored_file_object = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
# Case 1 : This is not a managed file
@@ -1621,30 +1927,57 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Case 2: Managed file and the file object exists in the database
# The stored file_object has the raw provider ID. Replace with the unified ID
# so callers see a consistent ID (matching Case 3 which does response.id = file_id).
- if stored_file_object and stored_file_object.file_object:
- # Use model_copy to ensure the ID update persists (Pydantic v2 compatibility)
- response = stored_file_object.file_object.model_copy(update={"id": file_id})
- return response
+ if stored_file_object.file_object is not None:
+ file_object: Final = stored_file_object.file_object
+ if file_object.litellm_details_fallback and stored_file_object.model_mappings:
+ try:
+ model_id, provider_file_id = next(iter(stored_file_object.model_mappings.items()))
+ refreshed_file_object, _ = await self._fetch_provider_file_object(
+ unified_file_id=file_id,
+ provider_file_id=provider_file_id,
+ model_id=model_id,
+ model_name=model_id,
+ llm_router=llm_router,
+ )
+ if refreshed_file_object is None:
+ return _public_file_object(file_object, file_id)
+ await self._save_refreshed_file_object(stored_file_object, refreshed_file_object)
+ return _public_file_object(refreshed_file_object, file_id)
+ except Exception as error:
+ verbose_logger.warning(
+ "Failed to refresh batch file object for "
+ f"file_id={sanitize_for_log(file_id)}: {sanitize_for_log(error)}"
+ )
+ return _public_file_object(file_object, file_id)
# Case 3: Managed file exists in the database but not the file object (for. e.g the batch task might not have run)
# So we fetch the file object from the provider. We deliberately do not store the result to avoid interfering with batch cost tracking code.
- if not llm_router:
+ model_mapping: Final = next(iter(stored_file_object.model_mappings.items()), None)
+ if model_mapping is None:
raise Exception(
- f"LiteLLM Managed File object with id={file_id} has no file_object "
- f"and llm_router is required to fetch from provider"
+ f"LiteLLM Managed File object with id={file_id} has no file_object and no provider route to fetch it"
)
+ model_id, model_file_id = model_mapping
try:
- model_id, model_file_id = next(iter(stored_file_object.model_mappings.items()))
- credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id) or {}
- response = await litellm.afile_retrieve(
- file_id=model_file_id,
- **cast(FileRetrieveCallOptions, credentials),
+ response, fetch_route_available = await self._fetch_provider_file_object(
+ unified_file_id=file_id,
+ provider_file_id=model_file_id,
+ model_id=model_id,
+ model_name=model_id,
+ llm_router=llm_router,
+ raise_on_failure=True,
+ allow_default_provider=True,
)
- response.id = file_id # Replace with unified ID
- return response
except Exception as e:
raise Exception(f"Failed to retrieve file {file_id} from provider: {str(e)}") from e
+ if not fetch_route_available:
+ raise Exception(
+ f"LiteLLM Managed File object with id={file_id} has no file_object and no provider route to fetch it"
+ )
+ if response is None:
+ raise ValueError("Provider file details could not be retrieved")
+ return _public_file_object(response, file_id)
async def afile_list(
self,
@@ -1707,7 +2040,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
**cursor_args,
)
matches.extend(
- parsed_file_object.model_copy(update={"id": row.unified_file_id})
+ _public_file_object(parsed_file_object, row.unified_file_id)
for row in chunk
if (parsed_file_object := _parse_managed_file_object(row.file_object, row.unified_file_id)) is not None
and (purpose is None or parsed_file_object.purpose == purpose)
diff --git a/litellm/constants.py b/litellm/constants.py
index b8048ad30b4..74273b9ecb9 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -1845,6 +1845,8 @@ RESET_BUDGET_JOB_LOCK_TTL_SECONDS: Final[int] = 900
PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600))
MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50)))
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7)))
+BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS: Final = 10.0
+BATCH_OUTPUT_FILE_FALLBACK_RETRY_AFTER_SECONDS: Final = 60
STALE_OBJECT_CLEANUP_BATCH_SIZE: Final = max(1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000)))
# Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and
# CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on
diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py
index ec0b28a63ad..cc36962da6b 100644
--- a/litellm/llms/base_llm/files/transformation.py
+++ b/litellm/llms/base_llm/files/transformation.py
@@ -135,6 +135,9 @@ class BaseFilesConfig(BaseConfig):
) -> OpenAIFileObject:
"""Transform file retrieve response into OpenAI format."""
+ def is_retrieve_file_response_successful(self, response: httpx.Response) -> bool:
+ return not httpx.codes.is_error(response.status_code)
+
@abstractmethod
def transform_delete_file_request(
self,
diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py
index e7ee2f89bbf..5c5d6e73417 100644
--- a/litellm/llms/bedrock/files/transformation.py
+++ b/litellm/llms/bedrock/files/transformation.py
@@ -8,6 +8,7 @@ from collections.abc import Iterable, Mapping, MutableMapping, Sequence
from contextlib import suppress
from dataclasses import dataclass
from datetime import datetime
+from email.utils import parsedate_to_datetime
from itertools import chain
from types import MappingProxyType
from typing import Any, Final, Literal, TypeAlias, TypedDict
@@ -16,7 +17,7 @@ from urllib.parse import quote, unquote, urlencode
import httpx
from httpx import Headers, Response
from openai.types.file_deleted import FileDeleted
-from pydantic import ConfigDict, Field
+from pydantic import ConfigDict, Field, TypeAdapter
from typing_extensions import ReadOnly
from litellm._logging import verbose_logger
@@ -70,12 +71,71 @@ from ..common_utils import (
)
S3_SIGNED_REQUEST_HEADERS_PARAM: Final = "_s3_signed_request_headers"
+S3_RETRIEVE_FILE_ID_PARAM: Final = "_s3_retrieve_file_id"
+S3_RETRIEVE_FILE_KEY_PARAM: Final = "_s3_retrieve_file_key"
+S3_RETRIEVE_FILE_RELATIVE_KEY_PARAM: Final = "_s3_retrieve_file_relative_key"
+_S3_SIGNED_REQUEST_HEADERS_ADAPTER: Final = TypeAdapter(
+ Mapping[str, str],
+ config=ConfigDict(strict=True),
+)
LIST_FILES_PURPOSE_PARAM: Final = "_s3_list_files_purpose"
LIST_FILES_LOCATION_PARAM: Final = "_s3_list_files_location"
+def _is_empty_s3_object_range_error(raw_response: Response) -> bool:
+ if raw_response.status_code != 416:
+ return False
+ if raw_response.headers.get("Content-Range") == "bytes */0":
+ return True
+ try:
+ error_xml: Final = ET.fromstring(raw_response.content)
+ except ET.ParseError:
+ return False
+ return error_xml.findtext("ActualObjectSize") == "0"
+
+
+def _retrieved_s3_file_size(raw_response: Response) -> int:
+ status_code: Final = raw_response.status_code
+ if _is_empty_s3_object_range_error(raw_response):
+ return 0
+ if status_code == 206:
+ content_range: Final = raw_response.headers.get("Content-Range", "")
+ range_parts: Final = content_range.removeprefix("bytes 0-0/")
+ if content_range.startswith("bytes 0-0/") and range_parts.isdigit():
+ return int(range_parts)
+ raise BedrockError(
+ status_code=status_code,
+ message=f"Invalid S3 Content-Range header: {content_range}",
+ headers=raw_response.headers,
+ response=raw_response,
+ )
+ if status_code == 200:
+ content_length: Final = raw_response.headers.get("Content-Length", "")
+ if content_length.isdigit():
+ return int(content_length)
+ raise BedrockError(
+ status_code=status_code,
+ message=f"Invalid S3 Content-Length header: {content_length}",
+ headers=raw_response.headers,
+ response=raw_response,
+ )
+ if status_code >= 400:
+ raise BedrockError(
+ status_code=status_code,
+ message=raw_response.text,
+ headers=raw_response.headers,
+ response=raw_response,
+ )
+ raise BedrockError(
+ status_code=status_code,
+ message=f"S3 file retrieval returned HTTP {status_code}",
+ headers=raw_response.headers,
+ response=raw_response,
+ )
+
+
class _S3DeleteContext(LiteLLMBaseModel):
file_id: str = Field(min_length=1)
@@ -280,6 +340,25 @@ def _resolve_managed_s3_object(file_id: str, litellm_params: Mapping[str, object
raise _rejected_file_id(reason) from reason
+def _relative_s3_object_key(
+ bucket_name: str,
+ object_key: str,
+ litellm_params: Mapping[str, object],
+) -> str:
+ configured_bucket_prefixes: Final = tuple(
+ split_configured_cloud_bucket_name(configured_bucket_name)
+ for configured_bucket_name in get_configured_s3_bucket_names(litellm_params)
+ )
+ matching_prefixes: Final = tuple(
+ configured_prefix
+ for configured_bucket, configured_prefix in configured_bucket_prefixes
+ if configured_bucket == bucket_name
+ and (not configured_prefix or object_key.startswith(f"{configured_prefix}/"))
+ )
+ configured_prefix: Final = max(matching_prefixes, key=len, default="")
+ return object_key[len(configured_prefix) + 1 :] if configured_prefix else object_key
+
+
_ANY_MANAGED_LISTING_PREFIX: Final = os.path.commonprefix(BEDROCK_MANAGED_S3_PREFIXES)
_MANAGED_LISTING_PREFIX_BY_PURPOSE: Final = MappingProxyType(
{
@@ -398,6 +477,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.BEDROCK
+ def is_retrieve_file_response_successful(self, response: httpx.Response) -> bool:
+ return not httpx.codes.is_error(response.status_code) or _is_empty_s3_object_range_error(response)
+
@property
def file_upload_http_method(self) -> str:
"""
@@ -1276,18 +1358,59 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
def transform_retrieve_file_request(
self,
file_id: str,
- optional_params: dict,
- litellm_params: dict,
- ) -> tuple[str, dict]:
- raise NotImplementedError("BedrockFilesConfig does not support file retrieval")
+ optional_params: Mapping[str, object],
+ litellm_params: MutableMapping[str, object],
+ ) -> tuple[str, dict[str, str]]:
+ """Prepare a ranged S3 GET for file retrieval."""
+ bucket_name, object_key = _resolve_managed_s3_object(file_id=file_id, litellm_params=litellm_params)
+ relative_key: Final = _relative_s3_object_key(
+ bucket_name=bucket_name,
+ object_key=object_key,
+ litellm_params=litellm_params,
+ )
+ url, params = self._transform_s3_file_request(
+ file_id=file_id,
+ method="GET",
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ )
+ signed_headers_object: Final = litellm_params.get(S3_SIGNED_REQUEST_HEADERS_PARAM)
+ if not isinstance(signed_headers_object, Mapping):
+ raise TypeError("S3 request signing did not produce request headers")
+ signed_headers: Final = _S3_SIGNED_REQUEST_HEADERS_ADAPTER.validate_python(signed_headers_object)
+ range_headers: Final = MappingProxyType({**signed_headers, "Range": "bytes=0-0"})
+ litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM] = range_headers # rebind-ok: handed to validate_environment
+ litellm_params[S3_RETRIEVE_FILE_ID_PARAM] = file_id # rebind-ok: required by response transform
+ litellm_params[S3_RETRIEVE_FILE_KEY_PARAM] = object_key # rebind-ok: required by response transform
+ litellm_params[S3_RETRIEVE_FILE_RELATIVE_KEY_PARAM] = relative_key # rebind-ok: required by response transform
+ return url, params
def transform_retrieve_file_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
- litellm_params: dict,
+ litellm_params: Mapping[str, object],
) -> OpenAIFileObject:
- raise NotImplementedError("BedrockFilesConfig does not support file retrieval")
+ """Build file metadata, accepting 416 only when S3 proves the object is empty."""
+ file_id: Final = litellm_params.get(S3_RETRIEVE_FILE_ID_PARAM)
+ object_key: Final = litellm_params.get(S3_RETRIEVE_FILE_KEY_PARAM)
+ relative_key: Final = litellm_params.get(S3_RETRIEVE_FILE_RELATIVE_KEY_PARAM)
+ if not isinstance(file_id, str) or not isinstance(object_key, str) or not isinstance(relative_key, str):
+ raise TypeError("S3 retrieve response is missing request context")
+
+ file_size: Final = _retrieved_s3_file_size(raw_response)
+
+ last_modified: Final = raw_response.headers.get("Last-Modified", "")
+ created_at: Final = int(parsedate_to_datetime(last_modified).timestamp()) if last_modified else 0
+ return OpenAIFileObject(
+ id=file_id,
+ bytes=file_size,
+ created_at=created_at,
+ filename=posixpath.basename(object_key),
+ object="file",
+ purpose="batch_output" if relative_key.startswith(BEDROCK_MANAGED_S3_OUTPUT_PREFIX) else "batch",
+ status="processed",
+ )
def transform_delete_file_request(
self,
diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py
index 124f45e9bf8..4a9b8fee2e8 100644
--- a/litellm/llms/custom_httpx/llm_http_handler.py
+++ b/litellm/llms/custom_httpx/llm_http_handler.py
@@ -4756,7 +4756,8 @@ class BaseLLMHTTPHandler:
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
- self._raise_for_provider_error_status(response=response, provider_config=provider_config)
+ if not provider_config.is_retrieve_file_response_successful(response):
+ self._raise_for_provider_error_status(response=response, provider_config=provider_config)
return provider_config.transform_retrieve_file_response(
raw_response=response,
logging_obj=logging_obj,
@@ -4814,7 +4815,8 @@ class BaseLLMHTTPHandler:
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
- self._raise_for_provider_error_status(response=response, provider_config=provider_config)
+ if not provider_config.is_retrieve_file_response_successful(response):
+ self._raise_for_provider_error_status(response=response, provider_config=provider_config)
return provider_config.transform_retrieve_file_response(
raw_response=response,
logging_obj=logging_obj,
diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py
index 413d5a7057f..104a6452f79 100644
--- a/litellm/proxy/batches_endpoints/endpoints.py
+++ b/litellm/proxy/batches_endpoints/endpoints.py
@@ -72,10 +72,10 @@ from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attributio
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy, is_known_model
from litellm.repositories.managed_batch_repository import ManagedBatchRepository
-from litellm.repositories.table_repositories import ManagedFileRepository
+from litellm.repositories.managed_file_repository import ManagedFileRepository
from litellm.router import Router
from litellm.types.llms.openai import LiteLLMBatchCreateRequest
-from litellm.types.utils import LiteLLMBatch
+from litellm.types.utils import LiteLLMBatch, LLMResponseTypes
if TYPE_CHECKING:
from prisma.models import LiteLLM_ManagedObjectTable
@@ -96,6 +96,12 @@ def _request_tags(data: Mapping[str, object]) -> tuple[str, ...] | None:
return request_tags_from_metadata(_METADATA_ADAPTER.validate_python(metadata))
+def _require_batch_response(response: LLMResponseTypes) -> LiteLLMBatch:
+ if not isinstance(response, LiteLLMBatch):
+ raise TypeError("Batch endpoint received a non-batch response")
+ return response
+
+
def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyLogging) -> LiteLLMExecutedBatchRunner:
from litellm.proxy.proxy_server import general_settings, prisma_client
@@ -652,8 +658,9 @@ async def retrieve_batch(
# The DB may store raw provider file IDs (before hooks translate them).
# Register any missing managed-file rows and return unified IDs.
if unified_batch_id:
+ terminal_batch_response: Final = _require_batch_response(response)
await ensure_batch_response_managed_file_ids(
- response=response,
+ response=terminal_batch_response,
managed_files_obj=managed_files_obj,
prisma_client=prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
@@ -799,8 +806,9 @@ async def retrieve_batch(
# Fix: bug_feb14_batch_retrieve_returns_raw_input_file_id
# Register any missing managed-file rows and return unified IDs.
if unified_batch_id:
+ retrieved_batch_response: Final = _require_batch_response(response)
await ensure_batch_response_managed_file_ids(
- response=response,
+ response=retrieved_batch_response,
managed_files_obj=managed_files_obj,
prisma_client=prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py
index bf67850ee9f..84e5e113e34 100644
--- a/litellm/proxy/litellm_pre_call_utils.py
+++ b/litellm/proxy/litellm_pre_call_utils.py
@@ -209,6 +209,10 @@ def _sanitize_for_log(value: object) -> str:
return text.replace("\r", "").replace("\n", "")
+def sanitize_for_log(value: object) -> str:
+ return _sanitize_for_log(value)
+
+
from litellm.router import Router
from litellm.secret_managers.main import get_secret_bool
from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS
diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py
index ffeca8d0257..6dd61a7bf14 100644
--- a/litellm/proxy/openai_files_endpoints/common_utils.py
+++ b/litellm/proxy/openai_files_endpoints/common_utils.py
@@ -1,4 +1,5 @@
import base64
+import logging
import mimetypes
import re
from collections.abc import Mapping, Sequence
@@ -18,8 +19,8 @@ from typing import (
from litellm.batches.batch_utils import batch_cost_is_final
from litellm.constants import MAX_FILE_LIST_LIMIT
from litellm.proxy._types import ProxyException
+from litellm.repositories.managed_file_repository import ManagedFileRepository
from litellm.repositories.table_repositories import (
- ManagedFileRepository,
ManagedObjectRepository,
)
from litellm.types.llms.openai import OpenAIFilesPurpose
@@ -27,6 +28,7 @@ from litellm.types.utils import SpecialEnums
if TYPE_CHECKING:
from fastapi import Request
+ from opentelemetry.trace import Span
from prisma.models import LiteLLM_ManagedObjectTable
from litellm.proxy._types import UserAPIKeyAuth
@@ -104,6 +106,31 @@ class ManagedFileIdResolver(Protocol):
) -> Mapping[str, str]: ...
+@runtime_checkable
+class ManagedBatchOutputFileWriter(Protocol):
+ def get_unified_output_file_id(
+ self,
+ output_file_id: str,
+ model_id: str,
+ model_name: str | None = None,
+ ) -> str: ...
+
+ async def store_batch_output_file(
+ self,
+ *,
+ unified_file_id: str,
+ provider_file_id: str,
+ model_id: str | None,
+ model_name: str | None = None,
+ owner: "UserAPIKeyAuth",
+ litellm_parent_otel_span: "Span | None",
+ size_bytes: int | None = None,
+ fetch_provider_details: bool = True,
+ ) -> None:
+ """Register file metadata, optionally fetching provider details."""
+ ...
+
+
def is_base64_encoded_unified_file_id(b64_uid: object) -> str | Literal[False]:
# Ensure b64_uid is a string and not a mock object
if not isinstance(b64_uid, str):
@@ -1282,19 +1309,21 @@ def apply_unified_file_ids(response: "LiteLLMBatch", unified_id_by_raw_id: Mappi
async def ensure_batch_response_managed_file_ids(
- response,
- managed_files_obj,
- prisma_client,
- verbose_proxy_logger,
- user_api_key_dict=None,
+ response: "LiteLLMBatch",
+ managed_files_obj: object | None,
+ prisma_client: "PrismaClient | None",
+ verbose_proxy_logger: logging.Logger,
+ user_api_key_dict: "UserAPIKeyAuth | None" = None,
db_batch_object: object | None = None,
unified_batch_id: str | Literal[False] | None = None,
+ *,
+ fetch_provider_details: bool = True,
) -> None:
- """Normalize batch file IDs to managed unified IDs before DB persistence."""
+ """Normalize batch file IDs and register output and error file metadata."""
await resolve_input_file_id_to_unified(response, prisma_client)
await resolve_output_file_ids_to_unified(response, prisma_client)
- if managed_files_obj is None:
+ if not isinstance(managed_files_obj, ManagedBatchOutputFileWriter):
return
model_id: Final = _model_id_for_batch_response(response, unified_batch_id)
@@ -1308,8 +1337,10 @@ async def ensure_batch_response_managed_file_ids(
if effective_auth is None:
return
- for file_attr in ("output_file_id", "error_file_id"):
- raw_file_id = getattr(response, file_attr, None)
+ for file_attr, raw_file_id in (
+ ("output_file_id", response.output_file_id),
+ ("error_file_id", response.error_file_id),
+ ):
if not raw_file_id or is_base64_encoded_unified_file_id(raw_file_id):
continue
try:
@@ -1318,12 +1349,15 @@ async def ensure_batch_response_managed_file_ids(
model_id=model_id,
model_name=model_name,
)
- await managed_files_obj.store_unified_file_id(
- file_id=new_unified_file_id,
- file_object=None,
- litellm_parent_otel_span=getattr(effective_auth, "parent_otel_span", None),
- model_mappings={model_id: raw_file_id},
- user_api_key_dict=effective_auth,
+ await managed_files_obj.store_batch_output_file(
+ unified_file_id=new_unified_file_id,
+ provider_file_id=raw_file_id,
+ model_id=model_id,
+ model_name=model_name,
+ owner=effective_auth,
+ litellm_parent_otel_span=effective_auth.parent_otel_span,
+ size_bytes=None,
+ fetch_provider_details=fetch_provider_details,
)
setattr(response, file_attr, new_unified_file_id)
verbose_proxy_logger.debug("Converted batch %s %r to managed ID before DB write", file_attr, raw_file_id)
diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py
index 4e31f9713d0..bdc71908c7c 100644
--- a/litellm/proxy/openai_files_endpoints/files_endpoints.py
+++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py
@@ -102,7 +102,7 @@ from litellm.proxy.openai_files_endpoints.general_upload_validation import (
raise_upload_validation_failure,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging, is_known_model
-from litellm.repositories.table_repositories import ManagedFileRepository
+from litellm.repositories.managed_file_repository import ManagedFileRepository
from litellm.router import Router
from litellm.types.llms.openai import (
CREATE_FILE_REQUESTS_PURPOSE,
diff --git a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py
index 94a75a9802e..a5a04fb68d2 100644
--- a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py
+++ b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py
@@ -54,10 +54,8 @@ from litellm.llms.base_llm.managed_resources.isolation import (
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.batches_endpoints.common_utils import validate_batch_list_limit
from litellm.proxy.common_utils.sse_keepalive import split_complete_sse_frames
-from litellm.repositories.table_repositories import (
- ManagedFileRepository,
- ManagedObjectRepository,
-)
+from litellm.repositories.managed_file_repository import ManagedFileRepository
+from litellm.repositories.table_repositories import ManagedObjectRepository
from litellm.types.llms.openai import BATCH_GUARDRAIL_RESPONSE_FIELD, OpenAIFileObject
from litellm.types.passthrough_endpoints.managed_id_rewriter import (
ManagedFileIdReader,
diff --git a/litellm/repositories/__init__.py b/litellm/repositories/__init__.py
index 7ffdcfa5ce6..f090911c416 100644
--- a/litellm/repositories/__init__.py
+++ b/litellm/repositories/__init__.py
@@ -6,6 +6,7 @@ from litellm.repositories.autorouter_session_repository import AutoRouterSession
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.config_repository import ConfigRepository
from litellm.repositories.credentials_repository import CredentialsRepository
+from litellm.repositories.managed_file_repository import ManagedFileRepository
from litellm.repositories.model_repository import ModelRepository
from litellm.repositories.object_permission_repository import (
ObjectPermissionRepository,
@@ -42,7 +43,6 @@ from litellm.repositories.table_repositories import (
HealthCheckRepository,
InvitationLinkRepository,
JWTKeyMappingRepository,
- ManagedFileRepository,
ManagedObjectRepository,
ManagedVectorStoreIndexRepository,
ManagedVectorStoresRepository,
diff --git a/litellm/repositories/managed_file_repository.py b/litellm/repositories/managed_file_repository.py
new file mode 100644
index 00000000000..e3cc4c1eac9
--- /dev/null
+++ b/litellm/repositories/managed_file_repository.py
@@ -0,0 +1,18 @@
+from typing import TYPE_CHECKING, Final
+
+from litellm.repositories.table_repositories import PrismaTableRepository
+from litellm.types.llms.openai import OpenAIFileObject
+
+if TYPE_CHECKING:
+ from prisma import models as prisma_models # noqa: F401 # used by quoted base-class subscripts
+
+
+class ManagedFileRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileTable"]):
+ table_name = "litellm_managedfiletable"
+
+ async def update_file_object(self, unified_file_id: str, file_object: OpenAIFileObject) -> bool:
+ updated_rows: Final = await self.table.update_many(
+ where={"unified_file_id": unified_file_id},
+ data={"file_object": file_object.model_dump_json()},
+ )
+ return updated_rows > 0
diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py
index 0f85818cba3..95b174f845f 100644
--- a/litellm/repositories/table_repositories.py
+++ b/litellm/repositories/table_repositories.py
@@ -131,10 +131,6 @@ class JWTKeyMappingRepository(PrismaTableRepository["prisma_models.LiteLLM_JWTKe
table_name = "litellm_jwtkeymapping"
-class ManagedFileRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileTable"]):
- table_name = "litellm_managedfiletable"
-
-
class MemoryRepository(PrismaTableRepository["prisma_models.LiteLLM_MemoryTable"]):
table_name = "litellm_memorytable"
diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py
index f7e3d644f34..cd77d6088f6 100644
--- a/litellm/types/llms/openai.py
+++ b/litellm/types/llms/openai.py
@@ -365,6 +365,7 @@ _JsonValue: TypeAlias = object
BATCH_GUARDRAIL_RESPONSE_FIELD: Final = "litellm_batch_guardrail"
+LITELLM_DETAILS_FALLBACK_RESPONSE_FIELD: Final = "litellm_details_fallback"
class OpenAIFileObject(LiteLLMBaseModel):
@@ -413,6 +414,9 @@ class OpenAIFileObject(LiteLLMBaseModel):
Absent on every other upload, so OpenAI-shaped clients see an unchanged response.
"""
+ litellm_details_fallback: bool | None = None
+ """Set by the LiteLLM proxy on a saved batch output file entry built without provider metadata; stripped from API responses."""
+
_hidden_params: dict = PrivateAttr(default={"response_cost": 0.0}) # no cost for writing a file
@property
@@ -424,13 +428,19 @@ class OpenAIFileObject(LiteLLMBaseModel):
self._hidden_params = hidden_params
@model_serializer(mode="wrap")
- def _omit_absent_batch_guardrail( # noqa: ANN202 # annotating it replaces the model's serialization schema
+ def _omit_absent_proxy_only_fields( # noqa: ANN202 # annotating it replaces the model's serialization schema
self, handler: SerializerFunctionWrapHandler
):
serialized: Final[Mapping[str, object]] = handler(self)
- if self.litellm_batch_guardrail is not None:
- return serialized
- return {key: value for key, value in serialized.items() if key != BATCH_GUARDRAIL_RESPONSE_FIELD}
+ fields_to_omit: Final = tuple(
+ field_name
+ for field_name, value in (
+ (BATCH_GUARDRAIL_RESPONSE_FIELD, self.litellm_batch_guardrail),
+ (LITELLM_DETAILS_FALLBACK_RESPONSE_FIELD, self.litellm_details_fallback),
+ )
+ if value is None
+ )
+ return {key: value for key, value in serialized.items() if key not in fields_to_omit}
def __contains__(self, key) -> bool:
# Define custom behavior for the 'in' operator
diff --git a/tests/integration/management/test_batch_output_file_listing.py b/tests/integration/management/test_batch_output_file_listing.py
new file mode 100644
index 00000000000..26d1ae8de46
--- /dev/null
+++ b/tests/integration/management/test_batch_output_file_listing.py
@@ -0,0 +1,651 @@
+from __future__ import annotations
+
+import contextlib
+import datetime
+import json
+import socket
+import socketserver
+import ssl
+import threading
+import uuid
+from collections.abc import Generator
+from contextlib import contextmanager
+from dataclasses import dataclass
+from pathlib import Path
+from queue import SimpleQueue
+from typing import Final
+
+import httpx
+from cryptography import x509
+from cryptography.hazmat.primitives import hashes, serialization
+from cryptography.hazmat.primitives.asymmetric import ec
+from cryptography.x509.oid import NameOID
+from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value
+from integration._support.database import read_rows
+from integration._support.process import owned_proxy
+from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
+from integration._support.wire import Reply, Request, Wire, wire_server
+from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse
+from pydantic import JsonValue
+
+from litellm.litellm_core_utils.cloud_storage_security import BEDROCK_MANAGED_S3_OUTPUT_PREFIX
+
+OUTPUT_BYTES: Final = 4096
+ERROR_BYTES: Final = 512
+BEDROCK_MODEL: Final = "bedrock/anthropic.claude-3-haiku-20240307-v1:0"
+BEDROCK_MODEL_ID: Final = "anthropic.claude-3-haiku-20240307-v1:0"
+BEDROCK_REGION: Final = "us-east-1"
+BEDROCK_AUTHORITY: Final = f"bedrock.{BEDROCK_REGION}.amazonaws.com:443"
+BEDROCK_BUCKET: Final = "integration-batch-listing-bucket"
+BEDROCK_ROLE_ARN: Final = "arn:aws:iam::123456789012:role/integration-batch-role"
+BEDROCK_JOB_ARN_PREFIX: Final = f"arn:aws:bedrock:{BEDROCK_REGION}:123456789012:model-invocation-job/"
+BEDROCK_LAST_MODIFIED: Final = "Thu, 02 Oct 2025 12:00:00 GMT"
+BEDROCK_OUTPUT_CONTENT: Final = b'{"recordId":"req-1","modelOutput":{}}\n'
+_BATCH_PROCESSED_SQL: Final = 'SELECT batch_processed FROM "LiteLLM_ManagedObjectTable" WHERE unified_object_id=%s'
+
+
+@dataclass(frozen=True, slots=True)
+class _OpenAIBatch:
+ model: str
+ owner_key: str
+ unrelated_key: str | None
+ batch_id: str
+ input_file_id: str
+ scenario: ScenarioHandle
+
+
+def _output_content(model: str) -> str:
+ return (
+ json.dumps(
+ {
+ "id": "batch_req_$REQUEST_ID",
+ "custom_id": "req-1",
+ "response": {
+ "status_code": 200,
+ "request_id": "$REQUEST_ID",
+ "body": {
+ "id": "chatcmpl-$REQUEST_ID",
+ "object": "chat.completion",
+ "model": model,
+ "choices": [
+ {
+ "index": 0,
+ "message": {"role": "assistant", "content": "ok"},
+ "finish_reason": "stop",
+ }
+ ],
+ "usage": {"prompt_tokens": 2, "completion_tokens": 1, "total_tokens": 3},
+ },
+ },
+ "error": None,
+ },
+ separators=(",", ":"),
+ )
+ + "\n"
+ )
+
+
+def _error_content() -> str:
+ return (
+ json.dumps(
+ {
+ "id": "batch_req_$REQUEST_ID",
+ "custom_id": "req-2",
+ "response": {"status_code": 400, "body": {"error": {"message": "rejected"}}},
+ "error": {"code": "bad_request", "message": "rejected"},
+ },
+ separators=(",", ":"),
+ )
+ + "\n"
+ )
+
+
+def _batch_routes(
+ model: str,
+ *,
+ output_file: bool = True,
+ metadata_fails: bool = False,
+) -> RoutedResponse:
+ completed: Final[dict[str, JsonValue]] = {
+ "id": "batch-$REQUEST_ID",
+ "object": "batch",
+ "endpoint": "/v1/chat/completions",
+ "errors": None,
+ "input_file_id": "file-in-$REQUEST_ID",
+ "completion_window": "24h",
+ "status": "completed",
+ "output_file_id": "file-out-$REQUEST_ID" if output_file else None,
+ "error_file_id": "file-err-$REQUEST_ID",
+ "created_at": 1,
+ "in_progress_at": 1,
+ "completed_at": 1,
+ "expires_at": 1,
+ "request_counts": {"total": 1 if not output_file else 2, "completed": 1 if output_file else 0, "failed": 1},
+ "metadata": None,
+ }
+ output_metadata: Final = JsonResponse(
+ content_type="application/json",
+ body=(
+ {"error": "provider metadata unavailable"}
+ if metadata_fails
+ else {
+ "id": "file-out-$REQUEST_ID",
+ "object": "file",
+ "purpose": "batch_output",
+ "bytes": OUTPUT_BYTES,
+ "created_at": 1,
+ "filename": "output.jsonl",
+ "status": "processed",
+ }
+ ),
+ status=500 if metadata_fails else 200,
+ )
+ return RoutedResponse(
+ content_type="application/x-routed",
+ routes={
+ "POST /files": JsonResponse(
+ content_type="application/json",
+ body={
+ "id": "file-in-$REQUEST_ID",
+ "object": "file",
+ "purpose": "batch",
+ "bytes": 100,
+ "created_at": 1,
+ "filename": "input.jsonl",
+ "status": "processed",
+ },
+ ),
+ "POST /batches": JsonResponse(
+ content_type="application/json",
+ body={**completed, "status": "validating", "output_file_id": None, "error_file_id": None},
+ ),
+ "GET /batches/batch-$REQUEST_ID": JsonResponse(content_type="application/json", body=completed),
+ "GET /files/file-out-$REQUEST_ID": output_metadata,
+ "GET /files/file-err-$REQUEST_ID": JsonResponse(
+ content_type="application/json",
+ body={
+ "id": "file-err-$REQUEST_ID",
+ "object": "file",
+ "purpose": "batch_output",
+ "bytes": ERROR_BYTES,
+ "created_at": 1,
+ "filename": "errors.jsonl",
+ "status": "processed",
+ },
+ ),
+ "GET /files/file-out-$REQUEST_ID/content": TextResponse(
+ content_type="application/jsonl",
+ body=_output_content(model),
+ ),
+ "GET /files/file-err-$REQUEST_ID/content": TextResponse(
+ content_type="application/jsonl",
+ body=_error_content(),
+ ),
+ },
+ )
+
+
+def _create_batch(
+ scenario: Scenario,
+ routes: RoutedResponse,
+ *,
+ unrelated_user: bool = False,
+) -> _OpenAIBatch:
+ handle: Final = register_scenario(f"batch-output-listing-{uuid.uuid4().hex}", routes)
+ scenario.cleanups.callback(delete_scenario, handle)
+ model: Final = scenario.model(api_base=handle.api_base())
+ owner_id: Final = scenario.user(user_role="internal_user")
+ owner_key: Final = scenario.key(user_id=owner_id, models=[model])
+ unrelated_key: Final = (
+ scenario.key(user_id=scenario.user(user_role="internal_user"), models=[model]) if unrelated_user else None
+ )
+ input_content: Final = json.dumps(
+ {
+ "custom_id": "req-1",
+ "method": "POST",
+ "url": "/v1/chat/completions",
+ "body": {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": 8},
+ }
+ ).encode()
+ uploaded: Final = scenario.gateway.request_multipart(
+ "/v1/files",
+ {"purpose": "batch", "target_model_names": model},
+ {"file": ("input.jsonl", input_content, "application/jsonl")},
+ key=owner_key,
+ )
+ assert uploaded.status_code == 200, uploaded.text
+ input_file_id: Final = string_value(JSON_OBJECT.validate_json(uploaded.content)["id"])
+ created: Final = scenario.gateway.request(
+ "POST",
+ "/v1/batches",
+ {
+ "input_file_id": input_file_id,
+ "endpoint": "/v1/chat/completions",
+ "completion_window": "24h",
+ "model": model,
+ },
+ key=owner_key,
+ )
+ assert created.status_code == 200, created.text
+ batch_id: Final = string_value(JSON_OBJECT.validate_json(created.content)["id"])
+ return _OpenAIBatch(model, owner_key, unrelated_key, batch_id, input_file_id, handle)
+
+
+def _retrieve_batch(gateway: Gateway, batch: _OpenAIBatch) -> dict[str, JsonValue]:
+ response: Final = gateway.request("GET", f"/v1/batches/{batch.batch_id}", key=batch.owner_key)
+ assert response.status_code == 200, response.text
+ return JSON_OBJECT.validate_json(response.content)
+
+
+def _list_files(
+ gateway: Gateway,
+ key: str,
+ *,
+ purpose: str | None = None,
+) -> tuple[dict[str, JsonValue], ...]:
+ response: Final = gateway.request(
+ "GET",
+ "/v1/files",
+ key=key,
+ params={"purpose": purpose} if purpose is not None else None,
+ )
+ assert response.status_code == 200, response.text
+ values: Final = JSON_OBJECT.validate_json(response.content)["data"]
+ assert isinstance(values, list), response.text
+ return tuple(object_value(value) for value in values)
+
+
+def _metadata_hit_count(gateway: Gateway, scenario: ScenarioHandle) -> int:
+ response: Final = httpx.get(
+ f"{gateway.upstream_url}/__observations",
+ timeout=15,
+ trust_env=False,
+ )
+ assert response.status_code == 200, response.text
+ requests: Final = JSON_OBJECT.validate_json(response.content)["requests"]
+ assert isinstance(requests, list), response.text
+ return sum(1 for request in requests if _is_metadata_hit(request, scenario.scenario_id))
+
+
+def _is_metadata_hit(request: JsonValue, scenario_id: str) -> bool:
+ if not isinstance(request, dict):
+ return False
+ path: Final = request.get("path")
+ return request.get("method") == "GET" and isinstance(path, str) and path.endswith(f"/files/file-out-{scenario_id}")
+
+
+def test_output_file_lists_after_owner_retrieves_completed_batch(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ batch: Final = _create_batch(scenario, _batch_routes(model="gpt-4o-mini"))
+ retrieved: Final = _retrieve_batch(gateway, batch)
+ assert retrieved["status"] == "completed", retrieved
+ output_id: Final = string_value(retrieved["output_file_id"])
+ error_id: Final = string_value(retrieved["error_file_id"])
+ files: Final = _list_files(gateway, batch.owner_key)
+ output: Final = next((file for file in files if file.get("id") == output_id), None)
+ assert output is not None, f"Completed batch output {output_id} is absent from GET /v1/files"
+ assert (output["purpose"], output["bytes"]) == ("batch_output", OUTPUT_BYTES), output
+ output_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output")
+ assert {string_value(file["id"]) for file in output_files} == {output_id, error_id}, output_files
+ assert all(file["purpose"] == "batch_output" for file in output_files), output_files
+ input_files: Final = _list_files(gateway, batch.owner_key, purpose="batch")
+ assert tuple(string_value(file["id"]) for file in input_files) == (batch.input_file_id,), input_files
+
+
+def test_poller_registers_listable_output_files_without_a_batch_retrieve(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ batch: Final = _create_batch(
+ scenario,
+ _batch_routes(model="gpt-4o-mini"),
+ unrelated_user=True,
+ )
+ assert batch.unrelated_key is not None
+ output_files: Final = eventually(
+ lambda: _list_files(gateway, batch.owner_key, purpose="batch_output"),
+ lambda values: len(values) == 2,
+ seconds=60,
+ )
+ assert {file["purpose"] for file in output_files} == {"batch_output"}, output_files
+ unrelated_files: Final = _list_files(gateway, batch.unrelated_key, purpose="batch_output")
+ assert unrelated_files == (), unrelated_files
+ retrieved: Final = _retrieve_batch(gateway, batch)
+ output_id: Final = string_value(retrieved["output_file_id"])
+ error_id: Final = string_value(retrieved["error_file_id"])
+ assert {string_value(file["id"]) for file in output_files} == {output_id, error_id}, output_files
+
+
+def test_error_file_only_batch_still_lists_its_error_file(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ batch: Final = _create_batch(
+ scenario,
+ _batch_routes(model="gpt-4o-mini", output_file=False),
+ )
+ error_files: Final = eventually(
+ lambda: _list_files(gateway, batch.owner_key, purpose="batch_output"),
+ lambda values: len(values) == 1,
+ seconds=60,
+ )
+ retrieved: Final = _retrieve_batch(gateway, batch)
+ assert retrieved["output_file_id"] is None, retrieved
+ error_id: Final = string_value(retrieved["error_file_id"])
+ assert tuple(string_value(file["id"]) for file in error_files) == (error_id,), error_files
+ assert error_files[0]["purpose"] == "batch_output", error_files
+
+
+def test_output_file_lists_with_basic_details_when_provider_file_lookup_fails(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ batch: Final = _create_batch(
+ scenario,
+ _batch_routes(model="gpt-4o-mini", metadata_fails=True),
+ )
+ retrieved: Final = _retrieve_batch(gateway, batch)
+ output_id: Final = string_value(retrieved["output_file_id"])
+ output: Final = next(
+ (file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id),
+ None,
+ )
+ assert output is not None, f"Completed batch output {output_id} is absent after provider metadata failure"
+ assert (output["purpose"], output["filename"]) == ("batch_output", f"file-out-{batch.scenario.scenario_id}"), (
+ output
+ )
+ assert "litellm_details_fallback" not in output, output
+
+
+def test_fallback_output_file_details_refresh_once_provider_recovers(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ batch: Final = _create_batch(
+ scenario,
+ _batch_routes(model="gpt-4o-mini", metadata_fails=True),
+ unrelated_user=True,
+ )
+ retrieved: Final = _retrieve_batch(gateway, batch)
+ output_id: Final = string_value(retrieved["output_file_id"])
+ listed_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output")
+ basic_output: Final = next((file for file in listed_files if file.get("id") == output_id), None)
+ assert basic_output is not None, f"Completed batch output {output_id} is absent after provider metadata failure"
+ assert (basic_output["purpose"], basic_output["filename"]) == (
+ "batch_output",
+ f"file-out-{batch.scenario.scenario_id}",
+ ), basic_output
+ assert basic_output["bytes"] != OUTPUT_BYTES, basic_output
+ assert "litellm_details_fallback" not in basic_output, basic_output
+
+ processed_batch_rows: Final = eventually(
+ lambda: read_rows(_BATCH_PROCESSED_SQL, (string_value(retrieved["id"]),)),
+ lambda rows: len(rows) == 1 and rows[0].get("batch_processed") is True,
+ seconds=60,
+ )
+ assert processed_batch_rows[0]["batch_processed"] is True, processed_batch_rows
+ metadata_hits_before_recovery: Final = _metadata_hit_count(gateway, batch.scenario)
+ assert metadata_hits_before_recovery >= 1, "The provider metadata route was not called before recovery"
+ register_scenario(
+ batch.scenario.scenario_id,
+ _batch_routes(model=batch.model),
+ control_url=batch.scenario.control_url,
+ )
+ details_response: Final = gateway.request("GET", f"/v1/files/{output_id}", key=batch.owner_key)
+ assert details_response.status_code == 200, details_response.text
+ details: Final = JSON_OBJECT.validate_json(details_response.content)
+ assert details["bytes"] == OUTPUT_BYTES, details
+ assert (details["filename"], details["purpose"]) == ("output.jsonl", "batch_output"), details
+ assert "litellm_details_fallback" not in details, details
+ assert _metadata_hit_count(gateway, batch.scenario) == 1
+
+ refreshed_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output")
+ refreshed_output: Final = next((file for file in refreshed_files if file.get("id") == output_id), None)
+ assert refreshed_output is not None, f"Refreshed batch output {output_id} is absent from the list"
+ assert (refreshed_output["bytes"], refreshed_output["filename"]) == (OUTPUT_BYTES, "output.jsonl"), (
+ refreshed_output
+ )
+ assert "litellm_details_fallback" not in refreshed_output, refreshed_output
+
+ assert batch.unrelated_key is not None
+ unrelated_details: Final = gateway.request("GET", f"/v1/files/{output_id}", key=batch.unrelated_key)
+ assert unrelated_details.status_code == 403, unrelated_details.text
+
+
+def _tls_context(directory: Path) -> ssl.SSLContext:
+ key: Final = ec.generate_private_key(ec.SECP256R1())
+ now: Final = datetime.datetime.now(datetime.timezone.utc)
+ name: Final = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, BEDROCK_AUTHORITY.split(":")[0])])
+ certificate: Final = (
+ x509.CertificateBuilder()
+ .subject_name(name)
+ .issuer_name(name)
+ .public_key(key.public_key())
+ .serial_number(x509.random_serial_number())
+ .not_valid_before(now - datetime.timedelta(days=1))
+ .not_valid_after(now + datetime.timedelta(days=1))
+ .sign(key, hashes.SHA256())
+ )
+ certificate_file: Final = directory / "bedrock.pem"
+ key_file: Final = directory / "bedrock.key"
+ certificate_file.write_bytes(certificate.public_bytes(serialization.Encoding.PEM))
+ key_file.write_bytes(
+ key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption())
+ )
+ context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
+ context.load_cert_chain(certificate_file, key_file)
+ return context
+
+
+def _pipe(source: socket.socket, sink: socket.socket) -> None:
+ with contextlib.suppress(OSError):
+ for chunk in iter(lambda: source.recv(65536), b""):
+ sink.sendall(chunk)
+ with contextlib.suppress(OSError):
+ sink.shutdown(socket.SHUT_WR)
+
+
+@dataclass(frozen=True, slots=True)
+class _ConnectProxy:
+ url: str
+ authorities: SimpleQueue[str]
+
+
+@contextmanager
+def _bedrock_tunnel(destination: Wire) -> Generator[_ConnectProxy, None, None]:
+ authorities: Final[SimpleQueue[str]] = SimpleQueue()
+ destination_port: Final = int(destination.url.rsplit(":", 1)[1])
+
+ class Tunnel(socketserver.StreamRequestHandler):
+ rbufsize = 0
+ request: socket.socket
+
+ def handle(self) -> None:
+ authority: Final = self.rfile.readline().decode().split()[1]
+ while self.rfile.readline() not in (b"\r\n", b""):
+ pass
+ authorities.put(authority)
+ if authority != BEDROCK_AUTHORITY:
+ self.wfile.write(b"HTTP/1.1 403 Forbidden\r\ncontent-length: 0\r\n\r\n")
+ return
+ self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n")
+ self.request.settimeout(10)
+ with socket.create_connection(("127.0.0.1", destination_port), timeout=10) as upstream:
+ outbound: Final = threading.Thread(target=_pipe, args=(self.request, upstream))
+ outbound.start()
+ _pipe(upstream, self.request)
+ outbound.join(timeout=12)
+
+ with socketserver.ThreadingTCPServer(("127.0.0.1", 0), Tunnel) as server:
+ thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05})
+ thread.start()
+ try:
+ yield _ConnectProxy(f"http://127.0.0.1:{server.server_address[1]}", authorities)
+ finally:
+ server.shutdown()
+ thread.join(timeout=6)
+
+
+@dataclass(frozen=True, slots=True)
+class _BedrockControlPlane:
+ job_locations: SimpleQueue[tuple[str, str, str]]
+ job_id: str
+ job_arn: str
+
+ def __call__(self, request: Request) -> Reply:
+ if request.method == "POST" and request.target == "/model-invocation-job":
+ body: Final = JSON_OBJECT.validate_json(request.body)
+ input_config: Final = object_value(object_value(body["inputDataConfig"])["s3InputDataConfig"])
+ output_config: Final = object_value(object_value(body["outputDataConfig"])["s3OutputDataConfig"])
+ job_name: Final = string_value(body["jobName"])
+ self.job_locations.put(
+ (string_value(input_config["s3Uri"]), string_value(output_config["s3Uri"]), job_name)
+ )
+ return Reply(body=json.dumps({"jobArn": self.job_arn}).encode())
+ if request.method == "GET" and request.target.endswith(self.job_id):
+ input_uri, output_uri, retrieved_job_name = self.job_locations.get()
+ self.job_locations.put((input_uri, output_uri, retrieved_job_name))
+ return Reply(
+ body=json.dumps(
+ {
+ "jobArn": self.job_arn,
+ "jobName": retrieved_job_name,
+ "modelId": BEDROCK_MODEL_ID,
+ "status": "Completed",
+ "submitTime": 1700000000,
+ "lastModifiedTime": 1700000001,
+ "endTime": 1700000002,
+ "inputDataConfig": {"s3InputDataConfig": {"s3Uri": input_uri}},
+ "outputDataConfig": {"s3OutputDataConfig": {"s3Uri": output_uri}},
+ "totalRecordCount": 1,
+ "successRecordCount": 1,
+ "errorRecordCount": 0,
+ }
+ ).encode()
+ )
+ return Reply(status=404, body=b'{"message":"not scripted"}')
+
+
+def _bedrock_s3_peer(request: Request) -> Reply:
+ if request.method == "PUT" and request.target.startswith(f"/{BEDROCK_BUCKET}/"):
+ return Reply(body=b"")
+ if request.method == "GET" and request.target.startswith(f"/{BEDROCK_BUCKET}/{BEDROCK_MANAGED_S3_OUTPUT_PREFIX}"):
+ if request.headers.get("range") == "bytes=0-0":
+ return Reply(
+ status=206,
+ body=BEDROCK_OUTPUT_CONTENT[:1],
+ headers={
+ "Content-Range": f"bytes 0-0/{len(BEDROCK_OUTPUT_CONTENT)}",
+ "Last-Modified": BEDROCK_LAST_MODIFIED,
+ },
+ )
+ return Reply(status=200, body=BEDROCK_OUTPUT_CONTENT, headers={"Last-Modified": BEDROCK_LAST_MODIFIED})
+ return Reply(status=404, body=b'{"message":"not scripted"}')
+
+
+def test_bedrock_batch_output_lists_and_retrieves_details(gateway: Gateway, tmp_path: Path) -> None:
+ job_id: Final = f"integration-batch-listing-{uuid.uuid4().hex}"
+ job_arn: Final = BEDROCK_JOB_ARN_PREFIX + job_id
+ environment: Final = {
+ "AWS_CA_BUNDLE": str(tmp_path / "bedrock.pem"),
+ "AWS_EC2_METADATA_DISABLED": "true",
+ "SSL_VERIFY": "False",
+ }
+ with (
+ wire_server(_bedrock_s3_peer) as s3,
+ wire_server(_BedrockControlPlane(SimpleQueue(), job_id, job_arn), tls=_tls_context(tmp_path)) as bedrock,
+ _bedrock_tunnel(bedrock) as tunnel,
+ owned_proxy(gateway, tmp_path, {**environment, "HTTPS_PROXY": tunnel.url}) as candidate,
+ candidate.scenario() as scenario,
+ ):
+ model: Final = scenario.model(
+ model=BEDROCK_MODEL,
+ api_key=None,
+ api_base=None,
+ aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
+ aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
+ aws_region_name=BEDROCK_REGION,
+ s3_bucket_name=BEDROCK_BUCKET,
+ s3_endpoint_url=s3.url,
+ aws_batch_role_arn=BEDROCK_ROLE_ARN,
+ )
+ owner_id: Final = scenario.user(user_role="internal_user")
+ owner_key: Final = scenario.key(user_id=owner_id, models=[model])
+ input_line: Final = {
+ "custom_id": "req-1",
+ "method": "POST",
+ "url": "/v1/chat/completions",
+ "body": {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": 8},
+ }
+ uploaded: Final = candidate.request_multipart(
+ "/v1/files",
+ {"purpose": "batch", "target_model_names": model},
+ {"file": ("input.jsonl", (json.dumps(input_line) + "\n").encode(), "application/jsonl")},
+ key=owner_key,
+ )
+ assert uploaded.status_code == 200, uploaded.text
+ created: Final = candidate.request(
+ "POST",
+ "/v1/batches",
+ {
+ "input_file_id": string_value(JSON_OBJECT.validate_json(uploaded.content)["id"]),
+ "endpoint": "/v1/chat/completions",
+ "completion_window": "24h",
+ },
+ key=owner_key,
+ )
+ assert created.status_code == 200, created.text
+ batch_id: Final = string_value(JSON_OBJECT.validate_json(created.content)["id"])
+ retrieved: Final = candidate.request("GET", f"/v1/batches/{batch_id}", key=owner_key)
+ assert retrieved.status_code == 200, retrieved.text
+ batch_object: Final = JSON_OBJECT.validate_json(retrieved.content)
+ assert batch_object["status"] == "completed", retrieved.text
+ output_id: Final = string_value(batch_object["output_file_id"])
+ output_files: Final = eventually(
+ lambda: _list_files(candidate, owner_key, purpose="batch_output"),
+ lambda values: any(file.get("id") == output_id for file in values),
+ seconds=30,
+ )
+ listed_output: Final = next(file for file in output_files if file.get("id") == output_id)
+ details: Final = candidate.request("GET", f"/v1/files/{output_id}", key=owner_key)
+ assert details.status_code == 200, details.text
+ detail_object: Final = JSON_OBJECT.validate_json(details.content)
+ assert (detail_object["id"], detail_object["bytes"], detail_object["purpose"]) == (
+ output_id,
+ len(BEDROCK_OUTPUT_CONTENT),
+ "batch_output",
+ ), detail_object
+ assert listed_output["id"] == output_id, listed_output
+ s3_requests: Final = s3.drain()
+ uploads: Final = tuple(request for request in s3_requests if request.method == "PUT")
+ assert len(uploads) == 1, f"Expected one input S3 upload, saw {[request.target for request in uploads]}"
+ ranged_metadata: Final = tuple(
+ request
+ for request in s3_requests
+ if request.method == "GET" and request.headers.get("range") == "bytes=0-0"
+ )
+ assert ranged_metadata, "Bedrock file metadata retrieval did not issue a ranged S3 GET"
+ assert all(
+ request.headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ") for request in ranged_metadata
+ ), ranged_metadata
+ authorities: Final = tuple(tunnel.authorities.get_nowait() for _ in range(tunnel.authorities.qsize()))
+ assert BEDROCK_AUTHORITY in authorities, authorities
+ bedrock_requests: Final = bedrock.drain()
+ assert any(request.method == "POST" for request in bedrock_requests), bedrock_requests
+ assert any(request.method == "GET" for request in bedrock_requests), bedrock_requests
+
+
+def test_repeated_batch_retrieve_does_not_refetch_saved_file_details(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ batch: Final = _create_batch(scenario, _batch_routes(model="gpt-4o-mini"))
+ first: Final = _retrieve_batch(gateway, batch)
+ output_id: Final = string_value(first["output_file_id"])
+ output_before: Final = next(
+ (file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id),
+ None,
+ )
+ assert output_before is not None, f"Completed batch output {output_id} is absent from GET /v1/files"
+ hits_before: Final = _metadata_hit_count(gateway, batch.scenario)
+ assert hits_before >= 1, "The provider file metadata route was not called for the batch output"
+ second: Final = _retrieve_batch(gateway, batch)
+ third: Final = _retrieve_batch(gateway, batch)
+ assert second["status"] == third["status"] == "completed", (second, third)
+ output_after: Final = next(
+ (file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id),
+ None,
+ )
+ assert output_after == output_before, output_after
+ additional_hits: Final = _metadata_hit_count(gateway, batch.scenario)
+ assert additional_hits == 0, f"Repeated batch retrieval fetched metadata {additional_hits} more times"
diff --git a/tests/unit/enterprise/proxy/hooks/test_managed_files.py b/tests/unit/enterprise/proxy/hooks/test_managed_files.py
index 11d9dbec49c..ed80bd9b928 100644
--- a/tests/unit/enterprise/proxy/hooks/test_managed_files.py
+++ b/tests/unit/enterprise/proxy/hooks/test_managed_files.py
@@ -1,19 +1,268 @@
import base64
import json
-from typing import cast
+import logging
+import time
+import asyncio
+from collections.abc import Awaitable, Callable, Mapping
+from types import SimpleNamespace
+from typing import TYPE_CHECKING, Final, cast
from unittest.mock import AsyncMock, MagicMock, patch
+import httpx
import pytest
from fastapi import HTTPException
-from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles
+from openai import APIConnectionError
+if TYPE_CHECKING:
+ from litellm.types.utils import LiteLLMBatch
+
+from litellm_enterprise.proxy.hooks.managed_files import (
+ PROXY_LiteLLMManagedFiles,
+ _provider_file_retrieve_credentials,
+)
from litellm.caching import DualCache
-from litellm.proxy._types import CallTypes
+from litellm.proxy._types import CallTypes, LiteLLM_ManagedFileTable
from litellm.proxy.openai_files_endpoints.common_utils import (
BATCH_CREATE_HIDDEN_PARAM,
is_base64_encoded_unified_file_id,
encode_file_id_with_model,
)
+from litellm.types.llms.openai import OpenAIFileObject
+
+
+class _InMemoryManagedFileTable:
+ def __init__(self, *rows: LiteLLM_ManagedFileTable) -> None:
+ self.rows: dict[str, LiteLLM_ManagedFileTable] = {
+ row.unified_file_id: row for row in rows
+ }
+ self.find_first_calls: list[Mapping[str, object]] = []
+ self.upsert_calls: list[
+ tuple[Mapping[str, object], Mapping[str, Mapping[str, object]]]
+ ] = []
+ self.update_many_calls: list[
+ tuple[Mapping[str, object], Mapping[str, object]]
+ ] = []
+
+ async def find_first(
+ self, where: Mapping[str, object]
+ ) -> LiteLLM_ManagedFileTable | None:
+ self.find_first_calls.append(where)
+ unified_file_id: Final = where.get("unified_file_id")
+ if isinstance(unified_file_id, str):
+ return self.rows.get(unified_file_id)
+ raw_file_filter: Final = where.get("flat_model_file_ids")
+ if isinstance(raw_file_filter, Mapping):
+ raw_file_id: Final = raw_file_filter.get("has")
+ if isinstance(raw_file_id, str):
+ return next(
+ (row for row in self.rows.values() if raw_file_id in row.flat_model_file_ids),
+ None,
+ )
+ return None
+
+ async def upsert(
+ self,
+ where: Mapping[str, object],
+ data: Mapping[str, Mapping[str, object]],
+ ) -> LiteLLM_ManagedFileTable:
+ self.upsert_calls.append((where, data))
+ unified_file_id: Final = cast(str, where["unified_file_id"])
+ previous_row: Final = self.rows.get(unified_file_id)
+ values: Final = {
+ **(
+ cast(dict[str, object], previous_row.model_dump())
+ if previous_row is not None
+ else {}
+ ),
+ **data["update" if previous_row is not None else "create"],
+ }
+ raw_model_mappings: Final = values.get("model_mappings", "{}")
+ model_mappings: Final = (
+ cast(dict[str, str], json.loads(raw_model_mappings))
+ if isinstance(raw_model_mappings, str)
+ else cast(dict[str, str], raw_model_mappings)
+ )
+ raw_file_object: Final = values.get("file_object")
+ file_object: Final = cast(
+ dict[str, object] | None,
+ json.loads(raw_file_object)
+ if isinstance(raw_file_object, str)
+ else raw_file_object,
+ )
+ row: Final = LiteLLM_ManagedFileTable.model_validate(
+ {
+ **values,
+ "file_object": file_object,
+ "model_mappings": model_mappings,
+ }
+ )
+ self.rows[unified_file_id] = row
+ return row
+
+ async def update_many(
+ self,
+ where: Mapping[str, object],
+ data: Mapping[str, object],
+ ) -> int:
+ self.update_many_calls.append((where, data))
+ unified_file_id: Final = cast(str, where["unified_file_id"])
+ existing_row: Final = self.rows.get(unified_file_id)
+ if existing_row is None:
+ return 0
+ file_object: Final = OpenAIFileObject.model_validate(
+ json.loads(cast(str, data["file_object"]))
+ )
+ self.rows[unified_file_id] = existing_row.model_copy(
+ update={"file_object": file_object}
+ )
+ return 1
+
+ async def delete(
+ self, where: Mapping[str, object]
+ ) -> LiteLLM_ManagedFileTable | None:
+ return self.rows.pop(cast(str, where["unified_file_id"]), None)
+
+ async def find_many(
+ self,
+ where: Mapping[str, object],
+ **_query: object,
+ ) -> list[LiteLLM_ManagedFileTable]:
+ return list(self.rows.values())
+
+
+class _InMemoryManagedObjectTable:
+ async def find_first(self, where: Mapping[str, object]) -> None:
+ return None
+
+ async def update_many(
+ self, where: Mapping[str, object], data: Mapping[str, object]
+ ) -> int:
+ return 0
+
+ async def upsert(
+ self,
+ where: Mapping[str, object],
+ data: Mapping[str, object],
+ ) -> None:
+ return None
+
+
+class _InMemoryManagedFilesDatabase:
+ def __init__(self, managed_file_table: _InMemoryManagedFileTable) -> None:
+ self.litellm_managedfiletable: Final = managed_file_table
+ self.litellm_managedobjecttable: Final = _InMemoryManagedObjectTable()
+
+
+class _InMemoryManagedFilesPrismaClient:
+ def __init__(self, managed_file_table: _InMemoryManagedFileTable) -> None:
+ self.db: Final = _InMemoryManagedFilesDatabase(managed_file_table)
+
+
+def _managed_files_with_fake_prisma(
+ *rows: LiteLLM_ManagedFileTable,
+ sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
+) -> tuple[PROXY_LiteLLMManagedFiles, _InMemoryManagedFileTable]:
+ managed_file_table: Final = _InMemoryManagedFileTable(*rows)
+ prisma_client: Final = _InMemoryManagedFilesPrismaClient(managed_file_table)
+ managed_files: Final = PROXY_LiteLLMManagedFiles(
+ DualCache(), prisma_client=prisma_client, sleep=sleep
+ )
+ return managed_files, managed_file_table
+
+
+def _marked_fallback_file_row(
+ *,
+ size_bytes: int = 0,
+ created_at: int = 123,
+) -> LiteLLM_ManagedFileTable:
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ file_object: Final = OpenAIFileObject(
+ id="unified-output",
+ object="file",
+ bytes=size_bytes,
+ created_at=created_at,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ litellm_details_fallback=True,
+ )
+ return LiteLLM_ManagedFileTable(
+ unified_file_id="unified-output",
+ file_object=file_object,
+ model_mappings={"model-123": "provider-output"},
+ flat_model_file_ids=["provider-output"],
+ created_by="user-123",
+ team_id="team-123",
+ )
+
+
+def _batch_response_for_output_listing() -> "LiteLLMBatch":
+ from openai.types.batch import BatchRequestCounts
+ from litellm.types.utils import LiteLLMBatch
+
+ response: Final = LiteLLMBatch(
+ id="batch-123",
+ completion_window="24h",
+ created_at=1,
+ endpoint="/v1/chat/completions",
+ input_file_id="input-file",
+ object="batch",
+ status="completed",
+ output_file_id="provider-output",
+ request_counts=BatchRequestCounts(completed=1, failed=0, total=1),
+ )
+ response.hidden_params = {
+ "model_id": "model-123",
+ "model_name": "bedrock/model-x",
+ }
+ return response
+
+
+async def _resolve_batch_for_output_listing(
+ managed_files: PROXY_LiteLLMManagedFiles,
+ response: "LiteLLMBatch",
+) -> "LiteLLMBatch | None":
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ return await managed_files._resolve_listed_batch(
+ row=SimpleNamespace(
+ unified_object_id="batch-row",
+ created_by="user-123",
+ team_id="team-123",
+ ),
+ batch_obj=response,
+ unified_id_by_raw_id={},
+ user_api_key_dict=UserAPIKeyAuth(user_id="user-123", team_id="team-123"),
+ )
+
+
+def test_provider_credentials_warning_sanitizes_newlines(
+ caplog: pytest.LogCaptureFixture,
+) -> None:
+ model_id: Final = "deployment-123\nforged warning"
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.side_effect = RuntimeError(
+ "credential lookup failed\nforged warning"
+ )
+
+ with caplog.at_level(logging.WARNING, logger="LiteLLM"):
+ credentials: Final = _provider_file_retrieve_credentials(
+ llm_router=router,
+ model_id=model_id,
+ )
+
+ warnings: Final = tuple(
+ record.getMessage()
+ for record in caplog.records
+ if "Failed to retrieve credentials for provider file" in record.getMessage()
+ )
+ assert credentials is None
+ assert warnings == (
+ "Failed to retrieve credentials for provider file "
+ "model_id=deployment-123forged warning: credential lookup failedforged warning",
+ )
+ assert all("\n" not in warning and "\r" not in warning for warning in warnings)
def test_get_file_ids_from_messages():
@@ -480,13 +729,11 @@ async def test_router_acreate_batch_only_selects_from_file_id_mapping(monkeypatc
@pytest.mark.asyncio
async def test_output_file_id_for_batch_retrieve():
- """
- Test that the output file id is the same as the input file id
- """
- from typing import cast
-
+ import litellm.proxy.proxy_server as proxy_server_module
from openai.types.batch import BatchRequestCounts
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.types.llms.openai import OpenAIFileObject
from litellm.types.utils import LiteLLMBatch
batch = LiteLLMBatch(
@@ -522,18 +769,42 @@ async def test_output_file_id_for_batch_retrieve():
"litellm_model_name": "gpt-5.5",
"unified_batch_id": "litellm_proxy;model_id:12345679;llm_batch_id:batch_685c5e5d63988190b85bdb2147ba131d",
}
- proxy_managed_files = PROXY_LiteLLMManagedFiles(
- DualCache(), prisma_client=AsyncMock()
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ provider_file_object = OpenAIFileObject(
+ id="file-provider-output",
+ object="file",
+ bytes=123,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
)
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
- response = await proxy_managed_files.async_post_call_success_hook(
- data={},
- user_api_key_dict=MagicMock(),
- response=batch,
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ return_value=provider_file_object,
+ ),
+ ):
+ response = await proxy_managed_files.async_post_call_success_hook(
+ data={},
+ user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
+ response=batch,
+ )
+
+ output_file_id = cast(str, cast(LiteLLMBatch, response).output_file_id)
+ assert not output_file_id.startswith("file-")
+ stored_file_object = managed_file_table.rows[output_file_id].file_object
+ assert stored_file_object is not None
+ assert stored_file_object.id == output_file_id
+ assert (
+ managed_file_table.upsert_calls[0][1]["create"]["file_object"]
+ == stored_file_object.model_dump_json()
)
- assert not cast(LiteLLMBatch, response).output_file_id.startswith("file-")
-
@pytest.mark.asyncio
async def test_output_file_id_preserves_target_model_names_when_model_name_missing():
@@ -542,6 +813,7 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi
(e.g. Vertex batch retrieve), unified output_file_id should still include
target_model_names from the managed input file ID.
"""
+ import litellm.proxy.proxy_server as proxy_server_module
from openai.types.batch import BatchRequestCounts
from litellm.proxy._types import UserAPIKeyAuth
@@ -582,10 +854,6 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi
# Intentionally omit model_name to mimic Vertex issue.
}
- proxy_managed_files = PROXY_LiteLLMManagedFiles(
- DualCache(), prisma_client=AsyncMock()
- )
-
provider_output_file = OpenAIFileObject(
id="file-provider-output-id",
object="file",
@@ -594,9 +862,18 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi
filename="predictions.jsonl",
purpose="batch_output",
)
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
- with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_retrieve:
- mock_retrieve.return_value = provider_output_file
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ return_value=provider_output_file,
+ ) as mock_retrieve,
+ ):
response = await proxy_managed_files.async_post_call_success_hook(
data={},
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
@@ -608,6 +885,15 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi
)
assert decoded_output_file_id
assert "target_model_names,gemini-2.5-pro" in cast(str, decoded_output_file_id)
+ mock_retrieve.assert_awaited_once()
+ output_file_id = cast(str, cast(LiteLLMBatch, response).output_file_id)
+ stored_file_object = managed_file_table.rows[output_file_id].file_object
+ assert stored_file_object is not None
+ assert stored_file_object.id == output_file_id
+ assert (
+ managed_file_table.upsert_calls[0][1]["create"]["file_object"]
+ == stored_file_object.model_dump_json()
+ )
@pytest.mark.asyncio
@@ -615,8 +901,7 @@ async def test_error_file_id_for_failed_batch():
"""
Test that the error_file_id is properly managed when a batch fails
"""
- from typing import cast
-
+ import litellm.proxy.proxy_server as proxy_server_module
from openai.types.batch import BatchRequestCounts
from litellm.proxy._types import UserAPIKeyAuth
@@ -658,9 +943,7 @@ async def test_error_file_id_for_failed_batch():
"unified_batch_id": "litellm_proxy;model_id:test-model-id;llm_batch_id:batch_abc123",
}
- proxy_managed_files = PROXY_LiteLLMManagedFiles(
- DualCache(), prisma_client=AsyncMock()
- )
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
# Create a proper OpenAIFileObject for the error file
error_file_object = OpenAIFileObject(
@@ -672,14 +955,20 @@ async def test_error_file_id_for_failed_batch():
purpose="batch_output",
status="processed",
)
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
# Mock the afile_retrieve to simulate retrieving error file metadata
- with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_retrieve:
- mock_retrieve.return_value = error_file_object
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ return_value=error_file_object,
+ ) as mock_retrieve,
+ ):
- user_api_key_dict = UserAPIKeyAuth(
- user_id="test-user-123", parent_otel_span=MagicMock()
- )
+ user_api_key_dict = UserAPIKeyAuth(user_id="test-user-123")
response = await proxy_managed_files.async_post_call_success_hook(
data={},
@@ -691,8 +980,15 @@ async def test_error_file_id_for_failed_batch():
assert cast(LiteLLMBatch, response).error_file_id is not None
assert not cast(LiteLLMBatch, response).error_file_id.startswith("error-")
# Verify it's a base64 encoded managed file ID
- assert is_base64_encoded_unified_file_id(
- cast(LiteLLMBatch, response).error_file_id
+ error_file_id = cast(str, cast(LiteLLMBatch, response).error_file_id)
+ assert is_base64_encoded_unified_file_id(error_file_id)
+ mock_retrieve.assert_awaited_once()
+ stored_file_object = managed_file_table.rows[error_file_id].file_object
+ assert stored_file_object is not None
+ assert stored_file_object.id == error_file_id
+ assert (
+ managed_file_table.upsert_calls[0][1]["create"]["file_object"]
+ == stored_file_object.model_dump_json()
)
@@ -1514,6 +1810,891 @@ async def test_store_unified_file_id_updates_file_metadata_on_existing_row():
assert second_update["storage_url"] == "s3://bucket/output.jsonl"
+@pytest.mark.asyncio
+async def test_store_batch_output_file_skips_existing_file_object():
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ existing_file_object = OpenAIFileObject(
+ id="unified-output",
+ object="file",
+ bytes=1,
+ created_at=1,
+ filename="output.jsonl",
+ purpose="batch_output",
+ )
+ existing_file_row = LiteLLM_ManagedFileTable(
+ unified_file_id="unified-output",
+ file_object=existing_file_object,
+ model_mappings={"model-123": "provider-output"},
+ flat_model_file_ids=["provider-output"],
+ created_by="user-123",
+ )
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ existing_file_row
+ )
+
+ with patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve:
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ retrieve.assert_not_called()
+ assert managed_file_table.find_first_calls == [
+ {"unified_file_id": "unified-output"}
+ ]
+ assert managed_file_table.upsert_calls == []
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_stores_provider_object_with_unified_id():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ provider_file_object = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=123,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ )
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ sleep_calls: list[float] = []
+
+ async def record_sleep(delay: float) -> None:
+ sleep_calls.append(delay)
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ sleep=record_sleep
+ )
+
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=[
+ HTTPException(status_code=500),
+ HTTPException(status_code=500),
+ provider_file_object,
+ ],
+ ) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ assert retrieve.await_count == 3
+ assert [call.kwargs["max_retries"] for call in retrieve.call_args_list] == [0, 0, 0]
+ assert sleep_calls == [0.5, 1.0]
+ stored_object = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert stored_object.id == "unified-output"
+ assert stored_object.bytes == 123
+ assert (
+ managed_file_table.upsert_calls[0][1]["create"]["file_object"]
+ == stored_object.model_dump_json()
+ )
+ assert managed_file_table.rows["unified-output"].model_mappings == {
+ "model-123": "provider-output"
+ }
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_retries_model_name_provider_fetch():
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ provider_file_object = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=123,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ )
+ sleep_calls: list[float] = []
+
+ async def record_sleep(delay: float) -> None:
+ sleep_calls.append(delay)
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ sleep=record_sleep
+ )
+ with patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=[
+ HTTPException(status_code=500),
+ HTTPException(status_code=500),
+ provider_file_object,
+ ],
+ ) as retrieve:
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id=None,
+ model_name="bedrock/model-x",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ assert retrieve.await_count == 3
+ assert all(
+ call.kwargs["custom_llm_provider"] == "bedrock"
+ and call.kwargs["max_retries"] == 0
+ and "_litellm_internal_model_credentials" not in call.kwargs
+ for call in retrieve.call_args_list
+ )
+ assert sleep_calls == [0.5, 1.0]
+ stored_object = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert stored_object.id == "unified-output"
+ assert stored_object.bytes == 123
+ assert stored_object.litellm_details_fallback is None
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_marks_model_name_fallback_after_four_transient_failures():
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ sleep_calls: list[float] = []
+
+ async def record_sleep(delay: float) -> None:
+ sleep_calls.append(delay)
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ sleep=record_sleep
+ )
+ with patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=[HTTPException(status_code=500) for _ in range(4)],
+ ) as retrieve:
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id=None,
+ model_name="bedrock/model-x",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ size_bytes=836,
+ )
+
+ assert retrieve.await_count == 4
+ assert all(
+ call.kwargs["custom_llm_provider"] == "bedrock"
+ and call.kwargs["max_retries"] == 0
+ and "_litellm_internal_model_credentials" not in call.kwargs
+ for call in retrieve.call_args_list
+ )
+ assert sleep_calls == [0.5, 1.0, 2.0]
+ stored_object = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert stored_object.bytes == 836
+ assert stored_object.litellm_details_fallback is True
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_falls_back_when_provider_retrieve_raises():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
+
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=RuntimeError("provider unavailable"),
+ ) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="s3://bucket/output.jsonl",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ size_bytes=321,
+ )
+
+ retrieve.assert_awaited_once()
+ stored_object = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert stored_object.id == "unified-output"
+ assert stored_object.filename == "output.jsonl"
+ assert stored_object.bytes == 321
+ assert stored_object.purpose == "batch_output"
+ assert (
+ managed_file_table.upsert_calls[0][1]["create"]["file_object"]
+ == stored_object.model_dump_json()
+ )
+ assert managed_file_table.rows["unified-output"].model_mappings == {
+ "model-123": "s3://bucket/output.jsonl"
+ }
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ ("model_id", "model_name"),
+ [("model-123", None), (None, "openai/gpt-4o-mini")],
+ ids=["deployment-credentials", "provider-from-model-name"],
+)
+async def test_store_batch_output_file_falls_back_when_provider_retrieve_hangs(
+ model_id: str | None, model_name: str | None
+):
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ async def never_answers(**_: object) -> None:
+ await asyncio.Event().wait()
+
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
+
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch("litellm.afile_retrieve", side_effect=never_answers),
+ patch("litellm_enterprise.proxy.hooks.managed_files.BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS", 0.01),
+ ):
+ await asyncio.wait_for(
+ proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="s3://bucket/output.jsonl",
+ model_id=model_id,
+ model_name=model_name,
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ size_bytes=321,
+ ),
+ timeout=5,
+ )
+
+ stored_object = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert (stored_object.id, stored_object.bytes, stored_object.purpose) == ("unified-output", 321, "batch_output")
+ assert stored_object.litellm_details_fallback is True
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_returns_marked_fallback_when_refresh_hangs():
+ async def never_answers(**_: object) -> None:
+ await asyncio.Event().wait()
+
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(_marked_fallback_file_row())
+
+ with (
+ patch("litellm.afile_retrieve", side_effect=never_answers),
+ patch("litellm_enterprise.proxy.hooks.managed_files.BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS", 0.01),
+ ):
+ response = await asyncio.wait_for(
+ proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ ),
+ timeout=5,
+ )
+
+ assert (response.id, response.bytes) == ("unified-output", 0)
+ assert managed_file_table.upsert_calls == []
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_falls_back_without_model_id():
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
+
+ with patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve:
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-error",
+ provider_file_id="error.jsonl",
+ model_id=None,
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ retrieve.assert_not_called()
+ stored_object = managed_file_table.rows["unified-error"].file_object
+ assert stored_object is not None
+ assert stored_object.id == "unified-error"
+ assert stored_object.filename == "error.jsonl"
+ assert stored_object.bytes == 0
+ assert stored_object.purpose == "batch_output"
+ assert (
+ managed_file_table.upsert_calls[0][1]["create"]["file_object"]
+ == stored_object.model_dump_json()
+ )
+ assert managed_file_table.rows["unified-error"].model_mappings == {}
+ assert stored_object.litellm_details_fallback is None
+ assert "litellm_details_fallback" not in stored_object.model_dump()
+
+
+@pytest.mark.asyncio
+async def test_listed_batch_saves_basic_output_without_provider_fetch():
+ import litellm.proxy.proxy_server as proxy_server_module
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
+ response: Final = _batch_response_for_output_listing()
+ with (
+ patch.object(proxy_server_module, "llm_router", None),
+ patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve,
+ ):
+ resolved: Final = await _resolve_batch_for_output_listing(
+ proxy_managed_files, response
+ )
+
+ assert resolved is response
+ retrieve.assert_not_called()
+ stored_file: Final = next(iter(managed_file_table.rows.values()))
+ assert stored_file.file_object is not None
+ assert stored_file.file_object.bytes == 0
+ assert stored_file.file_object.litellm_details_fallback is True
+ assert stored_file.model_mappings == {"model-123": "provider-output"}
+
+
+@pytest.mark.asyncio
+async def test_listed_batch_leaves_existing_marked_output_untouched():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ proxy_managed_files, _ = _managed_files_with_fake_prisma()
+ unified_file_id: Final = proxy_managed_files.get_unified_output_file_id(
+ output_file_id="provider-output",
+ model_id="model-123",
+ model_name="bedrock/model-x",
+ )
+ row: Final = _marked_fallback_file_row().model_copy(
+ update={"unified_file_id": unified_file_id}
+ )
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row)
+ response: Final = _batch_response_for_output_listing()
+ with (
+ patch.object(proxy_server_module, "llm_router", None),
+ patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve,
+ ):
+ resolved: Final = await _resolve_batch_for_output_listing(
+ proxy_managed_files, response
+ )
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id=unified_file_id,
+ provider_file_id="provider-output",
+ model_id="model-123",
+ model_name="bedrock/model-x",
+ owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"),
+ litellm_parent_otel_span=None,
+ fetch_provider_details=False,
+ )
+
+ assert resolved is response
+ retrieve.assert_not_called()
+ assert managed_file_table.rows[unified_file_id] is row
+ assert managed_file_table.upsert_calls == []
+
+
+@pytest.mark.asyncio
+async def test_batch_retrieve_registration_fetches_details_by_default():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.proxy.openai_files_endpoints.common_utils import (
+ ensure_batch_response_managed_file_ids,
+ )
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
+ response: Final = _batch_response_for_output_listing()
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {
+ "api_key": "key",
+ "custom_llm_provider": "bedrock",
+ }
+ provider_file_object: Final = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ return_value=provider_file_object,
+ ) as retrieve,
+ ):
+ await ensure_batch_response_managed_file_ids(
+ response=response,
+ managed_files_obj=proxy_managed_files,
+ prisma_client=proxy_managed_files.prisma_client,
+ verbose_proxy_logger=MagicMock(),
+ user_api_key_dict=UserAPIKeyAuth(user_id="user-123"),
+ db_batch_object=SimpleNamespace(created_by="user-123", team_id=None),
+ )
+
+ retrieve.assert_awaited_once()
+ output_file_id: Final = response.output_file_id
+ assert output_file_id is not None
+ stored_file: Final = managed_file_table.rows[output_file_id]
+ assert stored_file.file_object is not None
+ assert stored_file.file_object.bytes == 836
+
+
+@pytest.mark.parametrize(
+ ("error_factory", "expected_attempts"),
+ [
+ pytest.param(lambda: HTTPException(status_code=429), 2, id="429"),
+ pytest.param(lambda: HTTPException(status_code=408), 2, id="408"),
+ pytest.param(lambda: HTTPException(status_code=503), 2, id="503"),
+ pytest.param(
+ lambda: httpx.ConnectError(
+ "connection failed",
+ request=httpx.Request("GET", "https://api.openai.com/v1/files/file-1"),
+ ),
+ 2,
+ id="connect-error",
+ ),
+ pytest.param(
+ lambda: APIConnectionError(
+ message="connection failed",
+ request=httpx.Request("GET", "https://api.openai.com/v1/files/file-1"),
+ ),
+ 2,
+ id="openai-api-connection-error",
+ ),
+ pytest.param(lambda: asyncio.TimeoutError(), 2, id="async-timeout"),
+ pytest.param(lambda: HTTPException(status_code=400), 1, id="400"),
+ pytest.param(lambda: HTTPException(status_code=403), 1, id="403"),
+ pytest.param(lambda: HTTPException(status_code=404), 1, id="404"),
+ pytest.param(lambda: ValueError("invalid file"), 1, id="value-error"),
+ ],
+)
+@pytest.mark.asyncio
+async def test_batch_file_detail_retry_classification(
+ error_factory: Callable[[], Exception],
+ expected_attempts: int,
+) -> None:
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ sleep_calls: Final[list[float]] = []
+
+ async def record_sleep(delay: float) -> None:
+ sleep_calls.append(delay)
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ sleep=record_sleep
+ )
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ provider_file_object: Final = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+ retrieve_side_effect: Final = (
+ [error_factory(), provider_file_object]
+ if expected_attempts == 2
+ else error_factory()
+ )
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=retrieve_side_effect,
+ ) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ assert retrieve.await_count == expected_attempts
+ assert [call.kwargs["max_retries"] for call in retrieve.call_args_list] == [
+ 0
+ ] * expected_attempts
+ assert sleep_calls == ([0.5] if expected_attempts == 2 else [])
+ stored_object: Final = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert stored_object.bytes == (836 if expected_attempts == 2 else 0)
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_ignores_unmarked_row_with_provider_route():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ existing_file_object: Final = OpenAIFileObject(
+ id="unified-output",
+ object="file",
+ bytes=1,
+ created_at=1,
+ filename="output.jsonl",
+ purpose="batch_output",
+ )
+ existing_file_row: Final = LiteLLM_ManagedFileTable(
+ unified_file_id="unified-output",
+ file_object=existing_file_object,
+ model_mappings={"model-123": "provider-output"},
+ flat_model_file_ids=["provider-output"],
+ created_by="user-123",
+ )
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ existing_file_row
+ )
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ retrieve.assert_not_called()
+ assert managed_file_table.rows["unified-output"].file_object is existing_file_object
+ assert managed_file_table.upsert_calls == []
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_rejects_non_provider_model_prefix():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma()
+ with (
+ patch.object(proxy_server_module, "llm_router", None),
+ patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ model_name="my-team/gpt-4o",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ retrieve.assert_not_called()
+ stored_object: Final = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert stored_object.litellm_details_fallback is None
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_marks_fallback_after_four_transient_failures():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ sleep_calls: list[float] = []
+
+ async def record_sleep(delay: float) -> None:
+ sleep_calls.append(delay)
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ sleep=record_sleep
+ )
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=[HTTPException(status_code=500) for _ in range(4)],
+ ) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"),
+ litellm_parent_otel_span=None,
+ size_bytes=836,
+ )
+
+ assert retrieve.await_count == 4
+ assert [call.kwargs["max_retries"] for call in retrieve.call_args_list] == [0, 0, 0, 0]
+ assert sleep_calls == [0.5, 1.0, 2.0]
+ stored_object = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert stored_object.litellm_details_fallback is True
+ stored_file_object_json: Final = managed_file_table.upsert_calls[0][1]["create"]["file_object"]
+ assert isinstance(stored_file_object_json, str)
+ assert json.loads(stored_file_object_json)["litellm_details_fallback"] is True
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_skips_lookup_for_fallback_written_moments_ago():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ _marked_fallback_file_row(created_at=int(time.time()))
+ )
+
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"),
+ litellm_parent_otel_span=None,
+ size_bytes=836,
+ )
+
+ retrieve.assert_not_awaited()
+ stored_object = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert (stored_object.bytes, stored_object.litellm_details_fallback) == (836, True)
+ assert managed_file_table.upsert_calls == []
+ assert managed_file_table.update_many_calls[0][0] == {
+ "unified_file_id": "unified-output"
+ }
+ assert set(managed_file_table.update_many_calls[0][1]) == {"file_object"}
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_does_not_retry_not_found():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ sleep_calls: list[float] = []
+
+ async def record_sleep(delay: float) -> None:
+ sleep_calls.append(delay)
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ sleep=record_sleep
+ )
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=HTTPException(status_code=404),
+ ) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ retrieve.assert_awaited_once()
+ assert sleep_calls == []
+ stored_object = managed_file_table.rows["unified-output"].file_object
+ assert stored_object is not None
+ assert stored_object.litellm_details_fallback is True
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_refreshes_marked_fallback_on_later_write():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ row = _marked_fallback_file_row()
+ provider_object = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row)
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch("litellm.afile_retrieve", new_callable=AsyncMock, return_value=provider_object) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ retrieve.assert_awaited_once()
+ refreshed = managed_file_table.rows["unified-output"].file_object
+ assert refreshed is not None
+ assert refreshed.id == "unified-output"
+ assert refreshed.bytes == 836
+ assert refreshed.litellm_details_fallback is None
+ assert managed_file_table.update_many_calls == [
+ ({"unified_file_id": "unified-output"}, {"file_object": refreshed.model_dump_json()})
+ ]
+ assert managed_file_table.upsert_calls == []
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_does_not_recreate_deleted_row_after_refresh():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ provider_object: Final = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ _marked_fallback_file_row()
+ )
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+
+ async def delete_then_return_file(**_kwargs: object) -> OpenAIFileObject:
+ await managed_file_table.delete(where={"unified_file_id": "unified-output"})
+ return provider_object
+
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=delete_then_return_file,
+ ),
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ assert "unified-output" not in managed_file_table.rows
+ assert managed_file_table.upsert_calls == []
+ assert len(managed_file_table.update_many_calls) == 1
+ saved_file_object_json: Final = cast(
+ str, managed_file_table.update_many_calls[0][1]["file_object"]
+ )
+ saved_file_object: Final = json.loads(saved_file_object_json)
+ assert saved_file_object["id"] == "unified-output"
+ assert (
+ await proxy_managed_files.internal_usage_cache.async_get_cache(
+ key="unified-output",
+ litellm_parent_otel_span=None,
+ )
+ is None
+ )
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_keeps_marked_fallback_when_retryable_details_still_fail():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ _marked_fallback_file_row()
+ )
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=HTTPException(status_code=404),
+ ) as retrieve,
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"),
+ litellm_parent_otel_span=None,
+ )
+
+ retrieve.assert_awaited_once()
+ assert managed_file_table.upsert_calls == []
+ assert managed_file_table.rows["unified-output"].file_object.created_at == 123
+
+
+@pytest.mark.asyncio
+async def test_store_batch_output_file_updates_only_size_for_failed_marked_fallback():
+ import litellm.proxy.proxy_server as proxy_server_module
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ _marked_fallback_file_row()
+ )
+ with (
+ patch.object(proxy_server_module, "llm_router", router),
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=HTTPException(status_code=404),
+ ),
+ ):
+ await proxy_managed_files.store_batch_output_file(
+ unified_file_id="unified-output",
+ provider_file_id="provider-output",
+ model_id="model-123",
+ owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"),
+ litellm_parent_otel_span=None,
+ size_bytes=836,
+ )
+
+ updated = managed_file_table.rows["unified-output"].file_object
+ assert updated is not None
+ assert updated.bytes == 836
+ assert updated.created_at == 123
+ assert updated.litellm_details_fallback is True
+
+
@pytest.mark.asyncio
async def test_afile_delete_returns_provider_response_when_stored_file_object_none():
"""
@@ -1585,19 +2766,21 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none():
"""
from litellm.types.llms.openai import OpenAIFileObject
- prisma_client = AsyncMock()
- internal_usage_cache = MagicMock()
-
- proxy_managed_files = PROXY_LiteLLMManagedFiles(
- internal_usage_cache=internal_usage_cache,
- prisma_client=prisma_client,
+ stored_file = LiteLLM_ManagedFileTable(
+ unified_file_id="test-unified-file-id",
+ file_object=None,
+ model_mappings={"model-123": "file-provider-xyz"},
+ flat_model_file_ids=["file-provider-xyz"],
+ created_by="user-123",
)
+ sleep_calls: list[float] = []
- # Mock get_unified_file_id to return a stored object with file_object=None
- stored_file = MagicMock()
- stored_file.file_object = None
- stored_file.model_mappings = {"model-123": "file-provider-xyz"}
- proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file)
+ async def record_sleep(delay: float) -> None:
+ sleep_calls.append(delay)
+
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ stored_file, sleep=record_sleep
+ )
# Mock the router and provider response
provider_file_response = OpenAIFileObject(
@@ -1617,7 +2800,11 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none():
}
)
- with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_afile_retrieve:
+ with patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=[HTTPException(status_code=500), provider_file_response],
+ ) as mock_afile_retrieve:
mock_afile_retrieve.return_value = provider_file_response
unified_file_id = "test-unified-file-id"
@@ -1630,7 +2817,10 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none():
# Should return the provider response with the unified file ID
assert result is not None
assert result.id == unified_file_id
- mock_afile_retrieve.assert_called_once()
+ assert mock_afile_retrieve.await_count == 2
+ assert [call.kwargs["max_retries"] for call in mock_afile_retrieve.call_args_list] == [0, 0]
+ assert sleep_calls == [0.5]
+ assert managed_file_table.upsert_calls == []
@pytest.mark.asyncio
@@ -1639,19 +2829,14 @@ async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none()
Test that afile_retrieve raises an appropriate error when file_object is None
and no llm_router is provided to fetch from the provider.
"""
- prisma_client = AsyncMock()
- internal_usage_cache = MagicMock()
-
- proxy_managed_files = PROXY_LiteLLMManagedFiles(
- internal_usage_cache=internal_usage_cache,
- prisma_client=prisma_client,
+ stored_file = LiteLLM_ManagedFileTable(
+ unified_file_id="test-unified-file-id",
+ file_object=None,
+ model_mappings={"model-123": "file-provider-xyz"},
+ flat_model_file_ids=["file-provider-xyz"],
+ created_by="user-123",
)
-
- # Mock get_unified_file_id to return a stored object with file_object=None
- stored_file = MagicMock()
- stored_file.file_object = None
- stored_file.model_mappings = {"model-123": "file-provider-xyz"}
- proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file)
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file)
unified_file_id = "test-unified-file-id"
@@ -1662,7 +2847,7 @@ async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none()
llm_router=None,
)
- assert "llm_router is required" in str(exc_info.value)
+ assert "no provider route to fetch it" in str(exc_info.value)
@pytest.mark.asyncio
@@ -1673,15 +2858,6 @@ async def test_afile_retrieve_returns_stored_file_object_when_exists():
"""
from litellm.types.llms.openai import OpenAIFileObject
- prisma_client = AsyncMock()
- internal_usage_cache = MagicMock()
-
- proxy_managed_files = PROXY_LiteLLMManagedFiles(
- internal_usage_cache=internal_usage_cache,
- prisma_client=prisma_client,
- )
-
- # Mock get_unified_file_id to return a stored object WITH file_object
stored_file_object = OpenAIFileObject(
id="test-unified-file-id",
object="file",
@@ -1690,18 +2866,371 @@ async def test_afile_retrieve_returns_stored_file_object_when_exists():
filename="input.jsonl",
purpose="batch",
)
- stored_file = MagicMock()
- stored_file.file_object = stored_file_object
- proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file)
+ stored_file = LiteLLM_ManagedFileTable(
+ unified_file_id="test-unified-file-id",
+ file_object=stored_file_object,
+ model_mappings={"model-123": "provider-file-id"},
+ flat_model_file_ids=["provider-file-id"],
+ created_by="user-123",
+ )
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file)
+ router = MagicMock()
- result = await proxy_managed_files.afile_retrieve(
- file_id="test-unified-file-id",
+ with patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve:
+ result = await proxy_managed_files.afile_retrieve(
+ file_id="test-unified-file-id",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ )
+
+ assert result == stored_file_object
+ retrieve.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_refreshes_marked_fallback_and_preserves_ownership():
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ row = _marked_fallback_file_row()
+ provider_object = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row)
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+
+ with patch("litellm.afile_retrieve", new_callable=AsyncMock, return_value=provider_object) as retrieve:
+ response = await proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ )
+
+ assert response == provider_object.model_copy(update={"id": "unified-output"})
+ assert "litellm_details_fallback" not in response.model_dump()
+ retrieve.assert_awaited_once()
+ updated_row = managed_file_table.rows["unified-output"]
+ assert updated_row.file_object == response
+ assert updated_row.created_by == "user-123"
+ assert updated_row.team_id == "team-123"
+ assert managed_file_table.update_many_calls[0][1] == {
+ "file_object": response.model_dump_json()
+ }
+ assert managed_file_table.upsert_calls == []
+ cached_row = await proxy_managed_files.internal_usage_cache.async_get_cache(
+ key="unified-output",
litellm_parent_otel_span=None,
- llm_router=None,
+ )
+ assert cached_row["file_object"]["bytes"] == 836
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_refreshes_marked_fallback_without_router_from_model_name():
+ row: Final = _marked_fallback_file_row().model_copy(
+ update={"model_mappings": {"bedrock/model-x": "provider-output"}}
+ )
+ provider_object: Final = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row)
+
+ with patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ return_value=provider_object,
+ ) as retrieve:
+ response: Final = await proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=None,
+ )
+
+ assert response.bytes == 836
+ assert response.id == "unified-output"
+ assert response.litellm_details_fallback is None
+ retrieve.assert_awaited_once()
+ assert retrieve.await_args.kwargs["custom_llm_provider"] == "bedrock"
+ assert retrieve.await_args.kwargs["max_retries"] == 0
+ assert "_litellm_internal_model_credentials" not in retrieve.await_args.kwargs
+ updated_row: Final = managed_file_table.rows["unified-output"]
+ assert updated_row.file_object == response
+ assert updated_row.model_mappings == {"bedrock/model-x": "provider-output"}
+ assert updated_row.created_by == "user-123"
+ assert updated_row.team_id == "team-123"
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_uses_model_name_when_router_does_not_know_deployment():
+ row: Final = _marked_fallback_file_row().model_copy(
+ update={"model_mappings": {"bedrock/model-x": "provider-output"}}
+ )
+ provider_object: Final = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = None
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(row)
+
+ with patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ return_value=provider_object,
+ ) as retrieve:
+ response: Final = await proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ )
+
+ assert response.bytes == 836
+ retrieve.assert_awaited_once()
+ router.get_deployment_credentials_with_provider.assert_called_once_with(
+ "bedrock/model-x"
+ )
+ assert retrieve.await_args.kwargs["custom_llm_provider"] == "bedrock"
+ assert retrieve.await_args.kwargs["max_retries"] == 0
+ assert "_litellm_internal_model_credentials" not in retrieve.await_args.kwargs
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_does_not_recreate_deleted_row_after_refresh():
+ row: Final = _marked_fallback_file_row()
+ provider_object: Final = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row)
+
+ async def delete_row_then_return_file(
+ *, file_id: str, **_options: object
+ ) -> OpenAIFileObject:
+ assert file_id == "provider-output"
+ deleted_row: Final = await managed_file_table.delete(
+ where={"unified_file_id": "unified-output"}
+ )
+ assert deleted_row is not None
+ return provider_object
+
+ with patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=delete_row_then_return_file,
+ ):
+ response: Final = await proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ )
+
+ assert response.bytes == 836
+ assert "unified-output" not in managed_file_table.rows
+ assert managed_file_table.upsert_calls == []
+ cached_row: Final = await proxy_managed_files.internal_usage_cache.async_get_cache(
+ key="unified-output",
+ litellm_parent_otel_span=None,
+ )
+ assert cached_row is None
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_case3_includes_provider_error_text():
+ stored_file: Final = LiteLLM_ManagedFileTable(
+ unified_file_id="unified-output",
+ file_object=None,
+ model_mappings={"model-123": "provider-output"},
+ flat_model_file_ids=["provider-output"],
+ created_by="user-123",
+ )
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = None
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file)
+
+ with (
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=HTTPException(
+ status_code=404,
+ detail="provider file is missing",
+ ),
+ ),
+ pytest.raises(
+ Exception,
+ match="Failed to retrieve file unified-output from provider",
+ ) as error,
+ ):
+ await proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ )
+
+ assert type(error.value) is Exception
+ assert str(error.value) == (
+ "Failed to retrieve file unified-output from provider: "
+ "404: provider file is missing"
)
- # Should return the stored file object directly
- assert result == stored_file_object
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_case3_uses_default_provider_for_unknown_router_deployment():
+ stored_file: Final = LiteLLM_ManagedFileTable(
+ unified_file_id="unified-output",
+ file_object=None,
+ model_mappings={"model-123": "provider-output"},
+ flat_model_file_ids=["provider-output"],
+ created_by="user-123",
+ )
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = None
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file)
+ provider_file_object: Final = OpenAIFileObject(
+ id="provider-output",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+
+ with patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ return_value=provider_file_object,
+ ) as retrieve:
+ response: Final = await proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ )
+
+ assert response.bytes == 836
+ retrieve.assert_awaited_once()
+ assert retrieve.await_args.kwargs["file_id"] == "provider-output"
+ assert retrieve.await_args.kwargs["max_retries"] == 0
+ assert "custom_llm_provider" not in retrieve.await_args.kwargs
+ assert "_litellm_internal_model_credentials" not in retrieve.await_args.kwargs
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_case3_times_out_on_default_provider_route():
+ stored_file: Final = LiteLLM_ManagedFileTable(
+ unified_file_id="unified-output",
+ file_object=None,
+ model_mappings={"model-123": "provider-output"},
+ flat_model_file_ids=["provider-output"],
+ created_by="user-123",
+ )
+ router: Final = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = None
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file)
+
+ async def never_answers(**_options: object) -> None:
+ await asyncio.Event().wait()
+
+ with (
+ patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=never_answers,
+ ) as retrieve,
+ patch(
+ "litellm_enterprise.proxy.hooks.managed_files.BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS",
+ 0.01,
+ ),
+ pytest.raises(Exception, match="Provider file retrieve timed out"),
+ ):
+ await asyncio.wait_for(
+ proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ ),
+ timeout=1,
+ )
+
+ retrieve.assert_awaited_once()
+ assert retrieve.await_args.kwargs["file_id"] == "provider-output"
+ assert retrieve.await_args.kwargs["max_retries"] == 0
+ assert "custom_llm_provider" not in retrieve.await_args.kwargs
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_case3_without_route_raises_accurate_error():
+ import litellm.proxy.proxy_server as proxy_server_module
+
+ stored_file: Final = LiteLLM_ManagedFileTable(
+ unified_file_id="unified-output",
+ file_object=None,
+ model_mappings={"my-team/gpt-4o": "provider-output"},
+ flat_model_file_ids=["provider-output"],
+ created_by="user-123",
+ )
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file)
+ with (
+ patch.object(proxy_server_module, "llm_router", None),
+ patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve,
+ pytest.raises(Exception, match="no provider route to fetch it"),
+ ):
+ await proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=None,
+ )
+
+ retrieve.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_afile_retrieve_returns_marked_fallback_without_marker_when_refresh_fails():
+ router = MagicMock()
+ router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"}
+ proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(
+ _marked_fallback_file_row()
+ )
+
+ with patch(
+ "litellm.afile_retrieve",
+ new_callable=AsyncMock,
+ side_effect=HTTPException(status_code=404),
+ ) as retrieve:
+ response = await proxy_managed_files.afile_retrieve(
+ file_id="unified-output",
+ litellm_parent_otel_span=None,
+ llm_router=router,
+ )
+
+ assert response.bytes == 0
+ assert response.id == "unified-output"
+ assert "litellm_details_fallback" not in response.model_dump()
+ retrieve.assert_awaited_once()
+ assert managed_file_table.upsert_calls == []
@pytest.mark.asyncio
@@ -1710,16 +3239,7 @@ async def test_afile_retrieve_raises_error_for_non_managed_file():
Test that afile_retrieve raises an error when the file_id is not found
in the managed files table.
"""
- prisma_client = AsyncMock()
- internal_usage_cache = MagicMock()
-
- proxy_managed_files = PROXY_LiteLLMManagedFiles(
- internal_usage_cache=internal_usage_cache,
- prisma_client=prisma_client,
- )
-
- # Mock get_unified_file_id to return None (file not found)
- proxy_managed_files.get_unified_file_id = AsyncMock(return_value=None)
+ proxy_managed_files, _ = _managed_files_with_fake_prisma()
with pytest.raises(Exception, match='LiteLLM Managed File object with id=non-existent-file-id') as exc_info:
await proxy_managed_files.afile_retrieve(
@@ -1936,6 +3456,47 @@ async def test_list_batches_registers_and_returns_unified_output_file_ids():
assert c.kwargs["data"]["create"]["team_id"] == "owner-team"
+@pytest.mark.asyncio
+async def test_afile_list_returns_persisted_batch_output_file():
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(
+ _marked_fallback_file_row()
+ )
+
+ result = await proxy_managed_files.afile_list(
+ purpose="batch_output",
+ user_api_key_dict=UserAPIKeyAuth(user_id="owner-user"),
+ litellm_parent_otel_span=None,
+ limit=10,
+ )
+
+ listed = result.data[0]
+ assert listed.id == "unified-output"
+ assert listed.filename == "output.jsonl"
+ assert listed.bytes == 0
+ assert listed.purpose == "batch_output"
+ assert "litellm_details_fallback" not in listed.model_dump()
+
+
+@pytest.mark.asyncio
+async def test_get_user_created_file_ids_strips_fallback_marker():
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ proxy_managed_files, _ = _managed_files_with_fake_prisma(
+ _marked_fallback_file_row()
+ )
+
+ files = await proxy_managed_files.get_user_created_file_ids(
+ UserAPIKeyAuth(user_id="user-123"),
+ ["provider-output"],
+ )
+
+ assert len(files) == 1
+ assert files[0].id == "unified-output"
+ assert "litellm_details_fallback" not in files[0].model_dump()
+
+
@pytest.mark.asyncio
async def test_list_batches_resolves_existing_managed_rows_without_minting():
"""When the raw provider file IDs already have managed file rows, listing must
@@ -3453,8 +5014,10 @@ async def test_post_call_batch_sync_does_not_claim_ownership():
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 0
proxy_managed_files = PROXY_LiteLLMManagedFiles(
- MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
+ MagicMock(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()),
+ prisma_client=prisma_client,
)
+ prisma_client.db.litellm_managedfiletable.find_first.return_value = None
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
@@ -3477,22 +5040,16 @@ async def test_post_call_batch_sync_updates_existing_row():
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
- proxy_managed_files = PROXY_LiteLLMManagedFiles(
- MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
- )
+ proxy_managed_files = PROXY_LiteLLMManagedFiles(MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client)
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
- user_api_key_dict=UserAPIKeyAuth(
- user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
- ),
+ user_api_key_dict=UserAPIKeyAuth(user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()),
response=_batch_response(MODEL_ENCODED_BATCH_ID),
)
update_call = prisma_client.db.litellm_managedobjecttable.update_many.await_args
- assert update_call.kwargs["where"] == {
- "unified_object_id": MODEL_ENCODED_BATCH_ID
- }
+ assert update_call.kwargs["where"] == {"unified_object_id": MODEL_ENCODED_BATCH_ID}
assert update_call.kwargs["data"]["status"] == "completed"
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
@@ -3512,8 +5069,10 @@ async def test_post_call_batch_sync_stores_output_file_ownership_from_batch_row(
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = PROXY_LiteLLMManagedFiles(
- MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
+ MagicMock(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()),
+ prisma_client=prisma_client,
)
+ prisma_client.db.litellm_managedfiletable.find_first.return_value = None
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
diff --git a/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py b/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
index c48af1f6177..c6ab83a0d80 100644
--- a/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
+++ b/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
@@ -9,12 +9,12 @@ import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.openai_files_endpoints.common_utils import (
+ ManagedBatchOutputFileWriter,
ensure_batch_response_managed_file_ids,
get_batch_from_database,
)
from litellm.types.utils import LiteLLMBatch
-
UNIFIED_BATCH_ID = "litellm_proxy;model_id:my-model;llm_batch_id:batch-raw-123"
ENCODED_UNIFIED_BATCH_ID = (
base64.urlsafe_b64encode(UNIFIED_BATCH_ID.encode()).decode().rstrip("=")
@@ -40,9 +40,9 @@ def _build_batch_response(
def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ="):
- mock = MagicMock()
+ mock = MagicMock(spec=ManagedBatchOutputFileWriter)
mock.get_unified_output_file_id = MagicMock(return_value=unified_id)
- mock.store_unified_file_id = AsyncMock()
+ mock.store_batch_output_file = AsyncMock()
return mock
@@ -73,11 +73,12 @@ async def test_ensure_batch_response_derives_model_id_from_unified_batch_id():
)
assert response.output_file_id == unified_output_file_id
- mock_managed_files.store_unified_file_id.assert_called_once()
- store_kwargs = mock_managed_files.store_unified_file_id.call_args.kwargs
- assert store_kwargs["model_mappings"] == {"my-model": raw_output_file_id}
- assert store_kwargs["user_api_key_dict"].user_id == "batch-owner"
- assert store_kwargs["user_api_key_dict"].team_id == "team-owner"
+ mock_managed_files.store_batch_output_file.assert_awaited_once()
+ store_kwargs = mock_managed_files.store_batch_output_file.await_args.kwargs
+ assert store_kwargs["model_id"] == "my-model"
+ assert store_kwargs["provider_file_id"] == raw_output_file_id
+ assert store_kwargs["owner"].user_id == "batch-owner"
+ assert store_kwargs["owner"].team_id == "team-owner"
@pytest.mark.asyncio
@@ -100,13 +101,11 @@ async def test_ensure_batch_response_registers_output_and_error_file_ids():
assert response.output_file_id == unified_id
assert response.error_file_id == unified_id
- assert mock_managed_files.store_unified_file_id.call_count == 2
- mappings = [
- call.kwargs["model_mappings"]
- for call in mock_managed_files.store_unified_file_id.call_args_list
+ assert mock_managed_files.store_batch_output_file.await_count == 2
+ stored_provider_file_ids = [
+ call.kwargs["provider_file_id"] for call in mock_managed_files.store_batch_output_file.await_args_list
]
- assert {"my-model": "file-raw-output"} in mappings
- assert {"my-model": "file-raw-error"} in mappings
+ assert stored_provider_file_ids == ["file-raw-output", "file-raw-error"]
@pytest.mark.asyncio
@@ -148,10 +147,11 @@ async def test_get_batch_from_database_registers_missing_output_file_id():
assert response is not None
assert response.output_file_id == unified_output_file_id
- mock_managed_files.store_unified_file_id.assert_called_once()
- store_kwargs = mock_managed_files.store_unified_file_id.call_args.kwargs
- assert store_kwargs["model_mappings"] == {"my-model": raw_output_file_id}
- assert store_kwargs["user_api_key_dict"].user_id == "batch-owner"
+ mock_managed_files.store_batch_output_file.assert_awaited_once()
+ store_kwargs = mock_managed_files.store_batch_output_file.await_args.kwargs
+ assert store_kwargs["model_id"] == "my-model"
+ assert store_kwargs["provider_file_id"] == raw_output_file_id
+ assert store_kwargs["owner"].user_id == "batch-owner"
@pytest.mark.asyncio
@@ -172,9 +172,7 @@ async def test_ensure_batch_response_uses_batch_owner_when_db_batch_object_prese
)
# batch owner from db_batch_object wins over the caller auth context
- forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[
- "user_api_key_dict"
- ]
+ forwarded_auth = mock_managed_files.store_batch_output_file.await_args.kwargs["owner"]
assert forwarded_auth.user_id == "batch-owner"
assert forwarded_auth.team_id == "team-owner"
@@ -190,7 +188,10 @@ async def test_registered_output_file_row_denies_cross_user_access():
prisma.db.litellm_managedfiletable.upsert = AsyncMock()
prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
managed_files = PROXY_LiteLLMManagedFiles(
- internal_usage_cache=MagicMock(),
+ internal_usage_cache=MagicMock(
+ async_get_cache=AsyncMock(return_value=None),
+ async_set_cache=AsyncMock(),
+ ),
prisma_client=prisma,
)
response = _build_batch_response(output_file_id=raw_output_file_id)
diff --git a/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
index 2049ff58f95..cf305c660cd 100644
--- a/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
+++ b/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
@@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.openai_files_endpoints.common_utils import (
+ ManagedBatchOutputFileWriter,
ensure_batch_response_managed_file_ids,
update_batch_in_database,
)
@@ -39,9 +40,9 @@ def _build_batch_response(
def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ="):
- mock = MagicMock()
+ mock = MagicMock(spec=ManagedBatchOutputFileWriter)
mock.get_unified_output_file_id = MagicMock(return_value=unified_id)
- mock.store_unified_file_id = AsyncMock()
+ mock.store_batch_output_file = AsyncMock()
return mock
@@ -114,9 +115,7 @@ async def test_cancel_path_registers_output_file_under_batch_owner():
operation="cancel",
)
- forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[
- "user_api_key_dict"
- ]
+ forwarded_auth = mock_managed_files.store_batch_output_file.await_args.kwargs["owner"]
assert forwarded_auth.user_id == "batch-owner"
assert forwarded_auth.team_id == "batch-team"
stored = json.loads(
@@ -156,9 +155,7 @@ async def test_update_batch_skips_lookup_when_db_batch_object_supplied():
)
mock_prisma.db.litellm_managedobjecttable.find_first.assert_not_called()
- forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[
- "user_api_key_dict"
- ]
+ forwarded_auth = mock_managed_files.store_batch_output_file.await_args.kwargs["owner"]
assert forwarded_auth.user_id == "caller-owner"
assert forwarded_auth.team_id == "caller-team"
@@ -222,11 +219,11 @@ async def test_ensure_batch_response_swallows_conversion_errors():
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
)
- mock_managed_files = MagicMock()
+ mock_managed_files = MagicMock(spec=ManagedBatchOutputFileWriter)
mock_managed_files.get_unified_output_file_id = MagicMock(
side_effect=RuntimeError("boom")
)
- mock_managed_files.store_unified_file_id = AsyncMock()
+ mock_managed_files.store_batch_output_file = AsyncMock()
mock_logger = MagicMock()
await ensure_batch_response_managed_file_ids(
@@ -263,9 +260,7 @@ async def test_ensure_batch_response_builds_auth_from_db_batch_object():
db_batch_object=db_batch_object,
)
- forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[
- "user_api_key_dict"
- ]
+ forwarded_auth = mock_managed_files.store_batch_output_file.await_args.kwargs["owner"]
assert forwarded_auth.user_id == "user-from-db"
assert forwarded_auth.team_id == "team-from-db"
diff --git a/tests/unit/enterprise/proxy/test_managed_files_hook.py b/tests/unit/enterprise/proxy/test_managed_files_hook.py
index 3ea3b97e1fe..7bf53b11cdb 100644
--- a/tests/unit/enterprise/proxy/test_managed_files_hook.py
+++ b/tests/unit/enterprise/proxy/test_managed_files_hook.py
@@ -183,7 +183,9 @@ def _make_managed_files_instance():
)
mock_cache = MagicMock()
+ mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_prisma = MagicMock()
+ mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
instance = PROXY_LiteLLMManagedFiles(
internal_usage_cache=mock_cache,
@@ -1384,9 +1386,11 @@ def _make_real_managed_files_instance():
)
mock_cache = MagicMock()
+ mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.async_set_cache = AsyncMock()
mock_prisma = MagicMock()
+ mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
mock_prisma.db.litellm_managedfiletable.upsert = AsyncMock()
mock_prisma.db.litellm_managedfiletable.create = AsyncMock(
side_effect=AssertionError(
diff --git a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py
index 4bfac61545d..6685deccf65 100644
--- a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py
+++ b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py
@@ -2216,6 +2216,216 @@ class TestBedrockFileContentTransformation:
assert "x-amz-content-sha256" in authorization
assert "X-Amz-Date" in signed_headers
+ def test_transform_retrieve_file_request_adds_unsigned_range(self, monkeypatch):
+ from litellm.llms.bedrock.files.transformation import (
+ S3_SIGNED_REQUEST_HEADERS_PARAM,
+ BedrockFilesConfig,
+ )
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ litellm_params = self._litellm_params()
+ url, params = BedrockFilesConfig().transform_retrieve_file_request(
+ file_id=self.S3_URI,
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+
+ assert url == self.EXPECTED_URL
+ assert params == {}
+ assert litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]["Range"] == "bytes=0-0"
+
+ @pytest.mark.parametrize(
+ ("file_id", "purpose"),
+ [
+ ("s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl", "batch_output"),
+ ("s3://my-bucket/litellm-bedrock-files-job-123/input.jsonl", "batch"),
+ ],
+ )
+ def test_transform_retrieve_file_response_parses_metadata(
+ self, file_id: str, purpose: str, monkeypatch: pytest.MonkeyPatch
+ ) -> None:
+ import httpx
+
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ litellm_params = self._litellm_params()
+ BedrockFilesConfig().transform_retrieve_file_request(
+ file_id=file_id,
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+ response = BedrockFilesConfig().transform_retrieve_file_response(
+ raw_response=httpx.Response(
+ 206,
+ headers={
+ "Content-Range": "bytes 0-0/4321",
+ "Last-Modified": "Wed, 21 Oct 2015 07:28:00 GMT",
+ },
+ request=httpx.Request("GET", file_id),
+ ),
+ logging_obj=MagicMock(),
+ litellm_params=litellm_params,
+ )
+
+ assert response.id == file_id
+ assert response.bytes == 4321
+ assert response.created_at == 1445412480
+ assert response.filename == file_id.rsplit("/", 1)[-1]
+ assert response.purpose == purpose
+ assert response.status == "processed"
+ assert response.object == "file"
+
+ @pytest.mark.parametrize(
+ ("file_id", "purpose"),
+ [
+ pytest.param(
+ "s3://out-bucket/outpfx/litellm-batch-outputs/job-123/x.jsonl.out",
+ "batch_output",
+ id="output-bucket",
+ ),
+ pytest.param(
+ "s3://in-bucket/pfx/litellm-bedrock-files/job-123/input.jsonl",
+ "batch",
+ id="input-bucket-upload",
+ ),
+ pytest.param(
+ "s3://in-bucket/pfx/litellm-batch-outputs/job-123/x.jsonl.out",
+ "batch_output",
+ id="input-bucket-output",
+ ),
+ ],
+ )
+ def test_transform_retrieve_file_response_uses_the_retrieved_bucket_prefix(
+ self, file_id: str, purpose: str, monkeypatch: pytest.MonkeyPatch
+ ) -> None:
+ import httpx
+
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ monkeypatch.delenv("AWS_S3_BUCKET_NAME", raising=False)
+ monkeypatch.delenv("AWS_S3_OUTPUT_BUCKET_NAME", raising=False)
+ litellm_params: Final = _trusted_bucket_snapshot(
+ s3_bucket_name="in-bucket/pfx",
+ s3_output_bucket_name="out-bucket/outpfx",
+ )
+ config: Final = BedrockFilesConfig()
+ config.transform_retrieve_file_request(
+ file_id=file_id,
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+ response: Final = config.transform_retrieve_file_response(
+ raw_response=httpx.Response(
+ 206,
+ headers={"Content-Range": "bytes 0-0/1"},
+ request=httpx.Request("GET", file_id),
+ ),
+ logging_obj=MagicMock(),
+ litellm_params=litellm_params,
+ )
+
+ assert response.purpose == purpose
+
+ def test_transform_retrieve_file_response_accepts_verified_empty_object(self, monkeypatch):
+ import httpx
+
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ litellm_params = self._litellm_params()
+ BedrockFilesConfig().transform_retrieve_file_request(
+ file_id=self.S3_URI,
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+ response = BedrockFilesConfig().transform_retrieve_file_response(
+ raw_response=httpx.Response(
+ 416,
+ content=b"InvalidRange0",
+ request=httpx.Request("GET", self.S3_URI),
+ ),
+ logging_obj=MagicMock(),
+ litellm_params=litellm_params,
+ )
+
+ assert response.bytes == 0
+ assert response.filename == "input.jsonl.out"
+ assert response.purpose == "batch_output"
+
+ def test_transform_retrieve_file_response_rejects_unverified_empty_object(self, monkeypatch):
+ import httpx
+
+ from litellm.llms.bedrock.common_utils import BedrockError
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ litellm_params = self._litellm_params()
+ BedrockFilesConfig().transform_retrieve_file_request(
+ file_id=self.S3_URI,
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+ with pytest.raises(BedrockError):
+ BedrockFilesConfig().transform_retrieve_file_response(
+ raw_response=httpx.Response(
+ 416,
+ content=b"InvalidRange",
+ request=httpx.Request("GET", self.S3_URI),
+ ),
+ logging_obj=MagicMock(),
+ litellm_params=litellm_params,
+ )
+
+ def test_transform_retrieve_file_response_uses_content_length_when_range_is_ignored(self, monkeypatch):
+ import httpx
+
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ litellm_params = self._litellm_params()
+ BedrockFilesConfig().transform_retrieve_file_request(
+ file_id=self.S3_URI,
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+ response = BedrockFilesConfig().transform_retrieve_file_response(
+ raw_response=httpx.Response(
+ 200,
+ headers={"Content-Length": "4321"},
+ request=httpx.Request("GET", self.S3_URI),
+ ),
+ logging_obj=MagicMock(),
+ litellm_params=litellm_params,
+ )
+
+ assert response.bytes == 4321
+
+ def test_transform_retrieve_file_response_raises_on_s3_error(self, monkeypatch):
+ import httpx
+
+ from litellm.llms.bedrock.common_utils import BedrockError
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ litellm_params = self._litellm_params()
+ BedrockFilesConfig().transform_retrieve_file_request(
+ file_id=self.S3_URI,
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+
+ with pytest.raises(BedrockError, match="AccessDenied"):
+ BedrockFilesConfig().transform_retrieve_file_response(
+ raw_response=httpx.Response(
+ 403,
+ content=b"AccessDenied",
+ request=httpx.Request("GET", self.S3_URI),
+ ),
+ logging_obj=MagicMock(),
+ litellm_params=litellm_params,
+ )
+
def test_transform_file_content_request_decodes_unified_file_id(self, monkeypatch):
"""Base64 unified ids carrying llm_output_file_id must resolve to their S3 object."""
import base64
diff --git a/tests/unit/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py
index 82cc9e9d3ac..25d2d11b5ca 100644
--- a/tests/unit/llms/custom_httpx/test_llm_http_handler.py
+++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py
@@ -45,6 +45,7 @@ from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
+from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
from litellm.llms.anthropic.skills.transformation import AnthropicSkillsConfig
from litellm.llms.openai.evals.transformation import OpenAIEvalsConfig
from litellm.llms.mistral.files.transformation import MistralFilesConfig
@@ -4925,3 +4926,133 @@ async def test_lookup_handlers_raise_the_provider_error_status(name: str, is_asy
assert error.value.status_code == status_code
assert "No such object" in error.value.message
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("is_async", (False, True))
+@pytest.mark.parametrize(
+ ("status_code", "content", "headers", "expected_bytes"),
+ (
+ (
+ 416,
+ b"InvalidRange0",
+ {},
+ 0,
+ ),
+ (206, b"", {"Content-Range": "bytes 0-0/4321"}, 4321),
+ ),
+)
+async def test_retrieve_file_accepts_bedrock_successful_range_responses(
+ is_async: bool,
+ status_code: int,
+ content: bytes,
+ headers: dict[str, str],
+ expected_bytes: int,
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ file_id = "s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl"
+ transport = httpx.MockTransport(
+ lambda request: httpx.Response(
+ status_code,
+ content=content,
+ headers=headers,
+ request=request,
+ )
+ )
+ params = {
+ "aws_access_key_id": "AKIAEXAMPLE",
+ "aws_secret_access_key": "secret",
+ "aws_region_name": "us-west-2",
+ }
+ handler = BaseLLMHTTPHandler()
+
+ if is_async:
+ client = AsyncHTTPHandler()
+ await client.close()
+ async_client = httpx.AsyncClient(transport=transport)
+ client.client = async_client
+ try:
+ result = await handler.async_retrieve_file(
+ file_id=file_id,
+ provider_config=BedrockFilesConfig(),
+ litellm_params=params,
+ headers={},
+ logging_obj=Mock(),
+ client=client,
+ )
+ finally:
+ await async_client.aclose()
+ else:
+ sync_client = httpx.Client(transport=transport)
+ client = HTTPHandler(client=sync_client)
+ try:
+ result = handler.retrieve_file(
+ file_id=file_id,
+ provider_config=BedrockFilesConfig(),
+ litellm_params=params,
+ headers={},
+ logging_obj=Mock(),
+ client=client,
+ )
+ finally:
+ sync_client.close()
+
+ assert result.bytes == expected_bytes
+ assert result.filename == "output.jsonl"
+ assert result.purpose == "batch_output"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("is_async", (False, True))
+async def test_retrieve_file_rejects_unverified_bedrock_range_error(
+ is_async: bool, monkeypatch: pytest.MonkeyPatch
+) -> None:
+ monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
+ file_id = "s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl"
+ transport = httpx.MockTransport(
+ lambda request: httpx.Response(
+ 416,
+ content=b"InvalidRange",
+ request=request,
+ )
+ )
+ params = {
+ "aws_access_key_id": "AKIAEXAMPLE",
+ "aws_secret_access_key": "secret",
+ "aws_region_name": "us-west-2",
+ }
+ handler = BaseLLMHTTPHandler()
+
+ if is_async:
+ client = AsyncHTTPHandler()
+ await client.close()
+ async_client = httpx.AsyncClient(transport=transport)
+ client.client = async_client
+ try:
+ with pytest.raises(BaseLLMException, match="InvalidRange"):
+ await handler.async_retrieve_file(
+ file_id=file_id,
+ provider_config=BedrockFilesConfig(),
+ litellm_params=params,
+ headers={},
+ logging_obj=Mock(),
+ client=client,
+ )
+ finally:
+ await async_client.aclose()
+ else:
+ sync_client = httpx.Client(transport=transport)
+ client = HTTPHandler(client=sync_client)
+ try:
+ with pytest.raises(BaseLLMException, match="InvalidRange"):
+ handler.retrieve_file(
+ file_id=file_id,
+ provider_config=BedrockFilesConfig(),
+ litellm_params=params,
+ headers={},
+ logging_obj=Mock(),
+ client=client,
+ )
+ finally:
+ sync_client.close()
diff --git a/tests/unit/proxy/common_utils/test_check_batch_cost.py b/tests/unit/proxy/common_utils/test_check_batch_cost.py
index 87df82c01d0..58d6a44e4a4 100644
--- a/tests/unit/proxy/common_utils/test_check_batch_cost.py
+++ b/tests/unit/proxy/common_utils/test_check_batch_cost.py
@@ -23,6 +23,32 @@ _CLAIM_UNIFIED_BATCH_ID = "dW5pZmllZF9iYXRjaF9pZA=="
_CLAIM_OUTPUT_FILE_ID = "file-output-123"
+def test_batch_output_file_object_derives_metadata():
+ from litellm_enterprise.proxy.hooks.managed_files import _batch_output_file_object
+
+ output_file = _batch_output_file_object(
+ unified_file_id="unified-output",
+ raw_file_id="s3://bucket/path/to/output.jsonl.out",
+ size_bytes=4321,
+ fallback=False,
+ )
+ provider_file = _batch_output_file_object(
+ unified_file_id="unified-provider",
+ raw_file_id="file-abc",
+ size_bytes=0,
+ fallback=False,
+ )
+
+ assert output_file.id == "unified-output"
+ assert output_file.object == "file"
+ assert output_file.purpose == "batch_output"
+ assert output_file.filename == "output.jsonl.out"
+ assert output_file.bytes == 4321
+ assert output_file.status == "processed"
+ assert provider_file.filename == "file-abc"
+ assert provider_file.bytes == 0
+
+
def _batch_cost_result(
cost: float,
usage: dict,
@@ -1015,9 +1041,11 @@ class TestCheckBatchCost:
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=response)
- mock_hook = MagicMock()
+ from litellm.proxy.openai_files_endpoints.common_utils import ManagedBatchOutputFileWriter
+
+ mock_hook = MagicMock(spec=ManagedBatchOutputFileWriter)
mock_hook.get_unified_output_file_id.side_effect = [unified_error_file_id]
- mock_hook.store_unified_file_id = AsyncMock()
+ mock_hook.store_batch_output_file = AsyncMock()
check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = mock_hook
await check_batch_cost_instance.check_batch_cost()
@@ -1027,14 +1055,19 @@ class TestCheckBatchCost:
model_id="model-123",
model_name="gpt-5-batch",
)
- stored = {
- next(iter(c.kwargs["model_mappings"].values())): c.kwargs["file_id"]
- for c in mock_hook.store_unified_file_id.call_args_list
- }
- assert stored == {raw_error_file_id: unified_error_file_id}
- for store_call in mock_hook.store_unified_file_id.call_args_list:
- assert store_call.kwargs["user_api_key_dict"].user_id == "user-1"
- assert store_call.kwargs["user_api_key_dict"].team_id == "team-1"
+ mock_hook.store_batch_output_file.assert_awaited_once_with(
+ unified_file_id=unified_error_file_id,
+ provider_file_id=raw_error_file_id,
+ model_id="model-123",
+ model_name="gpt-5-batch",
+ owner=mock_hook.store_batch_output_file.await_args.kwargs["owner"],
+ litellm_parent_otel_span=None,
+ size_bytes=None,
+ fetch_provider_details=True,
+ )
+ owner = mock_hook.store_batch_output_file.await_args.kwargs["owner"]
+ assert owner.user_id == "user-1"
+ assert owner.team_id == "team-1"
assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
update_call = mock_prisma_client.db.litellm_managedobjecttable.update.call_args
@@ -1506,6 +1539,8 @@ class TestCheckBatchCost:
Without this, GET /batches/{id} returns a raw file ID that cannot be routed
through the proxy, causing API_KEY errors when clients call GET /files/{id}/content.
"""
+ from litellm.proxy.openai_files_endpoints.common_utils import ManagedBatchOutputFileWriter
+
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
@@ -1540,12 +1575,12 @@ class TestCheckBatchCost:
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
- mock_hook = MagicMock()
+ mock_hook = MagicMock(spec=ManagedBatchOutputFileWriter)
mock_hook.get_unified_output_file_id.side_effect = [
fake_managed_output_id,
fake_managed_error_id,
]
- mock_hook.store_unified_file_id = AsyncMock()
+ mock_hook.store_batch_output_file = AsyncMock()
check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = mock_hook
mock_file_content = MagicMock()
@@ -1608,16 +1643,25 @@ class TestCheckBatchCost:
model_id="model-123",
model_name="gpt-5-batch",
)
- assert mock_hook.store_unified_file_id.await_count == 2
- # {raw_file_id: managed_file_id} for each store call
+ assert mock_hook.store_batch_output_file.await_count == 2
+ assert {
+ call.kwargs["model_name"]
+ for call in mock_hook.store_batch_output_file.await_args_list
+ } == {"gpt-5-batch"}
stored = {
- next(iter(c[1]["model_mappings"].values())): c[1]["file_id"]
- for c in mock_hook.store_unified_file_id.call_args_list
+ c.kwargs["provider_file_id"]: c.kwargs["unified_file_id"]
+ for c in mock_hook.store_batch_output_file.call_args_list
}
assert stored == {
raw_output_file_id: fake_managed_output_id,
raw_error_file_id: fake_managed_error_id,
}
+ stored_sizes = {
+ c.kwargs["provider_file_id"]: c.kwargs["size_bytes"]
+ for c in mock_hook.store_batch_output_file.call_args_list
+ }
+ assert stored_sizes[raw_output_file_id] == len(mock_file_content.content)
+ assert stored_sizes[raw_error_file_id] is None
assert mock_response.output_file_id == fake_managed_output_id
assert mock_response.error_file_id == fake_managed_error_id
@@ -2192,10 +2236,10 @@ class TestManagedOutputFileIdEncodesPublicModelGroup:
from litellm_enterprise.proxy.common_utils.check_batch_cost import (
CheckBatchCost,
)
+
+ from enterprise.litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles
+ from litellm.proxy.openai_files_endpoints.common_utils import ManagedBatchOutputFileWriter
from litellm.types.utils import LiteLLMBatch
- from enterprise.litellm_enterprise.proxy.hooks.managed_files import (
- PROXY_LiteLLMManagedFiles,
- )
router = MagicMock()
router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"})
@@ -2206,13 +2250,13 @@ class TestManagedOutputFileIdEncodesPublicModelGroup:
deployment.model_info.model_dump.return_value = {}
router.get_deployment = MagicMock(return_value=deployment)
- hook = MagicMock()
+ hook = MagicMock(spec=ManagedBatchOutputFileWriter)
hook.get_unified_output_file_id = lambda output_file_id, model_id, model_name: (
PROXY_LiteLLMManagedFiles.get_unified_output_file_id(
None, output_file_id=output_file_id, model_id=model_id, model_name=model_name
)
)
- hook.store_unified_file_id = AsyncMock()
+ hook.store_batch_output_file = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.get_proxy_hook.return_value = hook
diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py
index 3aed934f40c..f2e033c9141 100644
--- a/tests/unit/repositories/test_repositories.py
+++ b/tests/unit/repositories/test_repositories.py
@@ -2221,6 +2221,37 @@ class TestPrismaTableRepository:
with pytest.raises(RuntimeError, match="No DB Connected"):
_ = repo.table
+ @pytest.mark.asyncio
+ async def test_managed_file_repository_updates_existing_file_object_only(self):
+ from litellm.repositories.managed_file_repository import ManagedFileRepository
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ class UpdateManyMockTable(MockTable):
+ async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> int:
+ return int(await self.update(where, data) is not None)
+
+ file_table = UpdateManyMockTable(pk_field="unified_file_id")
+ await file_table.create({"unified_file_id": "existing-file", "file_object": "{}"})
+ prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_managedfiletable=file_table))
+ repository = ManagedFileRepository(prisma_client)
+ file_object = OpenAIFileObject(
+ id="existing-file",
+ object="file",
+ bytes=836,
+ created_at=456,
+ filename="output.jsonl",
+ purpose="batch_output",
+ status="processed",
+ )
+
+ assert await repository.update_file_object("existing-file", file_object) is True
+ stored_row = await file_table.find_unique(where={"unified_file_id": "existing-file"})
+ assert stored_row is not None
+ assert stored_row.file_object == file_object.model_dump_json()
+
+ assert await repository.update_file_object("missing-file", file_object) is False
+ assert await file_table.find_unique(where={"unified_file_id": "missing-file"}) is None
+
CONFIG_SYNCED_TABLE_NAMES = frozenset(
{
"litellm_agentstable",
diff --git a/tests/unit/types/llms/test_types_llms_openai.py b/tests/unit/types/llms/test_types_llms_openai.py
index e59643f6509..871e4cc8cfb 100644
--- a/tests/unit/types/llms/test_types_llms_openai.py
+++ b/tests/unit/types/llms/test_types_llms_openai.py
@@ -552,6 +552,17 @@ class TestOpenAIFileObjectBatchGuardrailSerialization:
original = self._file_object(litellm_batch_guardrail=self._report())
assert OpenAIFileObject(**original.model_dump()) == original
+ def test_details_fallback_marker_is_omitted_when_unset_and_round_trips_when_set(self):
+ from litellm.types.llms.openai import OpenAIFileObject
+
+ without_marker = self._file_object()
+ assert "litellm_details_fallback" not in without_marker.model_dump()
+ assert "litellm_details_fallback" not in without_marker.model_dump_json()
+
+ with_marker = self._file_object(litellm_details_fallback=True)
+ assert with_marker.model_dump()["litellm_details_fallback"] is True
+ assert OpenAIFileObject.model_validate_json(with_marker.model_dump_json()) == with_marker
+
def test_serialization_json_schema_still_describes_the_model(self):
"""A return annotation on the wrap serializer would collapse this to a bare object."""
from litellm.types.llms.openai import OpenAIFileObject