mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(proxy): save file details for every batch output file so they list and retrieve (#41761)
* fix(proxy): register bedrock batch output files with a file object so they list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): guard output file size when content is missing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover managed batch output file listings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): persist metadata for batch output files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover batch output file listing and retrieval end to end Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): preserve provider metadata for managed batch files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): update managed batch output registration mocks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): inject fake prisma client in managed batch output file tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): retry provider file details and refresh fallback batch output entries Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover fallback batch output file details refresh after provider recovers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover model-name managed batch file retrieval Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(managed_files): bound the batch output file lookup so a slow provider cannot stall GET /v1/batches * fix(managed_files): skip the provider lookup for a fallback entry written moments ago * fix(proxy): safely refresh managed batch file details Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep batch listing provider-free Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): return metadata for empty S3 objects Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): strengthen batch fallback regression assertions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(proxy): describe provider reads and DB writes for batch output files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): sanitize caller-derived ids in managed file logs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(proxy): simplify batch and file endpoint docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(ui): regenerate API types from proxy OpenAPI spec Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(test): resolve managed files lint violations Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): save refreshed managed file details through ManagedFileRepository Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): classify retrieved file purpose by its own bucket prefix Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(proxy): move batch and file endpoint behavior notes to litellm-docs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): resolve strict lint violations in file retrieval Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): mark retrieval request metadata handoffs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): resolve managed batch type-check diagnostics Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): avoid duplicate final response bindings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mrinal <mrinal@berri.ai> Co-authored-by: Yucheng He <yucheng@berri.ai>
This commit is contained in:
parent
f48d837cd2
commit
390bea6a53
25 changed files with 3442 additions and 272 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
18
litellm/repositories/managed_file_repository.py
Normal file
18
litellm/repositories/managed_file_repository.py
Normal file
|
|
@ -0,0 +1,18 @@
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.repositories.table_repositories import PrismaTableRepository
|
||||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models # noqa: F401 # used by quoted base-class subscripts
|
||||
|
||||
|
||||
class ManagedFileRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileTable"]):
|
||||
table_name = "litellm_managedfiletable"
|
||||
|
||||
async def update_file_object(self, unified_file_id: str, file_object: OpenAIFileObject) -> bool:
|
||||
updated_rows: Final = await self.table.update_many(
|
||||
where={"unified_file_id": unified_file_id},
|
||||
data={"file_object": file_object.model_dump_json()},
|
||||
)
|
||||
return updated_rows > 0
|
||||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
651
tests/integration/management/test_batch_output_file_listing.py
Normal file
651
tests/integration/management/test_batch_output_file_listing.py
Normal file
|
|
@ -0,0 +1,651 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import datetime
|
||||
import json
|
||||
import socket
|
||||
import socketserver
|
||||
import ssl
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from cryptography import x509
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import ec
|
||||
from cryptography.x509.oid import NameOID
|
||||
from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm.litellm_core_utils.cloud_storage_security import BEDROCK_MANAGED_S3_OUTPUT_PREFIX
|
||||
|
||||
OUTPUT_BYTES: Final = 4096
|
||||
ERROR_BYTES: Final = 512
|
||||
BEDROCK_MODEL: Final = "bedrock/anthropic.claude-3-haiku-20240307-v1:0"
|
||||
BEDROCK_MODEL_ID: Final = "anthropic.claude-3-haiku-20240307-v1:0"
|
||||
BEDROCK_REGION: Final = "us-east-1"
|
||||
BEDROCK_AUTHORITY: Final = f"bedrock.{BEDROCK_REGION}.amazonaws.com:443"
|
||||
BEDROCK_BUCKET: Final = "integration-batch-listing-bucket"
|
||||
BEDROCK_ROLE_ARN: Final = "arn:aws:iam::123456789012:role/integration-batch-role"
|
||||
BEDROCK_JOB_ARN_PREFIX: Final = f"arn:aws:bedrock:{BEDROCK_REGION}:123456789012:model-invocation-job/"
|
||||
BEDROCK_LAST_MODIFIED: Final = "Thu, 02 Oct 2025 12:00:00 GMT"
|
||||
BEDROCK_OUTPUT_CONTENT: Final = b'{"recordId":"req-1","modelOutput":{}}\n'
|
||||
_BATCH_PROCESSED_SQL: Final = 'SELECT batch_processed FROM "LiteLLM_ManagedObjectTable" WHERE unified_object_id=%s'
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _OpenAIBatch:
|
||||
model: str
|
||||
owner_key: str
|
||||
unrelated_key: str | None
|
||||
batch_id: str
|
||||
input_file_id: str
|
||||
scenario: ScenarioHandle
|
||||
|
||||
|
||||
def _output_content(model: str) -> str:
|
||||
return (
|
||||
json.dumps(
|
||||
{
|
||||
"id": "batch_req_$REQUEST_ID",
|
||||
"custom_id": "req-1",
|
||||
"response": {
|
||||
"status_code": 200,
|
||||
"request_id": "$REQUEST_ID",
|
||||
"body": {
|
||||
"id": "chatcmpl-$REQUEST_ID",
|
||||
"object": "chat.completion",
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "ok"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 2, "completion_tokens": 1, "total_tokens": 3},
|
||||
},
|
||||
},
|
||||
"error": None,
|
||||
},
|
||||
separators=(",", ":"),
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
|
||||
def _error_content() -> str:
|
||||
return (
|
||||
json.dumps(
|
||||
{
|
||||
"id": "batch_req_$REQUEST_ID",
|
||||
"custom_id": "req-2",
|
||||
"response": {"status_code": 400, "body": {"error": {"message": "rejected"}}},
|
||||
"error": {"code": "bad_request", "message": "rejected"},
|
||||
},
|
||||
separators=(",", ":"),
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
|
||||
def _batch_routes(
|
||||
model: str,
|
||||
*,
|
||||
output_file: bool = True,
|
||||
metadata_fails: bool = False,
|
||||
) -> RoutedResponse:
|
||||
completed: Final[dict[str, JsonValue]] = {
|
||||
"id": "batch-$REQUEST_ID",
|
||||
"object": "batch",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"errors": None,
|
||||
"input_file_id": "file-in-$REQUEST_ID",
|
||||
"completion_window": "24h",
|
||||
"status": "completed",
|
||||
"output_file_id": "file-out-$REQUEST_ID" if output_file else None,
|
||||
"error_file_id": "file-err-$REQUEST_ID",
|
||||
"created_at": 1,
|
||||
"in_progress_at": 1,
|
||||
"completed_at": 1,
|
||||
"expires_at": 1,
|
||||
"request_counts": {"total": 1 if not output_file else 2, "completed": 1 if output_file else 0, "failed": 1},
|
||||
"metadata": None,
|
||||
}
|
||||
output_metadata: Final = JsonResponse(
|
||||
content_type="application/json",
|
||||
body=(
|
||||
{"error": "provider metadata unavailable"}
|
||||
if metadata_fails
|
||||
else {
|
||||
"id": "file-out-$REQUEST_ID",
|
||||
"object": "file",
|
||||
"purpose": "batch_output",
|
||||
"bytes": OUTPUT_BYTES,
|
||||
"created_at": 1,
|
||||
"filename": "output.jsonl",
|
||||
"status": "processed",
|
||||
}
|
||||
),
|
||||
status=500 if metadata_fails else 200,
|
||||
)
|
||||
return RoutedResponse(
|
||||
content_type="application/x-routed",
|
||||
routes={
|
||||
"POST /files": JsonResponse(
|
||||
content_type="application/json",
|
||||
body={
|
||||
"id": "file-in-$REQUEST_ID",
|
||||
"object": "file",
|
||||
"purpose": "batch",
|
||||
"bytes": 100,
|
||||
"created_at": 1,
|
||||
"filename": "input.jsonl",
|
||||
"status": "processed",
|
||||
},
|
||||
),
|
||||
"POST /batches": JsonResponse(
|
||||
content_type="application/json",
|
||||
body={**completed, "status": "validating", "output_file_id": None, "error_file_id": None},
|
||||
),
|
||||
"GET /batches/batch-$REQUEST_ID": JsonResponse(content_type="application/json", body=completed),
|
||||
"GET /files/file-out-$REQUEST_ID": output_metadata,
|
||||
"GET /files/file-err-$REQUEST_ID": JsonResponse(
|
||||
content_type="application/json",
|
||||
body={
|
||||
"id": "file-err-$REQUEST_ID",
|
||||
"object": "file",
|
||||
"purpose": "batch_output",
|
||||
"bytes": ERROR_BYTES,
|
||||
"created_at": 1,
|
||||
"filename": "errors.jsonl",
|
||||
"status": "processed",
|
||||
},
|
||||
),
|
||||
"GET /files/file-out-$REQUEST_ID/content": TextResponse(
|
||||
content_type="application/jsonl",
|
||||
body=_output_content(model),
|
||||
),
|
||||
"GET /files/file-err-$REQUEST_ID/content": TextResponse(
|
||||
content_type="application/jsonl",
|
||||
body=_error_content(),
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _create_batch(
|
||||
scenario: Scenario,
|
||||
routes: RoutedResponse,
|
||||
*,
|
||||
unrelated_user: bool = False,
|
||||
) -> _OpenAIBatch:
|
||||
handle: Final = register_scenario(f"batch-output-listing-{uuid.uuid4().hex}", routes)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
model: Final = scenario.model(api_base=handle.api_base())
|
||||
owner_id: Final = scenario.user(user_role="internal_user")
|
||||
owner_key: Final = scenario.key(user_id=owner_id, models=[model])
|
||||
unrelated_key: Final = (
|
||||
scenario.key(user_id=scenario.user(user_role="internal_user"), models=[model]) if unrelated_user else None
|
||||
)
|
||||
input_content: Final = json.dumps(
|
||||
{
|
||||
"custom_id": "req-1",
|
||||
"method": "POST",
|
||||
"url": "/v1/chat/completions",
|
||||
"body": {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": 8},
|
||||
}
|
||||
).encode()
|
||||
uploaded: Final = scenario.gateway.request_multipart(
|
||||
"/v1/files",
|
||||
{"purpose": "batch", "target_model_names": model},
|
||||
{"file": ("input.jsonl", input_content, "application/jsonl")},
|
||||
key=owner_key,
|
||||
)
|
||||
assert uploaded.status_code == 200, uploaded.text
|
||||
input_file_id: Final = string_value(JSON_OBJECT.validate_json(uploaded.content)["id"])
|
||||
created: Final = scenario.gateway.request(
|
||||
"POST",
|
||||
"/v1/batches",
|
||||
{
|
||||
"input_file_id": input_file_id,
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
"model": model,
|
||||
},
|
||||
key=owner_key,
|
||||
)
|
||||
assert created.status_code == 200, created.text
|
||||
batch_id: Final = string_value(JSON_OBJECT.validate_json(created.content)["id"])
|
||||
return _OpenAIBatch(model, owner_key, unrelated_key, batch_id, input_file_id, handle)
|
||||
|
||||
|
||||
def _retrieve_batch(gateway: Gateway, batch: _OpenAIBatch) -> dict[str, JsonValue]:
|
||||
response: Final = gateway.request("GET", f"/v1/batches/{batch.batch_id}", key=batch.owner_key)
|
||||
assert response.status_code == 200, response.text
|
||||
return JSON_OBJECT.validate_json(response.content)
|
||||
|
||||
|
||||
def _list_files(
|
||||
gateway: Gateway,
|
||||
key: str,
|
||||
*,
|
||||
purpose: str | None = None,
|
||||
) -> tuple[dict[str, JsonValue], ...]:
|
||||
response: Final = gateway.request(
|
||||
"GET",
|
||||
"/v1/files",
|
||||
key=key,
|
||||
params={"purpose": purpose} if purpose is not None else None,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
values: Final = JSON_OBJECT.validate_json(response.content)["data"]
|
||||
assert isinstance(values, list), response.text
|
||||
return tuple(object_value(value) for value in values)
|
||||
|
||||
|
||||
def _metadata_hit_count(gateway: Gateway, scenario: ScenarioHandle) -> int:
|
||||
response: Final = httpx.get(
|
||||
f"{gateway.upstream_url}/__observations",
|
||||
timeout=15,
|
||||
trust_env=False,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
requests: Final = JSON_OBJECT.validate_json(response.content)["requests"]
|
||||
assert isinstance(requests, list), response.text
|
||||
return sum(1 for request in requests if _is_metadata_hit(request, scenario.scenario_id))
|
||||
|
||||
|
||||
def _is_metadata_hit(request: JsonValue, scenario_id: str) -> bool:
|
||||
if not isinstance(request, dict):
|
||||
return False
|
||||
path: Final = request.get("path")
|
||||
return request.get("method") == "GET" and isinstance(path, str) and path.endswith(f"/files/file-out-{scenario_id}")
|
||||
|
||||
|
||||
def test_output_file_lists_after_owner_retrieves_completed_batch(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
batch: Final = _create_batch(scenario, _batch_routes(model="gpt-4o-mini"))
|
||||
retrieved: Final = _retrieve_batch(gateway, batch)
|
||||
assert retrieved["status"] == "completed", retrieved
|
||||
output_id: Final = string_value(retrieved["output_file_id"])
|
||||
error_id: Final = string_value(retrieved["error_file_id"])
|
||||
files: Final = _list_files(gateway, batch.owner_key)
|
||||
output: Final = next((file for file in files if file.get("id") == output_id), None)
|
||||
assert output is not None, f"Completed batch output {output_id} is absent from GET /v1/files"
|
||||
assert (output["purpose"], output["bytes"]) == ("batch_output", OUTPUT_BYTES), output
|
||||
output_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output")
|
||||
assert {string_value(file["id"]) for file in output_files} == {output_id, error_id}, output_files
|
||||
assert all(file["purpose"] == "batch_output" for file in output_files), output_files
|
||||
input_files: Final = _list_files(gateway, batch.owner_key, purpose="batch")
|
||||
assert tuple(string_value(file["id"]) for file in input_files) == (batch.input_file_id,), input_files
|
||||
|
||||
|
||||
def test_poller_registers_listable_output_files_without_a_batch_retrieve(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
batch: Final = _create_batch(
|
||||
scenario,
|
||||
_batch_routes(model="gpt-4o-mini"),
|
||||
unrelated_user=True,
|
||||
)
|
||||
assert batch.unrelated_key is not None
|
||||
output_files: Final = eventually(
|
||||
lambda: _list_files(gateway, batch.owner_key, purpose="batch_output"),
|
||||
lambda values: len(values) == 2,
|
||||
seconds=60,
|
||||
)
|
||||
assert {file["purpose"] for file in output_files} == {"batch_output"}, output_files
|
||||
unrelated_files: Final = _list_files(gateway, batch.unrelated_key, purpose="batch_output")
|
||||
assert unrelated_files == (), unrelated_files
|
||||
retrieved: Final = _retrieve_batch(gateway, batch)
|
||||
output_id: Final = string_value(retrieved["output_file_id"])
|
||||
error_id: Final = string_value(retrieved["error_file_id"])
|
||||
assert {string_value(file["id"]) for file in output_files} == {output_id, error_id}, output_files
|
||||
|
||||
|
||||
def test_error_file_only_batch_still_lists_its_error_file(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
batch: Final = _create_batch(
|
||||
scenario,
|
||||
_batch_routes(model="gpt-4o-mini", output_file=False),
|
||||
)
|
||||
error_files: Final = eventually(
|
||||
lambda: _list_files(gateway, batch.owner_key, purpose="batch_output"),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=60,
|
||||
)
|
||||
retrieved: Final = _retrieve_batch(gateway, batch)
|
||||
assert retrieved["output_file_id"] is None, retrieved
|
||||
error_id: Final = string_value(retrieved["error_file_id"])
|
||||
assert tuple(string_value(file["id"]) for file in error_files) == (error_id,), error_files
|
||||
assert error_files[0]["purpose"] == "batch_output", error_files
|
||||
|
||||
|
||||
def test_output_file_lists_with_basic_details_when_provider_file_lookup_fails(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
batch: Final = _create_batch(
|
||||
scenario,
|
||||
_batch_routes(model="gpt-4o-mini", metadata_fails=True),
|
||||
)
|
||||
retrieved: Final = _retrieve_batch(gateway, batch)
|
||||
output_id: Final = string_value(retrieved["output_file_id"])
|
||||
output: Final = next(
|
||||
(file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id),
|
||||
None,
|
||||
)
|
||||
assert output is not None, f"Completed batch output {output_id} is absent after provider metadata failure"
|
||||
assert (output["purpose"], output["filename"]) == ("batch_output", f"file-out-{batch.scenario.scenario_id}"), (
|
||||
output
|
||||
)
|
||||
assert "litellm_details_fallback" not in output, output
|
||||
|
||||
|
||||
def test_fallback_output_file_details_refresh_once_provider_recovers(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
batch: Final = _create_batch(
|
||||
scenario,
|
||||
_batch_routes(model="gpt-4o-mini", metadata_fails=True),
|
||||
unrelated_user=True,
|
||||
)
|
||||
retrieved: Final = _retrieve_batch(gateway, batch)
|
||||
output_id: Final = string_value(retrieved["output_file_id"])
|
||||
listed_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output")
|
||||
basic_output: Final = next((file for file in listed_files if file.get("id") == output_id), None)
|
||||
assert basic_output is not None, f"Completed batch output {output_id} is absent after provider metadata failure"
|
||||
assert (basic_output["purpose"], basic_output["filename"]) == (
|
||||
"batch_output",
|
||||
f"file-out-{batch.scenario.scenario_id}",
|
||||
), basic_output
|
||||
assert basic_output["bytes"] != OUTPUT_BYTES, basic_output
|
||||
assert "litellm_details_fallback" not in basic_output, basic_output
|
||||
|
||||
processed_batch_rows: Final = eventually(
|
||||
lambda: read_rows(_BATCH_PROCESSED_SQL, (string_value(retrieved["id"]),)),
|
||||
lambda rows: len(rows) == 1 and rows[0].get("batch_processed") is True,
|
||||
seconds=60,
|
||||
)
|
||||
assert processed_batch_rows[0]["batch_processed"] is True, processed_batch_rows
|
||||
metadata_hits_before_recovery: Final = _metadata_hit_count(gateway, batch.scenario)
|
||||
assert metadata_hits_before_recovery >= 1, "The provider metadata route was not called before recovery"
|
||||
register_scenario(
|
||||
batch.scenario.scenario_id,
|
||||
_batch_routes(model=batch.model),
|
||||
control_url=batch.scenario.control_url,
|
||||
)
|
||||
details_response: Final = gateway.request("GET", f"/v1/files/{output_id}", key=batch.owner_key)
|
||||
assert details_response.status_code == 200, details_response.text
|
||||
details: Final = JSON_OBJECT.validate_json(details_response.content)
|
||||
assert details["bytes"] == OUTPUT_BYTES, details
|
||||
assert (details["filename"], details["purpose"]) == ("output.jsonl", "batch_output"), details
|
||||
assert "litellm_details_fallback" not in details, details
|
||||
assert _metadata_hit_count(gateway, batch.scenario) == 1
|
||||
|
||||
refreshed_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output")
|
||||
refreshed_output: Final = next((file for file in refreshed_files if file.get("id") == output_id), None)
|
||||
assert refreshed_output is not None, f"Refreshed batch output {output_id} is absent from the list"
|
||||
assert (refreshed_output["bytes"], refreshed_output["filename"]) == (OUTPUT_BYTES, "output.jsonl"), (
|
||||
refreshed_output
|
||||
)
|
||||
assert "litellm_details_fallback" not in refreshed_output, refreshed_output
|
||||
|
||||
assert batch.unrelated_key is not None
|
||||
unrelated_details: Final = gateway.request("GET", f"/v1/files/{output_id}", key=batch.unrelated_key)
|
||||
assert unrelated_details.status_code == 403, unrelated_details.text
|
||||
|
||||
|
||||
def _tls_context(directory: Path) -> ssl.SSLContext:
|
||||
key: Final = ec.generate_private_key(ec.SECP256R1())
|
||||
now: Final = datetime.datetime.now(datetime.timezone.utc)
|
||||
name: Final = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, BEDROCK_AUTHORITY.split(":")[0])])
|
||||
certificate: Final = (
|
||||
x509.CertificateBuilder()
|
||||
.subject_name(name)
|
||||
.issuer_name(name)
|
||||
.public_key(key.public_key())
|
||||
.serial_number(x509.random_serial_number())
|
||||
.not_valid_before(now - datetime.timedelta(days=1))
|
||||
.not_valid_after(now + datetime.timedelta(days=1))
|
||||
.sign(key, hashes.SHA256())
|
||||
)
|
||||
certificate_file: Final = directory / "bedrock.pem"
|
||||
key_file: Final = directory / "bedrock.key"
|
||||
certificate_file.write_bytes(certificate.public_bytes(serialization.Encoding.PEM))
|
||||
key_file.write_bytes(
|
||||
key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption())
|
||||
)
|
||||
context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
context.load_cert_chain(certificate_file, key_file)
|
||||
return context
|
||||
|
||||
|
||||
def _pipe(source: socket.socket, sink: socket.socket) -> None:
|
||||
with contextlib.suppress(OSError):
|
||||
for chunk in iter(lambda: source.recv(65536), b""):
|
||||
sink.sendall(chunk)
|
||||
with contextlib.suppress(OSError):
|
||||
sink.shutdown(socket.SHUT_WR)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ConnectProxy:
|
||||
url: str
|
||||
authorities: SimpleQueue[str]
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _bedrock_tunnel(destination: Wire) -> Generator[_ConnectProxy, None, None]:
|
||||
authorities: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
destination_port: Final = int(destination.url.rsplit(":", 1)[1])
|
||||
|
||||
class Tunnel(socketserver.StreamRequestHandler):
|
||||
rbufsize = 0
|
||||
request: socket.socket
|
||||
|
||||
def handle(self) -> None:
|
||||
authority: Final = self.rfile.readline().decode().split()[1]
|
||||
while self.rfile.readline() not in (b"\r\n", b""):
|
||||
pass
|
||||
authorities.put(authority)
|
||||
if authority != BEDROCK_AUTHORITY:
|
||||
self.wfile.write(b"HTTP/1.1 403 Forbidden\r\ncontent-length: 0\r\n\r\n")
|
||||
return
|
||||
self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n")
|
||||
self.request.settimeout(10)
|
||||
with socket.create_connection(("127.0.0.1", destination_port), timeout=10) as upstream:
|
||||
outbound: Final = threading.Thread(target=_pipe, args=(self.request, upstream))
|
||||
outbound.start()
|
||||
_pipe(upstream, self.request)
|
||||
outbound.join(timeout=12)
|
||||
|
||||
with socketserver.ThreadingTCPServer(("127.0.0.1", 0), Tunnel) as server:
|
||||
thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05})
|
||||
thread.start()
|
||||
try:
|
||||
yield _ConnectProxy(f"http://127.0.0.1:{server.server_address[1]}", authorities)
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join(timeout=6)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BedrockControlPlane:
|
||||
job_locations: SimpleQueue[tuple[str, str, str]]
|
||||
job_id: str
|
||||
job_arn: str
|
||||
|
||||
def __call__(self, request: Request) -> Reply:
|
||||
if request.method == "POST" and request.target == "/model-invocation-job":
|
||||
body: Final = JSON_OBJECT.validate_json(request.body)
|
||||
input_config: Final = object_value(object_value(body["inputDataConfig"])["s3InputDataConfig"])
|
||||
output_config: Final = object_value(object_value(body["outputDataConfig"])["s3OutputDataConfig"])
|
||||
job_name: Final = string_value(body["jobName"])
|
||||
self.job_locations.put(
|
||||
(string_value(input_config["s3Uri"]), string_value(output_config["s3Uri"]), job_name)
|
||||
)
|
||||
return Reply(body=json.dumps({"jobArn": self.job_arn}).encode())
|
||||
if request.method == "GET" and request.target.endswith(self.job_id):
|
||||
input_uri, output_uri, retrieved_job_name = self.job_locations.get()
|
||||
self.job_locations.put((input_uri, output_uri, retrieved_job_name))
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"jobArn": self.job_arn,
|
||||
"jobName": retrieved_job_name,
|
||||
"modelId": BEDROCK_MODEL_ID,
|
||||
"status": "Completed",
|
||||
"submitTime": 1700000000,
|
||||
"lastModifiedTime": 1700000001,
|
||||
"endTime": 1700000002,
|
||||
"inputDataConfig": {"s3InputDataConfig": {"s3Uri": input_uri}},
|
||||
"outputDataConfig": {"s3OutputDataConfig": {"s3Uri": output_uri}},
|
||||
"totalRecordCount": 1,
|
||||
"successRecordCount": 1,
|
||||
"errorRecordCount": 0,
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
return Reply(status=404, body=b'{"message":"not scripted"}')
|
||||
|
||||
|
||||
def _bedrock_s3_peer(request: Request) -> Reply:
|
||||
if request.method == "PUT" and request.target.startswith(f"/{BEDROCK_BUCKET}/"):
|
||||
return Reply(body=b"")
|
||||
if request.method == "GET" and request.target.startswith(f"/{BEDROCK_BUCKET}/{BEDROCK_MANAGED_S3_OUTPUT_PREFIX}"):
|
||||
if request.headers.get("range") == "bytes=0-0":
|
||||
return Reply(
|
||||
status=206,
|
||||
body=BEDROCK_OUTPUT_CONTENT[:1],
|
||||
headers={
|
||||
"Content-Range": f"bytes 0-0/{len(BEDROCK_OUTPUT_CONTENT)}",
|
||||
"Last-Modified": BEDROCK_LAST_MODIFIED,
|
||||
},
|
||||
)
|
||||
return Reply(status=200, body=BEDROCK_OUTPUT_CONTENT, headers={"Last-Modified": BEDROCK_LAST_MODIFIED})
|
||||
return Reply(status=404, body=b'{"message":"not scripted"}')
|
||||
|
||||
|
||||
def test_bedrock_batch_output_lists_and_retrieves_details(gateway: Gateway, tmp_path: Path) -> None:
|
||||
job_id: Final = f"integration-batch-listing-{uuid.uuid4().hex}"
|
||||
job_arn: Final = BEDROCK_JOB_ARN_PREFIX + job_id
|
||||
environment: Final = {
|
||||
"AWS_CA_BUNDLE": str(tmp_path / "bedrock.pem"),
|
||||
"AWS_EC2_METADATA_DISABLED": "true",
|
||||
"SSL_VERIFY": "False",
|
||||
}
|
||||
with (
|
||||
wire_server(_bedrock_s3_peer) as s3,
|
||||
wire_server(_BedrockControlPlane(SimpleQueue(), job_id, job_arn), tls=_tls_context(tmp_path)) as bedrock,
|
||||
_bedrock_tunnel(bedrock) as tunnel,
|
||||
owned_proxy(gateway, tmp_path, {**environment, "HTTPS_PROXY": tunnel.url}) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(
|
||||
model=BEDROCK_MODEL,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
|
||||
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
|
||||
aws_region_name=BEDROCK_REGION,
|
||||
s3_bucket_name=BEDROCK_BUCKET,
|
||||
s3_endpoint_url=s3.url,
|
||||
aws_batch_role_arn=BEDROCK_ROLE_ARN,
|
||||
)
|
||||
owner_id: Final = scenario.user(user_role="internal_user")
|
||||
owner_key: Final = scenario.key(user_id=owner_id, models=[model])
|
||||
input_line: Final = {
|
||||
"custom_id": "req-1",
|
||||
"method": "POST",
|
||||
"url": "/v1/chat/completions",
|
||||
"body": {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": 8},
|
||||
}
|
||||
uploaded: Final = candidate.request_multipart(
|
||||
"/v1/files",
|
||||
{"purpose": "batch", "target_model_names": model},
|
||||
{"file": ("input.jsonl", (json.dumps(input_line) + "\n").encode(), "application/jsonl")},
|
||||
key=owner_key,
|
||||
)
|
||||
assert uploaded.status_code == 200, uploaded.text
|
||||
created: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/batches",
|
||||
{
|
||||
"input_file_id": string_value(JSON_OBJECT.validate_json(uploaded.content)["id"]),
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
key=owner_key,
|
||||
)
|
||||
assert created.status_code == 200, created.text
|
||||
batch_id: Final = string_value(JSON_OBJECT.validate_json(created.content)["id"])
|
||||
retrieved: Final = candidate.request("GET", f"/v1/batches/{batch_id}", key=owner_key)
|
||||
assert retrieved.status_code == 200, retrieved.text
|
||||
batch_object: Final = JSON_OBJECT.validate_json(retrieved.content)
|
||||
assert batch_object["status"] == "completed", retrieved.text
|
||||
output_id: Final = string_value(batch_object["output_file_id"])
|
||||
output_files: Final = eventually(
|
||||
lambda: _list_files(candidate, owner_key, purpose="batch_output"),
|
||||
lambda values: any(file.get("id") == output_id for file in values),
|
||||
seconds=30,
|
||||
)
|
||||
listed_output: Final = next(file for file in output_files if file.get("id") == output_id)
|
||||
details: Final = candidate.request("GET", f"/v1/files/{output_id}", key=owner_key)
|
||||
assert details.status_code == 200, details.text
|
||||
detail_object: Final = JSON_OBJECT.validate_json(details.content)
|
||||
assert (detail_object["id"], detail_object["bytes"], detail_object["purpose"]) == (
|
||||
output_id,
|
||||
len(BEDROCK_OUTPUT_CONTENT),
|
||||
"batch_output",
|
||||
), detail_object
|
||||
assert listed_output["id"] == output_id, listed_output
|
||||
s3_requests: Final = s3.drain()
|
||||
uploads: Final = tuple(request for request in s3_requests if request.method == "PUT")
|
||||
assert len(uploads) == 1, f"Expected one input S3 upload, saw {[request.target for request in uploads]}"
|
||||
ranged_metadata: Final = tuple(
|
||||
request
|
||||
for request in s3_requests
|
||||
if request.method == "GET" and request.headers.get("range") == "bytes=0-0"
|
||||
)
|
||||
assert ranged_metadata, "Bedrock file metadata retrieval did not issue a ranged S3 GET"
|
||||
assert all(
|
||||
request.headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ") for request in ranged_metadata
|
||||
), ranged_metadata
|
||||
authorities: Final = tuple(tunnel.authorities.get_nowait() for _ in range(tunnel.authorities.qsize()))
|
||||
assert BEDROCK_AUTHORITY in authorities, authorities
|
||||
bedrock_requests: Final = bedrock.drain()
|
||||
assert any(request.method == "POST" for request in bedrock_requests), bedrock_requests
|
||||
assert any(request.method == "GET" for request in bedrock_requests), bedrock_requests
|
||||
|
||||
|
||||
def test_repeated_batch_retrieve_does_not_refetch_saved_file_details(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
batch: Final = _create_batch(scenario, _batch_routes(model="gpt-4o-mini"))
|
||||
first: Final = _retrieve_batch(gateway, batch)
|
||||
output_id: Final = string_value(first["output_file_id"])
|
||||
output_before: Final = next(
|
||||
(file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id),
|
||||
None,
|
||||
)
|
||||
assert output_before is not None, f"Completed batch output {output_id} is absent from GET /v1/files"
|
||||
hits_before: Final = _metadata_hit_count(gateway, batch.scenario)
|
||||
assert hits_before >= 1, "The provider file metadata route was not called for the batch output"
|
||||
second: Final = _retrieve_batch(gateway, batch)
|
||||
third: Final = _retrieve_batch(gateway, batch)
|
||||
assert second["status"] == third["status"] == "completed", (second, third)
|
||||
output_after: Final = next(
|
||||
(file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id),
|
||||
None,
|
||||
)
|
||||
assert output_after == output_before, output_after
|
||||
additional_hits: Final = _metadata_hit_count(gateway, batch.scenario)
|
||||
assert additional_hits == 0, f"Repeated batch retrieval fetched metadata {additional_hits} more times"
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -2216,6 +2216,216 @@ class TestBedrockFileContentTransformation:
|
|||
assert "x-amz-content-sha256" in authorization
|
||||
assert "X-Amz-Date" in signed_headers
|
||||
|
||||
def test_transform_retrieve_file_request_adds_unsigned_range(self, monkeypatch):
|
||||
from litellm.llms.bedrock.files.transformation import (
|
||||
S3_SIGNED_REQUEST_HEADERS_PARAM,
|
||||
BedrockFilesConfig,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
litellm_params = self._litellm_params()
|
||||
url, params = BedrockFilesConfig().transform_retrieve_file_request(
|
||||
file_id=self.S3_URI,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert url == self.EXPECTED_URL
|
||||
assert params == {}
|
||||
assert litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]["Range"] == "bytes=0-0"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("file_id", "purpose"),
|
||||
[
|
||||
("s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl", "batch_output"),
|
||||
("s3://my-bucket/litellm-bedrock-files-job-123/input.jsonl", "batch"),
|
||||
],
|
||||
)
|
||||
def test_transform_retrieve_file_response_parses_metadata(
|
||||
self, file_id: str, purpose: str, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
litellm_params = self._litellm_params()
|
||||
BedrockFilesConfig().transform_retrieve_file_request(
|
||||
file_id=file_id,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
response = BedrockFilesConfig().transform_retrieve_file_response(
|
||||
raw_response=httpx.Response(
|
||||
206,
|
||||
headers={
|
||||
"Content-Range": "bytes 0-0/4321",
|
||||
"Last-Modified": "Wed, 21 Oct 2015 07:28:00 GMT",
|
||||
},
|
||||
request=httpx.Request("GET", file_id),
|
||||
),
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert response.id == file_id
|
||||
assert response.bytes == 4321
|
||||
assert response.created_at == 1445412480
|
||||
assert response.filename == file_id.rsplit("/", 1)[-1]
|
||||
assert response.purpose == purpose
|
||||
assert response.status == "processed"
|
||||
assert response.object == "file"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("file_id", "purpose"),
|
||||
[
|
||||
pytest.param(
|
||||
"s3://out-bucket/outpfx/litellm-batch-outputs/job-123/x.jsonl.out",
|
||||
"batch_output",
|
||||
id="output-bucket",
|
||||
),
|
||||
pytest.param(
|
||||
"s3://in-bucket/pfx/litellm-bedrock-files/job-123/input.jsonl",
|
||||
"batch",
|
||||
id="input-bucket-upload",
|
||||
),
|
||||
pytest.param(
|
||||
"s3://in-bucket/pfx/litellm-batch-outputs/job-123/x.jsonl.out",
|
||||
"batch_output",
|
||||
id="input-bucket-output",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_transform_retrieve_file_response_uses_the_retrieved_bucket_prefix(
|
||||
self, file_id: str, purpose: str, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.delenv("AWS_S3_BUCKET_NAME", raising=False)
|
||||
monkeypatch.delenv("AWS_S3_OUTPUT_BUCKET_NAME", raising=False)
|
||||
litellm_params: Final = _trusted_bucket_snapshot(
|
||||
s3_bucket_name="in-bucket/pfx",
|
||||
s3_output_bucket_name="out-bucket/outpfx",
|
||||
)
|
||||
config: Final = BedrockFilesConfig()
|
||||
config.transform_retrieve_file_request(
|
||||
file_id=file_id,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
response: Final = config.transform_retrieve_file_response(
|
||||
raw_response=httpx.Response(
|
||||
206,
|
||||
headers={"Content-Range": "bytes 0-0/1"},
|
||||
request=httpx.Request("GET", file_id),
|
||||
),
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert response.purpose == purpose
|
||||
|
||||
def test_transform_retrieve_file_response_accepts_verified_empty_object(self, monkeypatch):
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
litellm_params = self._litellm_params()
|
||||
BedrockFilesConfig().transform_retrieve_file_request(
|
||||
file_id=self.S3_URI,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
response = BedrockFilesConfig().transform_retrieve_file_response(
|
||||
raw_response=httpx.Response(
|
||||
416,
|
||||
content=b"<Error><Code>InvalidRange</Code><ActualObjectSize>0</ActualObjectSize></Error>",
|
||||
request=httpx.Request("GET", self.S3_URI),
|
||||
),
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert response.bytes == 0
|
||||
assert response.filename == "input.jsonl.out"
|
||||
assert response.purpose == "batch_output"
|
||||
|
||||
def test_transform_retrieve_file_response_rejects_unverified_empty_object(self, monkeypatch):
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
litellm_params = self._litellm_params()
|
||||
BedrockFilesConfig().transform_retrieve_file_request(
|
||||
file_id=self.S3_URI,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
with pytest.raises(BedrockError):
|
||||
BedrockFilesConfig().transform_retrieve_file_response(
|
||||
raw_response=httpx.Response(
|
||||
416,
|
||||
content=b"<Error><Code>InvalidRange</Code></Error>",
|
||||
request=httpx.Request("GET", self.S3_URI),
|
||||
),
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
def test_transform_retrieve_file_response_uses_content_length_when_range_is_ignored(self, monkeypatch):
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
litellm_params = self._litellm_params()
|
||||
BedrockFilesConfig().transform_retrieve_file_request(
|
||||
file_id=self.S3_URI,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
response = BedrockFilesConfig().transform_retrieve_file_response(
|
||||
raw_response=httpx.Response(
|
||||
200,
|
||||
headers={"Content-Length": "4321"},
|
||||
request=httpx.Request("GET", self.S3_URI),
|
||||
),
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert response.bytes == 4321
|
||||
|
||||
def test_transform_retrieve_file_response_raises_on_s3_error(self, monkeypatch):
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
litellm_params = self._litellm_params()
|
||||
BedrockFilesConfig().transform_retrieve_file_request(
|
||||
file_id=self.S3_URI,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
with pytest.raises(BedrockError, match="AccessDenied"):
|
||||
BedrockFilesConfig().transform_retrieve_file_response(
|
||||
raw_response=httpx.Response(
|
||||
403,
|
||||
content=b"<Error><Code>AccessDenied</Code></Error>",
|
||||
request=httpx.Request("GET", self.S3_URI),
|
||||
),
|
||||
logging_obj=MagicMock(),
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
def test_transform_file_content_request_decodes_unified_file_id(self, monkeypatch):
|
||||
"""Base64 unified ids carrying llm_output_file_id must resolve to their S3 object."""
|
||||
import base64
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.llms.azure.videos.transformation import AzureVideoConfig
|
|||
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeMessagesConfig,
|
||||
)
|
||||
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
|
||||
from litellm.llms.anthropic.skills.transformation import AnthropicSkillsConfig
|
||||
from litellm.llms.openai.evals.transformation import OpenAIEvalsConfig
|
||||
from litellm.llms.mistral.files.transformation import MistralFilesConfig
|
||||
|
|
@ -4925,3 +4926,133 @@ async def test_lookup_handlers_raise_the_provider_error_status(name: str, is_asy
|
|||
|
||||
assert error.value.status_code == status_code
|
||||
assert "No such object" in error.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("is_async", (False, True))
|
||||
@pytest.mark.parametrize(
|
||||
("status_code", "content", "headers", "expected_bytes"),
|
||||
(
|
||||
(
|
||||
416,
|
||||
b"<Error><Code>InvalidRange</Code><ActualObjectSize>0</ActualObjectSize></Error>",
|
||||
{},
|
||||
0,
|
||||
),
|
||||
(206, b"", {"Content-Range": "bytes 0-0/4321"}, 4321),
|
||||
),
|
||||
)
|
||||
async def test_retrieve_file_accepts_bedrock_successful_range_responses(
|
||||
is_async: bool,
|
||||
status_code: int,
|
||||
content: bytes,
|
||||
headers: dict[str, str],
|
||||
expected_bytes: int,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
file_id = "s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl"
|
||||
transport = httpx.MockTransport(
|
||||
lambda request: httpx.Response(
|
||||
status_code,
|
||||
content=content,
|
||||
headers=headers,
|
||||
request=request,
|
||||
)
|
||||
)
|
||||
params = {
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "secret",
|
||||
"aws_region_name": "us-west-2",
|
||||
}
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
if is_async:
|
||||
client = AsyncHTTPHandler()
|
||||
await client.close()
|
||||
async_client = httpx.AsyncClient(transport=transport)
|
||||
client.client = async_client
|
||||
try:
|
||||
result = await handler.async_retrieve_file(
|
||||
file_id=file_id,
|
||||
provider_config=BedrockFilesConfig(),
|
||||
litellm_params=params,
|
||||
headers={},
|
||||
logging_obj=Mock(),
|
||||
client=client,
|
||||
)
|
||||
finally:
|
||||
await async_client.aclose()
|
||||
else:
|
||||
sync_client = httpx.Client(transport=transport)
|
||||
client = HTTPHandler(client=sync_client)
|
||||
try:
|
||||
result = handler.retrieve_file(
|
||||
file_id=file_id,
|
||||
provider_config=BedrockFilesConfig(),
|
||||
litellm_params=params,
|
||||
headers={},
|
||||
logging_obj=Mock(),
|
||||
client=client,
|
||||
)
|
||||
finally:
|
||||
sync_client.close()
|
||||
|
||||
assert result.bytes == expected_bytes
|
||||
assert result.filename == "output.jsonl"
|
||||
assert result.purpose == "batch_output"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("is_async", (False, True))
|
||||
async def test_retrieve_file_rejects_unverified_bedrock_range_error(
|
||||
is_async: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket")
|
||||
file_id = "s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl"
|
||||
transport = httpx.MockTransport(
|
||||
lambda request: httpx.Response(
|
||||
416,
|
||||
content=b"<Error><Code>InvalidRange</Code></Error>",
|
||||
request=request,
|
||||
)
|
||||
)
|
||||
params = {
|
||||
"aws_access_key_id": "AKIAEXAMPLE",
|
||||
"aws_secret_access_key": "secret",
|
||||
"aws_region_name": "us-west-2",
|
||||
}
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
if is_async:
|
||||
client = AsyncHTTPHandler()
|
||||
await client.close()
|
||||
async_client = httpx.AsyncClient(transport=transport)
|
||||
client.client = async_client
|
||||
try:
|
||||
with pytest.raises(BaseLLMException, match="InvalidRange"):
|
||||
await handler.async_retrieve_file(
|
||||
file_id=file_id,
|
||||
provider_config=BedrockFilesConfig(),
|
||||
litellm_params=params,
|
||||
headers={},
|
||||
logging_obj=Mock(),
|
||||
client=client,
|
||||
)
|
||||
finally:
|
||||
await async_client.aclose()
|
||||
else:
|
||||
sync_client = httpx.Client(transport=transport)
|
||||
client = HTTPHandler(client=sync_client)
|
||||
try:
|
||||
with pytest.raises(BaseLLMException, match="InvalidRange"):
|
||||
handler.retrieve_file(
|
||||
file_id=file_id,
|
||||
provider_config=BedrockFilesConfig(),
|
||||
litellm_params=params,
|
||||
headers={},
|
||||
logging_obj=Mock(),
|
||||
client=client,
|
||||
)
|
||||
finally:
|
||||
sync_client.close()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue