fix(proxy): save file details for every batch output file so they list and retrieve (#41761)

* fix(proxy): register bedrock batch output files with a file object so they list

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): guard output file size when content is missing

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): cover managed batch output file listings

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): persist metadata for batch output files

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): cover batch output file listing and retrieval end to end

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): preserve provider metadata for managed batch files

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): update managed batch output registration mocks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): inject fake prisma client in managed batch output file tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): retry provider file details and refresh fallback batch output entries

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): cover fallback batch output file details refresh after provider recovers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): cover model-name managed batch file retrieval

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(managed_files): bound the batch output file lookup so a slow provider cannot stall GET /v1/batches

* fix(managed_files): skip the provider lookup for a fallback entry written moments ago

* fix(proxy): safely refresh managed batch file details

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep batch listing provider-free

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(bedrock): return metadata for empty S3 objects

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): strengthen batch fallback regression assertions

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* docs(proxy): describe provider reads and DB writes for batch output files

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): sanitize caller-derived ids in managed file logs

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* docs(proxy): simplify batch and file endpoint docstrings

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* chore(ui): regenerate API types from proxy OpenAPI spec

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(test): resolve managed files lint violations

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(proxy): save refreshed managed file details through ManagedFileRepository

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(bedrock): classify retrieved file purpose by its own bucket prefix

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* docs(proxy): move batch and file endpoint behavior notes to litellm-docs

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(bedrock): resolve strict lint violations in file retrieval

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(bedrock): mark retrieval request metadata handoffs

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): resolve managed batch type-check diagnostics

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): avoid duplicate final response bindings

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mrinal <mrinal@berri.ai>
Co-authored-by: Yucheng He <yucheng@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 18:51:45 -07:00 • committed by GitHub
parent f48d837cd2
commit 390bea6a53
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
25 changed files with 3442 additions and 272 deletions

View file

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

View file

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

View file

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

View file

@ -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,

View file

@ -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,

View file

@ -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,

View file

@ -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,

View file

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

View file

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

View file

@ -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,

View file

@ -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,

View file

@ -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,

View file

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

View file

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

View file

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

View file

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

File diff suppressed because it is too large Load diff

View file

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

View file

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

View file

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

View file

@ -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"<Error><Code>InvalidRange</Code><ActualObjectSize>0</ActualObjectSize></Error>",
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"<Error><Code>InvalidRange</Code></Error>",
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"<Error><Code>AccessDenied</Code></Error>",
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

View file

@ -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"<Error><Code>InvalidRange</Code><ActualObjectSize>0</ActualObjectSize></Error>",
{},
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"<Error><Code>InvalidRange</Code></Error>",
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()

View file

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

View file

@ -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",

View file

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