From fd45412c89df25928ff186f21b602b387df492f2 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 01:47:43 -0700 Subject: [PATCH 01/12] feat(batches): run hosted_vllm batches inside LiteLLM vLLM serves no /v1/files or /v1/batches, so a hosted_vllm deployment can never host a batch. Batch inputs for such a deployment now land in a LiteLLM-owned storage backend, the batch is executed line by line through the deployment's own chat, completion, embedding, or responses route, and the batch plus its output and error files are served back from the database under the creating key --- .../proxy/hooks/managed_files.py | 99 ++- .../migration.sql | 8 + litellm/constants.py | 1 + .../files/litellm_db_storage_backend.py | 65 ++ .../base_llm/files/storage_backend_factory.py | 26 +- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- litellm/proxy/batches_endpoints/endpoints.py | 116 ++- .../litellm_executed_batches.py | 562 ++++++++++++++ .../openai_files_endpoints/common_utils.py | 5 + .../openai_files_endpoints/files_endpoints.py | 67 +- .../storage_backend_service.py | 16 +- litellm/proxy/schema.prisma | 6 + litellm/types/llms/openai.py | 1 + litellm/types/utils.py | 2 + schema.prisma | 6 + tests/e2e/batches/test_batches_e2e.py | 158 +++- .../proxy/test_managed_files_hook.py | 135 +++- .../files/test_litellm_db_storage_backend.py | 92 +++ .../files/test_storage_backend_factory.py | 28 + .../proxy/batches_endpoints/test_endpoints.py | 189 ++++- .../test_litellm_executed_batches.py | 715 ++++++++++++++++++ .../test_files_common_utils.py | 13 + .../test_files_endpoint.py | 115 +++ .../test_storage_backend_service.py | 34 +- 24 files changed, 2339 insertions(+), 122 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260918000000_add_managed_file_content_table/migration.sql create mode 100644 litellm/llms/base_llm/files/litellm_db_storage_backend.py create mode 100644 litellm/proxy/batches_endpoints/litellm_executed_batches.py create mode 100644 tests/test_litellm/llms/base_llm/files/test_litellm_db_storage_backend.py create mode 100644 tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py create mode 100644 tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 4899b87da7a..8eef8a5f1ce 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -34,6 +34,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 +60,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, @@ -204,6 +206,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 +241,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 +251,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 +280,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 +326,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 +341,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 +396,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 +1361,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 +1815,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,6 +1828,39 @@ 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, @@ -1889,16 +1930,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)}" diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260918000000_add_managed_file_content_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260918000000_add_managed_file_content_table/migration.sql new file mode 100644 index 00000000000..bb1a3eab6ee --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260918000000_add_managed_file_content_table/migration.sql @@ -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") +); diff --git a/litellm/constants.py b/litellm/constants.py index d62cad74a36..bbeb4846e27 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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" diff --git a/litellm/llms/base_llm/files/litellm_db_storage_backend.py b/litellm/llms/base_llm/files/litellm_db_storage_backend.py new file mode 100644 index 00000000000..bca4b8f4c6f --- /dev/null +++ b/litellm/llms/base_llm/files/litellm_db_storage_backend.py @@ -0,0 +1,65 @@ +from collections.abc import Mapping +from typing import TYPE_CHECKING, Final + +from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend +from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.table_repositories import PrismaTableRepository + +if TYPE_CHECKING: + from prisma import models as prisma_models + + 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) + + +def _where_id(storage_url: str) -> Mapping[str, str]: + return {"id": storage_url_to_row_id(storage_url)} # mutable-ok: Prisma filter + + +class ManagedFileContentRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileContentTable"]): + table_name = "litellm_managedfilecontenttable" + + +class LiteLLMDbStorageBackend(BaseFileStorageBackend): + def __init__(self, prisma_client: "PrismaClient") -> None: + self._prisma_client = prisma_client + + @property + def _table(self) -> "TableActions[prisma_models.LiteLLM_ManagedFileContentTable]": + return ManagedFileContentRepository(self._prisma_client).table + + async def upload_file( + self, + file_content: bytes, + filename: str, + content_type: str, + path_prefix: str | None = None, + file_naming_strategy: str = "uuid", + ) -> str: + from prisma import Base64 + + data: Final = {"content": Base64.encode(file_content)} # mutable-ok: Prisma payload + row: Final = await self._table.create(data=data) + return f"{LITELLM_DB_STORAGE_URL_PREFIX}{row.id}" + + async def download_file(self, storage_url: str) -> bytes: + row: Final = await self._table.find_unique(where=_where_id(storage_url)) + if row is None: + raise ValueError(f"No stored file content for {storage_url}") + return row.content.decode() + + async def delete_file(self, storage_url: str) -> None: + from prisma.errors import RecordNotFoundError + + try: + await self._table.delete(where=_where_id(storage_url)) + except RecordNotFoundError: + return diff --git a/litellm/llms/base_llm/files/storage_backend_factory.py b/litellm/llms/base_llm/files/storage_backend_factory.py index 0cf8164bc4a..e126da44d0a 100644 --- a/litellm/llms/base_llm/files/storage_backend_factory.py +++ b/litellm/llms/base_llm/files/storage_backend_factory.py @@ -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}" + ) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 73ea8cf1991..18849ef5b64 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -19622,7 +19622,7 @@ } } }, - "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " + "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" }, "500": { "content": { diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 5d9ecddd4c2..cc698ad760e 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -11,6 +11,7 @@ from types import MappingProxyType from typing import Any, Final, 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 +19,14 @@ 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_http_error, + litellm_executed_provider_of, + resolve_litellm_executed_provider, +) from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, log_llm_api_exception, @@ -45,16 +54,55 @@ 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 ProxyLogging, handle_exception_on_proxy, is_known_model 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 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 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_http_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, + proxy_logging_obj=proxy_logging_obj, + ) + + +def _raise_when_input_file_must_be_managed(model: str, credentials: Mapping[str, object]) -> None: + if litellm_executed_provider_of(credentials) is None: + return + raise batch_http_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( @@ -99,6 +147,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)], @@ -292,24 +358,33 @@ async def create_batch( model: Final = target_model_names[0] _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 = 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 @@ -320,6 +395,7 @@ async def create_batch( model_id=model_param, operation_context="batch creation", ) + _raise_when_input_file_must_be_managed(model_param, credentials) prepare_data_with_credentials( data=_create_batch_data, @@ -478,15 +554,15 @@ async def retrieve_batch( verbose_proxy_logger=verbose_proxy_logger, ) + executed_batch: Final = isinstance(unified_batch_id, str) and is_litellm_executed_batch(unified_batch_id) + if executed_batch and response is None: + raise batch_http_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 @@ -989,6 +1065,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_http_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( diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py new file mode 100644 index 00000000000..f567f0e2263 --- /dev/null +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -0,0 +1,562 @@ +import asyncio +import json +import time +from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass +from itertools import pairwise +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable + +from fastapi import HTTPException +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, assert_never + +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.models.managed_files import LiteLLM_ManagedFileTable +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.openai_files_endpoints.common_utils import ( + LITELLM_EXECUTED_BATCH_ID_PREFIX, + convert_b64_uid_to_unified_uid, + get_batch_id_from_unified_batch_id, +) +from litellm.proxy.openai_files_endpoints.storage_backend_service import StorageBackendFileService +from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.repositories.table_repositories import ManagedObjectRepository +from litellm.types.llms.openai import LiteLLMBatchCreateRequest, OpenAIFileObject, OpenAIFilesPurpose +from litellm.types.utils import LITELLM_EXECUTED_BATCH_PROVIDERS, ExtractedFileData, LiteLLMBatch + +if TYPE_CHECKING: + from prisma import models as prisma_models + + 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"] + +TERMINAL_BATCH_STATUSES: Final[frozenset[str]] = frozenset({"completed", "failed", "cancelled", "expired"}) +_BATCH_ENDPOINT_ADAPTER: Final[TypeAdapter[BatchEndpoint]] = TypeAdapter(BatchEndpoint) +_CANCEL_POLL_SECONDS: Final = 1.0 +_COMPLETION_WINDOW_SECONDS: Final = 24 * 60 * 60 +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 _ResultLine(TypedDict): + id: ReadOnly[str] + custom_id: ReadOnly[str] + response: ReadOnly[_ResultResponse] + error: ReadOnly[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 _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, ...] + + +@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, + create_if_missing: bool = True, + 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 _router_method_name(endpoint: BatchEndpoint) -> str: + match endpoint: + case "/v1/chat/completions": + return "acompletion" + case "/v1/completions": + return "atext_completion" + case "/v1/embeddings": + return "aembedding" + case "/v1/responses": + return "aresponses" + case _: + assert_never(endpoint) + + +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 + + +def resolve_litellm_executed_provider(llm_router: "Router", model: str, team_id: str | None) -> 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 litellm_executed_provider_of(credentials) + + +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 _parse_line(line_number: int, raw: bytes, endpoint: BatchEndpoint) -> 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") + return line + + +def parse_batch_input(content: bytes, endpoint: BatchEndpoint) -> 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) 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_http_error(status_code: int, message: str) -> HTTPException: + detail: Final = {"error": message} # mutable-ok: HTTPException detail must be a plain mapping + return HTTPException(status_code=status_code, detail=detail) + + +def _validate_endpoint(endpoint: object) -> BatchEndpoint: + try: + return _BATCH_ENDPOINT_ADAPTER.validate_python(endpoint) + except ValidationError: + raise batch_http_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 _batch_of(blob: object) -> LiteLLMBatch: + return LiteLLMBatch.model_validate_json(blob) if isinstance(blob, str) else LiteLLMBatch.model_validate(blob) + + +def _error_body(error: Exception) -> _ErrorBody: + body: Final[_ErrorBody] = { + "error": {"message": str(error), "type": type(error).__name__, "param": None, "code": None} + } + return body + + +def _result_line(outcome: RowOutcome) -> _ResultLine: + line: Final[_ResultLine] = { + "id": f"batch_req_{uuid_module.uuid4().hex[:24]}", + "custom_id": outcome.custom_id, + "response": { + "status_code": outcome.status_code, + "request_id": f"req_{uuid_module.uuid4().hex[:24]}", + "body": outcome.body, + }, + "error": None, + } + 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 + match requested: + case "completed": + return "cancelled" + case "in_progress" | "finalizing": + return "cancelling" + case "failed" | "cancelling" | "cancelled": + return requested + case _: + assert_never(requested) + + +def _llm_batch_id_of(unified_batch_id: str) -> str: + return get_batch_id_from_unified_batch_id(convert_b64_uid_to_unified_uid(unified_batch_id)) + + +class _CancelWatch: + 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._cancelling = False + + async def cancelling(self) -> bool: + if self._cancelling: + return True + now: Final = time.monotonic() + if now - self._checked_at < self._interval_seconds: + return False + self._checked_at = now + self._cancelling = await self._load_status() == "cancelling" + return self._cancelling + + +class LiteLLMExecutedBatchRunner: + def __init__( + self, + llm_router: "Router", + prisma_client: PrismaClient, + managed_files: ManagedBatchStore, + proxy_logging_obj: ProxyLogging, + concurrency: int = LITELLM_EXECUTED_BATCH_CONCURRENCY, + 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.proxy_logging_obj = proxy_logging_obj + self.concurrency = concurrency + 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) + if isinstance(parsed, InvalidBatchInput): + raise batch_http_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) + created_at: Final = int(time.time()) + 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 + _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 ()), + ) + 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._load_batch(unified_batch_id) + if current is None: + raise batch_http_error(404, f"Batch {unified_batch_id} not found") + if current.status in TERMINAL_BATCH_STATUSES: + raise batch_http_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())}) + ) + await self._store(cancelling, user_api_key_dict) + return cancelling + + 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_http_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_http_error(400, str(e)) + + async def _run(self, run: _BatchRun) -> None: + 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 + ) + + async def _execute(self, run: _BatchRun) -> None: + await self._advance(run, "in_progress") + watch: Final = _CancelWatch(lambda: self._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) + await self._advance(run, "finalizing") + succeeded: Final = tuple(outcome for outcome in outcomes if outcome.succeeded) + failed: Final = tuple(outcome for outcome in outcomes if 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)) + await self._advance( + run, + "completed", + 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: _CancelWatch, semaphore: asyncio.Semaphore + ) -> RowOutcome | None: + async with semaphore: + if await watch.cancelling(): + return None + 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)}) + return _dump(await self._router_call(run.endpoint)(**params)) + + def _router_call(self, endpoint: BatchEndpoint) -> _RouterCall: + method: Final[object] = getattr(self.llm_router, _router_method_name(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] + ) -> 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) -> None: + current: Final = await self._load_batch(run.unified_batch_id) + if current is None: + raise RuntimeError(f"Batch {run.unified_batch_id} is no longer stored") + status: Final = _resolve_transition(current.status, requested) + updated: Final = current.model_copy( + update=MappingProxyType({**fields, "status": status, f"{status}_at": int(time.time())}) + ) + await self._store(updated, run.user_api_key_dict) + + async def _store(self, batch: LiteLLMBatch, user_api_key_dict: UserAPIKeyAuth) -> None: + await self.managed_files.store_unified_object_id( + unified_object_id=batch.id, + file_object=batch, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + model_object_id=_llm_batch_id_of(batch.id), + file_purpose="batch", + user_api_key_dict=user_api_key_dict, + create_if_missing=False, + ) + + async def _find_row(self, unified_batch_id: str) -> "prisma_models.LiteLLM_ManagedObjectTable | None": + return await ManagedObjectRepository(self.prisma_client).table.find_first( + where={"unified_object_id": unified_batch_id} # mutable-ok: Prisma filter + ) + + 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 + + +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 "", + ) diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 38a907892b4..c45d08c5546 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -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,10 @@ 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: + return get_batch_id_from_unified_batch_id(decoded_unified_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. diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index ae6e222a863..ad869100fb9 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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,12 @@ 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 resolve_litellm_executed_provider from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, @@ -86,7 +88,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 +101,39 @@ from litellm.types.llms.openai import ( router: Final = APIRouter() + +def _litellm_executed_batch_input_model( + llm_router: Router | None, + purpose: OpenAIFilesPurpose, + model: str | None, + target_model_names_list: Sequence[str], + team_id: str | None, +) -> str | None: + if purpose != "batch" or llm_router is None: + return None + candidates: Final = (model,) if model is not None else tuple(target_model_names_list) + executed: Final = tuple( + candidate + for candidate in candidates + if resolve_litellm_executed_provider(llm_router, candidate, team_id) is not None + ) + match executed: + case (): + return None + case (only,) if len(candidates) == 1: + return only + case _: + 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 +279,30 @@ 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": + executed_model: Final = _litellm_executed_batch_input_model( + llm_router, purpose, model, target_model_names_list, user_api_key_dict.team_id + ) + explicit_storage: Final = target_storage if target_storage and target_storage != "default" else None + 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 @@ -847,7 +882,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} ) @@ -862,7 +897,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 diff --git a/litellm/proxy/openai_files_endpoints/storage_backend_service.py b/litellm/proxy/openai_files_endpoints/storage_backend_service.py index e766f335071..b4a36336c22 100644 --- a/litellm/proxy/openai_files_endpoints/storage_backend_service.py +++ b/litellm/proxy/openai_files_endpoints/storage_backend_service.py @@ -7,7 +7,7 @@ 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 @@ -15,7 +15,7 @@ from litellm._uuid import uuid as uuid_module 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 +35,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 +61,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), @@ -164,7 +166,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 +196,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, diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 91b59e56906..c4606796ebf 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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 diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 632efcc3c4f..d9e8b545765 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -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): diff --git a/litellm/types/utils.py b/litellm/types/utils.py index c63d971b89b..12535493eb5 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -4143,6 +4143,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)) diff --git a/schema.prisma b/schema.prisma index 91b59e56906..c4606796ebf 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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 diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py index eff8f297f25..9b3c06d9a1b 100644 --- a/tests/e2e/batches/test_batches_e2e.py +++ b/tests/e2e/batches/test_batches_e2e.py @@ -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 diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index cb08e00ff65..0ae4e3a5fd8 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -1732,8 +1732,38 @@ 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 -async def test_afile_delete_passes_trusted_model_credentials_to_router(): +@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 """ afile_delete must hand the deployment's credential snapshot to the router call, since Bedrock validates the s3:// file id against the bucket in it. @@ -1743,6 +1773,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 +1840,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 +1859,104 @@ 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_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") diff --git a/tests/test_litellm/llms/base_llm/files/test_litellm_db_storage_backend.py b/tests/test_litellm/llms/base_llm/files/test_litellm_db_storage_backend.py new file mode 100644 index 00000000000..fabcb340a48 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/files/test_litellm_db_storage_backend.py @@ -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") diff --git a/tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py b/tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py new file mode 100644 index 00000000000..39b0adb56fc --- /dev/null +++ b/tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py @@ -0,0 +1,28 @@ +from unittest.mock import MagicMock + +import pytest + +from litellm.llms.base_llm.files.litellm_db_storage_backend import ( + LITELLM_DB_STORAGE_BACKEND_NAME, + LiteLLMDbStorageBackend, +) +from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend + + +def test_litellm_db_backend_is_built_on_the_given_prisma_client(): + prisma_client = MagicMock() + + backend = get_storage_backend(LITELLM_DB_STORAGE_BACKEND_NAME, prisma_client=prisma_client) + + assert isinstance(backend, LiteLLMDbStorageBackend) + assert backend._table is prisma_client.db.litellm_managedfilecontenttable + + +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()) diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index d9bfb3fe3da..e655438a672 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -51,7 +51,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( from litellm.proxy.utils import ProxyLogging from litellm.router import Router from litellm.types.llms.openai import BatchJobStatus -from litellm.types.utils import CredentialItem, LiteLLMBatch +from litellm.types.utils import CredentialItem, LiteLLMBatch, SpecialEnums from fastapi import Request, Response @@ -73,6 +73,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". @@ -161,9 +167,10 @@ 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 @@ -250,6 +257,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. @@ -761,6 +787,98 @@ 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_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 async def test_create__model_encoded_beats_unified(harness): """Precedence row: a file id that is BOTH model-encoded and (pretend) unified @@ -1141,6 +1259,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: @@ -1546,6 +1669,33 @@ 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() + 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_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. # --------------------------------------------------------------------------- # @@ -2257,6 +2407,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}. @@ -2774,8 +2953,6 @@ async def test_cancel__unified_batch_id_allowed_when_managed_files_required(canc assert cancel_harness.router_acancel.call_count == 1 - - @pytest.mark.asyncio async def test_retrieve__managed_batch_defers_cost_to_the_poller_when_it_is_running(retrieve_harness): with patch.object(endpoints, "batch_cost_poller_is_active", MagicMock(return_value=True)): diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py new file mode 100644 index 00000000000..0806dd3451e --- /dev/null +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -0,0 +1,715 @@ +import asyncio +import json +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from typing import Final, Literal, cast +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException +from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles +from openai.types.batch_request_counts import BatchRequestCounts + +from litellm.models.managed_files import LiteLLM_ManagedFileTable +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.batches_endpoints import litellm_executed_batches +from litellm.proxy.batches_endpoints.litellm_executed_batches import ( + BatchEndpoint, + BatchInputLine, + BatchStatus, + InvalidBatchInput, + LiteLLMExecutedBatchRunner, + _resolve_transition, + litellm_executed_provider_of, + parse_batch_input, + resolve_litellm_executed_provider, +) +from litellm.proxy.openai_files_endpoints.common_utils import ( + _is_base64_encoded_unified_file_id, + get_batch_id_from_unified_batch_id, + is_litellm_executed_batch, +) +from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.router import Router +from litellm.types.llms.openai import LiteLLMBatchCreateRequest, OpenAIFileObject, OpenAIFilesPurpose +from litellm.types.utils import EmbeddingResponse, LiteLLMBatch, ModelResponse, SpecialEnums + +BATCH_MODEL: Final = "batch-model" +DEPLOYMENT_ID: Final = "deployment-id-1" +INPUT_FILE_ID: Final = "unified-input-file" +STORAGE_BACKEND: Final = "s3" +STORAGE_URL: Final = "s3://bucket/input.jsonl" +CHAT_ENDPOINT: Final = "/v1/chat/completions" +ROUTER_METHODS: Final = ("acompletion", "atext_completion", "aembedding", "aresponses") +ALL_STATUSES: Final[tuple[BatchStatus, ...]] = ( + "in_progress", + "finalizing", + "completed", + "failed", + "cancelling", + "cancelled", +) + + +def chat_row(custom_id: str, content: str, **body_extra: object) -> dict[str, object]: + return { + "custom_id": custom_id, + "method": "POST", + "url": CHAT_ENDPOINT, + "body": {"model": "row-model", "messages": [{"role": "user", "content": content}], **body_extra}, + } + + +def jsonl(*rows: Mapping[str, object]) -> bytes: + return "".join(f"{json.dumps(row)}\n" for row in rows).encode() + + +TWO_CHAT_ROWS: Final = jsonl(chat_row("row-1", "hi 1"), chat_row("row-2", "hi 2")) + + +def chat_response(content: str) -> ModelResponse: + return ModelResponse( + id=f"chatcmpl-{content}", + model=BATCH_MODEL, + choices=[{"index": 0, "message": {"role": "assistant", "content": f"echo {content}"}, "finish_reason": "stop"}], + ) + + +def managed_input_file(storage_backend: str | None = STORAGE_BACKEND) -> LiteLLM_ManagedFileTable: + return LiteLLM_ManagedFileTable( + unified_file_id=INPUT_FILE_ID, + model_mappings={}, + flat_model_file_ids=[], + storage_backend=storage_backend, + storage_url=STORAGE_URL, + ) + + +def batch_request(endpoint: str) -> LiteLLMBatchCreateRequest: + return cast( + "LiteLLMBatchCreateRequest", + {"endpoint": endpoint, "input_file_id": INPUT_FILE_ID, "completion_window": "24h"}, + ) + + +class ProviderRateLimited(Exception): + status_code = 429 + + +@dataclass(frozen=True, slots=True) +class StoredObject: + file_object: str + status: str + + +@dataclass(frozen=True, slots=True) +class StoreCall: + unified_object_id: str + model_object_id: str + status: str + request_tags: tuple[str, ...] | None + persist_attribution: bool + create_if_missing: bool + batch_processed: bool + + +class FakeManagedBatchStore: + def __init__(self, files: Mapping[str, LiteLLM_ManagedFileTable]) -> None: + self.files = files + self.objects: dict[str, StoredObject] = {} + self.calls: list[StoreCall] = [] + + def get_unified_batch_id(self, batch_id: str, model_id: str) -> str: + return SpecialEnums.LITELLM_MANAGED_BATCH_COMPLETE_STR.value.format(model_id, batch_id) + + async def get_unified_file_id( + self, file_id: str, litellm_parent_otel_span: object | None = None + ) -> LiteLLM_ManagedFileTable | None: + return self.files.get(file_id) + + 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, + create_if_missing: bool = True, + batch_processed: bool = False, + ) -> None: + self.calls.append( + StoreCall( + unified_object_id=unified_object_id, + model_object_id=model_object_id, + status=file_object.status, + request_tags=tuple(request_tags) if request_tags is not None else None, + persist_attribution=persist_attribution, + create_if_missing=create_if_missing, + batch_processed=batch_processed, + ) + ) + if create_if_missing or unified_object_id in self.objects: + self.write(file_object) + + def write(self, batch: LiteLLMBatch) -> None: + self.objects[batch.id] = StoredObject(file_object=batch.model_dump_json(), status=batch.status) + + def batch(self, unified_batch_id: str) -> LiteLLMBatch: + return LiteLLMBatch.model_validate_json(self.objects[unified_batch_id].file_object) + + +REAL_HOOK: Final = _PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()) + + +class RealIdManagedBatchStore(FakeManagedBatchStore): + def get_unified_batch_id(self, batch_id: str, model_id: str) -> str: + return REAL_HOOK.get_unified_batch_id(batch_id=batch_id, model_id=model_id) + + +class FakeManagedObjectTable: + def __init__(self, objects: Mapping[str, StoredObject]) -> None: + self.objects = objects + + async def find_first(self, where: Mapping[str, str]) -> StoredObject | None: + return self.objects.get(where["unified_object_id"]) + + +class FakeDb: + def __init__(self, objects: Mapping[str, StoredObject]) -> None: + self.litellm_managedobjecttable = FakeManagedObjectTable(objects) + + +class FakePrismaClient: + def __init__(self, objects: Mapping[str, StoredObject]) -> None: + self.db = FakeDb(objects) + + +class FakeRouter: + def __init__(self) -> None: + self.acompletion = AsyncMock(return_value=chat_response("default")) + self.atext_completion = AsyncMock(return_value=chat_response("default")) + self.aembedding = AsyncMock( + return_value=EmbeddingResponse( + model=BATCH_MODEL, data=[{"embedding": [0.1], "index": 0, "object": "embedding"}] + ) + ) + self.aresponses = AsyncMock(return_value=chat_response("default")) + + def get_model_ids(self, model_name: str) -> list[str]: + return [DEPLOYMENT_ID] if model_name == BATCH_MODEL else [] + + +class FakeStorageBackend: + def __init__(self, contents: Mapping[str, bytes]) -> None: + self.contents = contents + self.downloads: list[str] = [] + + async def download_file(self, storage_url: str) -> bytes: + self.downloads.append(storage_url) + return self.contents[storage_url] + + +class FakeStorageBackendFactory: + def __init__(self, backend: FakeStorageBackend, error: ValueError | None) -> None: + self.backend = backend + self.error = error + self.calls: list[tuple[str, object]] = [] + + def __call__(self, backend_type: str, prisma_client: object = None) -> FakeStorageBackend: + self.calls.append((backend_type, prisma_client)) + if self.error is not None: + raise self.error + return self.backend + + +@dataclass(frozen=True, slots=True) +class UploadCall: + content: bytes + filename: str + target_storage: str + target_model_names: tuple[str, ...] + purpose: str + user_api_key_dict: UserAPIKeyAuth + prisma_client: object + + def lines(self) -> dict[str, dict[str, object]]: + parsed = tuple(json.loads(line) for line in self.content.decode().splitlines()) + return {str(line["custom_id"]): line for line in parsed} + + +class FakeResultFileUploader: + def __init__(self, error: Exception | None) -> None: + self.error = error + self.calls: list[UploadCall] = [] + + async def __call__( + self, + file_data: Mapping[str, object], + target_storage: str, + target_model_names: list[str], + purpose: OpenAIFilesPurpose, + proxy_logging_obj: ProxyLogging, + user_api_key_dict: UserAPIKeyAuth, + prisma_client: object = None, + ) -> OpenAIFileObject: + content = file_data["content"] + assert isinstance(content, bytes) + self.calls.append( + UploadCall( + content=content, + filename=str(file_data["filename"]), + target_storage=target_storage, + target_model_names=tuple(target_model_names), + purpose=purpose, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) + ) + if self.error is not None: + raise self.error + return OpenAIFileObject( + id=f"unified-output-{len(self.calls)}", + object="file", + bytes=len(content), + created_at=0, + filename=str(file_data["filename"]), + purpose=purpose, + status="uploaded", + ) + + +@dataclass(frozen=True, slots=True) +class Harness: + runner: LiteLLMExecutedBatchRunner + store: FakeManagedBatchStore + router: FakeRouter + uploads: FakeResultFileUploader + storage: FakeStorageBackend + storage_factory: FakeStorageBackendFactory + prisma: FakePrismaClient + user: UserAPIKeyAuth + + async def create(self, endpoint: str = CHAT_ENDPOINT) -> LiteLLMBatch: + return await self.runner.create( + create_request=batch_request(endpoint), + unified_input_file_id=INPUT_FILE_ID, + model=BATCH_MODEL, + provider="hosted_vllm", + user_api_key_dict=self.user, + request_tags=["tag-a"], + ) + + async def create_and_finish(self, endpoint: str = CHAT_ENDPOINT) -> tuple[LiteLLMBatch, LiteLLMBatch]: + created = await self.create(endpoint) + await asyncio.gather(*list(litellm_executed_batches._RUNNING_BATCHES)) + return created, self.store.batch(created.id) + + +def make_runner( + content: bytes = TWO_CHAT_ROWS, + concurrency: int = 4, + files: Mapping[str, LiteLLM_ManagedFileTable] | None = None, + upload_error: Exception | None = None, + storage_error: ValueError | None = None, + store_factory: Callable[[Mapping[str, LiteLLM_ManagedFileTable]], FakeManagedBatchStore] = FakeManagedBatchStore, +) -> Harness: + store = store_factory({INPUT_FILE_ID: managed_input_file()} if files is None else files) + router = FakeRouter() + uploads = FakeResultFileUploader(upload_error) + storage = FakeStorageBackend({STORAGE_URL: content}) + storage_factory = FakeStorageBackendFactory(storage, storage_error) + prisma = FakePrismaClient(store.objects) + user = UserAPIKeyAuth( + api_key="sk-batch-key", user_id="user-1", team_id="team-1", key_alias="alias-1", user_email="user@example.com" + ) + runner = LiteLLMExecutedBatchRunner( + llm_router=cast("Router", router), + prisma_client=cast("PrismaClient", prisma), + managed_files=store, + proxy_logging_obj=MagicMock(spec=ProxyLogging), + concurrency=concurrency, + storage_backend_factory=storage_factory, + upload_result_file=uploads, + ) + return Harness(runner, store, router, uploads, storage, storage_factory, prisma, user) + + +def seeded_batch(store: FakeManagedBatchStore, status: Literal["in_progress", "completed"]) -> LiteLLMBatch: + batch = LiteLLMBatch( + id=store.get_unified_batch_id(batch_id="litellm_batch_seed", model_id=DEPLOYMENT_ID), + object="batch", + endpoint=CHAT_ENDPOINT, + input_file_id=INPUT_FILE_ID, + completion_window="24h", + status=status, + created_at=1, + model=BATCH_MODEL, + ) + store.write(batch) + return batch + + +@pytest.mark.parametrize( + ("content", "line_number", "reason_fragment"), + [ + (b"", None, "no requests"), + (b"\n \n", None, "no requests"), + (b"{not json", 1, "JSON"), + (jsonl({"custom_id": "a", "method": "POST", "url": CHAT_ENDPOINT}), 1, "body"), + (jsonl({**chat_row("a", "hi"), "extra_field": 1}), 1, "extra_field"), + ( + jsonl(chat_row("a", "hi")) + b"\n" + jsonl({**chat_row("b", "hi"), "url": "/v1/embeddings"}), + 3, + "/v1/embeddings", + ), + (jsonl(chat_row("a", "hi", stream=True)), 1, "streaming"), + (jsonl(chat_row("a", "hi"), chat_row("a", "again")), None, "'a'"), + ], + ids=["empty", "blank lines", "not json", "missing body", "unknown field", "url mismatch", "stream", "duplicate id"], +) +def test_parse_batch_input_rejects(content: bytes, line_number: int | None, reason_fragment: str) -> None: + result = parse_batch_input(content, CHAT_ENDPOINT) + assert isinstance(result, InvalidBatchInput) + assert result.line_number == line_number + assert reason_fragment in result.reason + + +def test_parse_batch_input_keeps_every_request_and_skips_blank_lines() -> None: + content = b"\n" + jsonl(chat_row("a", "hi 1")) + b"\n" + jsonl(chat_row("b", "hi 2")) + b"\n\n" + lines = parse_batch_input(content, CHAT_ENDPOINT) + assert isinstance(lines, tuple) + assert [line.custom_id for line in lines] == ["a", "b"] + assert lines[1] == BatchInputLine( + custom_id="b", + method="POST", + url=CHAT_ENDPOINT, + body={"model": "row-model", "messages": [{"role": "user", "content": "hi 2"}]}, + ) + + +@pytest.mark.parametrize("current", ["validating", "in_progress", "finalizing"]) +@pytest.mark.parametrize("requested", ALL_STATUSES) +def test_resolve_transition_keeps_the_requested_status_unless_cancelling(current: str, requested: BatchStatus) -> None: + assert _resolve_transition(current, requested) == requested + + +@pytest.mark.parametrize( + ("requested", "expected"), + [ + ("completed", "cancelled"), + ("in_progress", "cancelling"), + ("finalizing", "cancelling"), + ("failed", "failed"), + ("cancelling", "cancelling"), + ("cancelled", "cancelled"), + ], +) +def test_resolve_transition_from_cancelling(requested: BatchStatus, expected: BatchStatus) -> None: + assert _resolve_transition("cancelling", requested) == expected + + +@pytest.mark.parametrize( + ("credentials", "expected"), + [ + ({"custom_llm_provider": "hosted_vllm", "model": "openai/gpt-4o"}, "hosted_vllm"), + ({"model": "hosted_vllm/qwen"}, "hosted_vllm"), + ({"custom_llm_provider": "openai", "model": "gpt-4o"}, None), + ({"model": "gpt-4o"}, None), + ], + ids=["explicit hosted_vllm", "model prefix", "explicit openai", "openai model"], +) +def test_litellm_executed_provider_of(credentials: Mapping[str, object], expected: str | None) -> None: + assert litellm_executed_provider_of(credentials) == expected + + +@pytest.mark.parametrize( + ("credentials", "expected"), [(None, None), ({"model": "hosted_vllm/qwen"}, "hosted_vllm")], ids=["unknown", "vllm"] +) +def test_resolve_litellm_executed_provider_asks_the_router_for_the_team_scoped_deployment( + credentials: Mapping[str, object] | None, expected: str | None +) -> None: + router = MagicMock(spec=Router) + router.get_deployment_credentials_with_provider.return_value = credentials + assert resolve_litellm_executed_provider(router, BATCH_MODEL, "team-1") == expected + router.get_deployment_credentials_with_provider.assert_called_once_with(model_id=BATCH_MODEL, team_id="team-1") + + +async def test_create_stores_a_validating_batch_and_completes_it_in_the_background() -> None: + harness = make_runner() + created, finished = await harness.create_and_finish() + + assert created.status == "validating" + assert is_litellm_executed_batch(created.id) + assert created.id.startswith(f"litellm_proxy;model_id:{DEPLOYMENT_ID};llm_batch_id:litellm_batch_") + assert (created.model, created.input_file_id) == (BATCH_MODEL, INPUT_FILE_ID) + assert created.request_counts == BatchRequestCounts(completed=0, failed=0, total=2) + first_write = harness.store.calls[0] + assert (first_write.unified_object_id, first_write.model_object_id) == ( + created.id, + get_batch_id_from_unified_batch_id(created.id), + ) + assert (first_write.persist_attribution, first_write.batch_processed, first_write.request_tags) == ( + True, + True, + ("tag-a",), + ) + assert harness.storage_factory.calls == [(STORAGE_BACKEND, harness.prisma)] + assert harness.storage.downloads == [STORAGE_URL] + + assert finished.status == "completed" + assert finished.request_counts == BatchRequestCounts(completed=2, failed=0, total=2) + assert (finished.output_file_id, finished.error_file_id) == ("unified-output-1", None) + assert finished.in_progress_at is not None + assert finished.completed_at is not None + + +async def test_create_dispatches_each_row_with_the_batch_model_and_the_key_metadata() -> None: + harness = make_runner() + created, _ = await harness.create_and_finish() + + calls = {call.kwargs["messages"][0]["content"]: call.kwargs for call in harness.router.acompletion.await_args_list} + assert set(calls) == {"hi 1", "hi 2"} + for content, kwargs in calls.items(): + assert kwargs["model"] == BATCH_MODEL + assert kwargs["messages"] == [{"role": "user", "content": content}] + metadata = kwargs["metadata"] + assert metadata["user_api_key"] == harness.user.api_key + assert metadata["tags"] == ["tag-a"] + assert metadata["batch_id"] == created.id + assert metadata["user_api_key_user_id"] == "user-1" + assert metadata["user_api_key_team_id"] == "team-1" + assert metadata["user_api_key_alias"] == "alias-1" + assert metadata["user_api_key_user_email"] == "user@example.com" + + +async def test_create_uploads_one_output_line_per_row_with_the_router_response() -> None: + harness = make_runner() + replies = {"hi 1": chat_response("hi 1"), "hi 2": chat_response("hi 2")} + harness.router.acompletion.side_effect = lambda **kwargs: replies[kwargs["messages"][0]["content"]] + created, _ = await harness.create_and_finish() + + assert len(harness.uploads.calls) == 1 + upload = harness.uploads.calls[0] + assert (upload.target_storage, upload.purpose, upload.target_model_names) == ( + "litellm_db", + "batch_output", + (BATCH_MODEL,), + ) + assert upload.filename == f"{get_batch_id_from_unified_batch_id(created.id)}_output.jsonl" + assert upload.user_api_key_dict is harness.user + assert upload.prisma_client is harness.prisma + lines = upload.lines() + assert set(lines) == {"row-1", "row-2"} + for custom_id, content in (("row-1", "hi 1"), ("row-2", "hi 2")): + line = lines[custom_id] + assert str(line["id"]).startswith("batch_req_") + assert line["error"] is None + response = line["response"] + assert isinstance(response, dict) + assert response["status_code"] == 200 + assert response["body"] == replies[content].model_dump(mode="json") + + +async def test_create_splits_failed_rows_into_the_error_file() -> None: + harness = make_runner() + failure = ProviderRateLimited("slow down") + reply = chat_response("hi 1") + + def dispatch(messages: Sequence[Mapping[str, str]], **_: object) -> ModelResponse: + if messages[0]["content"] == "hi 1": + return reply + raise failure + + harness.router.acompletion.side_effect = dispatch + created, finished = await harness.create_and_finish() + + assert finished.status == "completed" + assert finished.request_counts == BatchRequestCounts(completed=1, failed=1, total=2) + assert (finished.output_file_id, finished.error_file_id) == ("unified-output-1", "unified-output-2") + llm_batch_id = get_batch_id_from_unified_batch_id(created.id) + assert [call.filename for call in harness.uploads.calls] == [ + f"{llm_batch_id}_output.jsonl", + f"{llm_batch_id}_error.jsonl", + ] + assert set(harness.uploads.calls[0].lines()) == {"row-1"} + error_lines = harness.uploads.calls[1].lines() + assert set(error_lines) == {"row-2"} + response = error_lines["row-2"]["response"] + assert isinstance(response, dict) + assert response["status_code"] == 429 + assert response["body"] == { + "error": {"message": str(failure), "type": "ProviderRateLimited", "param": None, "code": None} + } + + +async def test_create_rejects_an_unsupported_endpoint() -> None: + harness = make_runner() + with pytest.raises(HTTPException) as raised: + await harness.create(endpoint="/v1/moderations") + assert raised.value.status_code == 400 + assert "/v1/moderations" in raised.value.detail["error"] + assert harness.store.calls == [] + assert harness.storage_factory.calls == [] + + +@pytest.mark.parametrize( + "files", + [{}, {INPUT_FILE_ID: managed_input_file(storage_backend=None)}], + ids=["unknown file", "no stored content"], +) +async def test_create_rejects_an_input_file_litellm_does_not_hold( + files: Mapping[str, LiteLLM_ManagedFileTable], +) -> None: + harness = make_runner(files=files) + with pytest.raises(HTTPException) as raised: + await harness.create() + assert raised.value.status_code == 400 + assert "POST /v1/files" in raised.value.detail["error"] + assert harness.storage_factory.calls == [] + assert harness.store.calls == [] + + +async def test_create_rejects_an_invalid_input_file() -> None: + harness = make_runner(content=jsonl(chat_row("a", "hi"), chat_row("a", "again"))) + with pytest.raises(HTTPException) as raised: + await harness.create() + assert raised.value.status_code == 400 + assert raised.value.detail["error"].startswith("Invalid batch input file:") + assert "'a'" in raised.value.detail["error"] + assert harness.store.calls == [] + + +async def test_create_surfaces_a_storage_backend_error_as_a_400() -> None: + harness = make_runner(storage_error=ValueError("Unknown storage backend 's3'")) + with pytest.raises(HTTPException) as raised: + await harness.create() + assert raised.value.status_code == 400 + assert raised.value.detail["error"] == "Unknown storage backend 's3'" + assert harness.store.calls == [] + + +@pytest.mark.parametrize( + ("endpoint", "body", "method"), + [ + ("/v1/chat/completions", {"messages": [{"role": "user", "content": "hi"}]}, "acompletion"), + ("/v1/completions", {"prompt": "hi"}, "atext_completion"), + ("/v1/embeddings", {"input": "hi"}, "aembedding"), + ("/v1/responses", {"input": "hi"}, "aresponses"), + ], +) +async def test_each_endpoint_awaits_only_its_router_method( + endpoint: BatchEndpoint, body: Mapping[str, object], method: str +) -> None: + row = {"custom_id": "a", "method": "POST", "url": endpoint, "body": {"model": "row-model", **body}} + harness = make_runner(content=jsonl(row)) + _, finished = await harness.create_and_finish(endpoint) + + assert finished.request_counts == BatchRequestCounts(completed=1, failed=0, total=1) + awaited = {name: getattr(harness.router, name).await_count for name in ROUTER_METHODS} + assert awaited == {name: int(name == method) for name in ROUTER_METHODS} + kwargs = getattr(harness.router, method).await_args.kwargs + assert kwargs["model"] == BATCH_MODEL + assert all(kwargs[key] == value for key, value in body.items()) + + +async def test_cancel_unknown_batch_is_404() -> None: + harness = make_runner() + with pytest.raises(HTTPException) as raised: + await harness.runner.cancel("missing-batch", harness.user) + assert raised.value.status_code == 404 + + +async def test_cancel_terminal_batch_is_400() -> None: + harness = make_runner() + batch = seeded_batch(harness.store, "completed") + with pytest.raises(HTTPException) as raised: + await harness.runner.cancel(batch.id, harness.user) + assert raised.value.status_code == 400 + assert "completed" in raised.value.detail["error"] + assert harness.store.calls == [] + + +async def test_cancel_marks_a_running_batch_cancelling_once() -> None: + harness = make_runner() + batch = seeded_batch(harness.store, "in_progress") + + cancelled = await harness.runner.cancel(batch.id, harness.user) + + assert cancelled.status == "cancelling" + assert cancelled.cancelling_at is not None + assert harness.store.batch(batch.id).status == "cancelling" + assert [(call.status, call.create_if_missing) for call in harness.store.calls] == [("cancelling", False)] + + again = await harness.runner.cancel(batch.id, harness.user) + + assert again.model_dump() == cancelled.model_dump() + assert len(harness.store.calls) == 1 + + +async def test_running_batch_skips_the_remaining_rows_after_an_operator_cancel( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm_executed_batches, "_CANCEL_POLL_SECONDS", 0.0) + rows = jsonl(chat_row("row-1", "hi 1"), chat_row("row-2", "hi 2"), chat_row("row-3", "hi 3")) + harness = make_runner(content=rows, concurrency=1) + reply = chat_response("hi 1") + + def dispatch(metadata: Mapping[str, object], **_: object) -> ModelResponse: + running = harness.store.batch(str(metadata["batch_id"])) + harness.store.write(running.model_copy(update={"status": "cancelling"})) + return reply + + harness.router.acompletion.side_effect = dispatch + _, finished = await harness.create_and_finish() + + assert harness.router.acompletion.await_count == 1 + assert finished.status == "cancelled" + assert finished.cancelled_at is not None + assert finished.request_counts == BatchRequestCounts(completed=1, failed=0, total=3) + assert (finished.output_file_id, finished.error_file_id) == ("unified-output-1", None) + + +async def test_upload_failure_marks_the_batch_failed() -> None: + harness = make_runner(upload_error=RuntimeError("storage exploded")) + _, finished = await harness.create_and_finish() + + assert finished.status == "failed" + assert finished.failed_at is not None + assert finished.output_file_id is None + assert finished.errors is not None + assert [(error.message, error.code) for error in finished.errors.data or []] == [ + ("storage exploded", "internal_error") + ] + + +async def test_only_the_create_write_carries_attribution_and_billing_flags() -> None: + harness = make_runner() + await harness.create_and_finish() + + assert [call.status for call in harness.store.calls] == ["validating", "in_progress", "finalizing", "completed"] + flags = [(call.persist_attribution, call.batch_processed, call.create_if_missing) for call in harness.store.calls] + assert flags[0] == (True, True, True) + assert flags[1:] == [(False, False, False)] * 3 + + +async def test_run_completes_under_the_real_hooks_base64_batch_id() -> None: + harness = make_runner(store_factory=RealIdManagedBatchStore) + created, finished = await harness.create_and_finish() + + assert _is_base64_encoded_unified_file_id(created.id) + assert finished.status == "completed" + llm_batch_id = harness.store.calls[0].model_object_id + assert llm_batch_id.startswith("litellm_batch_") + assert [call.model_object_id for call in harness.store.calls] == [llm_batch_id] * 4 + + +async def test_cancel_works_under_the_real_hooks_base64_batch_id() -> None: + harness = make_runner(store_factory=RealIdManagedBatchStore) + batch = seeded_batch(harness.store, "in_progress") + cancelled = await harness.runner.cancel(batch.id, harness.user) + + assert cancelled.status == "cancelling" + assert [call.model_object_id for call in harness.store.calls] == ["litellm_batch_seed"] diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index 87cd2aaff1f..f88d94d2c08 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -6,6 +6,7 @@ import pytest from litellm.proxy.openai_files_endpoints.common_utils import ( apply_unified_file_ids, + is_litellm_executed_batch, map_raw_file_ids_to_unified, ) from litellm.types.utils import LiteLLMBatch @@ -478,3 +479,15 @@ 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), + ], +) +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 diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index aa505c3019b..08d7bf8e0da 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -609,6 +609,121 @@ 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) + 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: + yield stored, provider_upload + 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" + + +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() + + +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 diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py b/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py index 07a85a70815..067826004e2 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py @@ -1,3 +1,5 @@ +from unittest.mock import MagicMock + import pytest from litellm.llms.base_llm.files.transformation import BaseFileEndpoints @@ -6,6 +8,7 @@ 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: @@ -57,7 +60,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 +83,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 +104,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 +128,28 @@ 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)] From b4f10e211c7655eeb5e0784c8d698a66f5149da7 Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" <41898282+github-actions[bot]@users.noreply.github.com> Date: Sat, 19 Sep 2026 08:49:34 +0000 Subject: [PATCH 02/12] chore: sync schema.prisma copies from root --- litellm-proxy-extras/litellm_proxy_extras/schema.prisma | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 91b59e56906..c4606796ebf 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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 From 167e3244ab7b2e904e6d7f3f193d3f2bf737a874 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 01:58:53 -0700 Subject: [PATCH 03/12] fix(batches): shape LiteLLM-executed batch errors like OpenAI errors --- litellm/proxy/batches_endpoints/endpoints.py | 10 ++--- .../litellm_executed_batches.py | 21 +++++----- .../test_litellm_executed_batches.py | 40 +++++++++---------- 3 files changed, 35 insertions(+), 36 deletions(-) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index cc698ad760e..7e49d937171 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -23,7 +23,7 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE, LiteLLMExecutedBatchRunner, ManagedBatchStore, - batch_http_error, + batch_error, litellm_executed_provider_of, resolve_litellm_executed_provider, ) @@ -83,7 +83,7 @@ def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyL managed_files: Final = proxy_logging_obj.get_proxy_hook("managed_files") if prisma_client is None or not isinstance(managed_files, ManagedBatchStore): - raise batch_http_error( + raise batch_error( 400, "LiteLLM-executed batches need a database: set DATABASE_URL so LiteLLM can keep the batch and its files", ) @@ -98,7 +98,7 @@ def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyL def _raise_when_input_file_must_be_managed(model: str, credentials: Mapping[str, object]) -> None: if litellm_executed_provider_of(credentials) is None: return - raise batch_http_error( + 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}", @@ -556,7 +556,7 @@ async def retrieve_batch( executed_batch: Final = isinstance(unified_batch_id, str) and is_litellm_executed_batch(unified_batch_id) if executed_batch and response is None: - raise batch_http_error(404, f"No batch found with id '{batch_id}'.") + 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"). @@ -1067,7 +1067,7 @@ 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_http_error(500, "LLM Router not initialized. Ensure models added to proxy.") + 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) diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index f567f0e2263..77e583781ea 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -7,7 +7,6 @@ from itertools import pairwise from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable -from fastapi import HTTPException from openai.types.batch import Errors from openai.types.batch_error import BatchError from openai.types.batch_request_counts import BatchRequestCounts @@ -23,7 +22,7 @@ from litellm.llms.base_llm.files.litellm_db_storage_backend import LITELLM_DB_ST 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.models.managed_files import LiteLLM_ManagedFileTable -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.openai_files_endpoints.common_utils import ( LITELLM_EXECUTED_BATCH_ID_PREFIX, @@ -234,16 +233,16 @@ def parse_batch_input(content: bytes, endpoint: BatchEndpoint) -> tuple[BatchInp return lines -def batch_http_error(status_code: int, message: str) -> HTTPException: - detail: Final = {"error": message} # mutable-ok: HTTPException detail must be a plain mapping - return HTTPException(status_code=status_code, detail=detail) +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_http_error(400, f"endpoint {endpoint!r} is not supported for a LiteLLM-executed batch") + raise batch_error(400, f"endpoint {endpoint!r} is not supported for a LiteLLM-executed batch") def _status_code_of(error: Exception) -> int: @@ -350,7 +349,7 @@ class LiteLLMExecutedBatchRunner: content: Final = await self._download_input(unified_input_file_id, user_api_key_dict) parsed: Final = parse_batch_input(content, endpoint) if isinstance(parsed, InvalidBatchInput): - raise batch_http_error(400, f"Invalid batch input file: {parsed.describe()}") + 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) @@ -397,9 +396,9 @@ class LiteLLMExecutedBatchRunner: async def cancel(self, unified_batch_id: str, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch: current: Final = await self._load_batch(unified_batch_id) if current is None: - raise batch_http_error(404, f"Batch {unified_batch_id} not found") + raise batch_error(404, f"Batch {unified_batch_id} not found") if current.status in TERMINAL_BATCH_STATUSES: - raise batch_http_error(400, f"Cannot cancel a batch with status '{current.status}'") + raise batch_error(400, f"Cannot cancel a batch with status '{current.status}'") if current.status == "cancelling": return current cancelling: Final = current.model_copy( @@ -413,7 +412,7 @@ class LiteLLMExecutedBatchRunner: 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_http_error( + raise batch_error( 400, f"LiteLLM does not hold the content of input file {unified_input_file_id}: " f"{LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE}", @@ -422,7 +421,7 @@ class LiteLLMExecutedBatchRunner: 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_http_error(400, str(e)) + raise batch_error(400, str(e)) async def _run(self, run: _BatchRun) -> None: try: diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index 0806dd3451e..9c775fd2c97 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -6,12 +6,11 @@ from typing import Final, Literal, cast from unittest.mock import AsyncMock, MagicMock import pytest -from fastapi import HTTPException from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles from openai.types.batch_request_counts import BatchRequestCounts from litellm.models.managed_files import LiteLLM_ManagedFileTable -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.batches_endpoints import litellm_executed_batches from litellm.proxy.batches_endpoints.litellm_executed_batches import ( BatchEndpoint, @@ -547,10 +546,11 @@ async def test_create_splits_failed_rows_into_the_error_file() -> None: async def test_create_rejects_an_unsupported_endpoint() -> None: harness = make_runner() - with pytest.raises(HTTPException) as raised: + with pytest.raises(ProxyException) as raised: await harness.create(endpoint="/v1/moderations") - assert raised.value.status_code == 400 - assert "/v1/moderations" in raised.value.detail["error"] + assert raised.value.code == "400" + assert raised.value.type == "invalid_request_error" + assert "/v1/moderations" in raised.value.message assert harness.store.calls == [] assert harness.storage_factory.calls == [] @@ -564,30 +564,30 @@ async def test_create_rejects_an_input_file_litellm_does_not_hold( files: Mapping[str, LiteLLM_ManagedFileTable], ) -> None: harness = make_runner(files=files) - with pytest.raises(HTTPException) as raised: + with pytest.raises(ProxyException) as raised: await harness.create() - assert raised.value.status_code == 400 - assert "POST /v1/files" in raised.value.detail["error"] + assert raised.value.code == "400" + assert "POST /v1/files" in raised.value.message assert harness.storage_factory.calls == [] assert harness.store.calls == [] async def test_create_rejects_an_invalid_input_file() -> None: harness = make_runner(content=jsonl(chat_row("a", "hi"), chat_row("a", "again"))) - with pytest.raises(HTTPException) as raised: + with pytest.raises(ProxyException) as raised: await harness.create() - assert raised.value.status_code == 400 - assert raised.value.detail["error"].startswith("Invalid batch input file:") - assert "'a'" in raised.value.detail["error"] + assert raised.value.code == "400" + assert raised.value.message.startswith("Invalid batch input file:") + assert "'a'" in raised.value.message assert harness.store.calls == [] async def test_create_surfaces_a_storage_backend_error_as_a_400() -> None: harness = make_runner(storage_error=ValueError("Unknown storage backend 's3'")) - with pytest.raises(HTTPException) as raised: + with pytest.raises(ProxyException) as raised: await harness.create() - assert raised.value.status_code == 400 - assert raised.value.detail["error"] == "Unknown storage backend 's3'" + assert raised.value.code == "400" + assert raised.value.message == "Unknown storage backend 's3'" assert harness.store.calls == [] @@ -617,18 +617,18 @@ async def test_each_endpoint_awaits_only_its_router_method( async def test_cancel_unknown_batch_is_404() -> None: harness = make_runner() - with pytest.raises(HTTPException) as raised: + with pytest.raises(ProxyException) as raised: await harness.runner.cancel("missing-batch", harness.user) - assert raised.value.status_code == 404 + assert raised.value.code == "404" async def test_cancel_terminal_batch_is_400() -> None: harness = make_runner() batch = seeded_batch(harness.store, "completed") - with pytest.raises(HTTPException) as raised: + with pytest.raises(ProxyException) as raised: await harness.runner.cancel(batch.id, harness.user) - assert raised.value.status_code == 400 - assert "completed" in raised.value.detail["error"] + assert raised.value.code == "400" + assert "completed" in raised.value.message assert harness.store.calls == [] From 9e8b686c7accbcf8e2cba87b982970dafbad94f2 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 02:53:46 -0700 Subject: [PATCH 04/12] fix(batches): run hosted_vllm batches in LiteLLM only when the server has no Files API --- litellm/proxy/batches_endpoints/endpoints.py | 12 ++- .../litellm_executed_batches.py | 67 +++++++++++- .../openai_files_endpoints/files_endpoints.py | 23 ++-- .../proxy/batches_endpoints/test_endpoints.py | 50 ++++++++- .../test_litellm_executed_batches.py | 101 +++++++++++++++++- .../test_files_endpoint.py | 52 ++++++++- 6 files changed, 283 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 7e49d937171..ba5853b7c46 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -24,7 +24,7 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( LiteLLMExecutedBatchRunner, ManagedBatchStore, batch_error, - litellm_executed_provider_of, + litellm_executed_provider_for, resolve_litellm_executed_provider, ) from litellm.proxy.common_request_processing import ( @@ -95,8 +95,8 @@ def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyL ) -def _raise_when_input_file_must_be_managed(model: str, credentials: Mapping[str, object]) -> None: - if litellm_executed_provider_of(credentials) is None: +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, @@ -364,7 +364,9 @@ async def create_batch( detail={"error": "LLM Router not initialized. Ensure models added to proxy."}, ) - executed_provider: Final = resolve_litellm_executed_provider(llm_router, model, user_api_key_dict.team_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, @@ -395,7 +397,7 @@ async def create_batch( model_id=model_param, operation_context="batch creation", ) - _raise_when_input_file_must_be_managed(model_param, credentials) + await _raise_when_input_file_must_be_managed(model_param, credentials) prepare_data_with_credentials( data=_create_batch_data, diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index 77e583781ea..642b2ea5f8c 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -7,6 +7,7 @@ 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 @@ -21,6 +22,7 @@ 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.litellm_pre_call_utils import LiteLLMProxyRequestSetup @@ -33,7 +35,7 @@ from litellm.proxy.openai_files_endpoints.storage_backend_service import Storage from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.table_repositories import ManagedObjectRepository from litellm.types.llms.openai import LiteLLMBatchCreateRequest, OpenAIFileObject, OpenAIFilesPurpose -from litellm.types.utils import LITELLM_EXECUTED_BATCH_PROVIDERS, ExtractedFileData, LiteLLMBatch +from litellm.types.utils import LITELLM_EXECUTED_BATCH_PROVIDERS, ExtractedFileData, LiteLLMBatch, LlmProviders if TYPE_CHECKING: from prisma import models as prisma_models @@ -46,6 +48,7 @@ BatchStatus: TypeAlias = Literal["in_progress", "finalizing", "completed", "fail TERMINAL_BATCH_STATUSES: Final[frozenset[str]] = frozenset({"completed", "failed", "cancelled", "expired"}) _BATCH_ENDPOINT_ADAPTER: Final[TypeAdapter[BatchEndpoint]] = TypeAdapter(BatchEndpoint) _CANCEL_POLL_SECONDS: Final = 1.0 +_FILES_API_PROBE_TIMEOUT_SECONDS: Final = 5.0 _COMPLETION_WINDOW_SECONDS: Final = 24 * 60 * 60 LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE: Final = ( "upload it through POST /v1/files with purpose=batch and either the x-litellm-model header or the " @@ -184,9 +187,67 @@ def litellm_executed_provider_of(credentials: Mapping[str, object]) -> str | Non return provider if provider in LITELLM_EXECUTED_BATCH_PROVIDERS else None -def resolve_litellm_executed_provider(llm_router: "Router", model: str, team_id: str | None) -> str | 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: ... + + +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}"} 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 litellm_executed_provider_of(credentials) + return None if credentials is None else await litellm_executed_provider_for(credentials, lacks_files_api) def _provider_of(model: object) -> str | None: diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index ad869100fb9..b8dcb89baef 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -102,24 +102,35 @@ from litellm.types.llms.openai import ( router: Final = APIRouter() -def _litellm_executed_batch_input_model( +async def _litellm_executed_batch_input_model( llm_router: Router | None, purpose: OpenAIFilesPurpose, model: str | None, target_model_names_list: Sequence[str], team_id: str | None, ) -> str | None: - if purpose != "batch" or llm_router is None: + if llm_router is None: return None candidates: Final = (model,) if model is not None else tuple(target_model_names_list) + providers: Final = await asyncio.gather( + *(resolve_litellm_executed_provider(llm_router, candidate, team_id) for candidate in candidates) + ) executed: Final = tuple( - candidate - for candidate in candidates - if resolve_litellm_executed_provider(llm_router, candidate, team_id) is not None + candidate for candidate, provider in zip(candidates, providers, strict=True) if provider is not None ) match executed: case (): return None + case _ 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, + ) case (only,) if len(candidates) == 1: return only case _: @@ -279,7 +290,7 @@ async def route_create_file( 5. Else -> use custom_llm_provider with files_settings """ - executed_model: Final = _litellm_executed_batch_input_model( + executed_model: Final = await _litellm_executed_batch_input_model( llm_router, purpose, model, target_model_names_list, user_api_key_dict.team_id ) explicit_storage: Final = target_storage if target_storage and target_storage != "default" else None diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index e655438a672..3a2ddf50143 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -37,7 +37,9 @@ from dataclasses import dataclass 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 @@ -152,6 +154,7 @@ class Harness: router: MagicMock logging: MagicMock creds_resolver: MagicMock + upstream_files_route: respx.Route @property def router_acreate(self) -> AsyncMock: @@ -174,7 +177,7 @@ def _creds_lookup(*, model_id: str, team_id: str | None = None) -> dict[str, str @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] = {} @@ -194,6 +197,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)) @@ -215,6 +219,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", {})) @@ -233,6 +241,7 @@ def harness(): router=router, logging=logging, creds_resolver=router.get_deployment_credentials_with_provider, + upstream_files_route=upstream_files_route, ) yield h @@ -843,6 +852,25 @@ async def test_create__unified_executed_provider_without_database_400(harness): 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 @@ -879,6 +907,26 @@ async def test_create__raw_file_with_executed_model_400_with_upload_guidance(har 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 diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index 9c775fd2c97..96e054272a4 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -5,6 +5,7 @@ from dataclasses import dataclass from typing import Final, Literal, cast from unittest.mock import AsyncMock, MagicMock +import httpx import pytest from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles from openai.types.batch_request_counts import BatchRequestCounts @@ -19,9 +20,11 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( InvalidBatchInput, LiteLLMExecutedBatchRunner, _resolve_transition, + litellm_executed_provider_for, litellm_executed_provider_of, parse_batch_input, resolve_litellm_executed_provider, + upstream_lacks_files_api, ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, @@ -424,15 +427,107 @@ def test_litellm_executed_provider_of(credentials: Mapping[str, object], expecte assert litellm_executed_provider_of(credentials) == expected +VLLM_CREDENTIALS: Final[Mapping[str, object]] = { + "model": "hosted_vllm/qwen", + "api_base": "http://vllm.test/v1/", + "api_key": "vllm-key", +} + + +@dataclass(slots=True) +class FakeFilesApiProbe: + lacks_files_api: bool + upstreams: list[tuple[str, str | None]] + + async def __call__(self, api_base: str, api_key: str | None) -> bool: + self.upstreams.append((api_base, api_key)) + return self.lacks_files_api + + +@dataclass(slots=True) +class FakeHttpGetter: + outcome: int | httpx.HTTPError + requests: list[tuple[str, dict[str, str] | None]] + + async def get( + self, url: str, *, headers: dict[str, str] | None = None, timeout: float | httpx.Timeout | None = None + ) -> httpx.Response: + self.requests.append((url, headers)) + if isinstance(self.outcome, httpx.HTTPError): + raise self.outcome + return httpx.Response(self.outcome) + + @pytest.mark.parametrize( - ("credentials", "expected"), [(None, None), ({"model": "hosted_vllm/qwen"}, "hosted_vllm")], ids=["unknown", "vllm"] + ("outcome", "expected"), + [ + (404, True), + (200, False), + (405, False), + (401, False), + (500, False), + (httpx.ConnectError("refused"), False), + (httpx.ReadTimeout("slow"), False), + ], + ids=["no files route", "lists files", "files route without list", "unauthorized", "server error", "down", "slow"], ) -def test_resolve_litellm_executed_provider_asks_the_router_for_the_team_scoped_deployment( +async def test_upstream_lacks_files_api_only_when_the_files_route_is_a_404( + outcome: int | httpx.HTTPError, expected: bool +) -> None: + assert await upstream_lacks_files_api("http://vllm.test/v1", "vllm-key", FakeHttpGetter(outcome, [])) is expected + + +@pytest.mark.parametrize( + ("api_base", "api_key", "expected_headers"), + [ + ("http://vllm.test/v1/", "vllm-key", {"Authorization": "Bearer vllm-key"}), + ("http://vllm.test/v1", None, None), + ], + ids=["trailing slash with key", "keyless"], +) +async def test_upstream_lacks_files_api_asks_the_files_route_under_the_api_base( + api_base: str, api_key: str | None, expected_headers: dict[str, str] | None +) -> None: + http_client = FakeHttpGetter(404, []) + await upstream_lacks_files_api(api_base, api_key, http_client) + assert http_client.requests == [("http://vllm.test/v1/files", expected_headers)] + + +@pytest.mark.parametrize( + ("lacks_files_api", "expected"), [(True, "hosted_vllm"), (False, None)], ids=["bare", "router"] +) +async def test_litellm_executed_provider_for_leaves_a_server_with_its_own_files_api_alone( + lacks_files_api: bool, expected: str | None +) -> None: + probe = FakeFilesApiProbe(lacks_files_api, []) + assert await litellm_executed_provider_for(VLLM_CREDENTIALS, probe) == expected + assert probe.upstreams == [("http://vllm.test/v1/", "vllm-key")] + + +@pytest.mark.parametrize( + "credentials", + [{"custom_llm_provider": "openai", "model": "gpt-4o", "api_base": "http://openai.test/v1"}, {"model": 7}], + ids=["provider runs its own batches", "no model to resolve an api_base from"], +) +async def test_litellm_executed_provider_for_never_probes_what_it_would_not_run( + credentials: Mapping[str, object], +) -> None: + probe = FakeFilesApiProbe(True, []) + assert await litellm_executed_provider_for(credentials, probe) is None + assert probe.upstreams == [] + + +@pytest.mark.parametrize( + ("credentials", "expected"), [(None, None), (VLLM_CREDENTIALS, "hosted_vllm")], ids=["unknown", "vllm"] +) +async def test_resolve_litellm_executed_provider_asks_the_router_for_the_team_scoped_deployment( credentials: Mapping[str, object] | None, expected: str | None ) -> None: router = MagicMock(spec=Router) router.get_deployment_credentials_with_provider.return_value = credentials - assert resolve_litellm_executed_provider(router, BATCH_MODEL, "team-1") == expected + assert ( + await resolve_litellm_executed_provider(router, BATCH_MODEL, "team-1", FakeFilesApiProbe(True, [])) == expected + ) router.get_deployment_credentials_with_provider.assert_called_once_with(model_id=BATCH_MODEL, team_id="team-1") diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 08d7bf8e0da..a8e0097b831 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -646,6 +646,7 @@ def batch_upload_seams(mocker: MockerFixture, monkeypatch): 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" ) @@ -666,7 +667,11 @@ def batch_upload_seams(mocker: MockerFixture, monkeypatch): "litellm.acreate_file", new=mocker.AsyncMock(return_value=uploaded) ) try: - yield stored, provider_upload + 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) @@ -688,7 +693,7 @@ def _upload_batch_file(headers: dict[str, str], form: dict[str, str]): 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 + stored, provider_upload, _ = batch_upload_seams response = _upload_batch_file(headers, form) @@ -702,7 +707,7 @@ def test_batch_upload_for_a_litellm_executed_model_is_kept_by_litellm( def test_batch_upload_naming_an_executed_and_a_provider_model_is_rejected(batch_upload_seams): - stored, provider_upload = batch_upload_seams + stored, provider_upload, _ = batch_upload_seams response = _upload_batch_file({}, {"target_model_names": "my-vllm,gemini-2.0-flash"}) @@ -713,8 +718,47 @@ def test_batch_upload_naming_an_executed_and_a_provider_model_is_rejected(batch_ 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" + + def test_batch_upload_for_a_provider_model_still_goes_to_the_provider(batch_upload_seams): - stored, provider_upload = batch_upload_seams + stored, provider_upload, _ = batch_upload_seams response = _upload_batch_file({"x-litellm-model": "gemini-2.0-flash"}, {}) From a0957edc9cf1794728afdf7094e0554f6ade5b88 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 04:07:34 -0700 Subject: [PATCH 05/12] fix(batches): gate row credentials, heartbeat executed batches, clean orphaned uploads --- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- litellm/proxy/batches_endpoints/endpoints.py | 47 +++++- .../litellm_executed_batches.py | 154 +++++++++++++----- .../openai_files_endpoints/files_endpoints.py | 48 +++--- .../storage_backend_service.py | 24 ++- .../proxy/batches_endpoints/test_endpoints.py | 37 +++++ .../test_litellm_executed_batches.py | 121 ++++++++++++++ .../test_storage_backend_service.py | 36 +++- 8 files changed, 391 insertions(+), 78 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 18849ef5b64..73ea8cf1991 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -19622,7 +19622,7 @@ } } }, - "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" + "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " }, "500": { "content": { diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index ba5853b7c46..19e56e97338 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -7,8 +7,9 @@ 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 @@ -24,6 +25,7 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( LiteLLMExecutedBatchRunner, ManagedBatchStore, batch_error, + executed_batch_runner_lost, litellm_executed_provider_for, resolve_litellm_executed_provider, ) @@ -61,12 +63,15 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( ) from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import request_tags_from_metadata from litellm.proxy.route_llm_request import raise_if_required_body_param_missing -from litellm.proxy.utils import ProxyLogging, handle_exception_on_proxy, is_known_model +from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy, is_known_model 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]) @@ -79,7 +84,7 @@ def _request_tags(data: Mapping[str, object]) -> tuple[str, ...] | None: def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyLogging) -> LiteLLMExecutedBatchRunner: - from litellm.proxy.proxy_server import prisma_client + 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): @@ -92,9 +97,36 @@ def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyL prisma_client=prisma_client, managed_files=managed_files, 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 @@ -548,15 +580,18 @@ 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, ) - executed_batch: Final = isinstance(unified_batch_id, str) and is_litellm_executed_batch(unified_batch_id) if executed_batch and response is None: raise batch_error(404, f"No batch found with id '{batch_id}'.") diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index 642b2ea5f8c..7bd4c678183 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -3,6 +3,7 @@ import json import time from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass +from datetime import datetime, timezone from itertools import pairwise from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable @@ -12,7 +13,7 @@ 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, assert_never +from typing_extensions import ReadOnly, TypedDict import litellm from litellm._logging import verbose_proxy_logger @@ -25,6 +26,7 @@ from litellm.llms.base_llm.files.storage_backend_factory import get_storage_back 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, @@ -44,12 +46,26 @@ if TYPE_CHECKING: BatchEndpoint: TypeAlias = Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"] BatchStatus: TypeAlias = Literal["in_progress", "finalizing", "completed", "failed", "cancelling", "cancelled"] - 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" +_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", "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" @@ -165,20 +181,6 @@ class _RouterCall(Protocol): def __call__(self, **params: object) -> Awaitable[object]: ... # kwargs-ok: the request body is passed as keywords -def _router_method_name(endpoint: BatchEndpoint) -> str: - match endpoint: - case "/v1/chat/completions": - return "acompletion" - case "/v1/completions": - return "atext_completion" - case "/v1/embeddings": - return "aembedding" - case "/v1/responses": - return "aresponses" - case _: - assert_never(endpoint) - - def litellm_executed_provider_of(credentials: Mapping[str, object]) -> str | None: explicit_provider: Final = credentials.get("custom_llm_provider") provider: Final = ( @@ -197,12 +199,20 @@ 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}"} if api_key else None, + 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: @@ -266,7 +276,13 @@ def _validation_reason(error: ValidationError) -> str: ) -def _parse_line(line_number: int, raw: bytes, endpoint: BatchEndpoint) -> BatchInputLine | InvalidBatchInput: +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: @@ -275,14 +291,19 @@ def _parse_line(line_number: int, raw: bytes, endpoint: BatchEndpoint) -> BatchI 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) -> tuple[BatchInputLine, ...] | InvalidBatchInput: +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) for number, raw in raw_lines) + 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 @@ -345,37 +366,35 @@ def _dump(response: object) -> Mapping[str, object]: def _resolve_transition(current_status: str, requested: BatchStatus) -> BatchStatus: if current_status != "cancelling": return requested - match requested: - case "completed": - return "cancelled" - case "in_progress" | "finalizing": - return "cancelling" - case "failed" | "cancelling" | "cancelled": - return requested - case _: - assert_never(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 def _llm_batch_id_of(unified_batch_id: str) -> str: return get_batch_id_from_unified_batch_id(convert_b64_uid_to_unified_uid(unified_batch_id)) -class _CancelWatch: +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._cancelling = False + self._stopped = False - async def cancelling(self) -> bool: - if self._cancelling: + 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._cancelling = await self._load_status() == "cancelling" - return self._cancelling + self._stopped = await self._load_status() in _STOP_STATUSES + return self._stopped class LiteLLMExecutedBatchRunner: @@ -385,7 +404,9 @@ class LiteLLMExecutedBatchRunner: prisma_client: PrismaClient, managed_files: ManagedBatchStore, proxy_logging_obj: ProxyLogging, + general_settings: Mapping[str, object], concurrency: int = LITELLM_EXECUTED_BATCH_CONCURRENCY, + heartbeat_seconds: float = _HEARTBEAT_SECONDS, storage_backend_factory: _StorageBackendFactory = get_storage_backend, upload_result_file: _ResultFileUploader = StorageBackendFileService.upload_file_to_storage_backend, ) -> None: @@ -393,7 +414,9 @@ class LiteLLMExecutedBatchRunner: self.prisma_client = prisma_client self.managed_files = managed_files self.proxy_logging_obj = proxy_logging_obj + self.general_settings = general_settings self.concurrency = concurrency + self.heartbeat_seconds = heartbeat_seconds self.storage_backend_factory = storage_backend_factory self.upload_result_file = upload_result_file @@ -408,7 +431,7 @@ class LiteLLMExecutedBatchRunner: ) -> 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) + 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}" @@ -468,6 +491,30 @@ class LiteLLMExecutedBatchRunner: await self._store(cancelling, user_api_key_dict) return cancelling + 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}) + ) + await self._store(failed, user_api_key_dict) + return failed + + 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 @@ -485,6 +532,7 @@ class LiteLLMExecutedBatchRunner: 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 @@ -497,14 +545,31 @@ class LiteLLMExecutedBatchRunner: 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 ManagedObjectRepository(self.prisma_client).table.update_many( + where={"unified_object_id": run.unified_batch_id}, # mutable-ok: Prisma filter + data={"updated_by": run.user_api_key_dict.user_id}, # mutable-ok: Prisma payload + ) async def _execute(self, run: _BatchRun) -> None: await self._advance(run, "in_progress") - watch: Final = _CancelWatch(lambda: self._load_status(run.unified_batch_id), _CANCEL_POLL_SECONDS) + watch: Final = _StopWatch(lambda: self._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) - await self._advance(run, "finalizing") + if await self._advance(run, "finalizing") is None: + return succeeded: Final = tuple(outcome for outcome in outcomes if outcome.succeeded) failed: Final = tuple(outcome for outcome in outcomes if not outcome.succeeded) output_file_id: Final = await self._upload_results(run, "output", succeeded) @@ -519,10 +584,10 @@ class LiteLLMExecutedBatchRunner: ) async def _run_row( - self, run: _BatchRun, line: BatchInputLine, watch: _CancelWatch, semaphore: asyncio.Semaphore + self, run: _BatchRun, line: BatchInputLine, watch: _StopWatch, semaphore: asyncio.Semaphore ) -> RowOutcome | None: async with semaphore: - if await watch.cancelling(): + if await watch.stopped(): return None try: body: Final = await self._dispatch(run, line) @@ -537,7 +602,7 @@ class LiteLLMExecutedBatchRunner: return _dump(await self._router_call(run.endpoint)(**params)) def _router_call(self, endpoint: BatchEndpoint) -> _RouterCall: - method: Final[object] = getattr(self.llm_router, _router_method_name(endpoint), None) + 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 @@ -574,15 +639,20 @@ class LiteLLMExecutedBatchRunner: ) return file_object.id - async def _advance(self, run: _BatchRun, requested: BatchStatus, fields: Mapping[str, object] = _NO_FIELDS) -> None: + async def _advance( + self, run: _BatchRun, requested: BatchStatus, fields: Mapping[str, object] = _NO_FIELDS + ) -> BatchStatus | None: current: Final = await self._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())}) ) await self._store(updated, run.user_api_key_dict) + return status async def _store(self, batch: LiteLLMBatch, user_api_key_dict: UserAPIKeyAuth) -> None: await self.managed_files.store_unified_object_id( diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index b8dcb89baef..7356d197be9 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -118,31 +118,29 @@ async def _litellm_executed_batch_input_model( executed: Final = tuple( candidate for candidate, provider in zip(candidates, providers, strict=True) if provider is not None ) - match executed: - case (): - return None - case _ 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, - ) - case (only,) if len(candidates) == 1: - return only - case _: - 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, - ) + 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) diff --git a/litellm/proxy/openai_files_endpoints/storage_backend_service.py b/litellm/proxy/openai_files_endpoints/storage_backend_service.py index b4a36336c22..66dbcd87c0b 100644 --- a/litellm/proxy/openai_files_endpoints/storage_backend_service.py +++ b/litellm/proxy/openai_files_endpoints/storage_backend_service.py @@ -12,6 +12,7 @@ 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 @@ -105,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, @@ -116,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, diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 3a2ddf50143..59c0e05d97b 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -34,6 +34,7 @@ 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 @@ -1722,6 +1723,7 @@ async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harn 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) @@ -1734,6 +1736,41 @@ async def test_retrieve__executed_batch_served_from_db_in_every_status(retrieve_ 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: diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index 96e054272a4..e5ed873a29d 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -2,6 +2,8 @@ import asyncio import json from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from types import MappingProxyType from typing import Final, Literal, cast from unittest.mock import AsyncMock, MagicMock @@ -20,6 +22,7 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( InvalidBatchInput, LiteLLMExecutedBatchRunner, _resolve_transition, + executed_batch_runner_lost, litellm_executed_provider_for, litellm_executed_provider_of, parse_batch_input, @@ -174,10 +177,15 @@ class RealIdManagedBatchStore(FakeManagedBatchStore): class FakeManagedObjectTable: def __init__(self, objects: Mapping[str, StoredObject]) -> None: self.objects = objects + self.touches: list[tuple[str, str | None]] = [] async def find_first(self, where: Mapping[str, str]) -> StoredObject | None: return self.objects.get(where["unified_object_id"]) + async def update_many(self, where: Mapping[str, str], data: Mapping[str, str | None]) -> int: + self.touches.append((where["unified_object_id"], data["updated_by"])) + return 1 + class FakeDb: def __init__(self, objects: Mapping[str, StoredObject]) -> None: @@ -203,6 +211,9 @@ class FakeRouter: def get_model_ids(self, model_name: str) -> list[str]: return [DEPLOYMENT_ID] if model_name == BATCH_MODEL else [] + def get_model_group_info(self, model_group: str) -> None: + return None + class FakeStorageBackend: def __init__(self, contents: Mapping[str, bytes]) -> None: @@ -317,6 +328,8 @@ def make_runner( upload_error: Exception | None = None, storage_error: ValueError | None = None, store_factory: Callable[[Mapping[str, LiteLLM_ManagedFileTable]], FakeManagedBatchStore] = FakeManagedBatchStore, + general_settings: Mapping[str, object] = MappingProxyType({}), + heartbeat_seconds: float = 30.0, ) -> Harness: store = store_factory({INPUT_FILE_ID: managed_input_file()} if files is None else files) router = FakeRouter() @@ -332,7 +345,9 @@ def make_runner( prisma_client=cast("PrismaClient", prisma), managed_files=store, proxy_logging_obj=MagicMock(spec=ProxyLogging), + general_settings=general_settings, concurrency=concurrency, + heartbeat_seconds=heartbeat_seconds, storage_backend_factory=storage_factory, upload_result_file=uploads, ) @@ -413,6 +428,27 @@ def test_resolve_transition_from_cancelling(requested: BatchStatus, expected: Ba assert _resolve_transition("cancelling", requested) == expected +@pytest.mark.parametrize( + ("status", "age_seconds", "lost"), + [ + ("validating", 200, True), + ("in_progress", 200, True), + ("in_progress", 100, False), + ("finalizing", 200, True), + ("cancelling", 200, True), + ("completed", 200, False), + ("failed", 200, False), + ("cancelled", 200, False), + ("expired", 200, False), + ], +) +def test_executed_batch_runner_lost_only_for_a_stale_non_terminal_batch( + status: str, age_seconds: int, lost: bool +) -> None: + updated_at = datetime.now(timezone.utc) - timedelta(seconds=age_seconds) + assert executed_batch_runner_lost(status, updated_at) is lost + + @pytest.mark.parametrize( ("credentials", "expected"), [ @@ -686,6 +722,91 @@ async def test_create_surfaces_a_storage_backend_error_as_a_400() -> None: assert harness.store.calls == [] +CREDENTIAL_ROWS: Final = jsonl(chat_row("row-1", "hi 1"), chat_row("row-2", "hi 2", api_base="https://evil.example")) + + +async def test_create_rejects_a_row_carrying_client_side_credentials() -> None: + harness = make_runner(content=CREDENTIAL_ROWS) + with pytest.raises(ProxyException) as raised: + await harness.create() + assert raised.value.code == "400" + assert raised.value.message.startswith("Invalid batch input file: line 2") + assert "api_base" in raised.value.message + assert "allow_client_side_credentials" in raised.value.message + assert harness.store.calls == [] + assert harness.router.acompletion.await_count == 0 + + +async def test_create_forwards_row_credentials_when_the_admin_opted_in() -> None: + harness = make_runner( + content=CREDENTIAL_ROWS, general_settings=MappingProxyType({"allow_client_side_credentials": True}) + ) + _, finished = await harness.create_and_finish() + + assert finished.status == "completed" + assert finished.request_counts == BatchRequestCounts(completed=2, failed=0, total=2) + by_content = {call.kwargs["messages"][0]["content"]: call.kwargs for call in harness.router.acompletion.await_args_list} + assert by_content["hi 2"]["api_base"] == "https://evil.example" + assert "api_base" not in by_content["hi 1"] + + +async def test_running_batch_touches_its_row_until_it_finishes() -> None: + harness = make_runner(heartbeat_seconds=0.01) + + async def slow_dispatch(**_: object) -> ModelResponse: + await asyncio.sleep(0.05) + return chat_response("slow") + + harness.router.acompletion.side_effect = slow_dispatch + created, finished = await harness.create_and_finish() + + touches = harness.prisma.db.litellm_managedobjecttable.touches + assert finished.status == "completed" + assert touches + assert set(touches) == {(created.id, "user-1")} + assert [call.status for call in harness.store.calls] == ["validating", "in_progress", "finalizing", "completed"] + beats_at_finish = len(touches) + await asyncio.sleep(0.05) + assert len(touches) == beats_at_finish + + +async def test_fail_abandoned_marks_the_batch_failed_with_the_runner_lost_error() -> None: + harness = make_runner() + batch = seeded_batch(harness.store, "in_progress") + + failed = await harness.runner.fail_abandoned(batch, harness.user) + + assert failed.status == "failed" + assert failed.failed_at is not None + assert failed.errors is not None + assert [(error.message, error.code) for error in failed.errors.data or []] == [ + (litellm_executed_batches._RUNNER_LOST_MESSAGE, "runner_lost") + ] + assert harness.store.batch(batch.id).status == "failed" + assert [(call.status, call.create_if_missing) for call in harness.store.calls] == [("failed", False)] + + +async def test_running_batch_stops_and_writes_nothing_once_a_retriever_marked_it_failed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm_executed_batches, "_CANCEL_POLL_SECONDS", 0.0) + rows = jsonl(chat_row("row-1", "hi 1"), chat_row("row-2", "hi 2"), chat_row("row-3", "hi 3")) + harness = make_runner(content=rows, concurrency=1) + + def dispatch(metadata: Mapping[str, object], **_: object) -> ModelResponse: + running = harness.store.batch(str(metadata["batch_id"])) + harness.store.write(running.model_copy(update={"status": "failed"})) + return chat_response("hi 1") + + harness.router.acompletion.side_effect = dispatch + _, finished = await harness.create_and_finish() + + assert harness.router.acompletion.await_count == 1 + assert finished.status == "failed" + assert [call.status for call in harness.store.calls] == ["validating", "in_progress"] + assert harness.uploads.calls == [] + + @pytest.mark.parametrize( ("endpoint", "body", "method"), [ diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py b/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py index 067826004e2..81c5803da33 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py @@ -12,13 +12,20 @@ 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): @@ -45,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 @@ -153,3 +165,25 @@ async def test_upload_hands_the_prisma_client_to_the_storage_backend_factory(mon ) 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"] From 42271b282a7d195b0f8b0ac324b852ab8109b8fa Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 04:44:26 -0700 Subject: [PATCH 06/12] fix(batches): guard batch status writes against stale reads and disable per-line fallbacks --- .../litellm_executed_batches.py | 65 +++--- .../test_litellm_executed_batches.py | 194 +++++++++++++++--- 2 files changed, 198 insertions(+), 61 deletions(-) diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index 7bd4c678183..b121ad1590e 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -3,7 +3,7 @@ import json import time from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass -from datetime import datetime, timezone +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 @@ -28,11 +28,7 @@ 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, - convert_b64_uid_to_unified_uid, - get_batch_id_from_unified_batch_id, -) +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.table_repositories import ManagedObjectRepository @@ -41,6 +37,7 @@ from litellm.types.utils import LITELLM_EXECUTED_BATCH_PROVIDERS, ExtractedFileD if TYPE_CHECKING: from prisma import models as prisma_models + from prisma import types as prisma_types from litellm.router import Router @@ -154,7 +151,6 @@ class ManagedBatchStore(Protocol): user_api_key_dict: UserAPIKeyAuth, request_tags: Sequence[str] | None = None, persist_attribution: bool = False, - create_if_missing: bool = True, batch_processed: bool = False, ) -> None: ... @@ -375,10 +371,6 @@ def executed_batch_runner_lost(status: str, updated_at: datetime) -> bool: return (datetime.now(timezone.utc) - updated_at).total_seconds() > _STALE_AFTER_SECONDS -def _llm_batch_id_of(unified_batch_id: str) -> str: - return get_batch_id_from_unified_batch_id(convert_b64_uid_to_unified_uid(unified_batch_id)) - - class _StopWatch: def __init__(self, load_status: Callable[[], Awaitable[str | None]], interval_seconds: float) -> None: self._load_status = load_status @@ -488,8 +480,10 @@ class LiteLLMExecutedBatchRunner: cancelling: Final = current.model_copy( update=MappingProxyType({"status": "cancelling", "cancelling_at": int(time.time())}) ) - await self._store(cancelling, user_api_key_dict) - return cancelling + unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status} + if await self._store_unless_changed(cancelling, unchanged, user_api_key_dict): + 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") @@ -497,8 +491,16 @@ class LiteLLMExecutedBatchRunner: failed: Final = batch.model_copy( update=MappingProxyType({"status": "failed", "failed_at": int(time.time()), "errors": errors}) ) - await self._store(failed, user_api_key_dict) - return failed + 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._store_unless_changed(failed, still_abandoned, user_api_key_dict): + return failed + return await self._load_batch(batch.id) or batch def _body_rejection(self, model: str) -> BodyRejection: def reject(body: Mapping[str, object]) -> str | None: @@ -598,7 +600,9 @@ class LiteLLMExecutedBatchRunner: 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)}) + 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: @@ -651,19 +655,26 @@ class LiteLLMExecutedBatchRunner: updated: Final = current.model_copy( update=MappingProxyType({**fields, "status": status, f"{status}_at": int(time.time())}) ) - await self._store(updated, run.user_api_key_dict) - return status + unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status} + if await self._store_unless_changed(updated, unchanged, run.user_api_key_dict): + return status + return await self._advance(run, requested, fields) - async def _store(self, batch: LiteLLMBatch, user_api_key_dict: UserAPIKeyAuth) -> None: - await self.managed_files.store_unified_object_id( - unified_object_id=batch.id, - file_object=batch, - litellm_parent_otel_span=user_api_key_dict.parent_otel_span, - model_object_id=_llm_batch_id_of(batch.id), - file_purpose="batch", - user_api_key_dict=user_api_key_dict, - create_if_missing=False, + async def _store_unless_changed( + self, + batch: LiteLLMBatch, + guard: "prisma_types.LiteLLM_ManagedObjectTableWhereInput", + user_api_key_dict: UserAPIKeyAuth, + ) -> bool: + updated_rows: Final = await ManagedObjectRepository(self.prisma_client).table.update_many( + where={"unified_object_id": batch.id, **guard}, # mutable-ok: Prisma filter + data={ # mutable-ok: Prisma payload + "file_object": batch.model_dump_json(), + "status": batch.status, + "updated_by": user_api_key_dict.user_id, + }, ) + return updated_rows > 0 async def _find_row(self, unified_batch_id: str) -> "prisma_models.LiteLLM_ManagedObjectTable | None": return await ManagedObjectRepository(self.prisma_client).table.find_first( diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index e5ed873a29d..cb566d4dca7 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -105,6 +105,10 @@ class ProviderRateLimited(Exception): class StoredObject: file_object: str status: str + updated_at: datetime + + def batch(self) -> LiteLLMBatch: + return LiteLLMBatch.model_validate_json(self.file_object) @dataclass(frozen=True, slots=True) @@ -114,10 +118,20 @@ class StoreCall: status: str request_tags: tuple[str, ...] | None persist_attribution: bool - create_if_missing: bool batch_processed: bool +@dataclass(frozen=True, slots=True) +class StatusWrite: + unified_object_id: str + status: str + columns: frozenset[str] + + +STATUS_WRITE_COLUMNS: Final = frozenset({"file_object", "status", "updated_by"}) +STALE: Final = timedelta(seconds=litellm_executed_batches._STALE_AFTER_SECONDS + 20) + + class FakeManagedBatchStore: def __init__(self, files: Mapping[str, LiteLLM_ManagedFileTable]) -> None: self.files = files @@ -142,7 +156,6 @@ class FakeManagedBatchStore: user_api_key_dict: UserAPIKeyAuth, request_tags: Sequence[str] | None = None, persist_attribution: bool = False, - create_if_missing: bool = True, batch_processed: bool = False, ) -> None: self.calls.append( @@ -152,18 +165,18 @@ class FakeManagedBatchStore: status=file_object.status, request_tags=tuple(request_tags) if request_tags is not None else None, persist_attribution=persist_attribution, - create_if_missing=create_if_missing, batch_processed=batch_processed, ) ) - if create_if_missing or unified_object_id in self.objects: - self.write(file_object) + self.write(file_object) - def write(self, batch: LiteLLMBatch) -> None: - self.objects[batch.id] = StoredObject(file_object=batch.model_dump_json(), status=batch.status) + def write(self, batch: LiteLLMBatch, age: timedelta = timedelta(0)) -> None: + self.objects[batch.id] = StoredObject( + file_object=batch.model_dump_json(), status=batch.status, updated_at=datetime.now(timezone.utc) - age + ) def batch(self, unified_batch_id: str) -> LiteLLMBatch: - return LiteLLMBatch.model_validate_json(self.objects[unified_batch_id].file_object) + return self.objects[unified_batch_id].batch() REAL_HOOK: Final = _PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()) @@ -174,26 +187,51 @@ class RealIdManagedBatchStore(FakeManagedBatchStore): return REAL_HOOK.get_unified_batch_id(batch_id=batch_id, model_id=model_id) +def row_matches(row: StoredObject, where: Mapping[str, object]) -> bool: + if "status" in where and row.status != where["status"]: + return False + match where.get("updated_at"): + case {"lt": datetime() as before}: + return row.updated_at < before + case _: + return True + + class FakeManagedObjectTable: - def __init__(self, objects: Mapping[str, StoredObject]) -> None: + def __init__(self, objects: dict[str, StoredObject]) -> None: self.objects = objects self.touches: list[tuple[str, str | None]] = [] + self.writes: list[StatusWrite] = [] + self.after_read: Callable[[StoredObject | None], None] | None = None async def find_first(self, where: Mapping[str, str]) -> StoredObject | None: - return self.objects.get(where["unified_object_id"]) + row = self.objects.get(where["unified_object_id"]) + if self.after_read is not None: + self.after_read(row) + return row - async def update_many(self, where: Mapping[str, str], data: Mapping[str, str | None]) -> int: - self.touches.append((where["unified_object_id"], data["updated_by"])) + async def update_many(self, where: Mapping[str, object], data: Mapping[str, str | None]) -> int: + unified_object_id = str(where["unified_object_id"]) + row = self.objects.get(unified_object_id) + if row is None or not row_matches(row, where): + return 0 + now = datetime.now(timezone.utc) + if "status" not in data: + self.touches.append((unified_object_id, data["updated_by"])) + self.objects[unified_object_id] = StoredObject(row.file_object, row.status, now) + return 1 + self.writes.append(StatusWrite(unified_object_id, str(data["status"]), frozenset(data))) + self.objects[unified_object_id] = StoredObject(str(data["file_object"]), str(data["status"]), now) return 1 class FakeDb: - def __init__(self, objects: Mapping[str, StoredObject]) -> None: + def __init__(self, objects: dict[str, StoredObject]) -> None: self.litellm_managedobjecttable = FakeManagedObjectTable(objects) class FakePrismaClient: - def __init__(self, objects: Mapping[str, StoredObject]) -> None: + def __init__(self, objects: dict[str, StoredObject]) -> None: self.db = FakeDb(objects) @@ -320,6 +358,13 @@ class Harness: await asyncio.gather(*list(litellm_executed_batches._RUNNING_BATCHES)) return created, self.store.batch(created.id) + @property + def table(self) -> FakeManagedObjectTable: + return self.prisma.db.litellm_managedobjecttable + + def written_statuses(self) -> list[str]: + return [write.status for write in self.table.writes] + def make_runner( content: bytes = TWO_CHAT_ROWS, @@ -354,7 +399,9 @@ def make_runner( return Harness(runner, store, router, uploads, storage, storage_factory, prisma, user) -def seeded_batch(store: FakeManagedBatchStore, status: Literal["in_progress", "completed"]) -> LiteLLMBatch: +def seeded_batch( + store: FakeManagedBatchStore, status: Literal["in_progress", "completed"], age: timedelta = timedelta(0) +) -> LiteLLMBatch: batch = LiteLLMBatch( id=store.get_unified_batch_id(batch_id="litellm_batch_seed", model_id=DEPLOYMENT_ID), object="batch", @@ -365,7 +412,7 @@ def seeded_batch(store: FakeManagedBatchStore, status: Literal["in_progress", "c created_at=1, model=BATCH_MODEL, ) - store.write(batch) + store.write(batch, age) return batch @@ -745,7 +792,9 @@ async def test_create_forwards_row_credentials_when_the_admin_opted_in() -> None assert finished.status == "completed" assert finished.request_counts == BatchRequestCounts(completed=2, failed=0, total=2) - by_content = {call.kwargs["messages"][0]["content"]: call.kwargs for call in harness.router.acompletion.await_args_list} + by_content = { + call.kwargs["messages"][0]["content"]: call.kwargs for call in harness.router.acompletion.await_args_list + } assert by_content["hi 2"]["api_base"] == "https://evil.example" assert "api_base" not in by_content["hi 1"] @@ -760,19 +809,20 @@ async def test_running_batch_touches_its_row_until_it_finishes() -> None: harness.router.acompletion.side_effect = slow_dispatch created, finished = await harness.create_and_finish() - touches = harness.prisma.db.litellm_managedobjecttable.touches + touches = harness.table.touches assert finished.status == "completed" assert touches assert set(touches) == {(created.id, "user-1")} - assert [call.status for call in harness.store.calls] == ["validating", "in_progress", "finalizing", "completed"] + assert [call.status for call in harness.store.calls] == ["validating"] + assert harness.written_statuses() == ["in_progress", "finalizing", "completed"] beats_at_finish = len(touches) await asyncio.sleep(0.05) assert len(touches) == beats_at_finish -async def test_fail_abandoned_marks_the_batch_failed_with_the_runner_lost_error() -> None: +async def test_fail_abandoned_marks_a_stale_batch_failed_with_the_runner_lost_error() -> None: harness = make_runner() - batch = seeded_batch(harness.store, "in_progress") + batch = seeded_batch(harness.store, "in_progress", age=STALE) failed = await harness.runner.fail_abandoned(batch, harness.user) @@ -783,7 +833,63 @@ async def test_fail_abandoned_marks_the_batch_failed_with_the_runner_lost_error( (litellm_executed_batches._RUNNER_LOST_MESSAGE, "runner_lost") ] assert harness.store.batch(batch.id).status == "failed" - assert [(call.status, call.create_if_missing) for call in harness.store.calls] == [("failed", False)] + assert harness.store.calls == [] + assert harness.table.writes == [StatusWrite(batch.id, "failed", STATUS_WRITE_COLUMNS)] + + +async def test_fail_abandoned_leaves_a_batch_that_finished_after_the_stale_read() -> None: + harness = make_runner() + stale_read = seeded_batch(harness.store, "in_progress", age=STALE) + harness.store.write(stale_read.model_copy(update={"status": "completed", "output_file_id": "out-1"}), age=STALE) + + current = await harness.runner.fail_abandoned(stale_read, harness.user) + + assert (current.status, current.output_file_id) == ("completed", "out-1") + assert harness.store.batch(stale_read.id).status == "completed" + assert harness.table.writes == [] + + +async def test_fail_abandoned_leaves_a_batch_its_runner_touched_since_the_read() -> None: + harness = make_runner() + batch = seeded_batch(harness.store, "in_progress", age=STALE) + harness.store.write(batch) + + current = await harness.runner.fail_abandoned(batch, harness.user) + + assert current.status == "in_progress" + assert harness.store.batch(batch.id).status == "in_progress" + assert harness.table.writes == [] + + +async def test_run_does_not_reverse_a_failure_written_between_its_read_and_its_completed_write() -> None: + harness = make_runner() + + def fail_once_finalizing_is_read(row: StoredObject | None) -> None: + if row is not None and row.status == "finalizing": + harness.store.write(row.batch().model_copy(update={"status": "failed"})) + + harness.table.after_read = fail_once_finalizing_is_read + _, finished = await harness.create_and_finish() + + assert finished.status == "failed" + assert finished.output_file_id is None + assert harness.written_statuses() == ["in_progress", "finalizing"] + + +async def test_run_honours_a_cancel_written_between_its_read_and_its_finalizing_write() -> None: + harness = make_runner(content=jsonl(chat_row("row-1", "hi 1"))) + + def cancel_once_the_row_is_dispatched(row: StoredObject | None) -> None: + if row is not None and row.status == "in_progress" and harness.router.acompletion.await_count == 1: + harness.store.write(row.batch().model_copy(update={"status": "cancelling"})) + + harness.table.after_read = cancel_once_the_row_is_dispatched + _, finished = await harness.create_and_finish() + + assert finished.status == "cancelled" + assert finished.request_counts == BatchRequestCounts(completed=1, failed=0, total=1) + assert finished.output_file_id == "unified-output-1" + assert harness.written_statuses() == ["in_progress", "cancelling", "cancelled"] async def test_running_batch_stops_and_writes_nothing_once_a_retriever_marked_it_failed( @@ -803,7 +909,8 @@ async def test_running_batch_stops_and_writes_nothing_once_a_retriever_marked_it assert harness.router.acompletion.await_count == 1 assert finished.status == "failed" - assert [call.status for call in harness.store.calls] == ["validating", "in_progress"] + assert [call.status for call in harness.store.calls] == ["validating"] + assert harness.written_statuses() == ["in_progress"] assert harness.uploads.calls == [] @@ -828,6 +935,7 @@ async def test_each_endpoint_awaits_only_its_router_method( assert awaited == {name: int(name == method) for name in ROUTER_METHODS} kwargs = getattr(harness.router, method).await_args.kwargs assert kwargs["model"] == BATCH_MODEL + assert kwargs["disable_fallbacks"] is True assert all(kwargs[key] == value for key, value in body.items()) @@ -845,7 +953,7 @@ async def test_cancel_terminal_batch_is_400() -> None: await harness.runner.cancel(batch.id, harness.user) assert raised.value.code == "400" assert "completed" in raised.value.message - assert harness.store.calls == [] + assert harness.table.writes == [] async def test_cancel_marks_a_running_batch_cancelling_once() -> None: @@ -857,12 +965,30 @@ async def test_cancel_marks_a_running_batch_cancelling_once() -> None: assert cancelled.status == "cancelling" assert cancelled.cancelling_at is not None assert harness.store.batch(batch.id).status == "cancelling" - assert [(call.status, call.create_if_missing) for call in harness.store.calls] == [("cancelling", False)] + assert harness.store.calls == [] + assert harness.table.writes == [StatusWrite(batch.id, "cancelling", STATUS_WRITE_COLUMNS)] again = await harness.runner.cancel(batch.id, harness.user) assert again.model_dump() == cancelled.model_dump() - assert len(harness.store.calls) == 1 + assert len(harness.table.writes) == 1 + + +async def test_cancel_racing_a_completion_is_400_and_leaves_the_batch_completed() -> None: + harness = make_runner() + batch = seeded_batch(harness.store, "in_progress") + + def complete_once_read(row: StoredObject | None) -> None: + if row is not None and row.status == "in_progress": + harness.store.write(row.batch().model_copy(update={"status": "completed"})) + + harness.table.after_read = complete_once_read + with pytest.raises(ProxyException) as raised: + await harness.runner.cancel(batch.id, harness.user) + + assert raised.value.code == "400" + assert harness.store.batch(batch.id).status == "completed" + assert harness.table.writes == [] async def test_running_batch_skips_the_remaining_rows_after_an_operator_cancel( @@ -905,10 +1031,10 @@ async def test_only_the_create_write_carries_attribution_and_billing_flags() -> harness = make_runner() await harness.create_and_finish() - assert [call.status for call in harness.store.calls] == ["validating", "in_progress", "finalizing", "completed"] - flags = [(call.persist_attribution, call.batch_processed, call.create_if_missing) for call in harness.store.calls] - assert flags[0] == (True, True, True) - assert flags[1:] == [(False, False, False)] * 3 + assert [(call.status, call.persist_attribution, call.batch_processed) for call in harness.store.calls] == [ + ("validating", True, True) + ] + assert [write.columns for write in harness.table.writes] == [STATUS_WRITE_COLUMNS] * 3 async def test_run_completes_under_the_real_hooks_base64_batch_id() -> None: @@ -917,9 +1043,8 @@ async def test_run_completes_under_the_real_hooks_base64_batch_id() -> None: assert _is_base64_encoded_unified_file_id(created.id) assert finished.status == "completed" - llm_batch_id = harness.store.calls[0].model_object_id - assert llm_batch_id.startswith("litellm_batch_") - assert [call.model_object_id for call in harness.store.calls] == [llm_batch_id] * 4 + assert [call.model_object_id.startswith("litellm_batch_") for call in harness.store.calls] == [True] + assert [write.unified_object_id for write in harness.table.writes] == [created.id] * 3 async def test_cancel_works_under_the_real_hooks_base64_batch_id() -> None: @@ -928,4 +1053,5 @@ async def test_cancel_works_under_the_real_hooks_base64_batch_id() -> None: cancelled = await harness.runner.cancel(batch.id, harness.user) assert cancelled.status == "cancelling" - assert [call.model_object_id for call in harness.store.calls] == ["litellm_batch_seed"] + assert harness.store.batch(batch.id).status == "cancelling" + assert [write.unified_object_id for write in harness.table.writes] == [batch.id] From 2ee8c1bd0e6ac9125befdcb8fb479773d7c5dd07 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 06:25:42 -0700 Subject: [PATCH 07/12] fix(batches): enforce the completion window and guard executed-batch id parsing --- .../litellm_executed_batches.py | 82 ++++++++++++++----- .../openai_files_endpoints/common_utils.py | 3 +- .../test_litellm_executed_batches.py | 44 ++++++++++ .../test_files_common_utils.py | 2 + 4 files changed, 111 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index b121ad1590e..5a7061d9ab1 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -42,7 +42,9 @@ if TYPE_CHECKING: 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"] +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) @@ -52,6 +54,7 @@ _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", @@ -61,7 +64,7 @@ _ROUTER_METHODS: Final[Mapping[BatchEndpoint, str]] = MappingProxyType( } ) _CANCELLING_TRANSITIONS: Final[Mapping[BatchStatus, BatchStatus]] = MappingProxyType( - {"completed": "cancelled", "in_progress": "cancelling", "finalizing": "cancelling"} + {"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 " @@ -89,11 +92,16 @@ class _ResultResponse(TypedDict): 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] - error: ReadOnly[None] + response: ReadOnly[_ResultResponse | None] + error: ReadOnly[_LineError | None] class BatchInputLine(BaseModel): @@ -122,6 +130,11 @@ class RowOutcome: succeeded: bool +@dataclass(frozen=True, slots=True) +class ExpiredRow: + custom_id: str + + @dataclass(frozen=True, slots=True) class _BatchRun: unified_batch_id: str @@ -131,6 +144,7 @@ class _BatchRun: lines: tuple[BatchInputLine, ...] user_api_key_dict: UserAPIKeyAuth request_tags: tuple[str, ...] + deadline: float @runtime_checkable @@ -339,16 +353,30 @@ def _error_body(error: Exception) -> _ErrorBody: return body -def _result_line(outcome: RowOutcome) -> _ResultLine: +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": { - "status_code": outcome.status_code, - "request_id": f"req_{uuid_module.uuid4().hex[:24]}", - "body": outcome.body, - }, - "error": None, + "response": _line_response(outcome), + "error": _line_error(outcome), } return line @@ -399,6 +427,7 @@ class LiteLLMExecutedBatchRunner: 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: @@ -409,6 +438,7 @@ class LiteLLMExecutedBatchRunner: 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 @@ -429,7 +459,8 @@ class LiteLLMExecutedBatchRunner: 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) - created_at: Final = int(time.time()) + now: Final = time.time() + created_at: Final = int(now) batch: Final = LiteLLMBatch( id=unified_batch_id, object="batch", @@ -438,7 +469,7 @@ class LiteLLMExecutedBatchRunner: completion_window="24h", status="validating", created_at=created_at, - expires_at=created_at + _COMPLETION_WINDOW_SECONDS, + 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)), @@ -463,6 +494,7 @@ class LiteLLMExecutedBatchRunner: 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) @@ -572,14 +604,21 @@ class LiteLLMExecutedBatchRunner: 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 outcome.succeeded) - failed: Final = tuple(outcome for outcome in outcomes if not outcome.succeeded) + 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, - "completed", + final_status, MappingProxyType( {"output_file_id": output_file_id, "error_file_id": error_file_id, "request_counts": request_counts} ), @@ -587,12 +626,17 @@ class LiteLLMExecutedBatchRunner: async def _run_row( self, run: _BatchRun, line: BatchInputLine, watch: _StopWatch, semaphore: asyncio.Semaphore - ) -> RowOutcome | None: + ) -> 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: - body: Final = await self._dispatch(run, line) + body: Final = await asyncio.wait_for(self._dispatch(run, line), timeout=remaining) + except asyncio.TimeoutError: + return ExpiredRow(custom_id=line.custom_id) 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 @@ -621,7 +665,7 @@ class LiteLLMExecutedBatchRunner: } async def _upload_results( - self, run: _BatchRun, kind: Literal["output", "error"], outcomes: Sequence[RowOutcome] + self, run: _BatchRun, kind: Literal["output", "error"], outcomes: Sequence[RowOutcome | ExpiredRow] ) -> str | None: if not outcomes: return None diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 25365d34187..ccc3b5b7f5f 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -181,7 +181,8 @@ def get_batch_id_from_unified_batch_id(file_id: str) -> str: def is_litellm_executed_batch(decoded_unified_batch_id: str) -> bool: - return get_batch_id_from_unified_batch_id(decoded_unified_batch_id).startswith(LITELLM_EXECUTED_BATCH_ID_PREFIX) + _, 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: diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index cb566d4dca7..860827e8fbd 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -53,6 +53,7 @@ ALL_STATUSES: Final[tuple[BatchStatus, ...]] = ( "failed", "cancelling", "cancelled", + "expired", ) @@ -375,6 +376,7 @@ def make_runner( store_factory: Callable[[Mapping[str, LiteLLM_ManagedFileTable]], FakeManagedBatchStore] = FakeManagedBatchStore, general_settings: Mapping[str, object] = MappingProxyType({}), heartbeat_seconds: float = 30.0, + completion_window_seconds: float = 24 * 60 * 60, ) -> Harness: store = store_factory({INPUT_FILE_ID: managed_input_file()} if files is None else files) router = FakeRouter() @@ -393,6 +395,7 @@ def make_runner( general_settings=general_settings, concurrency=concurrency, heartbeat_seconds=heartbeat_seconds, + completion_window_seconds=completion_window_seconds, storage_backend_factory=storage_factory, upload_result_file=uploads, ) @@ -464,6 +467,7 @@ def test_resolve_transition_keeps_the_requested_status_unless_cancelling(current ("requested", "expected"), [ ("completed", "cancelled"), + ("expired", "cancelled"), ("in_progress", "cancelling"), ("finalizing", "cancelling"), ("failed", "failed"), @@ -1014,6 +1018,46 @@ async def test_running_batch_skips_the_remaining_rows_after_an_operator_cancel( assert (finished.output_file_id, finished.error_file_id) == ("unified-output-1", None) +async def test_batch_expires_at_the_completion_window_and_keeps_what_finished() -> None: + rows = jsonl(chat_row("row-1", "hi 1"), chat_row("row-2", "hi 2"), chat_row("row-3", "hi 3")) + harness = make_runner(content=rows, concurrency=1, completion_window_seconds=0.2) + reply = chat_response("hi 1") + + async def dispatch(messages: Sequence[Mapping[str, str]], **_: object) -> ModelResponse: + if messages[0]["content"] == "hi 1": + return reply + await asyncio.Event().wait() + raise AssertionError("a row still running at the completion window must be cut off") + + harness.router.acompletion.side_effect = dispatch + created, finished = await harness.create_and_finish() + + assert created.expires_at == created.created_at + assert finished.status == "expired" + assert finished.expired_at is not None + assert finished.request_counts == BatchRequestCounts(completed=1, failed=2, total=3) + assert (finished.output_file_id, finished.error_file_id) == ("unified-output-1", "unified-output-2") + assert set(harness.uploads.calls[0].lines()) == {"row-1"} + error_lines = harness.uploads.calls[1].lines() + assert set(error_lines) == {"row-2", "row-3"} + for line in error_lines.values(): + assert line["response"] is None + error = line["error"] + assert isinstance(error, dict) + assert error["code"] == "batch_expired" + + +async def test_batch_created_past_its_window_dispatches_nothing() -> None: + harness = make_runner(completion_window_seconds=0) + _, finished = await harness.create_and_finish() + + assert harness.router.acompletion.await_count == 0 + assert finished.status == "expired" + assert finished.request_counts == BatchRequestCounts(completed=0, failed=2, total=2) + assert (finished.output_file_id, finished.error_file_id) == (None, "unified-output-1") + assert set(harness.uploads.calls[0].lines()) == {"row-1", "row-2"} + + async def test_upload_failure_marks_the_batch_failed() -> None: harness = make_runner(upload_error=RuntimeError("storage exploded")) _, finished = await harness.create_and_finish() diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index f88d94d2c08..b7e3088ff01 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -487,6 +487,8 @@ class TestCompletedBatchSafeToRetire: ("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): From 03a63db1fded5690687f2fcc24f199e9f3a11df8 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 11:54:57 -0700 Subject: [PATCH 08/12] fix(batches): keep provider timeouts as failed rows and move batch rows behind a repository --- litellm/proxy/batches_endpoints/endpoints.py | 2 + .../litellm_executed_batches.py | 73 ++++++------------- .../repositories/managed_batch_repository.py | 48 ++++++++++++ .../test_litellm_executed_batches.py | 29 ++++++++ 4 files changed, 100 insertions(+), 52 deletions(-) create mode 100644 litellm/repositories/managed_batch_repository.py diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 1284690172a..3f6c9f4d6ed 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -66,6 +66,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import request_tags_from_metadata from litellm.proxy.route_llm_request import raise_if_required_body_param_missing from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy, is_known_model +from litellm.repositories.managed_batch_repository import ManagedBatchRepository from litellm.repositories.table_repositories import ManagedFileRepository from litellm.router import Router from litellm.types.llms.openai import LiteLLMBatchCreateRequest @@ -98,6 +99,7 @@ def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyL 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, ) diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index 5a7061d9ab1..67201d99422 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -31,12 +31,11 @@ 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.table_repositories import ManagedObjectRepository +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 models as prisma_models from prisma import types as prisma_types from litellm.router import Router @@ -342,10 +341,6 @@ def _status_code_of(error: Exception) -> int: return status_code if isinstance(status_code, int) else 500 -def _batch_of(blob: object) -> LiteLLMBatch: - return LiteLLMBatch.model_validate_json(blob) if isinstance(blob, str) else LiteLLMBatch.model_validate(blob) - - def _error_body(error: Exception) -> _ErrorBody: body: Final[_ErrorBody] = { "error": {"message": str(error), "type": type(error).__name__, "param": None, "code": None} @@ -423,6 +418,7 @@ class LiteLLMExecutedBatchRunner: 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, @@ -434,6 +430,7 @@ class LiteLLMExecutedBatchRunner: 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 @@ -502,7 +499,7 @@ class LiteLLMExecutedBatchRunner: return batch async def cancel(self, unified_batch_id: str, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch: - current: Final = await self._load_batch(unified_batch_id) + 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: @@ -513,7 +510,7 @@ class LiteLLMExecutedBatchRunner: update=MappingProxyType({"status": "cancelling", "cancelling_at": int(time.time())}) ) unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status} - if await self._store_unless_changed(cancelling, unchanged, user_api_key_dict): + 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) @@ -530,9 +527,9 @@ class LiteLLMExecutedBatchRunner: "status": batch.status, "updated_at": untouched, } - if await self._store_unless_changed(failed, still_abandoned, user_api_key_dict): + if await self.batches.compare_and_set(failed, still_abandoned, user_api_key_dict.user_id): return failed - return await self._load_batch(batch.id) or batch + 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: @@ -591,14 +588,11 @@ class LiteLLMExecutedBatchRunner: verbose_proxy_logger.warning("LiteLLM-executed batch %s heartbeat failed: %s", run.unified_batch_id, e) async def _touch(self, run: _BatchRun) -> None: - await ManagedObjectRepository(self.prisma_client).table.update_many( - where={"unified_object_id": run.unified_batch_id}, # mutable-ok: Prisma filter - data={"updated_by": run.user_api_key_dict.user_id}, # mutable-ok: Prisma payload - ) + 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._load_status(run.unified_batch_id), _CANCEL_POLL_SECONDS) + 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) @@ -634,14 +628,18 @@ class LiteLLMExecutedBatchRunner: if remaining <= 0: return ExpiredRow(custom_id=line.custom_id) try: - body: Final = await asyncio.wait_for(self._dispatch(run, line), timeout=remaining) + return await asyncio.wait_for(self._row_outcome(run, line), timeout=remaining) except asyncio.TimeoutError: return ExpiredRow(custom_id=line.custom_id) - 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 _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( @@ -690,7 +688,7 @@ class LiteLLMExecutedBatchRunner: async def _advance( self, run: _BatchRun, requested: BatchStatus, fields: Mapping[str, object] = _NO_FIELDS ) -> BatchStatus | None: - current: Final = await self._load_batch(run.unified_batch_id) + 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: @@ -700,39 +698,10 @@ class LiteLLMExecutedBatchRunner: update=MappingProxyType({**fields, "status": status, f"{status}_at": int(time.time())}) ) unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status} - if await self._store_unless_changed(updated, unchanged, run.user_api_key_dict): + 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) - async def _store_unless_changed( - self, - batch: LiteLLMBatch, - guard: "prisma_types.LiteLLM_ManagedObjectTableWhereInput", - user_api_key_dict: UserAPIKeyAuth, - ) -> bool: - updated_rows: Final = await ManagedObjectRepository(self.prisma_client).table.update_many( - where={"unified_object_id": batch.id, **guard}, # mutable-ok: Prisma filter - data={ # mutable-ok: Prisma payload - "file_object": batch.model_dump_json(), - "status": batch.status, - "updated_by": user_api_key_dict.user_id, - }, - ) - return updated_rows > 0 - - async def _find_row(self, unified_batch_id: str) -> "prisma_models.LiteLLM_ManagedObjectTable | None": - return await ManagedObjectRepository(self.prisma_client).table.find_first( - where={"unified_object_id": unified_batch_id} # mutable-ok: Prisma filter - ) - - 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 - def _record_batch_created(model: str, provider: str, user_api_key_dict: UserAPIKeyAuth) -> None: prometheus_logger: Final = PrometheusLogger.get_instance() diff --git a/litellm/repositories/managed_batch_repository.py b/litellm/repositories/managed_batch_repository.py new file mode 100644 index 00000000000..3f85251fdbd --- /dev/null +++ b/litellm/repositories/managed_batch_repository.py @@ -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 + ) diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index 860827e8fbd..6f2341a578c 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -35,6 +35,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( is_litellm_executed_batch, ) from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.repositories.managed_batch_repository import ManagedBatchRepository from litellm.router import Router from litellm.types.llms.openai import LiteLLMBatchCreateRequest, OpenAIFileObject, OpenAIFilesPurpose from litellm.types.utils import EmbeddingResponse, LiteLLMBatch, ModelResponse, SpecialEnums @@ -391,6 +392,7 @@ def make_runner( llm_router=cast("Router", router), prisma_client=cast("PrismaClient", prisma), managed_files=store, + batches=ManagedBatchRepository(prisma), proxy_logging_obj=MagicMock(spec=ProxyLogging), general_settings=general_settings, concurrency=concurrency, @@ -1047,6 +1049,33 @@ async def test_batch_expires_at_the_completion_window_and_keeps_what_finished() assert error["code"] == "batch_expired" +async def test_a_provider_timeout_fails_its_row_without_expiring_the_batch() -> None: + harness = make_runner() + reply = chat_response("hi 2") + + def dispatch(messages: Sequence[Mapping[str, str]], **_: object) -> ModelResponse: + if messages[0]["content"] == "hi 1": + raise asyncio.TimeoutError("the provider took too long") + return reply + + harness.router.acompletion.side_effect = dispatch + _, finished = await harness.create_and_finish() + + assert finished.status == "completed" + assert finished.expired_at is None + assert finished.request_counts == BatchRequestCounts(completed=1, failed=1, total=2) + assert set(harness.uploads.calls[0].lines()) == {"row-2"} + error_lines = harness.uploads.calls[1].lines() + assert set(error_lines) == {"row-1"} + assert error_lines["row-1"]["error"] is None + response = error_lines["row-1"]["response"] + assert isinstance(response, dict) + assert response["status_code"] == 500 + assert response["body"] == { + "error": {"message": "the provider took too long", "type": "TimeoutError", "param": None, "code": None} + } + + async def test_batch_created_past_its_window_dispatches_nothing() -> None: harness = make_runner(completion_window_seconds=0) _, finished = await harness.create_and_finish() From 9f0eb5082ae4e2360f68b7cba19478023fc1ccbb Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 12:26:11 -0700 Subject: [PATCH 09/12] fix(batches): authorize executed upload targets before the files api probe A batch upload naming a model on a LiteLLM-executed provider now checks that the key may call that model before the upstream server is probed for a Files API, matching the order batch create already uses. Only targets on an executed provider are checked here, so provider-model uploads keep their existing behavior. File content reads and writes move out of the storage backend into ManagedFileContentRepository, so the backend no longer queries Prisma directly. --- .../files/litellm_db_storage_backend.py | 39 ++++--------------- .../openai_files_endpoints/files_endpoints.py | 23 +++++++++-- .../managed_file_content_repository.py | 30 ++++++++++++++ .../files/test_storage_backend_factory.py | 14 +++++-- .../test_files_endpoint.py | 24 ++++++++++++ 5 files changed, 91 insertions(+), 39 deletions(-) create mode 100644 litellm/repositories/managed_file_content_repository.py diff --git a/litellm/llms/base_llm/files/litellm_db_storage_backend.py b/litellm/llms/base_llm/files/litellm_db_storage_backend.py index bca4b8f4c6f..a686062b2f7 100644 --- a/litellm/llms/base_llm/files/litellm_db_storage_backend.py +++ b/litellm/llms/base_llm/files/litellm_db_storage_backend.py @@ -1,13 +1,9 @@ -from collections.abc import Mapping from typing import TYPE_CHECKING, Final from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend -from litellm.repositories.prisma_protocols import TableActions -from litellm.repositories.table_repositories import PrismaTableRepository +from litellm.repositories.managed_file_content_repository import ManagedFileContentRepository if TYPE_CHECKING: - from prisma import models as prisma_models - from litellm.proxy.utils import PrismaClient LITELLM_DB_STORAGE_BACKEND_NAME: Final = "litellm_db" @@ -20,21 +16,9 @@ def storage_url_to_row_id(storage_url: str) -> str: return storage_url.removeprefix(LITELLM_DB_STORAGE_URL_PREFIX) -def _where_id(storage_url: str) -> Mapping[str, str]: - return {"id": storage_url_to_row_id(storage_url)} # mutable-ok: Prisma filter - - -class ManagedFileContentRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileContentTable"]): - table_name = "litellm_managedfilecontenttable" - - class LiteLLMDbStorageBackend(BaseFileStorageBackend): def __init__(self, prisma_client: "PrismaClient") -> None: - self._prisma_client = prisma_client - - @property - def _table(self) -> "TableActions[prisma_models.LiteLLM_ManagedFileContentTable]": - return ManagedFileContentRepository(self._prisma_client).table + self._contents = ManagedFileContentRepository(prisma_client) async def upload_file( self, @@ -44,22 +28,13 @@ class LiteLLMDbStorageBackend(BaseFileStorageBackend): path_prefix: str | None = None, file_naming_strategy: str = "uuid", ) -> str: - from prisma import Base64 - - data: Final = {"content": Base64.encode(file_content)} # mutable-ok: Prisma payload - row: Final = await self._table.create(data=data) - return f"{LITELLM_DB_STORAGE_URL_PREFIX}{row.id}" + return f"{LITELLM_DB_STORAGE_URL_PREFIX}{await self._contents.store(file_content)}" async def download_file(self, storage_url: str) -> bytes: - row: Final = await self._table.find_unique(where=_where_id(storage_url)) - if row is None: + 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 row.content.decode() + return content async def delete_file(self, storage_url: str) -> None: - from prisma.errors import RecordNotFoundError - - try: - await self._table.delete(where=_where_id(storage_url)) - except RecordNotFoundError: - return + await self._contents.delete(storage_url_to_row_id(storage_url)) diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 6d60bc3fda5..a5921cc6380 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -37,7 +37,10 @@ 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 resolve_litellm_executed_provider +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, @@ -69,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, @@ -102,16 +106,29 @@ 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], - team_id: str | None, + user_api_key_dict: UserAPIKeyAuth, ) -> 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) + ) + ) providers: Final = await asyncio.gather( *(resolve_litellm_executed_provider(llm_router, candidate, team_id) for candidate in candidates) ) @@ -289,7 +306,7 @@ async def route_create_file( """ executed_model: Final = await _litellm_executed_batch_input_model( - llm_router, purpose, model, target_model_names_list, user_api_key_dict.team_id + llm_router, purpose, model, target_model_names_list, user_api_key_dict ) explicit_storage: Final = target_storage if target_storage and target_storage != "default" else None storage: Final = explicit_storage or (LITELLM_DB_STORAGE_BACKEND_NAME if executed_model is not None else None) diff --git a/litellm/repositories/managed_file_content_repository.py b/litellm/repositories/managed_file_content_repository.py new file mode 100644 index 00000000000..8810269279c --- /dev/null +++ b/litellm/repositories/managed_file_content_repository.py @@ -0,0 +1,30 @@ +from typing import TYPE_CHECKING, Final + +from litellm.repositories.table_repositories import PrismaTableRepository + +if TYPE_CHECKING: + from prisma import models as prisma_models # noqa: F401 # used by the quoted base-class subscript + + +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 = 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 diff --git a/tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py b/tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py index 39b0adb56fc..945691c5b98 100644 --- a/tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py +++ b/tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py @@ -1,21 +1,27 @@ -from unittest.mock import MagicMock +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 -def test_litellm_db_backend_is_built_on_the_given_prisma_client(): - prisma_client = MagicMock() +@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) - assert backend._table is prisma_client.db.litellm_managedfilecontenttable + 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(): diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 6c0f4c012ba..6787aaa3525 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -706,6 +706,30 @@ def test_batch_upload_for_a_litellm_executed_model_is_kept_by_litellm( 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 From 020cbba4ddc4c336b4fd6a39de146aa1392f1e2d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 12:39:44 -0700 Subject: [PATCH 10/12] refactor(batches): annotate the stored file row so its model import is a real use --- litellm/repositories/managed_file_content_repository.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/litellm/repositories/managed_file_content_repository.py b/litellm/repositories/managed_file_content_repository.py index 8810269279c..c55d0060080 100644 --- a/litellm/repositories/managed_file_content_repository.py +++ b/litellm/repositories/managed_file_content_repository.py @@ -3,7 +3,7 @@ from typing import TYPE_CHECKING, Final from litellm.repositories.table_repositories import PrismaTableRepository if TYPE_CHECKING: - from prisma import models as prisma_models # noqa: F401 # used by the quoted base-class subscript + from prisma import models as prisma_models class ManagedFileContentRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileContentTable"]): @@ -18,7 +18,9 @@ class ManagedFileContentRepository(PrismaTableRepository["prisma_models.LiteLLM_ return row.id async def load(self, row_id: str) -> bytes | None: - row: Final = await self.table.find_unique(where={"id": row_id}) # mutable-ok: prisma filters are plain dicts + 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 549548de62454448b1b89ea3b362f9d0e980ee3d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 13:20:02 -0700 Subject: [PATCH 11/12] fix(files): keep an explicit target_storage on its old path and refuse litellm_db as a caller choice An explicit target_storage=litellm_db upload was accepted for any model, so an OpenAI model's litellm_db:// id was sent to OpenAI as input_file_id and a model-less upload left a content row nothing can read; it now answers 400 on target_storage. An explicit target_storage skips the files api probe and the purpose and single-target gates, which only decide whether LiteLLM keeps the file itself, so an azure_storage user_data upload for a vLLM model reaches the storage path again as it did before this branch. cancel_batch authorizes the model of every LiteLLM-managed batch id before it branches, the way retrieve_batch already does, so the LiteLLM-executed branch gets the check its provider sibling had. Restores the test_afile_delete_passes_trusted_model_credentials_to_router definition line an earlier commit dropped --- litellm/proxy/batches_endpoints/endpoints.py | 16 ++++-- .../openai_files_endpoints/files_endpoints.py | 20 ++++++- .../proxy/test_managed_files_hook.py | 4 ++ .../proxy/batches_endpoints/test_endpoints.py | 13 +++++ .../test_files_endpoint.py | 57 +++++++++++++++++++ 5 files changed, 102 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 3f6c9f4d6ed..6d5f7a65855 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -1093,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( @@ -1143,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) diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index a5921cc6380..9f12b6faa61 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -117,6 +117,7 @@ async def _litellm_executed_batch_input_model( 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 @@ -129,6 +130,8 @@ async def _litellm_executed_batch_input_model( 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) ) @@ -305,10 +308,21 @@ async def route_create_file( 5. Else -> use custom_llm_provider with files_settings """ - executed_model: Final = await _litellm_executed_batch_input_model( - llm_router, purpose, model, target_model_names_list, user_api_key_dict - ) 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 ( diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index 0ae4e3a5fd8..419a460d098 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -1764,6 +1764,10 @@ async def test_post_call_hook_leaves_litellm_executed_batches_untouched(llm_batc 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(): """ afile_delete must hand the deployment's credential snapshot to the router call, since Bedrock validates the s3:// file id against the bucket in it. diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index c2101abd350..8571ff20e57 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -3220,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() diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 6787aaa3525..48699b47e7f 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -781,6 +781,63 @@ def test_upload_for_a_litellm_executed_model_goes_to_the_provider_unless_the_ser 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 From a61bceb0cf00dd05be46387e8d56a3dfb2daf013 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 14:00:51 -0700 Subject: [PATCH 12/12] fix(files): read storage-backed managed files from their storage backend The managed files hook's content read looped the file's model mappings and asked each deployment for the file. A file LiteLLM stored itself maps every model to its storage url, so the read sent that internal id to the upstream server, failed, and the batch rate limiter failed open: a key's TPM limit did not apply to a LiteLLM-executed batch. The hook now returns the stored bytes from the file's storage backend before it consults any deployment --- .../proxy/hooks/managed_files.py | 17 +++++-- .../proxy/test_managed_files_hook.py | 46 +++++++++++++++++++ 2 files changed, 58 insertions(+), 5 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 8eef8a5f1ce..09cd0ed192f 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -20,6 +20,7 @@ from typing import ( ) from uuid import NAMESPACE_URL, uuid5 +import httpx from fastapi import HTTPException from pydantic import ValidationError @@ -77,6 +78,7 @@ from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccess CreateFileRequest, FileListPage, FileObject, + HttpxBinaryResponseContent, OpenAIFileObject, ResponsesAPIResponse, ) @@ -88,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 ( @@ -1867,10 +1865,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): 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 @@ -1900,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], diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index 419a460d098..74bd67efaf2 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -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}} ) @@ -1908,6 +1911,49 @@ async def test_afile_delete_storage_backed_row_deletes_stored_content_not_provid 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()