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