chore: merge latest main before JWT regression validation

This commit is contained in:
Joshua Valluru 2026-09-19 15:45:52 -07:00
commit dbb1f1285a
26 changed files with 3411 additions and 137 deletions

View file

@ -20,6 +20,7 @@ from typing import (
)
from uuid import NAMESPACE_URL, uuid5
import httpx
from fastapi import HTTPException
from pydantic import ValidationError
@ -34,6 +35,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from openai.types.file_deleted import FileDeleted
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.llms.base_llm.managed_resources.isolation import (
build_list_page,
@ -59,6 +61,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
get_content_type_from_file_object,
get_model_id_from_unified_batch_id,
get_original_file_id,
is_litellm_executed_batch,
map_raw_file_ids_to_unified,
normalize_mime_type_for_provider,
resolve_managed_output_file_model_name,
@ -75,6 +78,7 @@ from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccess
CreateFileRequest,
FileListPage,
FileObject,
HttpxBinaryResponseContent,
OpenAIFileObject,
ResponsesAPIResponse,
)
@ -86,10 +90,6 @@ from litellm.types.utils import (
SpecialEnums,
)
if TYPE_CHECKING:
from litellm.types.llms.openai import HttpxBinaryResponseContent
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from prisma.models import (
@ -204,6 +204,19 @@ def _managed_object_table(prisma_client: PrismaClient) -> _ManagedObjectTableAct
return prisma_client.db.litellm_managedobjecttable
def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, str]:
hidden_params: Final = cast( # cast-ok: _hidden_params is an untyped attribute the upload path sets
"Mapping[str, object]", getattr(file_object, "_hidden_params", None) or {}
)
return MappingProxyType(
{
key: value
for key in ("storage_backend", "storage_url")
if isinstance(value := hidden_params.get(key), str)
}
)
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Class variables or attributes
def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient):
@ -226,6 +239,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
user_api_key_dict: UserAPIKeyAuth,
) -> None:
verbose_logger.info(f"Storing LiteLLM Managed File object with id={file_id} in cache")
storage_metadata: Final = _storage_metadata_of(file_object)
if file_object is not None:
litellm_managed_file_object = LiteLLM_ManagedFileTable(
unified_file_id=file_id,
@ -235,6 +249,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
created_by=resolve_resource_owner_id(user_api_key_dict),
team_id=user_api_key_dict.team_id,
updated_by=user_api_key_dict.user_id,
storage_backend=storage_metadata.get("storage_backend"),
storage_url=storage_metadata.get("storage_url"),
)
await self.internal_usage_cache.async_set_cache(
key=file_id,
@ -262,14 +278,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
file_object_json = file_object.model_dump_json()
db_data["file_object"] = file_object_json
update_data["file_object"] = file_object_json
# Extract storage metadata from hidden params if present
hidden_params = getattr(file_object, "_hidden_params", {}) or {}
if "storage_backend" in hidden_params:
db_data["storage_backend"] = hidden_params["storage_backend"]
update_data["storage_backend"] = hidden_params["storage_backend"]
if "storage_url" in hidden_params:
db_data["storage_url"] = hidden_params["storage_url"]
update_data["storage_url"] = hidden_params["storage_url"]
db_data.update(storage_metadata)
update_data.update(storage_metadata)
verbose_logger.debug(
f"Storage metadata: storage_backend={db_data.get('storage_backend')}, "
@ -314,6 +324,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
request_tags: Sequence[str] | None = None,
persist_attribution: bool = False,
create_if_missing: bool = True,
batch_processed: bool = False,
) -> None:
"""Persist a managed object row, caching it and upserting it in the DB.
@ -328,6 +339,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
row absent from the table is left absent rather than created with the
observer as its creator, because created_by and team_id are written from
whoever calls the create branch.
batch_processed is set by callers that have already billed the batch
themselves, so CheckBatchCost skips the row instead of billing it twice.
It is written only in the upsert create branch.
"""
verbose_logger.info(f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache")
litellm_managed_object = LiteLLM_ManagedObjectTable(
@ -379,6 +394,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"updated_by": user_api_key_dict.user_id,
"status": file_object.status,
**attribution_columns,
"batch_processed": batch_processed,
},
"update": update_columns,
},
@ -1343,6 +1359,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
self, data: Dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
) -> LLMResponseTypes:
if isinstance(response, LiteLLMBatch):
decoded_batch_id: Final = _is_base64_encoded_unified_file_id(response.id)
if decoded_batch_id and is_litellm_executed_batch(decoded_batch_id):
return response
## Check if unified_file_id is in the response
unified_file_id = response._hidden_params.get("unified_file_id") # managed file id
unified_batch_id = response._hidden_params.get("unified_batch_id") # managed batch id
@ -1794,24 +1813,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Check if file deletion should be blocked due to batch references
await self._check_file_deletion_allowed(file_id)
# file_id = convert_b64_uid_to_unified_uid(file_id)
model_file_id_mapping = await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
if specific_model_file_id_mapping:
# Remove conflicting keys from data to avoid duplicate keyword arguments
filtered_data = {k: v for k, v in data.items() if k not in ("model", "file_id")}
for model_id, model_file_id in specific_model_file_id_mapping.items():
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
delete_data = {
**{k: v for k, v in filtered_data.items() if k != "_litellm_internal_model_credentials"},
**(
{"_litellm_internal_model_credentials": MappingProxyType(dict(credentials))}
if credentials is not None
else {}
),
}
await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
managed_file: Final = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
if managed_file is not None and managed_file.storage_backend and managed_file.storage_url:
await self._delete_storage_backend_content(managed_file.storage_backend, managed_file.storage_url)
else:
await self._delete_provider_files(file_id, litellm_parent_otel_span, llm_router, data)
await self.delete_unified_file_id(file_id, litellm_parent_otel_span)
@ -1820,16 +1826,53 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
prom_logger.record_managed_file_deleted(result="success")
return FileDeleted(id=file_id, object="file", deleted=True)
async def _delete_storage_backend_content(self, storage_backend_name: str, storage_url: str) -> None:
try:
storage_backend: Final = get_storage_backend(storage_backend_name, prisma_client=self.prisma_client)
except ValueError as e:
raise HTTPException(status_code=400, detail=f"Cannot delete the stored file content: {e}") from e
await storage_backend.delete_file(storage_url)
async def _delete_provider_files(
self,
file_id: str,
litellm_parent_otel_span: Span | None,
llm_router: Router,
data: Mapping[str, object],
) -> None:
model_file_id_mapping: Final = await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
specific_model_file_id_mapping: Final = model_file_id_mapping.get(file_id)
if not specific_model_file_id_mapping:
return
filtered_data: Final = {
k: v for k, v in data.items() if k not in ("model", "file_id", "_litellm_internal_model_credentials")
}
for model_id, model_file_id in specific_model_file_id_mapping.items():
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
delete_data = {
**filtered_data,
**(
{"_litellm_internal_model_credentials": MappingProxyType(dict(credentials))}
if credentials is not None
else {}
),
}
await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
async def afile_content(
self,
file_id: str,
litellm_parent_otel_span: Optional[Span],
llm_router: Router,
**data: Dict,
) -> "HttpxBinaryResponseContent":
) -> HttpxBinaryResponseContent:
"""
Get the content of a file from first model that has it
"""
managed_file: Final = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
if managed_file is not None and managed_file.storage_backend and managed_file.storage_url:
return await self._storage_backend_content(managed_file.storage_backend, managed_file.storage_url)
model_file_id_mapping = data.pop("model_file_id_mapping", None)
model_file_id_mapping = model_file_id_mapping or await self.get_model_file_id_mapping(
[file_id], litellm_parent_otel_span
@ -1859,6 +1902,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
else:
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
async def _storage_backend_content(self, storage_backend_name: str, storage_url: str) -> HttpxBinaryResponseContent:
storage_backend: Final = get_storage_backend(storage_backend_name, prisma_client=self.prisma_client)
content: Final = await storage_backend.download_file(storage_url)
return HttpxBinaryResponseContent(response=httpx.Response(status_code=httpx.codes.OK, content=content))
async def _convert_storage_files_to_base64(
self,
messages: List[AllMessageValues],
@ -1889,16 +1937,12 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# File is stored in a storage backend, download and convert to base64
try:
from litellm.llms.base_llm.files.storage_backend_factory import (
get_storage_backend,
)
storage_backend_name = db_file.storage_backend
storage_url = db_file.storage_url
# Get storage backend (uses same env vars as callback)
try:
storage_backend = get_storage_backend(storage_backend_name)
storage_backend = get_storage_backend(storage_backend_name, prisma_client=self.prisma_client)
except ValueError as e:
verbose_logger.warning(
f"Storage backend '{storage_backend_name}' error for file {file_id}: {str(e)}"

View file

@ -0,0 +1,8 @@
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_ManagedFileContentTable" (
"id" TEXT NOT NULL,
"content" BYTEA NOT NULL,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_ManagedFileContentTable_pkey" PRIMARY KEY ("id")
);

View file

@ -1107,6 +1107,12 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t
@@index([team_id, created_at(sort: Desc)])
}
model LiteLLM_ManagedFileContentTable {
id String @id @default(uuid())
content Bytes
created_at DateTime @default(now())
}
model LiteLLM_ManagedVectorStoreTable {
id String @id @default(uuid())
unified_resource_id String @unique // The base64 encoded unified vector store ID

View file

@ -1694,6 +1694,7 @@ LOGIN_THROTTLE_NOT_BLOCKED: Final = (0, 0)
LITELLM_PROXY_ADMIN_NAME: Final = "default_user_id"
LITELLM_PROXY_BUDGET_NAME: Final = "litellm-proxy-budget"
GLOBAL_PROXY_SPEND_CACHE_KEY: Final = f"{LITELLM_PROXY_ADMIN_NAME}:spend"
LITELLM_EXECUTED_BATCH_CONCURRENCY: Final = max(1, int(os.getenv("LITELLM_EXECUTED_BATCH_CONCURRENCY", "4")))
########################### CLI SSO AUTHENTICATION CONSTANTS ###########################
LITELLM_CLI_SOURCE_IDENTIFIER: Final = "litellm-cli"

View file

@ -0,0 +1,40 @@
from typing import TYPE_CHECKING, Final
from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend
from litellm.repositories.managed_file_content_repository import ManagedFileContentRepository
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
LITELLM_DB_STORAGE_BACKEND_NAME: Final = "litellm_db"
LITELLM_DB_STORAGE_URL_PREFIX: Final = f"{LITELLM_DB_STORAGE_BACKEND_NAME}://"
def storage_url_to_row_id(storage_url: str) -> str:
if not storage_url.startswith(LITELLM_DB_STORAGE_URL_PREFIX):
raise ValueError(f"Not a {LITELLM_DB_STORAGE_BACKEND_NAME} storage url: {storage_url}")
return storage_url.removeprefix(LITELLM_DB_STORAGE_URL_PREFIX)
class LiteLLMDbStorageBackend(BaseFileStorageBackend):
def __init__(self, prisma_client: "PrismaClient") -> None:
self._contents = ManagedFileContentRepository(prisma_client)
async def upload_file(
self,
file_content: bytes,
filename: str,
content_type: str,
path_prefix: str | None = None,
file_naming_strategy: str = "uuid",
) -> str:
return f"{LITELLM_DB_STORAGE_URL_PREFIX}{await self._contents.store(file_content)}"
async def download_file(self, storage_url: str) -> bytes:
content: Final = await self._contents.load(storage_url_to_row_id(storage_url))
if content is None:
raise ValueError(f"No stored file content for {storage_url}")
return content
async def delete_file(self, storage_url: str) -> None:
await self._contents.delete(storage_url_to_row_id(storage_url))

View file

@ -6,32 +6,46 @@ based on the backend type. Backends use the same configuration as their correspo
callbacks (e.g., azure_storage uses the same env vars as AzureBlobStorageLogger).
"""
from typing import TYPE_CHECKING
from litellm._logging import verbose_logger
from .azure_blob_storage_backend import AzureBlobStorageBackend
from .litellm_db_storage_backend import LITELLM_DB_STORAGE_BACKEND_NAME, LiteLLMDbStorageBackend
from .storage_backend import BaseFileStorageBackend
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
def get_storage_backend(backend_type: str) -> BaseFileStorageBackend:
def get_storage_backend(backend_type: str, prisma_client: "PrismaClient | None" = None) -> BaseFileStorageBackend:
"""
Factory function to create a storage backend instance.
Backends are configured using the same environment variables as their
corresponding callbacks. For example, "azure_storage" uses the same
env vars as AzureBlobStorageLogger.
env vars as AzureBlobStorageLogger. "litellm_db" stores file bytes in the
proxy's own database and needs the connected Prisma client.
Args:
backend_type: Backend type identifier (e.g., "azure_storage")
backend_type: Backend type identifier (e.g., "azure_storage", "litellm_db")
prisma_client: The proxy's database client, required by "litellm_db"
Returns:
BaseFileStorageBackend: Instance of the appropriate storage backend
Raises:
ValueError: If backend_type is not supported
ValueError: If backend_type is not supported, or "litellm_db" is asked for without a database
"""
verbose_logger.debug("Creating storage backend: type=%s", backend_type)
if backend_type == "azure_storage":
return AzureBlobStorageBackend()
else:
raise ValueError(f"Unsupported storage backend type: {backend_type}. Supported types: azure_storage")
if backend_type == LITELLM_DB_STORAGE_BACKEND_NAME:
if prisma_client is None:
raise ValueError(f"Storage backend {LITELLM_DB_STORAGE_BACKEND_NAME} requires a database-connected proxy")
return LiteLLMDbStorageBackend(prisma_client)
raise ValueError(
f"Unsupported storage backend type: {backend_type}. "
f"Supported types: azure_storage, {LITELLM_DB_STORAGE_BACKEND_NAME}"
)

View file

@ -7,10 +7,12 @@
import asyncio
import os
from collections.abc import Mapping
from datetime import datetime
from types import MappingProxyType
from typing import Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response
from pydantic import TypeAdapter
import litellm
from litellm._logging import verbose_proxy_logger
@ -18,6 +20,15 @@ from litellm.batches.main import CancelBatchRequest, RetrieveBatchRequest
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.batches_endpoints.common_utils import validate_batch_list_limit
from litellm.proxy.batches_endpoints.litellm_executed_batches import (
LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE,
LiteLLMExecutedBatchRunner,
ManagedBatchStore,
batch_error,
executed_batch_runner_lost,
litellm_executed_provider_for,
resolve_litellm_executed_provider,
)
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
log_llm_api_exception,
@ -47,16 +58,87 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
get_model_id_from_unified_batch_id,
get_models_from_unified_file_id,
get_original_file_id,
is_litellm_executed_batch,
prepare_data_with_credentials,
update_batch_in_database,
validate_managed_id_requirement,
)
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import request_tags_from_metadata
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
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.router import Router
from litellm.types.llms.openai import LiteLLMBatchCreateRequest
from litellm.types.utils import LiteLLMBatch
if TYPE_CHECKING:
from prisma.models import LiteLLM_ManagedObjectTable
router: Final = APIRouter()
_METADATA_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
def _request_tags(data: Mapping[str, object]) -> tuple[str, ...] | None:
metadata: Final = data.get("litellm_metadata")
if metadata is None:
return None
return request_tags_from_metadata(_METADATA_ADAPTER.validate_python(metadata))
def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyLogging) -> LiteLLMExecutedBatchRunner:
from litellm.proxy.proxy_server import general_settings, prisma_client
managed_files: Final = proxy_logging_obj.get_proxy_hook("managed_files")
if prisma_client is None or not isinstance(managed_files, ManagedBatchStore):
raise batch_error(
400,
"LiteLLM-executed batches need a database: set DATABASE_URL so LiteLLM can keep the batch and its files",
)
return LiteLLMExecutedBatchRunner(
llm_router=llm_router,
prisma_client=prisma_client,
managed_files=managed_files,
batches=ManagedBatchRepository(prisma_client),
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
)
async def _batch_from_database(
batch_id: str,
unified_batch_id: str | Literal[False],
executed_batch: bool,
managed_files_obj: object,
prisma_client: PrismaClient | None,
llm_router: Router | None,
proxy_logging_obj: ProxyLogging,
user_api_key_dict: UserAPIKeyAuth,
) -> tuple["LiteLLM_ManagedObjectTable | None", LiteLLMBatch | None]:
row, batch = await get_batch_from_database(
batch_id=batch_id,
unified_batch_id=unified_batch_id,
managed_files_obj=managed_files_obj,
prisma_client=prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
)
updated_at: Final[object] = getattr(row, "updated_at", None)
if not executed_batch or batch is None or llm_router is None or not isinstance(updated_at, datetime):
return row, batch
if not executed_batch_runner_lost(batch.status, updated_at):
return row, batch
runner: Final = _litellm_executed_batch_runner(llm_router, proxy_logging_obj)
return row, await runner.fail_abandoned(batch, user_api_key_dict)
async def _raise_when_input_file_must_be_managed(model: str, credentials: Mapping[str, object]) -> None:
if await litellm_executed_provider_for(credentials) is None:
return
raise batch_error(
400,
f"Batches for {model} run inside LiteLLM, so the input file must be a LiteLLM managed file: "
f"{LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE}",
)
def _raise_not_found_when_openai_fallback_unservable(
@ -101,6 +183,24 @@ async def _resolve_managed_input_file_storage_url(input_file_id: str) -> "str |
return db_file.storage_url or None
async def _create_provider_batch_for_managed_file(
llm_router: Router,
create_batch_data: LiteLLMBatchCreateRequest,
input_file_id: str,
unified_file_id: str,
) -> LiteLLMBatch:
resolved_storage_url: Final = await _resolve_managed_input_file_storage_url(input_file_id)
request: Final[LiteLLMBatchCreateRequest] = {
**create_batch_data,
"input_file_id": resolved_storage_url or input_file_id,
"disable_fallbacks": True,
}
response: Final = await llm_router.acreate_batch(**request)
response.input_file_id = input_file_id
response._hidden_params["unified_file_id"] = unified_file_id
return response
@router.post(
"/{provider}/v1/batches",
dependencies=[Depends(user_api_key_auth)],
@ -296,24 +396,35 @@ async def create_batch(
await authorize_model_for_key(model_id=model, llm_router=llm_router, user_api_key_dict=user_api_key_dict)
_create_batch_data["model"] = model
resolved_storage_url: Final = await _resolve_managed_input_file_storage_url(input_file_id)
if resolved_storage_url is not None:
_create_batch_data["input_file_id"] = resolved_storage_url
if llm_router is None:
raise HTTPException(
status_code=500,
detail={"error": "LLM Router not initialized. Ensure models added to proxy."},
)
_create_batch_data.update(disable_fallbacks=True) # pyright: ignore[reportCallIssue] # router flag
response = await llm_router.acreate_batch(**_create_batch_data)
response.input_file_id = input_file_id
response._hidden_params["unified_file_id"] = unified_file_id
executed_provider: Final = await resolve_litellm_executed_provider(
llm_router, model, user_api_key_dict.team_id
)
response = (
await _litellm_executed_batch_runner(llm_router, proxy_logging_obj).create(
create_request=_create_batch_data,
unified_input_file_id=input_file_id,
model=model,
provider=executed_provider,
user_api_key_dict=user_api_key_dict,
request_tags=_request_tags(_create_batch_data),
)
if executed_provider is not None
else await _create_provider_batch_for_managed_file(
llm_router, _create_batch_data, input_file_id, unified_file_id
)
)
else:
# Check if model specified via header/query/body param
model_param: Final = (
data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model")
_create_batch_data.get("model")
or request.query_params.get("model")
or request.headers.get("x-litellm-model")
)
# SCENARIO 2 & 3: Model from header/query OR custom_llm_provider fallback
@ -325,6 +436,7 @@ async def create_batch(
user_api_key_dict=user_api_key_dict,
operation_context="batch creation",
)
await _raise_when_input_file_must_be_managed(model_param, credentials)
prepare_data_with_credentials(
data=_create_batch_data,
@ -486,23 +598,26 @@ async def retrieve_batch(
managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files")
from litellm.proxy.proxy_server import prisma_client
db_batch_object, response = await get_batch_from_database(
executed_batch: Final = isinstance(unified_batch_id, str) and is_litellm_executed_batch(unified_batch_id)
db_batch_object, response = await _batch_from_database(
batch_id=batch_id,
unified_batch_id=unified_batch_id,
executed_batch=executed_batch,
managed_files_obj=managed_files_obj,
prisma_client=prisma_client,
verbose_proxy_logger=verbose_proxy_logger,
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
)
if executed_batch and response is None:
raise batch_error(404, f"No batch found with id '{batch_id}'.")
# If batch is in a terminal state, return immediately.
# Include "complete" (DB-normalized form of "completed").
if response is not None and response.status in [
"completed",
"complete",
"failed",
"cancelled",
"expired",
]:
if response is not None and (
response.status in ("completed", "complete", "failed", "cancelled", "expired") or executed_batch
):
# Call hooks and return
response = await proxy_logging_obj.post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
@ -978,6 +1093,17 @@ async def cancel_batch(
proxy_config=proxy_config,
)
unified_model_id: Final = get_model_id_from_unified_batch_id(unified_batch_id) if unified_batch_id else None
if unified_model_id is not None:
resolved_unified_model: Final = (
llm_router.resolve_model_name_from_model_id(unified_model_id) if llm_router is not None else None
)
await authorize_model_for_key(
model_id=resolved_unified_model or unified_model_id,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
)
# SCENARIO 1: Batch ID is encoded with model info
if model_from_id is not None:
credentials: Final = await get_authorized_credentials_for_model(
@ -1009,6 +1135,12 @@ async def cancel_batch(
)
# SCENARIO 2: target_model_names based routing
elif unified_batch_id and is_litellm_executed_batch(unified_batch_id):
if llm_router is None:
raise batch_error(500, "LLM Router not initialized. Ensure models added to proxy.")
response = await _litellm_executed_batch_runner( # rebind-ok: each cancel path sets the route's response
llm_router, proxy_logging_obj
).cancel(batch_id, user_api_key_dict)
elif unified_batch_id:
if llm_router is None:
raise HTTPException(
@ -1022,11 +1154,6 @@ async def cancel_batch(
status_code=400,
detail={"error": "Invalid LiteLLM managed batch ID. Missing model_id."},
)
await authorize_model_for_key(
model_id=llm_router.resolve_model_name_from_model_id(model_id_from_batch) or model_id_from_batch,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
)
data["model"] = model_id_from_batch
data["batch_id"] = get_batch_id_from_unified_batch_id(unified_batch_id)
response = await llm_router.acancel_batch(**data)

View file

@ -0,0 +1,716 @@
import asyncio
import json
import time
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from itertools import pairwise
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable
import httpx
from openai.types.batch import Errors
from openai.types.batch_error import BatchError
from openai.types.batch_request_counts import BatchRequestCounts
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid as uuid_module
from litellm.constants import LITELLM_EXECUTED_BATCH_CONCURRENCY
from litellm.integrations.prometheus import PrometheusLogger
from litellm.llms.base_llm.files.litellm_db_storage_backend import LITELLM_DB_STORAGE_BACKEND_NAME
from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.models.managed_files import LiteLLM_ManagedFileTable
from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import is_request_body_safe
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.openai_files_endpoints.common_utils import LITELLM_EXECUTED_BATCH_ID_PREFIX
from litellm.proxy.openai_files_endpoints.storage_backend_service import StorageBackendFileService
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.managed_batch_repository import ManagedBatchRepository
from litellm.types.llms.openai import LiteLLMBatchCreateRequest, OpenAIFileObject, OpenAIFilesPurpose
from litellm.types.utils import LITELLM_EXECUTED_BATCH_PROVIDERS, ExtractedFileData, LiteLLMBatch, LlmProviders
if TYPE_CHECKING:
from prisma import types as prisma_types
from litellm.router import Router
BatchEndpoint: TypeAlias = Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"]
BatchStatus: TypeAlias = Literal[
"in_progress", "finalizing", "completed", "failed", "cancelling", "cancelled", "expired"
]
TERMINAL_BATCH_STATUSES: Final[frozenset[str]] = frozenset({"completed", "failed", "cancelled", "expired"})
_STOP_STATUSES: Final[frozenset[str]] = TERMINAL_BATCH_STATUSES | frozenset({"cancelling"})
_BATCH_ENDPOINT_ADAPTER: Final[TypeAdapter[BatchEndpoint]] = TypeAdapter(BatchEndpoint)
_CANCEL_POLL_SECONDS: Final = 1.0
_HEARTBEAT_SECONDS: Final = 30.0
_STALE_AFTER_SECONDS: Final = 180.0
_FILES_API_PROBE_TIMEOUT_SECONDS: Final = 5.0
_COMPLETION_WINDOW_SECONDS: Final = 24 * 60 * 60
_RUNNER_LOST_MESSAGE: Final = "the proxy replica running this batch stopped before it finished; resubmit the batch"
_EXPIRED_MESSAGE: Final = "This request could not be executed before the completion window expired."
_ROUTER_METHODS: Final[Mapping[BatchEndpoint, str]] = MappingProxyType(
{
"/v1/chat/completions": "acompletion",
"/v1/completions": "atext_completion",
"/v1/embeddings": "aembedding",
"/v1/responses": "aresponses",
}
)
_CANCELLING_TRANSITIONS: Final[Mapping[BatchStatus, BatchStatus]] = MappingProxyType(
{"completed": "cancelled", "expired": "cancelled", "in_progress": "cancelling", "finalizing": "cancelling"}
)
LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE: Final = (
"upload it through POST /v1/files with purpose=batch and either the x-litellm-model header or the "
"target_model_names form field naming the model, so LiteLLM keeps the file and runs the batch itself"
)
_RUNNING_BATCHES: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong references keep running batch tasks alive
_NO_FIELDS: Final[Mapping[str, object]] = MappingProxyType({})
_NO_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
class _ErrorDetail(TypedDict):
message: ReadOnly[str]
type: ReadOnly[str]
param: ReadOnly[None]
code: ReadOnly[None]
class _ErrorBody(TypedDict):
error: ReadOnly[_ErrorDetail]
class _ResultResponse(TypedDict):
status_code: ReadOnly[int]
request_id: ReadOnly[str]
body: ReadOnly[Mapping[str, object]]
class _LineError(TypedDict):
code: ReadOnly[str]
message: ReadOnly[str]
class _ResultLine(TypedDict):
id: ReadOnly[str]
custom_id: ReadOnly[str]
response: ReadOnly[_ResultResponse | None]
error: ReadOnly[_LineError | None]
class BatchInputLine(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
custom_id: str
method: Literal["POST"]
url: str
body: Mapping[str, object]
@dataclass(frozen=True, slots=True)
class InvalidBatchInput:
line_number: int | None
reason: str
def describe(self) -> str:
return f"line {self.line_number}: {self.reason}" if self.line_number is not None else self.reason
@dataclass(frozen=True, slots=True)
class RowOutcome:
custom_id: str
status_code: int
body: Mapping[str, object]
succeeded: bool
@dataclass(frozen=True, slots=True)
class ExpiredRow:
custom_id: str
@dataclass(frozen=True, slots=True)
class _BatchRun:
unified_batch_id: str
llm_batch_id: str
model: str
endpoint: BatchEndpoint
lines: tuple[BatchInputLine, ...]
user_api_key_dict: UserAPIKeyAuth
request_tags: tuple[str, ...]
deadline: float
@runtime_checkable
class ManagedBatchStore(Protocol):
def get_unified_batch_id(self, batch_id: str, model_id: str) -> str: ...
async def get_unified_file_id(
self, file_id: str, litellm_parent_otel_span: object | None = None
) -> LiteLLM_ManagedFileTable | None: ...
async def store_unified_object_id(
self,
unified_object_id: str,
file_object: LiteLLMBatch,
litellm_parent_otel_span: object | None,
model_object_id: str,
file_purpose: Literal["batch", "fine-tune", "response"],
user_api_key_dict: UserAPIKeyAuth,
request_tags: Sequence[str] | None = None,
persist_attribution: bool = False,
batch_processed: bool = False,
) -> None: ...
class _StorageBackendFactory(Protocol):
def __call__(self, backend_type: str, prisma_client: PrismaClient | None = None) -> BaseFileStorageBackend: ...
class _ResultFileUploader(Protocol):
def __call__(
self,
file_data: Mapping[str, object],
target_storage: str,
target_model_names: Sequence[str],
purpose: OpenAIFilesPurpose,
proxy_logging_obj: ProxyLogging,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient | None = None,
) -> Awaitable[OpenAIFileObject]: ...
@runtime_checkable
class _RouterCall(Protocol):
def __call__(self, **params: object) -> Awaitable[object]: ... # kwargs-ok: the request body is passed as keywords
def litellm_executed_provider_of(credentials: Mapping[str, object]) -> str | None:
explicit_provider: Final = credentials.get("custom_llm_provider")
provider: Final = (
explicit_provider if isinstance(explicit_provider, str) else _provider_of(credentials.get("model"))
)
return provider if provider in LITELLM_EXECUTED_BATCH_PROVIDERS else None
class _HttpGetter(Protocol):
async def get(
self, url: str, *, headers: dict[str, str] | None = None, timeout: float | httpx.Timeout | None = None
) -> httpx.Response: ...
class FilesApiProbe(Protocol):
async def __call__(self, api_base: str, api_key: str | None) -> bool: ...
class BodyRejection(Protocol):
def __call__(self, body: Mapping[str, object], /) -> str | None: ...
async def upstream_lacks_files_api(api_base: str, api_key: str | None, http_client: _HttpGetter | None = None) -> bool:
client: Final = http_client or get_async_httpx_client(llm_provider=LlmProviders.HOSTED_VLLM)
try:
response: Final = await client.get(
f"{api_base.rstrip('/')}/files",
headers=(
{"Authorization": f"Bearer {api_key}"} # mutable-ok: AsyncHTTPHandler.get wants a plain dict
if api_key
else None
),
timeout=_FILES_API_PROBE_TIMEOUT_SECONDS,
)
except httpx.HTTPError:
return False
return response.status_code == httpx.codes.NOT_FOUND
def _upstream_of(credentials: Mapping[str, object], provider: str) -> tuple[str, str | None] | None:
model: Final = credentials.get("model")
api_base: Final = credentials.get("api_base")
api_key: Final = credentials.get("api_key")
if not isinstance(model, str):
return None
try:
_, _, resolved_api_key, resolved_api_base = litellm.get_llm_provider(
model=model,
custom_llm_provider=provider,
api_base=api_base if isinstance(api_base, str) else None,
api_key=api_key if isinstance(api_key, str) else None,
)
except Exception: # noqa: BLE001 # get_llm_provider raises on a model it cannot map, which means nothing to probe
return None
return None if resolved_api_base is None else (resolved_api_base, resolved_api_key)
async def litellm_executed_provider_for(
credentials: Mapping[str, object], lacks_files_api: FilesApiProbe = upstream_lacks_files_api
) -> str | None:
provider: Final = litellm_executed_provider_of(credentials)
if provider is None:
return None
upstream: Final = _upstream_of(credentials, provider)
if upstream is None:
return None
return provider if await lacks_files_api(*upstream) else None
async def resolve_litellm_executed_provider(
llm_router: "Router",
model: str,
team_id: str | None,
lacks_files_api: FilesApiProbe = upstream_lacks_files_api,
) -> str | None:
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model, team_id=team_id)
return None if credentials is None else await litellm_executed_provider_for(credentials, lacks_files_api)
def _provider_of(model: object) -> str | None:
if not isinstance(model, str):
return None
try:
return litellm.get_llm_provider(model=model)[1]
except Exception: # noqa: BLE001 # get_llm_provider raises on an unknown model, which means no provider
return None
def _validation_reason(error: ValidationError) -> str:
return "; ".join(
f"{'.'.join(str(part) for part in item['loc'])}: {item['msg']}" if item["loc"] else item["msg"]
for item in error.errors()
)
def _accept_every_body(_body: Mapping[str, object]) -> str | None:
return None
def _parse_line(
line_number: int, raw: bytes, endpoint: BatchEndpoint, reject_body: BodyRejection
) -> BatchInputLine | InvalidBatchInput:
try:
line: Final = BatchInputLine.model_validate_json(raw)
except ValidationError as e:
return InvalidBatchInput(line_number, _validation_reason(e))
if line.url != endpoint:
return InvalidBatchInput(line_number, f"url {line.url!r} does not match the batch endpoint {endpoint!r}")
if line.body.get("stream"):
return InvalidBatchInput(line_number, "streaming requests are not supported in a batch")
rejection: Final = reject_body(line.body)
if rejection is not None:
return InvalidBatchInput(line_number, rejection)
return line
def parse_batch_input(
content: bytes, endpoint: BatchEndpoint, reject_body: BodyRejection = _accept_every_body
) -> tuple[BatchInputLine, ...] | InvalidBatchInput:
raw_lines: Final = tuple((number, raw) for number, raw in enumerate(content.splitlines(), start=1) if raw.strip())
if not raw_lines:
return InvalidBatchInput(None, "the input file has no requests")
parsed: Final = tuple(_parse_line(number, raw, endpoint, reject_body) for number, raw in raw_lines)
first_invalid: Final = next((item for item in parsed if isinstance(item, InvalidBatchInput)), None)
if first_invalid is not None:
return first_invalid
lines: Final = tuple(item for item in parsed if isinstance(item, BatchInputLine))
custom_ids: Final = sorted(line.custom_id for line in lines)
duplicate: Final = next((first for first, second in pairwise(custom_ids) if first == second), None)
if duplicate is not None:
return InvalidBatchInput(None, f"custom_id {duplicate!r} is used more than once")
return lines
def batch_error(status_code: int, message: str) -> ProxyException:
error_type: Final = "invalid_request_error" if status_code < 500 else ProxyErrorTypes.internal_server_error.value
return ProxyException(message=message, type=error_type, param=None, code=status_code)
def _validate_endpoint(endpoint: object) -> BatchEndpoint:
try:
return _BATCH_ENDPOINT_ADAPTER.validate_python(endpoint)
except ValidationError:
raise batch_error(400, f"endpoint {endpoint!r} is not supported for a LiteLLM-executed batch")
def _status_code_of(error: Exception) -> int:
status_code: Final[object] = getattr(error, "status_code", None)
return status_code if isinstance(status_code, int) else 500
def _error_body(error: Exception) -> _ErrorBody:
body: Final[_ErrorBody] = {
"error": {"message": str(error), "type": type(error).__name__, "param": None, "code": None}
}
return body
def _line_response(outcome: RowOutcome | ExpiredRow) -> _ResultResponse | None:
if isinstance(outcome, ExpiredRow):
return None
response: Final[_ResultResponse] = {
"status_code": outcome.status_code,
"request_id": f"req_{uuid_module.uuid4().hex[:24]}",
"body": outcome.body,
}
return response
def _line_error(outcome: RowOutcome | ExpiredRow) -> _LineError | None:
if isinstance(outcome, RowOutcome):
return None
error: Final[_LineError] = {"code": "batch_expired", "message": _EXPIRED_MESSAGE}
return error
def _result_line(outcome: RowOutcome | ExpiredRow) -> _ResultLine:
line: Final[_ResultLine] = {
"id": f"batch_req_{uuid_module.uuid4().hex[:24]}",
"custom_id": outcome.custom_id,
"response": _line_response(outcome),
"error": _line_error(outcome),
}
return line
def _dump(response: object) -> Mapping[str, object]:
if isinstance(response, BaseModel):
return response.model_dump(mode="json")
raise TypeError(f"Batch rows must return a single response object, got {type(response).__name__}")
def _resolve_transition(current_status: str, requested: BatchStatus) -> BatchStatus:
if current_status != "cancelling":
return requested
return _CANCELLING_TRANSITIONS.get(requested, requested)
def executed_batch_runner_lost(status: str, updated_at: datetime) -> bool:
if status in TERMINAL_BATCH_STATUSES:
return False
return (datetime.now(timezone.utc) - updated_at).total_seconds() > _STALE_AFTER_SECONDS
class _StopWatch:
def __init__(self, load_status: Callable[[], Awaitable[str | None]], interval_seconds: float) -> None:
self._load_status = load_status
self._interval_seconds = interval_seconds
self._checked_at = float("-inf")
self._stopped = False
async def stopped(self) -> bool:
if self._stopped:
return True
now: Final = time.monotonic()
if now - self._checked_at < self._interval_seconds:
return False
self._checked_at = now
self._stopped = await self._load_status() in _STOP_STATUSES
return self._stopped
class LiteLLMExecutedBatchRunner:
def __init__(
self,
llm_router: "Router",
prisma_client: PrismaClient,
managed_files: ManagedBatchStore,
batches: ManagedBatchRepository,
proxy_logging_obj: ProxyLogging,
general_settings: Mapping[str, object],
concurrency: int = LITELLM_EXECUTED_BATCH_CONCURRENCY,
heartbeat_seconds: float = _HEARTBEAT_SECONDS,
completion_window_seconds: float = _COMPLETION_WINDOW_SECONDS,
storage_backend_factory: _StorageBackendFactory = get_storage_backend,
upload_result_file: _ResultFileUploader = StorageBackendFileService.upload_file_to_storage_backend,
) -> None:
self.llm_router = llm_router
self.prisma_client = prisma_client
self.managed_files = managed_files
self.batches = batches
self.proxy_logging_obj = proxy_logging_obj
self.general_settings = general_settings
self.concurrency = concurrency
self.heartbeat_seconds = heartbeat_seconds
self.completion_window_seconds = completion_window_seconds
self.storage_backend_factory = storage_backend_factory
self.upload_result_file = upload_result_file
async def create(
self,
create_request: LiteLLMBatchCreateRequest,
unified_input_file_id: str,
model: str,
provider: str,
user_api_key_dict: UserAPIKeyAuth,
request_tags: Sequence[str] | None,
) -> LiteLLMBatch:
endpoint: Final = _validate_endpoint(create_request.get("endpoint"))
content: Final = await self._download_input(unified_input_file_id, user_api_key_dict)
parsed: Final = parse_batch_input(content, endpoint, self._body_rejection(model))
if isinstance(parsed, InvalidBatchInput):
raise batch_error(400, f"Invalid batch input file: {parsed.describe()}")
llm_batch_id: Final = f"{LITELLM_EXECUTED_BATCH_ID_PREFIX}{uuid_module.uuid4().hex}"
model_id: Final = next(iter(self.llm_router.get_model_ids(model_name=model)), model)
unified_batch_id: Final = self.managed_files.get_unified_batch_id(batch_id=llm_batch_id, model_id=model_id)
now: Final = time.time()
created_at: Final = int(now)
batch: Final = LiteLLMBatch(
id=unified_batch_id,
object="batch",
endpoint=endpoint,
input_file_id=unified_input_file_id,
completion_window="24h",
status="validating",
created_at=created_at,
expires_at=created_at + int(self.completion_window_seconds),
metadata=create_request.get("metadata"),
model=model,
request_counts=BatchRequestCounts(completed=0, failed=0, total=len(parsed)),
)
await self.managed_files.store_unified_object_id(
unified_object_id=unified_batch_id,
file_object=batch,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
model_object_id=llm_batch_id,
file_purpose="batch",
user_api_key_dict=user_api_key_dict,
request_tags=request_tags,
persist_attribution=True,
batch_processed=True,
)
_record_batch_created(model, provider, user_api_key_dict)
run: Final = _BatchRun(
unified_batch_id=unified_batch_id,
llm_batch_id=llm_batch_id,
model=model,
endpoint=endpoint,
lines=parsed,
user_api_key_dict=user_api_key_dict,
request_tags=tuple(request_tags or ()),
deadline=now + self.completion_window_seconds,
)
task: Final = asyncio.create_task(self._run(run))
_RUNNING_BATCHES.add(task)
task.add_done_callback(_RUNNING_BATCHES.discard)
return batch
async def cancel(self, unified_batch_id: str, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch:
current: Final = await self.batches.load_batch(unified_batch_id)
if current is None:
raise batch_error(404, f"Batch {unified_batch_id} not found")
if current.status in TERMINAL_BATCH_STATUSES:
raise batch_error(400, f"Cannot cancel a batch with status '{current.status}'")
if current.status == "cancelling":
return current
cancelling: Final = current.model_copy(
update=MappingProxyType({"status": "cancelling", "cancelling_at": int(time.time())})
)
unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status}
if await self.batches.compare_and_set(cancelling, unchanged, user_api_key_dict.user_id):
return cancelling
return await self.cancel(unified_batch_id, user_api_key_dict)
async def fail_abandoned(self, batch: LiteLLMBatch, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch:
error: Final = BatchError(message=_RUNNER_LOST_MESSAGE, code="runner_lost")
errors: Final = Errors(data=[error], object="list") # mutable-ok: Errors.data is typed as a list
failed: Final = batch.model_copy(
update=MappingProxyType({"status": "failed", "failed_at": int(time.time()), "errors": errors})
)
untouched: Final[prisma_types.DateTimeFilter] = {
"lt": datetime.now(timezone.utc) - timedelta(seconds=_STALE_AFTER_SECONDS)
}
still_abandoned: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {
"status": batch.status,
"updated_at": untouched,
}
if await self.batches.compare_and_set(failed, still_abandoned, user_api_key_dict.user_id):
return failed
return await self.batches.load_batch(batch.id) or batch
def _body_rejection(self, model: str) -> BodyRejection:
def reject(body: Mapping[str, object]) -> str | None:
try:
is_request_body_safe(
request_body=dict(body), # mutable-ok: is_request_body_safe takes a dict
general_settings=dict(self.general_settings), # mutable-ok: is_request_body_safe takes a dict
llm_router=self.llm_router,
model=model,
)
except ValueError as e:
return str(e)
return None
return reject
async def _download_input(self, unified_input_file_id: str, user_api_key_dict: UserAPIKeyAuth) -> bytes:
stored: Final = await self.managed_files.get_unified_file_id(
unified_input_file_id, litellm_parent_otel_span=user_api_key_dict.parent_otel_span
)
if stored is None or not stored.storage_backend or not stored.storage_url:
raise batch_error(
400,
f"LiteLLM does not hold the content of input file {unified_input_file_id}: "
f"{LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE}",
)
try:
backend: Final = self.storage_backend_factory(stored.storage_backend, prisma_client=self.prisma_client)
return await backend.download_file(stored.storage_url)
except ValueError as e:
raise batch_error(400, str(e))
async def _run(self, run: _BatchRun) -> None:
heartbeat: Final = asyncio.create_task(self._heartbeat(run))
try:
await self._execute(run)
except Exception as e: # noqa: BLE001 # whatever fails, the batch must end up marked failed
verbose_proxy_logger.exception("LiteLLM-executed batch %s failed: %s", run.unified_batch_id, e)
error: Final = BatchError(message=str(e), code="internal_error")
errors: Final = Errors(data=[error], object="list") # mutable-ok: Errors.data is typed as a list
try:
await self._advance(run, "failed", MappingProxyType({"errors": errors}))
except Exception as advance_error: # noqa: BLE001 # a failed status write is logged, never raised
verbose_proxy_logger.exception(
"LiteLLM-executed batch %s could not be marked failed: %s", run.unified_batch_id, advance_error
)
finally:
heartbeat.cancel()
async def _heartbeat(self, run: _BatchRun) -> None:
while True:
await asyncio.sleep(self.heartbeat_seconds)
try:
await self._touch(run)
except Exception as e: # noqa: BLE001 # a missed beat is logged and the next one retries
verbose_proxy_logger.warning("LiteLLM-executed batch %s heartbeat failed: %s", run.unified_batch_id, e)
async def _touch(self, run: _BatchRun) -> None:
await self.batches.touch(run.unified_batch_id, run.user_api_key_dict.user_id)
async def _execute(self, run: _BatchRun) -> None:
await self._advance(run, "in_progress")
watch: Final = _StopWatch(lambda: self.batches.load_status(run.unified_batch_id), _CANCEL_POLL_SECONDS)
semaphore: Final = asyncio.Semaphore(self.concurrency)
results: Final = await asyncio.gather(*(self._run_row(run, line, watch, semaphore) for line in run.lines))
outcomes: Final = tuple(outcome for outcome in results if outcome is not None)
if await self._advance(run, "finalizing") is None:
return
succeeded: Final = tuple(
outcome for outcome in outcomes if isinstance(outcome, RowOutcome) and outcome.succeeded
)
failed: Final = tuple(
outcome for outcome in outcomes if isinstance(outcome, ExpiredRow) or not outcome.succeeded
)
output_file_id: Final = await self._upload_results(run, "output", succeeded)
error_file_id: Final = await self._upload_results(run, "error", failed)
request_counts: Final = BatchRequestCounts(completed=len(succeeded), failed=len(failed), total=len(run.lines))
final_status: Final[BatchStatus] = (
"expired" if any(isinstance(outcome, ExpiredRow) for outcome in outcomes) else "completed"
)
await self._advance(
run,
final_status,
MappingProxyType(
{"output_file_id": output_file_id, "error_file_id": error_file_id, "request_counts": request_counts}
),
)
async def _run_row(
self, run: _BatchRun, line: BatchInputLine, watch: _StopWatch, semaphore: asyncio.Semaphore
) -> RowOutcome | ExpiredRow | None:
async with semaphore:
if await watch.stopped():
return None
remaining: Final = run.deadline - time.time()
if remaining <= 0:
return ExpiredRow(custom_id=line.custom_id)
try:
return await asyncio.wait_for(self._row_outcome(run, line), timeout=remaining)
except asyncio.TimeoutError:
return ExpiredRow(custom_id=line.custom_id)
async def _row_outcome(self, run: _BatchRun, line: BatchInputLine) -> RowOutcome:
try:
body: Final = await self._dispatch(run, line)
except Exception as e: # noqa: BLE001 # a provider error becomes the row's error line, never a crashed batch
return RowOutcome(
custom_id=line.custom_id, status_code=_status_code_of(e), body=_error_body(e), succeeded=False
)
return RowOutcome(custom_id=line.custom_id, status_code=200, body=body, succeeded=True)
async def _dispatch(self, run: _BatchRun, line: BatchInputLine) -> Mapping[str, object]:
params: Final = MappingProxyType(
{**line.body, "model": run.model, "metadata": self._row_metadata(run), "disable_fallbacks": True}
)
return _dump(await self._router_call(run.endpoint)(**params))
def _router_call(self, endpoint: BatchEndpoint) -> _RouterCall:
method: Final[object] = getattr(self.llm_router, _ROUTER_METHODS[endpoint], None)
if not isinstance(method, _RouterCall):
raise TypeError(f"the router has no callable for {endpoint}")
return method
def _row_metadata(self, run: _BatchRun) -> dict[str, object]: # mutable-ok: router updates metadata in place
return { # mutable-ok: the router updates request metadata in place
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(run.user_api_key_dict),
"user_api_key": run.user_api_key_dict.api_key,
"user_api_end_user_max_budget": run.user_api_key_dict.end_user_max_budget,
"tags": list(run.request_tags), # mutable-ok: litellm types request tags as a list
"batch_id": run.unified_batch_id,
}
async def _upload_results(
self, run: _BatchRun, kind: Literal["output", "error"], outcomes: Sequence[RowOutcome | ExpiredRow]
) -> str | None:
if not outcomes:
return None
content: Final = "".join(f"{json.dumps(_result_line(outcome))}\n" for outcome in outcomes).encode()
file_data: Final[ExtractedFileData] = {
"filename": f"{run.llm_batch_id}_{kind}.jsonl",
"content": content,
"content_type": "application/jsonl",
"headers": _NO_HEADERS,
}
file_object: Final = await self.upload_result_file(
file_data=file_data,
target_storage=LITELLM_DB_STORAGE_BACKEND_NAME,
target_model_names=(run.model,),
purpose="batch_output",
proxy_logging_obj=self.proxy_logging_obj,
user_api_key_dict=run.user_api_key_dict,
prisma_client=self.prisma_client,
)
return file_object.id
async def _advance(
self, run: _BatchRun, requested: BatchStatus, fields: Mapping[str, object] = _NO_FIELDS
) -> BatchStatus | None:
current: Final = await self.batches.load_batch(run.unified_batch_id)
if current is None:
raise RuntimeError(f"Batch {run.unified_batch_id} is no longer stored")
if current.status in TERMINAL_BATCH_STATUSES:
return None
status: Final = _resolve_transition(current.status, requested)
updated: Final = current.model_copy(
update=MappingProxyType({**fields, "status": status, f"{status}_at": int(time.time())})
)
unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status}
if await self.batches.compare_and_set(updated, unchanged, run.user_api_key_dict.user_id):
return status
return await self._advance(run, requested, fields)
def _record_batch_created(model: str, provider: str, user_api_key_dict: UserAPIKeyAuth) -> None:
prometheus_logger: Final = PrometheusLogger.get_instance()
if prometheus_logger is None:
return
prometheus_logger.record_managed_batch_created(
model=model,
api_provider=provider,
user=user_api_key_dict.user_id or "",
user_email=user_api_key_dict.user_email or "",
api_key_alias=user_api_key_dict.key_alias or "",
)

View file

@ -38,6 +38,7 @@ if TYPE_CHECKING:
FILE_LIST_CONTINUATION_CHUNK_SIZE: Final = 500
BATCH_CREATE_HIDDEN_PARAM: Final = "batch_create"
LITELLM_EXECUTED_BATCH_ID_PREFIX: Final = "litellm_batch_"
def validate_file_list_limit(limit: int | None) -> None:
@ -179,6 +180,11 @@ def get_batch_id_from_unified_batch_id(file_id: str) -> str:
return re.split(r"[;,]", batch_id, maxsplit=1)[0]
def is_litellm_executed_batch(decoded_unified_batch_id: str) -> bool:
_, marker, batch_id = decoded_unified_batch_id.partition("llm_batch_id:")
return bool(marker) and batch_id.startswith(LITELLM_EXECUTED_BATCH_ID_PREFIX)
def encode_file_id_with_model(file_id: str, model: str, id_type: Literal["file", "batch"] = "file") -> str:
"""
Encode a file/batch ID with model routing information.

View file

@ -7,7 +7,7 @@
import asyncio
import traceback
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from typing import Any, BinaryIO, Final, TypedDict, cast, get_args
import httpx
@ -32,10 +32,15 @@ from litellm.litellm_core_utils.cloud_storage_security import (
is_managed_cloud_storage_uri,
)
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.llms.base_llm.files.litellm_db_storage_backend import LITELLM_DB_STORAGE_BACKEND_NAME
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.llms.base_llm.managed_resources.isolation import build_list_page
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.batches_endpoints.litellm_executed_batches import (
litellm_executed_provider_of,
resolve_litellm_executed_provider,
)
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
@ -67,6 +72,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
add_internal_model_credentials,
apply_team_provider_credentials,
authorize_model_for_key,
encode_file_id_with_model,
extract_file_creation_params,
get_authorized_credentials_for_model,
@ -86,7 +92,7 @@ from litellm.proxy.openai_files_endpoints.general_upload_validation import (
coerce_optional_str_list_setting,
raise_upload_validation_failure,
)
from litellm.proxy.utils import ProxyLogging, is_known_model
from litellm.proxy.utils import PrismaClient, ProxyLogging, is_known_model
from litellm.repositories.table_repositories import ManagedFileRepository
from litellm.router import Router
from litellm.types.llms.openai import (
@ -99,6 +105,64 @@ from litellm.types.llms.openai import (
router: Final = APIRouter()
def _names_a_litellm_executed_provider(llm_router: Router, candidate: str, team_id: str | None) -> bool:
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=candidate, team_id=team_id)
return credentials is not None and litellm_executed_provider_of(credentials) is not None
async def _litellm_executed_batch_input_model(
llm_router: Router | None,
purpose: OpenAIFilesPurpose,
model: str | None,
target_model_names_list: Sequence[str],
user_api_key_dict: UserAPIKeyAuth,
explicit_storage: str | None,
) -> str | None:
if llm_router is None:
return None
candidates: Final = (model,) if model is not None else tuple(target_model_names_list)
team_id: Final = user_api_key_dict.team_id
await asyncio.gather(
*(
authorize_model_for_key(model_id=candidate, llm_router=llm_router, user_api_key_dict=user_api_key_dict)
for candidate in candidates
if _names_a_litellm_executed_provider(llm_router, candidate, team_id)
)
)
if explicit_storage is not None:
return None
providers: Final = await asyncio.gather(
*(resolve_litellm_executed_provider(llm_router, candidate, team_id) for candidate in candidates)
)
executed: Final = tuple(
candidate for candidate, provider in zip(candidates, providers, strict=True) if provider is not None
)
if not executed:
return None
if purpose != "batch":
raise ProxyException(
message=(
f"The server behind {', '.join(executed)} has no Files API, so LiteLLM keeps only batch input "
f"files for it and runs the batch itself: upload with purpose=batch; got purpose={purpose}"
),
type="invalid_request_error",
param="purpose",
code=400,
)
if len(candidates) == 1:
return executed[0]
raise ProxyException(
message=(
f"LiteLLM runs batches for {', '.join(executed)} itself and keeps their input files, so a batch "
f"input file can target only that one model; got target_model_names={', '.join(candidates)}"
),
type="invalid_request_error",
param="target_model_names",
code=400,
)
_MAX_BATCH_FILE_SIZE_MB_ADAPTER: Final = TypeAdapter(int | None)
_LISTED_FILES_ADAPTER: Final = TypeAdapter(list[OpenAIFileObject])
@ -244,30 +308,41 @@ async def route_create_file(
5. Else -> use custom_llm_provider with files_settings
"""
# Handle custom storage backend
if target_storage and target_storage != "default":
explicit_storage: Final = target_storage if target_storage and target_storage != "default" else None
if explicit_storage == LITELLM_DB_STORAGE_BACKEND_NAME:
raise ProxyException(
message=(
f"target_storage={LITELLM_DB_STORAGE_BACKEND_NAME} is not a storage a caller can pick: LiteLLM "
"chooses it on its own for the batch input files of a model whose batches it runs itself, so "
"upload with purpose=batch and name that model instead of target_storage"
),
type="invalid_request_error",
param="target_storage",
code=400,
)
executed_model: Final = await _litellm_executed_batch_input_model(
llm_router, purpose, model, target_model_names_list, user_api_key_dict, explicit_storage
)
storage: Final = explicit_storage or (LITELLM_DB_STORAGE_BACKEND_NAME if executed_model is not None else None)
if storage is not None:
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_file_data,
)
from litellm.proxy.openai_files_endpoints.storage_backend_service import (
StorageBackendFileService,
)
from litellm.proxy.proxy_server import prisma_client
# Extract file data
file_data: Final = extract_file_data(cast(Any, _create_file_request.get("file")))
# Use storage backend service to handle upload
file_object: Final = await StorageBackendFileService.upload_file_to_storage_backend(
file_data=file_data,
target_storage=target_storage,
target_model_names=target_model_names_list,
return await StorageBackendFileService.upload_file_to_storage_backend(
file_data=extract_file_data(cast(Any, _create_file_request.get("file"))),
target_storage=storage,
target_model_names=(executed_model,) if executed_model is not None else target_model_names_list,
purpose=purpose,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
return file_object
# NEW: Handle model-based routing (no DB required)
if model is not None:
# Get credentials from model_list via router
@ -848,7 +923,7 @@ async def get_file_content(
# Check if file is stored in a storage backend (check DB)
if hasattr(managed_files_obj, "prisma_client") and getattr(managed_files_obj, "prisma_client", None):
prisma_client: Final = getattr(managed_files_obj, "prisma_client")
prisma_client: Final[PrismaClient] = getattr(managed_files_obj, "prisma_client")
db_file: Final = await ManagedFileRepository(prisma_client).table.find_first(
where={"unified_file_id": file_id}
)
@ -863,7 +938,7 @@ async def get_file_content(
try:
# Get storage backend (uses same env vars as callback)
storage_backend: Final = get_storage_backend(storage_backend_name)
storage_backend: Final = get_storage_backend(storage_backend_name, prisma_client=prisma_client)
file_content: Final = await storage_backend.download_file(storage_url)
# Return file content

View file

@ -7,15 +7,16 @@ storage backends (e.g., Azure Blob Storage) and managing associated metadata.
import base64
import time
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from typing import Any, Final, cast
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid as uuid_module
from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.types.llms.openai import OpenAIFileObject, OpenAIFilesPurpose
from litellm.types.utils import SpecialEnums
@ -35,21 +36,23 @@ class StorageBackendFileService:
async def upload_file_to_storage_backend(
file_data: Mapping[str, Any],
target_storage: str,
target_model_names: list[str],
target_model_names: Sequence[str],
purpose: OpenAIFilesPurpose,
proxy_logging_obj: ProxyLogging,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient | None = None,
) -> OpenAIFileObject:
"""
Upload a file to a storage backend and create a file object.
Args:
file_data: File data dictionary from extract_file_data()
target_storage: Storage backend name (e.g., "azure_storage")
target_storage: Storage backend name (e.g., "azure_storage", "litellm_db")
target_model_names: List of model names for managed files
purpose: File purpose (e.g., "user_data", "batch")
proxy_logging_obj: Proxy logging object for accessing hooks
user_api_key_dict: User API key authentication data
prisma_client: The proxy's database client, required by the "litellm_db" backend
Returns:
OpenAIFileObject: Created file object with storage metadata
@ -59,7 +62,7 @@ class StorageBackendFileService:
"""
# Get storage backend instance
try:
storage_backend: Final = get_storage_backend(target_storage)
storage_backend: Final = get_storage_backend(target_storage, prisma_client=prisma_client)
except ValueError as e:
raise ProxyException(
message=str(e),
@ -103,8 +106,9 @@ class StorageBackendFileService:
storage_url=storage_url,
)
# Store in managed files if target_model_names provided
if target_model_names:
if not target_model_names:
return file_object
try:
await StorageBackendFileService._store_in_managed_files(
file_object=file_object,
file_data=file_data,
@ -114,9 +118,25 @@ class StorageBackendFileService:
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
)
except Exception:
await StorageBackendFileService._discard_orphaned_content(storage_backend, storage_url, target_storage)
raise
return file_object
@staticmethod
async def _discard_orphaned_content(
storage_backend: BaseFileStorageBackend, storage_url: str, target_storage: str
) -> None:
try:
await storage_backend.delete_file(storage_url)
except Exception as e: # noqa: BLE001 # the metadata failure is what surfaces; a failed cleanup is only logged
verbose_proxy_logger.warning(
"Could not delete orphaned content at %s on %s after its metadata write failed: %s",
storage_url,
target_storage,
e,
)
@staticmethod
def _create_file_object_with_storage_metadata(
file_content: bytes,
@ -164,7 +184,7 @@ class StorageBackendFileService:
@staticmethod
def _create_unified_file_id(
file_type: str,
target_model_names: list[str],
target_model_names: Sequence[str],
file_id: str,
) -> str:
"""
@ -194,7 +214,7 @@ class StorageBackendFileService:
async def _store_in_managed_files(
file_object: OpenAIFileObject,
file_data: Mapping[str, Any],
target_model_names: list[str],
target_model_names: Sequence[str],
target_storage: str,
storage_url: str,
proxy_logging_obj: ProxyLogging,

View file

@ -1107,6 +1107,12 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t
@@index([team_id, created_at(sort: Desc)])
}
model LiteLLM_ManagedFileContentTable {
id String @id @default(uuid())
content Bytes
created_at DateTime @default(now())
}
model LiteLLM_ManagedVectorStoreTable {
id String @id @default(uuid())
unified_resource_id String @unique // The base64 encoded unified vector store ID

View file

@ -0,0 +1,48 @@
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final
from litellm.repositories.table_repositories import PrismaTableRepository
from litellm.types.utils import LiteLLMBatch
if TYPE_CHECKING:
from prisma import models as prisma_models
def _batch_of(blob: object) -> LiteLLMBatch:
return LiteLLMBatch.model_validate_json(blob) if isinstance(blob, str) else LiteLLMBatch.model_validate(blob)
class ManagedBatchRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedObjectTable"]):
table_name = "litellm_managedobjecttable"
async def load_batch(self, unified_batch_id: str) -> LiteLLMBatch | None:
row: Final = await self._find_row(unified_batch_id)
return None if row is None or not row.file_object else _batch_of(row.file_object)
async def load_status(self, unified_batch_id: str) -> str | None:
row: Final = await self._find_row(unified_batch_id)
return row.status if row is not None else None
async def compare_and_set(
self, batch: LiteLLMBatch, unchanged: Mapping[str, object], updated_by: str | None
) -> bool:
updated_rows: Final = await self.table.update_many(
where={"unified_object_id": batch.id, **unchanged}, # mutable-ok: prisma filters are plain dicts
data={ # mutable-ok: prisma payloads are plain dicts
"file_object": batch.model_dump_json(),
"status": batch.status,
"updated_by": updated_by,
},
)
return updated_rows > 0
async def touch(self, unified_batch_id: str, updated_by: str | None) -> None:
await self.table.update_many(
where={"unified_object_id": unified_batch_id}, # mutable-ok: prisma filters are plain dicts
data={"updated_by": updated_by}, # mutable-ok: prisma payloads are plain dicts
)
async def _find_row(self, unified_batch_id: str) -> "prisma_models.LiteLLM_ManagedObjectTable | None":
return await self.table.find_first(
where={"unified_object_id": unified_batch_id} # mutable-ok: prisma filters are plain dicts
)

View file

@ -0,0 +1,32 @@
from typing import TYPE_CHECKING, Final
from litellm.repositories.table_repositories import PrismaTableRepository
if TYPE_CHECKING:
from prisma import models as prisma_models
class ManagedFileContentRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileContentTable"]):
table_name = "litellm_managedfilecontenttable"
async def store(self, content: bytes) -> str:
from prisma import Base64
row: Final = await self.table.create(
data={"content": Base64.encode(content)} # mutable-ok: prisma payloads are plain dicts
)
return row.id
async def load(self, row_id: str) -> bytes | None:
row: Final[prisma_models.LiteLLM_ManagedFileContentTable | None] = await self.table.find_unique(
where={"id": row_id} # mutable-ok: prisma filters are plain dicts
)
return None if row is None else row.content.decode()
async def delete(self, row_id: str) -> None:
from prisma.errors import RecordNotFoundError
try:
await self.table.delete(where={"id": row_id}) # mutable-ok: prisma filters are plain dicts
except RecordNotFoundError:
return

View file

@ -512,6 +512,7 @@ class CreateBatchRequest(TypedDict, total=False):
class LiteLLMBatchCreateRequest(CreateBatchRequest, total=False):
model: str
disable_fallbacks: ReadOnly[bool]
class RetrieveBatchRequest(TypedDict, total=False):

View file

@ -4178,6 +4178,8 @@ FILE_CONTENT_STREAMING_PROVIDERS: Final[frozenset[str]] = frozenset(
{*OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS, LlmProviders.VERTEX_AI.value}
)
LITELLM_EXECUTED_BATCH_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.HOSTED_VLLM.value})
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai"]
LIST_BATCHES_SUPPORTED_PROVIDERS: Final[frozenset[str]] = frozenset(get_args(ListBatchesSupportedProvider))

View file

@ -1107,6 +1107,12 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t
@@index([team_id, created_at(sort: Desc)])
}
model LiteLLM_ManagedFileContentTable {
id String @id @default(uuid())
content Bytes
created_at DateTime @default(now())
}
model LiteLLM_ManagedVectorStoreTable {
id String @id @default(uuid())
unified_resource_id String @unique // The base64 encoded unified vector store ID

View file

@ -1160,62 +1160,151 @@ def _vllm_params(api_base: str, api_key: str | None, model_id: str) -> LiteLLMPa
)
class TestHostedVllmBatch:
"""hosted_vllm file upload + batch create (OpenAI-compatible path, LIT-3266).
HOSTED_VLLM_DEFAULT_MODEL = "Qwen/Qwen2.5-0.5B-Instruct"
HOSTED_VLLM_BAD_LINE_CUSTOM_ID = "req-bad"
hosted_vllm is in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS, so /v1/files
and /v1/batches route through the OpenAI handler against the deployment's
api_base. Skipped for now: it needs a live vLLM (or OpenAI-compatible) server
exposing the files/batches APIs (HOSTED_VLLM_API_BASE), which the e2e
environment does not currently provision.
def _hosted_vllm_deployment(client: BatchClient, resources: ResourceManager) -> str:
api_base = os.environ.get("HOSTED_VLLM_API_BASE")
if api_base is None:
pytest.skip("set HOSTED_VLLM_API_BASE (the live vLLM server this deployment targets)")
api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None
model_id = (os.environ.get("HOSTED_VLLM_MODEL") or HOSTED_VLLM_DEFAULT_MODEL).strip()
proxy_name = batch_model_name("hosted-vllm-batch")
model_row_id = client.create_model(proxy_name, _vllm_params(api_base, api_key, model_id))
resources.defer(lambda: client.delete_model(model_row_id))
return proxy_name
def _upload_hosted_vllm_input(
client: BatchClient, content: bytes, *, proxy_name: str, key: str, upload_route: str
) -> Result[FileObject]:
if upload_route == "model_query":
return client.upload_file(content=content, form=FileUploadForm(purpose="batch"), model=proxy_name, key=key)
return client.upload_file(
content=content, form=FileUploadForm(purpose="batch", target_model_names=proxy_name), key=key
)
def _jsonl_with_a_failing_line(model: str) -> bytes:
bad_line = {
"custom_id": HOSTED_VLLM_BAD_LINE_CUSTOM_ID,
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": -1},
}
return render_jsonl(model) + (json.dumps(bad_line) + "\n").encode()
def _download_managed_file(client: BatchClient, file_id: str, *, key: str) -> list[str]:
downloaded = client.proxy.transport.download(
f"/v1/files/{file_id}/content", headers=client.proxy.transport.bearer(key)
)
assert downloaded.status_code == 200, (
f"file content must be 200, got {downloaded.status_code}: {downloaded.body[:300]}"
)
return downloaded.body.strip().splitlines()
class TestHostedVllmBatch:
"""hosted_vllm file upload + batch execution (LIT-5739).
vLLM implements neither /v1/files nor /v1/batches, so LiteLLM keeps the batch
input in its own database, runs every line through the deployment's
/v1/chat/completions itself, and serves the batch plus its output and error
files from that database under the creating key. Needs a live vLLM server
(HOSTED_VLLM_API_BASE), which the default e2e stack does not provision, so
the cases skip without it.
"""
@pytest.mark.skip(
reason="hosted_vllm batch/files needs a live vLLM server (HOSTED_VLLM_API_BASE) "
"not provisioned in the e2e environment; re-enable when available (LIT-3266)"
)
@pytest.mark.parametrize("upload_route", ["target_model_names", "model_query"])
@pytest.mark.covers(
"llm.batches.hosted_vllm.basic.nonstream.works",
"llm.files.hosted_vllm.upload.nonstream.works",
exercised_on=["batches", "files"],
)
def test_unified_file_and_batch_create(
self, client: BatchClient, resources: ResourceManager
def test_batch_runs_to_completion_with_a_downloadable_output(
self, client: BatchClient, resources: ResourceManager, upload_route: str
) -> None:
api_base = os.environ["HOSTED_VLLM_API_BASE"]
api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None
model_id = (
os.environ.get("HOSTED_VLLM_MODEL") or "meta-llama/Llama-3.2-3B-Instruct"
).strip()
proxy_name = batch_model_name("hosted-vllm-batch")
model_row_id = client.create_model(
proxy_name, _vllm_params(api_base, api_key, model_id)
)
resources.defer(lambda: client.delete_model(model_row_id))
proxy_name = _hosted_vllm_deployment(client, resources)
key = resources.key()
file = unwrap(
client.upload_file(
content=render_jsonl(model_id),
form=FileUploadForm(purpose="batch", target_model_names=proxy_name),
key=key,
_upload_hosted_vllm_input(
client, render_jsonl(proxy_name), proxy_name=proxy_name, key=key, upload_route=upload_route
)
)
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="hosted_vllm")
assert is_managed_id(file.id), f"hosted_vllm batch input must stay in LiteLLM, got file id {file.id!r}"
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert batch.id, f"hosted_vllm create returned no batch id: {created.body[:200]}"
assert batch.status in CREATED_BATCH_STATUSES, (
f"hosted_vllm batch has non-transitional status {batch.status!r}"
)
resources.defer(lambda: cleanup_batch(client, batch.id, key=key, delete_output_files=True))
assert is_managed_id(batch.id), f"hosted_vllm batch must be LiteLLM-managed, got {batch.id!r}"
assert batch.status in CREATED_BATCH_STATUSES, f"hosted_vllm batch has non-transitional status {batch.status!r}"
assert_batch_object(batch)
finished = _poll_until_terminal(client, batch.id, key)
assert finished.status == "completed", f"hosted_vllm batch ended {finished.status!r}: {finished.errors!r}"
assert finished.output_file_id, "completed hosted_vllm batch has no output_file_id"
assert finished.error_file_id is None, f"all lines succeeded but error_file_id={finished.error_file_id!r}"
output_lines = _download_managed_file(client, finished.output_file_id, key=key)
assert len(output_lines) == 1, f"one input line must yield one output line, got {output_lines!r}"
first_line = BatchOutputLine.model_validate_json(output_lines[0])
assert first_line.custom_id == "req-1", f"output line lost its custom_id: {output_lines[0][:300]}"
assert first_line.response.status_code == 200, f"batch output line reports failure: {output_lines[0][:400]}"
assert first_line.response.body is not None and first_line.response.body.choices, (
"batch output line has no choices"
)
rows = client.proxy.poll_logs_for_key(
key, predicate=lambda found: any(row.call_type == "acompletion" for row in found)
)
line_rows = [row for row in rows if row.call_type == "acompletion"]
assert line_rows, f"the batch line's chat call was not logged under the creating key: {rows!r}"
assert all(row.custom_llm_provider == "hosted_vllm" for row in line_rows), (
f"batch line rows must be attributed to hosted_vllm: {line_rows!r}"
)
@pytest.mark.covers("llm.batches.hosted_vllm.basic.nonstream.works", exercised_on=["batches", "files"])
def test_failing_line_lands_in_the_error_file_not_the_batch_status(
self, client: BatchClient, resources: ResourceManager
) -> None:
proxy_name = _hosted_vllm_deployment(client, resources)
key = resources.key()
file = unwrap(
_upload_hosted_vllm_input(
client,
_jsonl_with_a_failing_line(proxy_name),
proxy_name=proxy_name,
key=key,
upload_route="target_model_names",
)
)
resources.defer(lambda: cleanup_file(client, file.id, key=key))
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(lambda: cleanup_batch(client, batch.id, key=key, delete_output_files=True))
finished = _poll_until_terminal(client, batch.id, key)
assert finished.status == "completed", f"a failing line must not fail the batch, got {finished.status!r}"
assert finished.output_file_id, "the good line must still produce an output file"
assert finished.error_file_id, "the failing line must produce an error file"
output_lines = _download_managed_file(client, finished.output_file_id, key=key)
error_lines = _download_managed_file(client, finished.error_file_id, key=key)
assert [BatchOutputLine.model_validate_json(line).custom_id for line in output_lines] == ["req-1"]
assert len(error_lines) == 1, f"one failing line must yield one error line, got {error_lines!r}"
error_line = BatchOutputLine.model_validate_json(error_lines[0])
assert error_line.custom_id == HOSTED_VLLM_BAD_LINE_CUSTOM_ID
assert error_line.response.status_code == 400, f"error line must carry the provider's 4xx: {error_lines[0][:400]}"
BATCH_TERMINAL_STATUSES = frozenset({"completed", "failed", "expired", "cancelled"})
FAILED_BATCH_POLL_SECONDS = 120.0
@ -1443,6 +1532,7 @@ class BatchOutputResponse(BaseModel):
class BatchOutputLine(BaseModel):
custom_id: str | None = None
response: BatchOutputResponse

View file

@ -1067,6 +1067,7 @@ async def test_afile_content_passes_trusted_model_credentials_to_router():
managed_files = _make_managed_files_instance()
unified_file_id = "unified-file-id"
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
managed_files.get_unified_file_id = AsyncMock(return_value=None)
managed_files.get_model_file_id_mapping = AsyncMock(
return_value={unified_file_id: {"model-123": s3_uri}}
)
@ -1238,6 +1239,7 @@ async def test_afile_content_bedrock_unified_id_end_to_end(monkeypatch):
managed_files = _make_managed_files_instance()
unified_file_id = "unified-file-id"
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
managed_files.get_unified_file_id = AsyncMock(return_value=None)
managed_files.get_model_file_id_mapping = AsyncMock(
return_value={unified_file_id: {"model-123": s3_uri}}
)
@ -1268,6 +1270,7 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri():
managed_files = _make_managed_files_instance()
unified_file_id = "litellm_proxy_unified_id_abc"
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
managed_files.get_unified_file_id = AsyncMock(return_value=None)
managed_files.get_model_file_id_mapping = AsyncMock(
return_value={unified_file_id: {"model-123": s3_uri}}
)
@ -1732,6 +1735,40 @@ async def test_batch_retrieve_hook_does_not_claim_attribution():
assert managed_files.store_unified_object_id.await_args.kwargs["persist_attribution"] is False
def _unified_batch_id(llm_batch_id: str) -> str:
decoded = f"litellm_proxy;model_id:my-vllm;llm_batch_id:{llm_batch_id}"
return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=")
@pytest.mark.asyncio
@pytest.mark.parametrize(
"llm_batch_id, stores",
[("litellm_batch_abc", False), ("batch_abc", True)],
ids=["litellm-executed batch is left alone", "provider batch is still stored"],
)
async def test_post_call_hook_leaves_litellm_executed_batches_untouched(llm_batch_id: str, stores: bool):
managed_files = _make_managed_files_instance()
response = _make_batch_response(status="in_progress", output_file_id=None)
response.id = _unified_batch_id(llm_batch_id)
response._hidden_params = {
"unified_batch_id": response.id,
"model_id": "my-vllm",
"model_name": "hosted_vllm/qwen",
}
original_id = response.id
returned = await managed_files.async_post_call_success_hook(
data={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-the-poller", user_id="bob", parent_otel_span=None),
response=response,
)
assert returned is response
assert managed_files.store_unified_object_id.await_count == (1 if stores else 0)
if not stores:
assert response.id == original_id
@pytest.mark.asyncio
async def test_afile_delete_passes_trusted_model_credentials_to_router():
"""
@ -1743,6 +1780,7 @@ async def test_afile_delete_passes_trusted_model_credentials_to_router():
managed_files = _make_managed_files_instance()
unified_file_id = "unified-file-id"
s3_uri = "s3://my-bucket/litellm-bedrock-files/job-123/input.jsonl"
managed_files.get_unified_file_id = AsyncMock(return_value=None)
managed_files.get_model_file_id_mapping = AsyncMock(return_value={unified_file_id: {"model-123": s3_uri}})
managed_files.delete_unified_file_id = AsyncMock(return_value=_make_file_object(unified_file_id))
@ -1809,6 +1847,7 @@ async def test_afile_delete_bedrock_unified_id_end_to_end(monkeypatch):
managed_files = _make_managed_files_instance()
unified_file_id = "unified-file-id"
s3_uri = "s3://my-bucket/litellm-bedrock-files/job-123/input.jsonl"
managed_files.get_unified_file_id = AsyncMock(return_value=None)
managed_files.get_model_file_id_mapping = AsyncMock(return_value={unified_file_id: {"model-123": s3_uri}})
managed_files.delete_unified_file_id = AsyncMock(return_value=_make_file_object(unified_file_id))
@ -1827,3 +1866,147 @@ async def test_afile_delete_bedrock_unified_id_end_to_end(monkeypatch):
assert response.id == unified_file_id
assert response.model_dump() == {"id": unified_file_id, "object": "file", "deleted": True}
managed_files.delete_unified_file_id.assert_awaited_once_with(unified_file_id, None)
@pytest.mark.asyncio
async def test_afile_delete_storage_backed_row_deletes_stored_content_not_provider_files():
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
from openai.types import FileDeleted
from litellm.caching import DualCache
from litellm.models.managed_files import LiteLLM_ManagedFileTable
storage_url = "litellm_db://content-row-1"
unified_file_id = _managed_deletion_file_id(storage_url)
row = LiteLLM_ManagedFileTable(
unified_file_id=unified_file_id,
model_mappings={"vllm-batch": storage_url},
flat_model_file_ids=[storage_url],
file_object=_make_file_object(unified_file_id),
storage_backend="litellm_db",
storage_url=storage_url,
)
file_table = MagicMock(find_first=AsyncMock(return_value=row), delete=AsyncMock())
content_table = MagicMock(delete=AsyncMock())
managed_files = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=DualCache(),
prisma_client=MagicMock(
db=MagicMock(litellm_managedfiletable=file_table, litellm_managedfilecontenttable=content_table)
),
)
router = MagicMock(
get_deployment_credentials_with_provider=MagicMock(return_value=None),
afile_delete=AsyncMock(),
)
response = await managed_files.afile_delete(
file_id=unified_file_id,
litellm_parent_otel_span=None,
llm_router=router,
)
content_table.delete.assert_awaited_once_with(where={"id": "content-row-1"})
router.afile_delete.assert_not_awaited()
file_table.delete.assert_awaited_once_with(where={"unified_file_id": unified_file_id})
assert response == FileDeleted(id=unified_file_id, object="file", deleted=True)
@pytest.mark.asyncio
async def test_afile_content_storage_backed_row_returns_stored_bytes_not_provider_content():
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
from prisma import Base64
from litellm.caching import DualCache
from litellm.models.managed_files import LiteLLM_ManagedFileTable
storage_url = "litellm_db://content-row-1"
unified_file_id = _managed_deletion_file_id(storage_url)
stored_bytes = b'{"custom_id": "line-1", "method": "POST", "url": "/v1/chat/completions", "body": {}}\n'
row = LiteLLM_ManagedFileTable(
unified_file_id=unified_file_id,
model_mappings={"vllm-batch": storage_url},
flat_model_file_ids=[storage_url],
file_object=_make_file_object(unified_file_id),
storage_backend="litellm_db",
storage_url=storage_url,
)
file_table = MagicMock(find_first=AsyncMock(return_value=row))
content_table = MagicMock(find_unique=AsyncMock(return_value=MagicMock(content=Base64.encode(stored_bytes))))
managed_files = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=DualCache(),
prisma_client=MagicMock(
db=MagicMock(litellm_managedfiletable=file_table, litellm_managedfilecontenttable=content_table)
),
)
router = MagicMock(
get_deployment_credentials_with_provider=MagicMock(return_value=None),
afile_content=AsyncMock(),
)
response = await managed_files.afile_content(
file_id=unified_file_id,
litellm_parent_otel_span=None,
llm_router=router,
)
assert response.content == stored_bytes
content_table.find_unique.assert_awaited_once_with(where={"id": "content-row-1"})
router.afile_content.assert_not_awaited()
@pytest.mark.asyncio
async def test_store_unified_object_id_batch_processed_is_written_only_when_asked():
managed_files, mock_prisma = _make_object_store_instance()
upsert = mock_prisma.db.litellm_managedobjecttable.upsert
creator = UserAPIKeyAuth(api_key="sk-creator", user_id="alice", team_id="team-alpha", parent_otel_span=None)
await managed_files.store_unified_object_id(
unified_object_id="uoi-processed",
file_object=_make_batch_response(status="completed"),
litellm_parent_otel_span=None,
model_object_id="batch-processed",
file_purpose="batch",
user_api_key_dict=creator,
batch_processed=True,
)
await managed_files.store_unified_object_id(
unified_object_id="uoi-default",
file_object=_make_batch_response(status="completed"),
litellm_parent_otel_span=None,
model_object_id="batch-default",
file_purpose="batch",
user_api_key_dict=creator,
)
processed_create, default_create = (call.kwargs["data"]["create"] for call in upsert.await_args_list)
assert processed_create["batch_processed"] is True
assert default_create["batch_processed"] is False
@pytest.mark.asyncio
async def test_store_unified_file_id_caches_the_storage_location_the_db_row_gets():
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
from litellm.caching import DualCache
file_table = MagicMock(upsert=AsyncMock(), find_first=AsyncMock(side_effect=AssertionError("cache miss")))
managed_files = _PROXY_LiteLLMManagedFiles(
internal_usage_cache=DualCache(),
prisma_client=MagicMock(db=MagicMock(litellm_managedfiletable=file_table)),
)
stored = _make_file_object("file-kept").model_copy(update={"purpose": "batch"})
stored._hidden_params = {"storage_backend": "litellm_db", "storage_url": "litellm_db://content-row-1"}
await managed_files.store_unified_file_id(
file_id="unified-kept",
file_object=stored,
litellm_parent_otel_span=None,
model_mappings={"vllm-batch": "litellm_db://content-row-1"},
user_api_key_dict=_make_user_api_key_dict(),
)
cached = await managed_files.get_unified_file_id("unified-kept")
assert cached is not None
assert (cached.storage_backend, cached.storage_url) == ("litellm_db", "litellm_db://content-row-1")
create_data = file_table.upsert.await_args.kwargs["data"]["create"]
assert (create_data["storage_backend"], create_data["storage_url"]) == ("litellm_db", "litellm_db://content-row-1")

View file

@ -0,0 +1,92 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from prisma import Base64
from prisma.errors import RecordNotFoundError
from litellm.llms.base_llm.files.litellm_db_storage_backend import (
LITELLM_DB_STORAGE_URL_PREFIX,
LiteLLMDbStorageBackend,
storage_url_to_row_id,
)
def _backend_with_table():
table = MagicMock(create=AsyncMock(), find_unique=AsyncMock(), delete=AsyncMock())
prisma_client = MagicMock(db=MagicMock(litellm_managedfilecontenttable=table))
return LiteLLMDbStorageBackend(prisma_client), table
@pytest.mark.asyncio
async def test_upload_stores_bytes_and_returns_prefixed_row_id():
backend, table = _backend_with_table()
table.create.return_value = SimpleNamespace(id="row-1")
content = b"\x00\x01binary jsonl\n"
storage_url = await backend.upload_file(file_content=content, filename="input.jsonl", content_type="text/plain")
assert storage_url == f"{LITELLM_DB_STORAGE_URL_PREFIX}row-1"
stored = table.create.await_args.kwargs["data"]["content"]
assert isinstance(stored, Base64)
assert stored.decode() == content
@pytest.mark.asyncio
async def test_download_returns_exact_bytes_of_the_row():
backend, table = _backend_with_table()
content = b'{"custom_id": "1"}\n'
table.find_unique.return_value = SimpleNamespace(id="row-1", content=Base64.encode(content))
downloaded = await backend.download_file(f"{LITELLM_DB_STORAGE_URL_PREFIX}row-1")
assert downloaded == content
table.find_unique.assert_awaited_once_with(where={"id": "row-1"})
@pytest.mark.asyncio
async def test_download_missing_row_raises_value_error_naming_the_url():
backend, table = _backend_with_table()
table.find_unique.return_value = None
storage_url = f"{LITELLM_DB_STORAGE_URL_PREFIX}missing"
with pytest.raises(ValueError, match="missing"):
await backend.download_file(storage_url)
@pytest.mark.asyncio
async def test_download_rejects_url_without_prefix_before_touching_the_db():
backend, table = _backend_with_table()
with pytest.raises(ValueError, match="https://elsewhere/blob"):
await backend.download_file("https://elsewhere/blob")
table.find_unique.assert_not_awaited()
@pytest.mark.asyncio
async def test_delete_removes_the_parsed_row():
backend, table = _backend_with_table()
await backend.delete_file(f"{LITELLM_DB_STORAGE_URL_PREFIX}row-1")
table.delete.assert_awaited_once_with(where={"id": "row-1"})
@pytest.mark.asyncio
async def test_delete_tolerates_a_row_that_is_already_gone():
backend, table = _backend_with_table()
table.delete.side_effect = RecordNotFoundError({"user_facing_error": {"message": "gone"}})
await backend.delete_file(f"{LITELLM_DB_STORAGE_URL_PREFIX}row-1")
table.delete.assert_awaited_once_with(where={"id": "row-1"})
def test_storage_url_to_row_id_round_trips():
assert storage_url_to_row_id(f"{LITELLM_DB_STORAGE_URL_PREFIX}abc-123") == "abc-123"
def test_storage_url_to_row_id_rejects_foreign_urls():
with pytest.raises(ValueError, match="s3://bucket/key"):
storage_url_to_row_id("s3://bucket/key")

View file

@ -0,0 +1,34 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.llms.base_llm.files.litellm_db_storage_backend import (
LITELLM_DB_STORAGE_BACKEND_NAME,
LITELLM_DB_STORAGE_URL_PREFIX,
LiteLLMDbStorageBackend,
)
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
@pytest.mark.asyncio
async def test_litellm_db_backend_stores_through_the_given_prisma_client():
table = MagicMock(create=AsyncMock(return_value=SimpleNamespace(id="row-1")))
prisma_client = MagicMock(db=MagicMock(litellm_managedfilecontenttable=table))
backend = get_storage_backend(LITELLM_DB_STORAGE_BACKEND_NAME, prisma_client=prisma_client)
assert isinstance(backend, LiteLLMDbStorageBackend)
stored_at = await backend.upload_file(file_content=b"line\n", filename="input.jsonl", content_type="text/plain")
assert stored_at == f"{LITELLM_DB_STORAGE_URL_PREFIX}row-1"
table.create.assert_awaited_once()
def test_litellm_db_backend_without_a_database_is_rejected():
with pytest.raises(ValueError, match="database-connected proxy"):
get_storage_backend(LITELLM_DB_STORAGE_BACKEND_NAME)
def test_unknown_backend_is_still_rejected():
with pytest.raises(ValueError, match="Unsupported storage backend type: nope"):
get_storage_backend("nope", prisma_client=MagicMock())

View file

@ -34,10 +34,13 @@ import json
import logging
from contextlib import ExitStack
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import respx
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
import litellm
@ -52,7 +55,7 @@ from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
from litellm.types.llms.openai import BatchJobStatus
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
from litellm.types.utils import CredentialItem, LiteLLMBatch
from litellm.types.utils import CredentialItem, LiteLLMBatch, SpecialEnums
from fastapi import Request, Response
@ -74,6 +77,12 @@ CREDS: Dict[str, Dict[str, str]] = {
"api_base": "https://vertex.test",
"model": "vertex_ai/gemini-2.0",
},
"my-vllm": {
"custom_llm_provider": "hosted_vllm",
"api_key": "sk-vllm",
"api_base": "http://vllm.test/v1",
"model": "hosted_vllm/qwen",
},
}
# A real model-encoded file id: decodes to "azure/gpt-4o", strips to "file-original123".
@ -147,6 +156,7 @@ class Harness:
router: MagicMock
logging: MagicMock
creds_resolver: MagicMock
upstream_files_route: respx.Route
@property
def router_acreate(self) -> AsyncMock:
@ -162,13 +172,14 @@ class Harness:
return dict(self.router_acreate.call_args.kwargs)
def _creds_lookup(*, model_id: str) -> Dict[str, str]:
# KeyError on an unknown/hardcoded model_id - the bug cannot hide.
return dict(CREDS[model_id])
def _creds_lookup(*, model_id: str, team_id: str | None = None) -> dict[str, str] | None:
# An unknown/hardcoded model_id resolves to None exactly like the real router,
# which the endpoint turns into a 400 and a missing dispatch - the bug cannot hide.
return dict(CREDS[model_id]) if model_id in CREDS else None
@pytest.fixture
def harness():
def harness(monkeypatch: pytest.MonkeyPatch):
"""Seam harness. Patches only true I/O boundaries; pure encode/decode/merge
helpers run for real. Object mocks are spec'd so unknown method calls raise."""
body_holder: Dict[str, Any] = {}
@ -192,6 +203,7 @@ def harness():
provider_from_headers = MagicMock(return_value=None)
is_known_model = MagicMock(return_value=False)
litellm_acreate = AsyncMock(return_value=make_batch())
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
with ExitStack() as stack:
stack.enter_context(patch.object(endpoints, "_read_request_body", read_body))
@ -213,6 +225,10 @@ def harness():
stack.enter_context(patch.object(endpoints, "is_known_model", is_known_model))
stack.enter_context(patch.object(litellm, "acreate_batch", litellm_acreate))
stack.enter_context(patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False))
upstream = stack.enter_context(respx.mock(assert_all_called=False))
upstream_files_route = upstream.get(f"{CREDS['my-vllm']['api_base']}/files").mock(
return_value=httpx.Response(404, json={"detail": "Not Found"})
)
stack.enter_context(patch.object(proxy_server, "llm_router", router))
stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging))
stack.enter_context(patch.object(proxy_server, "general_settings", {}))
@ -231,6 +247,7 @@ def harness():
router=router,
logging=logging,
creds_resolver=router.get_deployment_credentials_with_provider,
upstream_files_route=upstream_files_route,
)
yield h
@ -255,6 +272,25 @@ async def call_create(
)
@pytest.fixture
def executed_runner():
runner = MagicMock(spec=endpoints.LiteLLMExecutedBatchRunner)
runner.create = AsyncMock(return_value=make_batch(id="litellm-executed-batch"))
runner.cancel = AsyncMock(return_value=make_batch(id="litellm-executed-batch", status="cancelling"))
factory = MagicMock(return_value=runner)
with patch.object( # test-quality-ok: the route builds its runner from proxy_server globals; the factory is the only seam
endpoints, "_litellm_executed_batch_runner", factory
):
yield runner, factory
def _managed_input_file_id(model: str) -> str:
unified = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
"application/jsonl", "managed-id", model, "file-id", "file-model-id"
)
return base64.urlsafe_b64encode(unified.encode()).decode().rstrip("=")
# =========================================================================== #
# SCENARIO 1 - input_file_id encoded with model. The full showcase: every
# assertion type from the design lives here.
@ -766,6 +802,137 @@ async def test_create__unified_file_id_legacy_row_without_storage_url_dispatches
assert harness.router_kwargs()["input_file_id"] == "litellm_proxy_unified_id"
# --------------------------------------------------------------------------- #
# LiteLLM-executed batches: a unified file targeting a provider whose API has
# no /v1/batches (hosted_vllm) runs inside LiteLLM instead of being forwarded.
# --------------------------------------------------------------------------- #
@pytest.mark.asyncio
async def test_create__unified_executed_provider_runs_inside_litellm(harness, executed_runner):
runner, factory = executed_runner
caller = UserAPIKeyAuth(api_key="sk-test", team_id="team-vllm")
input_file_id = _managed_input_file_id("my-vllm")
set_body(
harness,
{
"input_file_id": input_file_id,
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"litellm_metadata": {"tags": ["batch-tag"]},
},
)
resp = await call_create(harness, user=caller)
harness.router_acreate.assert_not_called()
harness.litellm_acreate.assert_not_called()
harness.creds_resolver.assert_called_once_with(model_id="my-vllm", team_id="team-vllm")
factory.assert_called_once_with(harness.router, harness.logging)
runner.create.assert_awaited_once()
create_kwargs = runner.create.call_args.kwargs
assert create_kwargs["unified_input_file_id"] == input_file_id
assert create_kwargs["model"] == "my-vllm"
assert create_kwargs["provider"] == "hosted_vllm"
assert create_kwargs["request_tags"] == ("batch-tag",)
assert create_kwargs["user_api_key_dict"] is caller
assert create_kwargs["create_request"]["model"] == "my-vllm"
assert resp.id == "litellm-executed-batch"
@pytest.mark.asyncio
async def test_create__unified_executed_provider_without_database_400(harness):
set_body(
harness,
{
"input_file_id": _managed_input_file_id("my-vllm"),
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
)
with pytest.raises(ProxyException) as exc:
await call_create(harness)
assert exc.value.code == "400"
assert "need a database" in exc.value.message
harness.router_acreate.assert_not_called()
harness.litellm_acreate.assert_not_called()
@pytest.mark.asyncio
async def test_create__unified_executed_provider_with_its_own_files_api_goes_to_the_provider(harness, executed_runner):
runner, factory = executed_runner
harness.upstream_files_route.mock(return_value=httpx.Response(200, json={"object": "list", "data": []}))
set_body(
harness,
{
"input_file_id": _managed_input_file_id("my-vllm"),
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
)
await call_create(harness)
factory.assert_not_called()
runner.create.assert_not_called()
assert harness.router_kwargs()["model"] == "my-vllm"
@pytest.mark.asyncio
async def test_create__unified_provider_model_never_touches_executed_runner(harness, executed_runner):
runner, factory = executed_runner
set_body(
harness,
{
"input_file_id": _managed_input_file_id("azure/gpt-4o"),
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
)
await call_create(harness)
factory.assert_not_called()
runner.create.assert_not_called()
harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o", team_id=None)
assert harness.router_kwargs()["model"] == "azure/gpt-4o"
@pytest.mark.asyncio
@pytest.mark.parametrize("via", ["body", "header"])
async def test_create__raw_file_with_executed_model_400_with_upload_guidance(harness, via):
body = {"input_file_id": "file-plain", "endpoint": "/v1/chat/completions", "completion_window": "24h"}
set_body(harness, {**body, "model": "my-vllm"} if via == "body" else body)
headers = {"x-litellm-model": "my-vllm"} if via == "header" else None
with pytest.raises(ProxyException) as exc:
await call_create(harness, headers=headers)
assert exc.value.code == "400"
assert "POST /v1/files" in exc.value.message
assert "x-litellm-model" in exc.value.message
harness.litellm_acreate.assert_not_called()
harness.router_acreate.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"upstream_answer",
[httpx.Response(200, json={"object": "list", "data": []}), httpx.Response(405), httpx.ConnectError("refused")],
ids=["lists files", "files route without list", "unreachable"],
)
async def test_create__raw_file_with_executed_model_is_forwarded_unless_the_server_lacks_a_files_api(
harness, upstream_answer
):
harness.upstream_files_route.mock(side_effect=[upstream_answer])
set_body(harness, {"input_file_id": "file-plain", "endpoint": "/v1/chat/completions", "completion_window": "24h"})
await call_create(harness, headers={"x-litellm-model": "my-vllm"})
forwarded = harness.acreate_kwargs()
assert forwarded["input_file_id"] == "file-plain"
assert forwarded["custom_llm_provider"] == "hosted_vllm"
assert forwarded["api_base"] == CREDS["my-vllm"]["api_base"]
@pytest.mark.asyncio
async def test_create__model_encoded_beats_unified(harness):
"""Precedence row: a file id that is BOTH model-encoded and (pretend) unified
@ -1146,6 +1313,11 @@ AZURE_BATCH_ID = encode_file_id_with_model("batch_orig123", "azure/gpt-4o", id_t
# returns). model_id / llm_batch_id are parsed out of this by the real helpers.
UNIFIED_BATCH_ID = "litellm_proxy;model_id:gpt-4o-mini;llm_batch_id:batch-raw-xyz"
# A decoded unified id of a batch LiteLLM runs itself: the llm_batch_id carries
# the litellm_batch_ prefix, so no provider holds a batch to sync with.
EXECUTED_BATCH_ID = "litellm_proxy;model_id:my-vllm;llm_batch_id:litellm_batch_abc"
EXECUTED_BATCH_B64 = base64.urlsafe_b64encode(EXECUTED_BATCH_ID.encode()).decode().rstrip("=")
@dataclass
class RetrieveHarness:
@ -1580,6 +1752,69 @@ async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harn
assert retrieve_harness.update_batch_in_db.call_count == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["validating", "in_progress", "finalizing", "cancelling"])
async def test_retrieve__executed_batch_served_from_db_in_every_status(retrieve_harness, status):
db_response = make_batch(id="litellm-executed-batch", status=status)
db_batch_object = MagicMock()
db_batch_object.updated_at = datetime.now(timezone.utc)
retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response)
resp = await call_retrieve(retrieve_harness, EXECUTED_BATCH_B64)
assert resp is db_response
retrieve_harness.litellm_aretrieve.assert_not_called()
retrieve_harness.router_aretrieve.assert_not_called()
retrieve_harness.update_batch_in_db.assert_not_called()
retrieve_harness.ensure_managed_files.assert_called_once()
assert retrieve_harness.ensure_managed_files.call_args.kwargs["unified_batch_id"] == EXECUTED_BATCH_ID
@pytest.mark.asyncio
async def test_retrieve__executed_batch_abandoned_by_its_runner_is_served_failed(retrieve_harness, executed_runner):
runner, _ = executed_runner
failed = make_batch(id="litellm-executed-batch", status="failed")
runner.fail_abandoned = AsyncMock(return_value=failed)
db_response = make_batch(id="litellm-executed-batch", status="in_progress")
db_batch_object = MagicMock()
db_batch_object.updated_at = datetime.now(timezone.utc) - timedelta(minutes=10)
retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response)
user = UserAPIKeyAuth(api_key="sk-test", user_id="user-1")
resp = await call_retrieve(retrieve_harness, EXECUTED_BATCH_B64, user=user)
assert resp is failed
runner.fail_abandoned.assert_awaited_once_with(db_response, user)
retrieve_harness.litellm_aretrieve.assert_not_called()
retrieve_harness.router_aretrieve.assert_not_called()
retrieve_harness.ensure_managed_files.assert_called_once()
@pytest.mark.asyncio
async def test_retrieve__executed_batch_with_a_fresh_heartbeat_is_left_running(retrieve_harness, executed_runner):
runner, _ = executed_runner
runner.fail_abandoned = AsyncMock()
db_response = make_batch(id="litellm-executed-batch", status="in_progress")
db_batch_object = MagicMock()
db_batch_object.updated_at = datetime.now(timezone.utc) - timedelta(seconds=30)
retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response)
resp = await call_retrieve(retrieve_harness, EXECUTED_BATCH_B64)
assert resp is db_response
runner.fail_abandoned.assert_not_awaited()
@pytest.mark.asyncio
async def test_retrieve__executed_batch_without_db_row_404(retrieve_harness):
with pytest.raises(ProxyException) as exc:
await call_retrieve(retrieve_harness, EXECUTED_BATCH_B64)
assert exc.value.code == "404"
retrieve_harness.litellm_aretrieve.assert_not_called()
retrieve_harness.router_aretrieve.assert_not_called()
# --------------------------------------------------------------------------- #
# Cross-cutting: enrichment route_type and failure-hook on provider error.
# --------------------------------------------------------------------------- #
@ -2299,6 +2534,35 @@ async def test_cancel__unified_no_router_500(cancel_harness):
assert exc.value.code == "500"
@pytest.mark.asyncio
async def test_cancel__executed_batch_routes_to_runner(cancel_harness, executed_runner):
runner, factory = executed_runner
caller = UserAPIKeyAuth(api_key="sk-test", user_id="user-cancel-2")
resp = await call_cancel(cancel_harness, EXECUTED_BATCH_B64, user=caller)
runner.cancel.assert_awaited_once_with(EXECUTED_BATCH_B64, caller)
factory.assert_called_once_with(cancel_harness.router, cancel_harness.logging)
cancel_harness.router_acancel.assert_not_called()
cancel_harness.litellm_acancel.assert_not_called()
cancel_harness.creds_resolver.assert_not_called()
assert resp is runner.cancel.return_value
assert cancel_harness.update_batch_in_db.call_args.kwargs["operation"] == "cancel"
@pytest.mark.asyncio
async def test_cancel__executed_batch_no_router_500(cancel_harness, executed_runner):
runner, factory = executed_runner
with patch.object( # test-quality-ok: proxy_server module global is the endpoint's only injection point
proxy_server, "llm_router", None
):
with pytest.raises(ProxyException) as exc:
await call_cancel(cancel_harness, EXECUTED_BATCH_B64)
assert exc.value.code == "500"
factory.assert_not_called()
runner.cancel.assert_not_called()
# --------------------------------------------------------------------------- #
# SCENARIO 3 - fallback to custom_llm_provider. Rebuilds a CancelBatchRequest
# and forwards only {custom_llm_provider, batch_id}.
@ -2956,3 +3220,16 @@ async def test_cancel__unified_batch_id_rejects_key_without_model_grant(cancel_h
assert exc_info.value.code == "403"
cancel_harness.router_acancel.assert_not_called()
@pytest.mark.asyncio
async def test_cancel__executed_batch_rejects_key_without_model_grant(cancel_harness, executed_runner):
runner, factory = executed_runner
with pytest.raises(ProxyException) as exc_info:
await call_cancel(cancel_harness, EXECUTED_BATCH_B64, user=_key_restricted_to("vertex-model"))
assert exc_info.value.code == "403"
factory.assert_not_called()
runner.cancel.assert_not_called()
cancel_harness.router_acancel.assert_not_called()

File diff suppressed because it is too large Load diff

View file

@ -4,10 +4,10 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.proxy.openai_files_endpoints.common_utils import (
apply_unified_file_ids,
get_credentials_for_model,
is_litellm_executed_batch,
map_raw_file_ids_to_unified,
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
@ -500,3 +500,17 @@ class TestCompletedBatchSafeToRetire:
def test_no_output_and_unknown_counts_is_not_safe(self):
assert _completed_batch_safe_to_retire(_completed_batch_for_retire(None)) is False
@pytest.mark.parametrize(
"decoded_unified_batch_id, executed",
[
("litellm_proxy;model_id:my-vllm;llm_batch_id:litellm_batch_0123abcd", True),
("litellm_proxy;model_id:my-vllm;llm_batch_id:batch_0123abcd", False),
("litellm_proxy;model_id:my-vllm;generic_response_id:resp_0123abcd", False),
("litellm_proxy;model_id:my-vllm;llm_output_file_id:file-0123abcd", False),
("batch_0123abcd", False),
],
)
def test_is_litellm_executed_batch_reads_the_llm_batch_id_prefix(decoded_unified_batch_id: str, executed: bool):
assert is_litellm_executed_batch(decoded_unified_batch_id) is executed

View file

@ -609,6 +609,246 @@ def test_target_storage_with_target_models(
app.dependency_overrides.pop(ps.user_api_key_auth, None)
BATCH_JSONL_LINE = (
b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions", '
b'"body": {"model": "my-vllm", "messages": [{"role": "user", "content": "hi"}]}}\n'
)
def _router_with_executed_batch_model() -> Router:
return Router(
model_list=[
{
"model_name": "my-vllm",
"litellm_params": {
"model": "hosted_vllm/qwen",
"api_key": "sk-vllm",
"api_base": "http://vllm.test/v1",
},
"model_info": {"id": "my-vllm-id"},
},
{
"model_name": "gemini-2.0-flash",
"litellm_params": {"model": "gemini/gemini-2.0-flash"},
"model_info": {"id": "gemini-2.0-flash-id"},
},
]
)
@pytest.fixture
def batch_upload_seams(mocker: MockerFixture, monkeypatch):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
llm_router = _router_with_executed_batch_model()
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
setup_proxy_logging_object(monkeypatch, llm_router)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"
)
uploaded = OpenAIFileObject(
id="file-kept",
object="file",
purpose="batch",
created_at=0,
bytes=len(BATCH_JSONL_LINE),
filename="batch.jsonl",
status="uploaded",
)
stored = mocker.patch( # test-quality-ok: the route calls the storage service directly with no injection seam
"litellm.proxy.openai_files_endpoints.storage_backend_service.StorageBackendFileService.upload_file_to_storage_backend",
new=mocker.AsyncMock(return_value=uploaded),
)
provider_upload = mocker.patch( # test-quality-ok: the route calls litellm.acreate_file directly with no injection seam
"litellm.acreate_file", new=mocker.AsyncMock(return_value=uploaded)
)
try:
with respx.mock(assert_all_called=False) as upstream:
upstream_files_route = upstream.get("http://vllm.test/v1/files").mock(
return_value=httpx.Response(404, json={"detail": "Not Found"})
)
yield stored, provider_upload, upstream_files_route
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
def _upload_batch_file(headers: dict[str, str], form: dict[str, str]):
return client.post(
"/v1/files",
files={"file": ("batch.jsonl", BATCH_JSONL_LINE, "application/jsonl")},
data={"purpose": "batch", **form},
headers={"Authorization": "Bearer test-key", **headers},
)
@pytest.mark.parametrize(
"headers, form",
[({"x-litellm-model": "my-vllm"}, {}), ({}, {"target_model_names": "my-vllm"})],
ids=["x-litellm-model header", "target_model_names form field"],
)
def test_batch_upload_for_a_litellm_executed_model_is_kept_by_litellm(
batch_upload_seams, headers: dict[str, str], form: dict[str, str]
):
stored, provider_upload, _ = batch_upload_seams
response = _upload_batch_file(headers, form)
assert response.status_code == 200, response.text
provider_upload.assert_not_awaited()
stored.assert_awaited_once()
kwargs = stored.call_args.kwargs
assert kwargs["target_storage"] == "litellm_db"
assert tuple(kwargs["target_model_names"]) == ("my-vllm",)
assert kwargs["purpose"] == "batch"
@pytest.mark.parametrize(
"headers, form",
[({"x-litellm-model": "my-vllm"}, {}), ({}, {"target_model_names": "my-vllm"})],
ids=["x-litellm-model header", "target_model_names form field"],
)
def test_batch_upload_for_a_litellm_executed_model_the_key_cannot_call_is_refused_before_the_server_is_probed(
batch_upload_seams, headers: dict[str, str], form: dict[str, str]
):
import litellm.proxy.proxy_server as ps
stored, provider_upload, upstream_files_route = batch_upload_seams
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="restricted-user", models=["gemini-2.0-flash"]
)
response = _upload_batch_file(headers, form)
assert response.status_code == 403, response.text
assert "my-vllm" in response.text
assert upstream_files_route.call_count == 0
stored.assert_not_awaited()
provider_upload.assert_not_awaited()
def test_batch_upload_naming_an_executed_and_a_provider_model_is_rejected(batch_upload_seams):
stored, provider_upload, _ = batch_upload_seams
response = _upload_batch_file({}, {"target_model_names": "my-vllm,gemini-2.0-flash"})
assert response.status_code == 400, response.text
assert "my-vllm" in response.text
assert "target_model_names" in response.text
stored.assert_not_awaited()
provider_upload.assert_not_awaited()
@pytest.mark.parametrize("purpose", ["assistants", "user_data"])
def test_non_batch_upload_for_a_litellm_executed_model_is_rejected_with_the_purpose_to_use(
batch_upload_seams, purpose: str
):
stored, provider_upload, _ = batch_upload_seams
response = _upload_batch_file({"x-litellm-model": "my-vllm"}, {"purpose": purpose})
assert response.status_code == 400, response.text
error = response.json()["error"]
assert error["type"] == "invalid_request_error"
assert error["param"] == "purpose"
assert "purpose=batch" in error["message"]
assert f"purpose={purpose}" in error["message"]
stored.assert_not_awaited()
provider_upload.assert_not_awaited()
@pytest.mark.parametrize("purpose", ["batch", "assistants"])
@pytest.mark.parametrize(
"upstream_answer",
[httpx.Response(200, json={"object": "list", "data": []}), httpx.Response(405), httpx.ConnectError("refused")],
ids=["lists files", "files route without list", "unreachable"],
)
def test_upload_for_a_litellm_executed_model_goes_to_the_provider_unless_the_server_lacks_a_files_api(
batch_upload_seams, upstream_answer: httpx.Response | httpx.ConnectError, purpose: str
):
stored, provider_upload, upstream_files_route = batch_upload_seams
upstream_files_route.mock(side_effect=[upstream_answer])
response = _upload_batch_file({"x-litellm-model": "my-vllm"}, {"purpose": purpose})
assert response.status_code == 200, response.text
stored.assert_not_awaited()
provider_upload.assert_awaited_once()
assert provider_upload.call_args.kwargs["custom_llm_provider"] == "hosted_vllm"
assert provider_upload.call_args.kwargs["api_base"] == "http://vllm.test/v1"
@pytest.mark.parametrize(
"form",
[{}, {"target_model_names": "my-vllm"}, {"target_model_names": "gemini-2.0-flash"}],
ids=["no model", "litellm-executed model", "provider model"],
)
def test_upload_naming_litellm_db_as_target_storage_is_rejected(batch_upload_seams, form: dict[str, str]):
stored, provider_upload, upstream_files_route = batch_upload_seams
response = _upload_batch_file({}, {**form, "target_storage": "litellm_db"})
assert response.status_code == 400, response.text
error = response.json()["error"]
assert error["type"] == "invalid_request_error"
assert error["param"] == "target_storage"
assert "litellm_db" in error["message"]
assert upstream_files_route.call_count == 0
stored.assert_not_awaited()
provider_upload.assert_not_awaited()
@pytest.mark.parametrize("purpose", ["user_data", "batch"])
def test_upload_with_an_explicit_target_storage_goes_where_the_caller_said_without_probing_the_server(
batch_upload_seams, purpose: str
):
stored, provider_upload, upstream_files_route = batch_upload_seams
response = _upload_batch_file(
{}, {"purpose": purpose, "target_model_names": "my-vllm", "target_storage": "azure_storage"}
)
assert response.status_code == 200, response.text
assert upstream_files_route.call_count == 0
provider_upload.assert_not_awaited()
stored.assert_awaited_once()
kwargs = stored.call_args.kwargs
assert kwargs["target_storage"] == "azure_storage"
assert tuple(kwargs["target_model_names"]) == ("my-vllm",)
assert kwargs["purpose"] == purpose
def test_upload_with_an_explicit_target_storage_still_refuses_a_key_without_the_executed_model(batch_upload_seams):
import litellm.proxy.proxy_server as ps
stored, provider_upload, upstream_files_route = batch_upload_seams
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="restricted-user", models=["gemini-2.0-flash"]
)
response = _upload_batch_file({}, {"target_model_names": "my-vllm", "target_storage": "azure_storage"})
assert response.status_code == 403, response.text
assert "my-vllm" in response.text
assert upstream_files_route.call_count == 0
stored.assert_not_awaited()
provider_upload.assert_not_awaited()
def test_batch_upload_for_a_provider_model_still_goes_to_the_provider(batch_upload_seams):
stored, provider_upload, _ = batch_upload_seams
response = _upload_batch_file({"x-litellm-model": "gemini-2.0-flash"}, {})
assert response.status_code == 200, response.text
stored.assert_not_awaited()
provider_upload.assert_awaited_once()
assert provider_upload.call_args.kwargs["custom_llm_provider"] == "gemini"
@pytest.mark.skip(reason="mock respx fails on ci/cd - unclear why")
def test_create_file_and_call_chat_completion_e2e(
mocker: MockerFixture, monkeypatch, llm_router: Router

View file

@ -1,3 +1,5 @@
from unittest.mock import MagicMock
import pytest
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
@ -6,16 +8,24 @@ from litellm.proxy.openai_files_endpoints import storage_backend_service
from litellm.proxy.openai_files_endpoints.storage_backend_service import (
StorageBackendFileService,
)
from litellm.proxy.utils import PrismaClient
class _RecordingStorageBackend:
def __init__(self):
def __init__(self, delete_error: Exception | None = None):
self.upload_calls = []
self.delete_calls: list[str] = []
self.delete_error = delete_error
async def upload_file(self, **kwargs):
self.upload_calls.append(kwargs)
return "https://storage.example/blob-1"
async def delete_file(self, storage_url: str) -> None:
self.delete_calls.append(storage_url)
if self.delete_error is not None:
raise self.delete_error
class _FakeManagedFilesHook(BaseFileEndpoints):
def __init__(self):
@ -42,6 +52,11 @@ class _FakeManagedFilesHook(BaseFileEndpoints):
self.stored.append(kwargs)
class _FailingManagedFilesHook(_FakeManagedFilesHook):
async def store_unified_file_id(self, **kwargs):
raise RuntimeError("db down")
class _FakeProxyLogging:
def __init__(self, hook):
self._hook = hook
@ -57,7 +72,7 @@ def _file_data():
@pytest.mark.asyncio
async def test_upload_with_target_model_names_but_no_hook_raises_before_uploading(monkeypatch):
backend = _RecordingStorageBackend()
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend)
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name, prisma_client=None: backend)
with pytest.raises(ProxyException) as exc_info:
await StorageBackendFileService.upload_file_to_storage_backend(
@ -80,7 +95,7 @@ async def test_upload_with_target_model_names_but_no_hook_raises_before_uploadin
@pytest.mark.asyncio
async def test_upload_without_target_model_names_skips_hook_requirement(monkeypatch):
backend = _RecordingStorageBackend()
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend)
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name, prisma_client=None: backend)
file_object = await StorageBackendFileService.upload_file_to_storage_backend(
file_data=_file_data(),
@ -101,7 +116,7 @@ async def test_upload_without_target_model_names_skips_hook_requirement(monkeypa
@pytest.mark.asyncio
async def test_upload_with_target_model_names_and_hook_stores_unified_id(monkeypatch):
backend = _RecordingStorageBackend()
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend)
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name, prisma_client=None: backend)
hook = _FakeManagedFilesHook()
file_object = await StorageBackendFileService.upload_file_to_storage_backend(
@ -125,3 +140,50 @@ async def test_upload_with_target_model_names_and_hook_stores_unified_id(monkeyp
"stored_id_matches_response": True,
"model_mappings": {"gpt-x": "https://storage.example/blob-1"},
}
@pytest.mark.asyncio
async def test_upload_hands_the_prisma_client_to_the_storage_backend_factory(monkeypatch: pytest.MonkeyPatch):
backend = _RecordingStorageBackend()
factory_calls: list[tuple[str, PrismaClient | None]] = []
def _factory(name: str, prisma_client: PrismaClient | None = None) -> _RecordingStorageBackend:
factory_calls.append((name, prisma_client))
return backend
monkeypatch.setattr(storage_backend_service, "get_storage_backend", _factory)
prisma_client = MagicMock()
await StorageBackendFileService.upload_file_to_storage_backend(
file_data=_file_data(),
target_storage="litellm_db",
target_model_names=[],
purpose="batch",
proxy_logging_obj=_FakeProxyLogging(hook=None),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
prisma_client=prisma_client,
)
assert factory_calls == [("litellm_db", prisma_client)]
@pytest.mark.asyncio
@pytest.mark.parametrize("delete_error", [None, OSError("blob locked")], ids=["delete succeeds", "delete fails"])
async def test_upload_deletes_the_uploaded_content_when_the_metadata_write_fails(
monkeypatch: pytest.MonkeyPatch, delete_error: Exception | None
):
backend = _RecordingStorageBackend(delete_error=delete_error)
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name, prisma_client=None: backend)
with pytest.raises(RuntimeError, match="db down"):
await StorageBackendFileService.upload_file_to_storage_backend(
file_data=_file_data(),
target_storage="azure_storage",
target_model_names=["gpt-x"],
purpose="batch",
proxy_logging_obj=_FakeProxyLogging(hook=_FailingManagedFilesHook()),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
)
assert len(backend.upload_calls) == 1
assert backend.delete_calls == ["https://storage.example/blob-1"]