mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
chore: merge latest main before JWT regression validation
This commit is contained in:
commit
dbb1f1285a
26 changed files with 3411 additions and 137 deletions
|
|
@ -20,6 +20,7 @@ from typing import (
|
|||
)
|
||||
from uuid import NAMESPACE_URL, uuid5
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
|
@ -34,6 +35,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
|
||||
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
|
||||
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
|
||||
from litellm.llms.base_llm.managed_resources.isolation import (
|
||||
build_list_page,
|
||||
|
|
@ -59,6 +61,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
get_content_type_from_file_object,
|
||||
get_model_id_from_unified_batch_id,
|
||||
get_original_file_id,
|
||||
is_litellm_executed_batch,
|
||||
map_raw_file_ids_to_unified,
|
||||
normalize_mime_type_for_provider,
|
||||
resolve_managed_output_file_model_name,
|
||||
|
|
@ -75,6 +78,7 @@ from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccess
|
|||
CreateFileRequest,
|
||||
FileListPage,
|
||||
FileObject,
|
||||
HttpxBinaryResponseContent,
|
||||
OpenAIFileObject,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
|
|
@ -86,10 +90,6 @@ from litellm.types.utils import (
|
|||
SpecialEnums,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
from prisma.models import (
|
||||
|
|
@ -204,6 +204,19 @@ def _managed_object_table(prisma_client: PrismaClient) -> _ManagedObjectTableAct
|
|||
return prisma_client.db.litellm_managedobjecttable
|
||||
|
||||
|
||||
def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, str]:
|
||||
hidden_params: Final = cast( # cast-ok: _hidden_params is an untyped attribute the upload path sets
|
||||
"Mapping[str, object]", getattr(file_object, "_hidden_params", None) or {}
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key in ("storage_backend", "storage_url")
|
||||
if isinstance(value := hidden_params.get(key), str)
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
||||
# Class variables or attributes
|
||||
def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient):
|
||||
|
|
@ -226,6 +239,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
verbose_logger.info(f"Storing LiteLLM Managed File object with id={file_id} in cache")
|
||||
storage_metadata: Final = _storage_metadata_of(file_object)
|
||||
if file_object is not None:
|
||||
litellm_managed_file_object = LiteLLM_ManagedFileTable(
|
||||
unified_file_id=file_id,
|
||||
|
|
@ -235,6 +249,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
created_by=resolve_resource_owner_id(user_api_key_dict),
|
||||
team_id=user_api_key_dict.team_id,
|
||||
updated_by=user_api_key_dict.user_id,
|
||||
storage_backend=storage_metadata.get("storage_backend"),
|
||||
storage_url=storage_metadata.get("storage_url"),
|
||||
)
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=file_id,
|
||||
|
|
@ -262,14 +278,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
file_object_json = file_object.model_dump_json()
|
||||
db_data["file_object"] = file_object_json
|
||||
update_data["file_object"] = file_object_json
|
||||
# Extract storage metadata from hidden params if present
|
||||
hidden_params = getattr(file_object, "_hidden_params", {}) or {}
|
||||
if "storage_backend" in hidden_params:
|
||||
db_data["storage_backend"] = hidden_params["storage_backend"]
|
||||
update_data["storage_backend"] = hidden_params["storage_backend"]
|
||||
if "storage_url" in hidden_params:
|
||||
db_data["storage_url"] = hidden_params["storage_url"]
|
||||
update_data["storage_url"] = hidden_params["storage_url"]
|
||||
db_data.update(storage_metadata)
|
||||
update_data.update(storage_metadata)
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Storage metadata: storage_backend={db_data.get('storage_backend')}, "
|
||||
|
|
@ -314,6 +324,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
request_tags: Sequence[str] | None = None,
|
||||
persist_attribution: bool = False,
|
||||
create_if_missing: bool = True,
|
||||
batch_processed: bool = False,
|
||||
) -> None:
|
||||
"""Persist a managed object row, caching it and upserting it in the DB.
|
||||
|
||||
|
|
@ -328,6 +339,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
row absent from the table is left absent rather than created with the
|
||||
observer as its creator, because created_by and team_id are written from
|
||||
whoever calls the create branch.
|
||||
|
||||
batch_processed is set by callers that have already billed the batch
|
||||
themselves, so CheckBatchCost skips the row instead of billing it twice.
|
||||
It is written only in the upsert create branch.
|
||||
"""
|
||||
verbose_logger.info(f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache")
|
||||
litellm_managed_object = LiteLLM_ManagedObjectTable(
|
||||
|
|
@ -379,6 +394,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
"updated_by": user_api_key_dict.user_id,
|
||||
"status": file_object.status,
|
||||
**attribution_columns,
|
||||
"batch_processed": batch_processed,
|
||||
},
|
||||
"update": update_columns,
|
||||
},
|
||||
|
|
@ -1343,6 +1359,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
self, data: Dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
|
||||
) -> LLMResponseTypes:
|
||||
if isinstance(response, LiteLLMBatch):
|
||||
decoded_batch_id: Final = _is_base64_encoded_unified_file_id(response.id)
|
||||
if decoded_batch_id and is_litellm_executed_batch(decoded_batch_id):
|
||||
return response
|
||||
## Check if unified_file_id is in the response
|
||||
unified_file_id = response._hidden_params.get("unified_file_id") # managed file id
|
||||
unified_batch_id = response._hidden_params.get("unified_batch_id") # managed batch id
|
||||
|
|
@ -1794,24 +1813,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
# Check if file deletion should be blocked due to batch references
|
||||
await self._check_file_deletion_allowed(file_id)
|
||||
|
||||
# file_id = convert_b64_uid_to_unified_uid(file_id)
|
||||
model_file_id_mapping = await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
|
||||
|
||||
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
|
||||
if specific_model_file_id_mapping:
|
||||
# Remove conflicting keys from data to avoid duplicate keyword arguments
|
||||
filtered_data = {k: v for k, v in data.items() if k not in ("model", "file_id")}
|
||||
for model_id, model_file_id in specific_model_file_id_mapping.items():
|
||||
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
|
||||
delete_data = {
|
||||
**{k: v for k, v in filtered_data.items() if k != "_litellm_internal_model_credentials"},
|
||||
**(
|
||||
{"_litellm_internal_model_credentials": MappingProxyType(dict(credentials))}
|
||||
if credentials is not None
|
||||
else {}
|
||||
),
|
||||
}
|
||||
await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
|
||||
managed_file: Final = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
|
||||
if managed_file is not None and managed_file.storage_backend and managed_file.storage_url:
|
||||
await self._delete_storage_backend_content(managed_file.storage_backend, managed_file.storage_url)
|
||||
else:
|
||||
await self._delete_provider_files(file_id, litellm_parent_otel_span, llm_router, data)
|
||||
|
||||
await self.delete_unified_file_id(file_id, litellm_parent_otel_span)
|
||||
|
||||
|
|
@ -1820,16 +1826,53 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
prom_logger.record_managed_file_deleted(result="success")
|
||||
return FileDeleted(id=file_id, object="file", deleted=True)
|
||||
|
||||
async def _delete_storage_backend_content(self, storage_backend_name: str, storage_url: str) -> None:
|
||||
try:
|
||||
storage_backend: Final = get_storage_backend(storage_backend_name, prisma_client=self.prisma_client)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=f"Cannot delete the stored file content: {e}") from e
|
||||
await storage_backend.delete_file(storage_url)
|
||||
|
||||
async def _delete_provider_files(
|
||||
self,
|
||||
file_id: str,
|
||||
litellm_parent_otel_span: Span | None,
|
||||
llm_router: Router,
|
||||
data: Mapping[str, object],
|
||||
) -> None:
|
||||
model_file_id_mapping: Final = await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
|
||||
specific_model_file_id_mapping: Final = model_file_id_mapping.get(file_id)
|
||||
if not specific_model_file_id_mapping:
|
||||
return
|
||||
filtered_data: Final = {
|
||||
k: v for k, v in data.items() if k not in ("model", "file_id", "_litellm_internal_model_credentials")
|
||||
}
|
||||
for model_id, model_file_id in specific_model_file_id_mapping.items():
|
||||
credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id)
|
||||
delete_data = {
|
||||
**filtered_data,
|
||||
**(
|
||||
{"_litellm_internal_model_credentials": MappingProxyType(dict(credentials))}
|
||||
if credentials is not None
|
||||
else {}
|
||||
),
|
||||
}
|
||||
await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data)
|
||||
|
||||
async def afile_content(
|
||||
self,
|
||||
file_id: str,
|
||||
litellm_parent_otel_span: Optional[Span],
|
||||
llm_router: Router,
|
||||
**data: Dict,
|
||||
) -> "HttpxBinaryResponseContent":
|
||||
) -> HttpxBinaryResponseContent:
|
||||
"""
|
||||
Get the content of a file from first model that has it
|
||||
"""
|
||||
managed_file: Final = await self.get_unified_file_id(file_id, litellm_parent_otel_span)
|
||||
if managed_file is not None and managed_file.storage_backend and managed_file.storage_url:
|
||||
return await self._storage_backend_content(managed_file.storage_backend, managed_file.storage_url)
|
||||
|
||||
model_file_id_mapping = data.pop("model_file_id_mapping", None)
|
||||
model_file_id_mapping = model_file_id_mapping or await self.get_model_file_id_mapping(
|
||||
[file_id], litellm_parent_otel_span
|
||||
|
|
@ -1859,6 +1902,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
else:
|
||||
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
||||
|
||||
async def _storage_backend_content(self, storage_backend_name: str, storage_url: str) -> HttpxBinaryResponseContent:
|
||||
storage_backend: Final = get_storage_backend(storage_backend_name, prisma_client=self.prisma_client)
|
||||
content: Final = await storage_backend.download_file(storage_url)
|
||||
return HttpxBinaryResponseContent(response=httpx.Response(status_code=httpx.codes.OK, content=content))
|
||||
|
||||
async def _convert_storage_files_to_base64(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
|
|
@ -1889,16 +1937,12 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
|
||||
# File is stored in a storage backend, download and convert to base64
|
||||
try:
|
||||
from litellm.llms.base_llm.files.storage_backend_factory import (
|
||||
get_storage_backend,
|
||||
)
|
||||
|
||||
storage_backend_name = db_file.storage_backend
|
||||
storage_url = db_file.storage_url
|
||||
|
||||
# Get storage backend (uses same env vars as callback)
|
||||
try:
|
||||
storage_backend = get_storage_backend(storage_backend_name)
|
||||
storage_backend = get_storage_backend(storage_backend_name, prisma_client=self.prisma_client)
|
||||
except ValueError as e:
|
||||
verbose_logger.warning(
|
||||
f"Storage backend '{storage_backend_name}' error for file {file_id}: {str(e)}"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
);
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
40
litellm/llms/base_llm/files/litellm_db_storage_backend.py
Normal file
40
litellm/llms/base_llm/files/litellm_db_storage_backend.py
Normal file
|
|
@ -0,0 +1,40 @@
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend
|
||||
from litellm.repositories.managed_file_content_repository import ManagedFileContentRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
LITELLM_DB_STORAGE_BACKEND_NAME: Final = "litellm_db"
|
||||
LITELLM_DB_STORAGE_URL_PREFIX: Final = f"{LITELLM_DB_STORAGE_BACKEND_NAME}://"
|
||||
|
||||
|
||||
def storage_url_to_row_id(storage_url: str) -> str:
|
||||
if not storage_url.startswith(LITELLM_DB_STORAGE_URL_PREFIX):
|
||||
raise ValueError(f"Not a {LITELLM_DB_STORAGE_BACKEND_NAME} storage url: {storage_url}")
|
||||
return storage_url.removeprefix(LITELLM_DB_STORAGE_URL_PREFIX)
|
||||
|
||||
|
||||
class LiteLLMDbStorageBackend(BaseFileStorageBackend):
|
||||
def __init__(self, prisma_client: "PrismaClient") -> None:
|
||||
self._contents = ManagedFileContentRepository(prisma_client)
|
||||
|
||||
async def upload_file(
|
||||
self,
|
||||
file_content: bytes,
|
||||
filename: str,
|
||||
content_type: str,
|
||||
path_prefix: str | None = None,
|
||||
file_naming_strategy: str = "uuid",
|
||||
) -> str:
|
||||
return f"{LITELLM_DB_STORAGE_URL_PREFIX}{await self._contents.store(file_content)}"
|
||||
|
||||
async def download_file(self, storage_url: str) -> bytes:
|
||||
content: Final = await self._contents.load(storage_url_to_row_id(storage_url))
|
||||
if content is None:
|
||||
raise ValueError(f"No stored file content for {storage_url}")
|
||||
return content
|
||||
|
||||
async def delete_file(self, storage_url: str) -> None:
|
||||
await self._contents.delete(storage_url_to_row_id(storage_url))
|
||||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -7,10 +7,12 @@
|
|||
import asyncio
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -18,6 +20,15 @@ from litellm.batches.main import CancelBatchRequest, RetrieveBatchRequest
|
|||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.batches_endpoints.common_utils import validate_batch_list_limit
|
||||
from litellm.proxy.batches_endpoints.litellm_executed_batches import (
|
||||
LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE,
|
||||
LiteLLMExecutedBatchRunner,
|
||||
ManagedBatchStore,
|
||||
batch_error,
|
||||
executed_batch_runner_lost,
|
||||
litellm_executed_provider_for,
|
||||
resolve_litellm_executed_provider,
|
||||
)
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
log_llm_api_exception,
|
||||
|
|
@ -47,16 +58,87 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
get_model_id_from_unified_batch_id,
|
||||
get_models_from_unified_file_id,
|
||||
get_original_file_id,
|
||||
is_litellm_executed_batch,
|
||||
prepare_data_with_credentials,
|
||||
update_batch_in_database,
|
||||
validate_managed_id_requirement,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import request_tags_from_metadata
|
||||
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
|
||||
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy, is_known_model
|
||||
from litellm.repositories.managed_batch_repository import ManagedBatchRepository
|
||||
from litellm.repositories.table_repositories import ManagedFileRepository
|
||||
from litellm.router import Router
|
||||
from litellm.types.llms.openai import LiteLLMBatchCreateRequest
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.models import LiteLLM_ManagedObjectTable
|
||||
|
||||
router: Final = APIRouter()
|
||||
_METADATA_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _request_tags(data: Mapping[str, object]) -> tuple[str, ...] | None:
|
||||
metadata: Final = data.get("litellm_metadata")
|
||||
if metadata is None:
|
||||
return None
|
||||
return request_tags_from_metadata(_METADATA_ADAPTER.validate_python(metadata))
|
||||
|
||||
|
||||
def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyLogging) -> LiteLLMExecutedBatchRunner:
|
||||
from litellm.proxy.proxy_server import general_settings, prisma_client
|
||||
|
||||
managed_files: Final = proxy_logging_obj.get_proxy_hook("managed_files")
|
||||
if prisma_client is None or not isinstance(managed_files, ManagedBatchStore):
|
||||
raise batch_error(
|
||||
400,
|
||||
"LiteLLM-executed batches need a database: set DATABASE_URL so LiteLLM can keep the batch and its files",
|
||||
)
|
||||
return LiteLLMExecutedBatchRunner(
|
||||
llm_router=llm_router,
|
||||
prisma_client=prisma_client,
|
||||
managed_files=managed_files,
|
||||
batches=ManagedBatchRepository(prisma_client),
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
general_settings=general_settings,
|
||||
)
|
||||
|
||||
|
||||
async def _batch_from_database(
|
||||
batch_id: str,
|
||||
unified_batch_id: str | Literal[False],
|
||||
executed_batch: bool,
|
||||
managed_files_obj: object,
|
||||
prisma_client: PrismaClient | None,
|
||||
llm_router: Router | None,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple["LiteLLM_ManagedObjectTable | None", LiteLLMBatch | None]:
|
||||
row, batch = await get_batch_from_database(
|
||||
batch_id=batch_id,
|
||||
unified_batch_id=unified_batch_id,
|
||||
managed_files_obj=managed_files_obj,
|
||||
prisma_client=prisma_client,
|
||||
verbose_proxy_logger=verbose_proxy_logger,
|
||||
)
|
||||
updated_at: Final[object] = getattr(row, "updated_at", None)
|
||||
if not executed_batch or batch is None or llm_router is None or not isinstance(updated_at, datetime):
|
||||
return row, batch
|
||||
if not executed_batch_runner_lost(batch.status, updated_at):
|
||||
return row, batch
|
||||
runner: Final = _litellm_executed_batch_runner(llm_router, proxy_logging_obj)
|
||||
return row, await runner.fail_abandoned(batch, user_api_key_dict)
|
||||
|
||||
|
||||
async def _raise_when_input_file_must_be_managed(model: str, credentials: Mapping[str, object]) -> None:
|
||||
if await litellm_executed_provider_for(credentials) is None:
|
||||
return
|
||||
raise batch_error(
|
||||
400,
|
||||
f"Batches for {model} run inside LiteLLM, so the input file must be a LiteLLM managed file: "
|
||||
f"{LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE}",
|
||||
)
|
||||
|
||||
|
||||
def _raise_not_found_when_openai_fallback_unservable(
|
||||
|
|
@ -101,6 +183,24 @@ async def _resolve_managed_input_file_storage_url(input_file_id: str) -> "str |
|
|||
return db_file.storage_url or None
|
||||
|
||||
|
||||
async def _create_provider_batch_for_managed_file(
|
||||
llm_router: Router,
|
||||
create_batch_data: LiteLLMBatchCreateRequest,
|
||||
input_file_id: str,
|
||||
unified_file_id: str,
|
||||
) -> LiteLLMBatch:
|
||||
resolved_storage_url: Final = await _resolve_managed_input_file_storage_url(input_file_id)
|
||||
request: Final[LiteLLMBatchCreateRequest] = {
|
||||
**create_batch_data,
|
||||
"input_file_id": resolved_storage_url or input_file_id,
|
||||
"disable_fallbacks": True,
|
||||
}
|
||||
response: Final = await llm_router.acreate_batch(**request)
|
||||
response.input_file_id = input_file_id
|
||||
response._hidden_params["unified_file_id"] = unified_file_id
|
||||
return response
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{provider}/v1/batches",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
@ -296,24 +396,35 @@ async def create_batch(
|
|||
await authorize_model_for_key(model_id=model, llm_router=llm_router, user_api_key_dict=user_api_key_dict)
|
||||
_create_batch_data["model"] = model
|
||||
|
||||
resolved_storage_url: Final = await _resolve_managed_input_file_storage_url(input_file_id)
|
||||
if resolved_storage_url is not None:
|
||||
_create_batch_data["input_file_id"] = resolved_storage_url
|
||||
|
||||
if llm_router is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": "LLM Router not initialized. Ensure models added to proxy."},
|
||||
)
|
||||
|
||||
_create_batch_data.update(disable_fallbacks=True) # pyright: ignore[reportCallIssue] # router flag
|
||||
response = await llm_router.acreate_batch(**_create_batch_data)
|
||||
response.input_file_id = input_file_id
|
||||
response._hidden_params["unified_file_id"] = unified_file_id
|
||||
executed_provider: Final = await resolve_litellm_executed_provider(
|
||||
llm_router, model, user_api_key_dict.team_id
|
||||
)
|
||||
response = (
|
||||
await _litellm_executed_batch_runner(llm_router, proxy_logging_obj).create(
|
||||
create_request=_create_batch_data,
|
||||
unified_input_file_id=input_file_id,
|
||||
model=model,
|
||||
provider=executed_provider,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_tags=_request_tags(_create_batch_data),
|
||||
)
|
||||
if executed_provider is not None
|
||||
else await _create_provider_batch_for_managed_file(
|
||||
llm_router, _create_batch_data, input_file_id, unified_file_id
|
||||
)
|
||||
)
|
||||
else:
|
||||
# Check if model specified via header/query/body param
|
||||
model_param: Final = (
|
||||
data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model")
|
||||
_create_batch_data.get("model")
|
||||
or request.query_params.get("model")
|
||||
or request.headers.get("x-litellm-model")
|
||||
)
|
||||
|
||||
# SCENARIO 2 & 3: Model from header/query OR custom_llm_provider fallback
|
||||
|
|
@ -325,6 +436,7 @@ async def create_batch(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
operation_context="batch creation",
|
||||
)
|
||||
await _raise_when_input_file_must_be_managed(model_param, credentials)
|
||||
|
||||
prepare_data_with_credentials(
|
||||
data=_create_batch_data,
|
||||
|
|
@ -486,23 +598,26 @@ async def retrieve_batch(
|
|||
managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files")
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
db_batch_object, response = await get_batch_from_database(
|
||||
executed_batch: Final = isinstance(unified_batch_id, str) and is_litellm_executed_batch(unified_batch_id)
|
||||
db_batch_object, response = await _batch_from_database(
|
||||
batch_id=batch_id,
|
||||
unified_batch_id=unified_batch_id,
|
||||
executed_batch=executed_batch,
|
||||
managed_files_obj=managed_files_obj,
|
||||
prisma_client=prisma_client,
|
||||
verbose_proxy_logger=verbose_proxy_logger,
|
||||
llm_router=llm_router,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
if executed_batch and response is None:
|
||||
raise batch_error(404, f"No batch found with id '{batch_id}'.")
|
||||
|
||||
# If batch is in a terminal state, return immediately.
|
||||
# Include "complete" (DB-normalized form of "completed").
|
||||
if response is not None and response.status in [
|
||||
"completed",
|
||||
"complete",
|
||||
"failed",
|
||||
"cancelled",
|
||||
"expired",
|
||||
]:
|
||||
if response is not None and (
|
||||
response.status in ("completed", "complete", "failed", "cancelled", "expired") or executed_batch
|
||||
):
|
||||
# Call hooks and return
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
|
|
@ -978,6 +1093,17 @@ async def cancel_batch(
|
|||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
unified_model_id: Final = get_model_id_from_unified_batch_id(unified_batch_id) if unified_batch_id else None
|
||||
if unified_model_id is not None:
|
||||
resolved_unified_model: Final = (
|
||||
llm_router.resolve_model_name_from_model_id(unified_model_id) if llm_router is not None else None
|
||||
)
|
||||
await authorize_model_for_key(
|
||||
model_id=resolved_unified_model or unified_model_id,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# SCENARIO 1: Batch ID is encoded with model info
|
||||
if model_from_id is not None:
|
||||
credentials: Final = await get_authorized_credentials_for_model(
|
||||
|
|
@ -1009,6 +1135,12 @@ async def cancel_batch(
|
|||
)
|
||||
|
||||
# SCENARIO 2: target_model_names based routing
|
||||
elif unified_batch_id and is_litellm_executed_batch(unified_batch_id):
|
||||
if llm_router is None:
|
||||
raise batch_error(500, "LLM Router not initialized. Ensure models added to proxy.")
|
||||
response = await _litellm_executed_batch_runner( # rebind-ok: each cancel path sets the route's response
|
||||
llm_router, proxy_logging_obj
|
||||
).cancel(batch_id, user_api_key_dict)
|
||||
elif unified_batch_id:
|
||||
if llm_router is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -1022,11 +1154,6 @@ async def cancel_batch(
|
|||
status_code=400,
|
||||
detail={"error": "Invalid LiteLLM managed batch ID. Missing model_id."},
|
||||
)
|
||||
await authorize_model_for_key(
|
||||
model_id=llm_router.resolve_model_name_from_model_id(model_id_from_batch) or model_id_from_batch,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
data["model"] = model_id_from_batch
|
||||
data["batch_id"] = get_batch_id_from_unified_batch_id(unified_batch_id)
|
||||
response = await llm_router.acancel_batch(**data)
|
||||
|
|
|
|||
716
litellm/proxy/batches_endpoints/litellm_executed_batches.py
Normal file
716
litellm/proxy/batches_endpoints/litellm_executed_batches.py
Normal file
|
|
@ -0,0 +1,716 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from itertools import pairwise
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from openai.types.batch import Errors
|
||||
from openai.types.batch_error import BatchError
|
||||
from openai.types.batch_request_counts import BatchRequestCounts
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid as uuid_module
|
||||
from litellm.constants import LITELLM_EXECUTED_BATCH_CONCURRENCY
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.llms.base_llm.files.litellm_db_storage_backend import LITELLM_DB_STORAGE_BACKEND_NAME
|
||||
from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend
|
||||
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.models.managed_files import LiteLLM_ManagedFileTable
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_utils import is_request_body_safe
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import LITELLM_EXECUTED_BATCH_ID_PREFIX
|
||||
from litellm.proxy.openai_files_endpoints.storage_backend_service import StorageBackendFileService
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.repositories.managed_batch_repository import ManagedBatchRepository
|
||||
from litellm.types.llms.openai import LiteLLMBatchCreateRequest, OpenAIFileObject, OpenAIFilesPurpose
|
||||
from litellm.types.utils import LITELLM_EXECUTED_BATCH_PROVIDERS, ExtractedFileData, LiteLLMBatch, LlmProviders
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import types as prisma_types
|
||||
|
||||
from litellm.router import Router
|
||||
|
||||
BatchEndpoint: TypeAlias = Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"]
|
||||
BatchStatus: TypeAlias = Literal[
|
||||
"in_progress", "finalizing", "completed", "failed", "cancelling", "cancelled", "expired"
|
||||
]
|
||||
TERMINAL_BATCH_STATUSES: Final[frozenset[str]] = frozenset({"completed", "failed", "cancelled", "expired"})
|
||||
_STOP_STATUSES: Final[frozenset[str]] = TERMINAL_BATCH_STATUSES | frozenset({"cancelling"})
|
||||
_BATCH_ENDPOINT_ADAPTER: Final[TypeAdapter[BatchEndpoint]] = TypeAdapter(BatchEndpoint)
|
||||
_CANCEL_POLL_SECONDS: Final = 1.0
|
||||
_HEARTBEAT_SECONDS: Final = 30.0
|
||||
_STALE_AFTER_SECONDS: Final = 180.0
|
||||
_FILES_API_PROBE_TIMEOUT_SECONDS: Final = 5.0
|
||||
_COMPLETION_WINDOW_SECONDS: Final = 24 * 60 * 60
|
||||
_RUNNER_LOST_MESSAGE: Final = "the proxy replica running this batch stopped before it finished; resubmit the batch"
|
||||
_EXPIRED_MESSAGE: Final = "This request could not be executed before the completion window expired."
|
||||
_ROUTER_METHODS: Final[Mapping[BatchEndpoint, str]] = MappingProxyType(
|
||||
{
|
||||
"/v1/chat/completions": "acompletion",
|
||||
"/v1/completions": "atext_completion",
|
||||
"/v1/embeddings": "aembedding",
|
||||
"/v1/responses": "aresponses",
|
||||
}
|
||||
)
|
||||
_CANCELLING_TRANSITIONS: Final[Mapping[BatchStatus, BatchStatus]] = MappingProxyType(
|
||||
{"completed": "cancelled", "expired": "cancelled", "in_progress": "cancelling", "finalizing": "cancelling"}
|
||||
)
|
||||
LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE: Final = (
|
||||
"upload it through POST /v1/files with purpose=batch and either the x-litellm-model header or the "
|
||||
"target_model_names form field naming the model, so LiteLLM keeps the file and runs the batch itself"
|
||||
)
|
||||
_RUNNING_BATCHES: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong references keep running batch tasks alive
|
||||
_NO_FIELDS: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
_NO_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
class _ErrorDetail(TypedDict):
|
||||
message: ReadOnly[str]
|
||||
type: ReadOnly[str]
|
||||
param: ReadOnly[None]
|
||||
code: ReadOnly[None]
|
||||
|
||||
|
||||
class _ErrorBody(TypedDict):
|
||||
error: ReadOnly[_ErrorDetail]
|
||||
|
||||
|
||||
class _ResultResponse(TypedDict):
|
||||
status_code: ReadOnly[int]
|
||||
request_id: ReadOnly[str]
|
||||
body: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class _LineError(TypedDict):
|
||||
code: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
|
||||
|
||||
class _ResultLine(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
custom_id: ReadOnly[str]
|
||||
response: ReadOnly[_ResultResponse | None]
|
||||
error: ReadOnly[_LineError | None]
|
||||
|
||||
|
||||
class BatchInputLine(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
custom_id: str
|
||||
method: Literal["POST"]
|
||||
url: str
|
||||
body: Mapping[str, object]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InvalidBatchInput:
|
||||
line_number: int | None
|
||||
reason: str
|
||||
|
||||
def describe(self) -> str:
|
||||
return f"line {self.line_number}: {self.reason}" if self.line_number is not None else self.reason
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RowOutcome:
|
||||
custom_id: str
|
||||
status_code: int
|
||||
body: Mapping[str, object]
|
||||
succeeded: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ExpiredRow:
|
||||
custom_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BatchRun:
|
||||
unified_batch_id: str
|
||||
llm_batch_id: str
|
||||
model: str
|
||||
endpoint: BatchEndpoint
|
||||
lines: tuple[BatchInputLine, ...]
|
||||
user_api_key_dict: UserAPIKeyAuth
|
||||
request_tags: tuple[str, ...]
|
||||
deadline: float
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ManagedBatchStore(Protocol):
|
||||
def get_unified_batch_id(self, batch_id: str, model_id: str) -> str: ...
|
||||
|
||||
async def get_unified_file_id(
|
||||
self, file_id: str, litellm_parent_otel_span: object | None = None
|
||||
) -> LiteLLM_ManagedFileTable | None: ...
|
||||
|
||||
async def store_unified_object_id(
|
||||
self,
|
||||
unified_object_id: str,
|
||||
file_object: LiteLLMBatch,
|
||||
litellm_parent_otel_span: object | None,
|
||||
model_object_id: str,
|
||||
file_purpose: Literal["batch", "fine-tune", "response"],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_tags: Sequence[str] | None = None,
|
||||
persist_attribution: bool = False,
|
||||
batch_processed: bool = False,
|
||||
) -> None: ...
|
||||
|
||||
|
||||
class _StorageBackendFactory(Protocol):
|
||||
def __call__(self, backend_type: str, prisma_client: PrismaClient | None = None) -> BaseFileStorageBackend: ...
|
||||
|
||||
|
||||
class _ResultFileUploader(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
file_data: Mapping[str, object],
|
||||
target_storage: str,
|
||||
target_model_names: Sequence[str],
|
||||
purpose: OpenAIFilesPurpose,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient | None = None,
|
||||
) -> Awaitable[OpenAIFileObject]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _RouterCall(Protocol):
|
||||
def __call__(self, **params: object) -> Awaitable[object]: ... # kwargs-ok: the request body is passed as keywords
|
||||
|
||||
|
||||
def litellm_executed_provider_of(credentials: Mapping[str, object]) -> str | None:
|
||||
explicit_provider: Final = credentials.get("custom_llm_provider")
|
||||
provider: Final = (
|
||||
explicit_provider if isinstance(explicit_provider, str) else _provider_of(credentials.get("model"))
|
||||
)
|
||||
return provider if provider in LITELLM_EXECUTED_BATCH_PROVIDERS else None
|
||||
|
||||
|
||||
class _HttpGetter(Protocol):
|
||||
async def get(
|
||||
self, url: str, *, headers: dict[str, str] | None = None, timeout: float | httpx.Timeout | None = None
|
||||
) -> httpx.Response: ...
|
||||
|
||||
|
||||
class FilesApiProbe(Protocol):
|
||||
async def __call__(self, api_base: str, api_key: str | None) -> bool: ...
|
||||
|
||||
|
||||
class BodyRejection(Protocol):
|
||||
def __call__(self, body: Mapping[str, object], /) -> str | None: ...
|
||||
|
||||
|
||||
async def upstream_lacks_files_api(api_base: str, api_key: str | None, http_client: _HttpGetter | None = None) -> bool:
|
||||
client: Final = http_client or get_async_httpx_client(llm_provider=LlmProviders.HOSTED_VLLM)
|
||||
try:
|
||||
response: Final = await client.get(
|
||||
f"{api_base.rstrip('/')}/files",
|
||||
headers=(
|
||||
{"Authorization": f"Bearer {api_key}"} # mutable-ok: AsyncHTTPHandler.get wants a plain dict
|
||||
if api_key
|
||||
else None
|
||||
),
|
||||
timeout=_FILES_API_PROBE_TIMEOUT_SECONDS,
|
||||
)
|
||||
except httpx.HTTPError:
|
||||
return False
|
||||
return response.status_code == httpx.codes.NOT_FOUND
|
||||
|
||||
|
||||
def _upstream_of(credentials: Mapping[str, object], provider: str) -> tuple[str, str | None] | None:
|
||||
model: Final = credentials.get("model")
|
||||
api_base: Final = credentials.get("api_base")
|
||||
api_key: Final = credentials.get("api_key")
|
||||
if not isinstance(model, str):
|
||||
return None
|
||||
try:
|
||||
_, _, resolved_api_key, resolved_api_base = litellm.get_llm_provider(
|
||||
model=model,
|
||||
custom_llm_provider=provider,
|
||||
api_base=api_base if isinstance(api_base, str) else None,
|
||||
api_key=api_key if isinstance(api_key, str) else None,
|
||||
)
|
||||
except Exception: # noqa: BLE001 # get_llm_provider raises on a model it cannot map, which means nothing to probe
|
||||
return None
|
||||
return None if resolved_api_base is None else (resolved_api_base, resolved_api_key)
|
||||
|
||||
|
||||
async def litellm_executed_provider_for(
|
||||
credentials: Mapping[str, object], lacks_files_api: FilesApiProbe = upstream_lacks_files_api
|
||||
) -> str | None:
|
||||
provider: Final = litellm_executed_provider_of(credentials)
|
||||
if provider is None:
|
||||
return None
|
||||
upstream: Final = _upstream_of(credentials, provider)
|
||||
if upstream is None:
|
||||
return None
|
||||
return provider if await lacks_files_api(*upstream) else None
|
||||
|
||||
|
||||
async def resolve_litellm_executed_provider(
|
||||
llm_router: "Router",
|
||||
model: str,
|
||||
team_id: str | None,
|
||||
lacks_files_api: FilesApiProbe = upstream_lacks_files_api,
|
||||
) -> str | None:
|
||||
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model, team_id=team_id)
|
||||
return None if credentials is None else await litellm_executed_provider_for(credentials, lacks_files_api)
|
||||
|
||||
|
||||
def _provider_of(model: object) -> str | None:
|
||||
if not isinstance(model, str):
|
||||
return None
|
||||
try:
|
||||
return litellm.get_llm_provider(model=model)[1]
|
||||
except Exception: # noqa: BLE001 # get_llm_provider raises on an unknown model, which means no provider
|
||||
return None
|
||||
|
||||
|
||||
def _validation_reason(error: ValidationError) -> str:
|
||||
return "; ".join(
|
||||
f"{'.'.join(str(part) for part in item['loc'])}: {item['msg']}" if item["loc"] else item["msg"]
|
||||
for item in error.errors()
|
||||
)
|
||||
|
||||
|
||||
def _accept_every_body(_body: Mapping[str, object]) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _parse_line(
|
||||
line_number: int, raw: bytes, endpoint: BatchEndpoint, reject_body: BodyRejection
|
||||
) -> BatchInputLine | InvalidBatchInput:
|
||||
try:
|
||||
line: Final = BatchInputLine.model_validate_json(raw)
|
||||
except ValidationError as e:
|
||||
return InvalidBatchInput(line_number, _validation_reason(e))
|
||||
if line.url != endpoint:
|
||||
return InvalidBatchInput(line_number, f"url {line.url!r} does not match the batch endpoint {endpoint!r}")
|
||||
if line.body.get("stream"):
|
||||
return InvalidBatchInput(line_number, "streaming requests are not supported in a batch")
|
||||
rejection: Final = reject_body(line.body)
|
||||
if rejection is not None:
|
||||
return InvalidBatchInput(line_number, rejection)
|
||||
return line
|
||||
|
||||
|
||||
def parse_batch_input(
|
||||
content: bytes, endpoint: BatchEndpoint, reject_body: BodyRejection = _accept_every_body
|
||||
) -> tuple[BatchInputLine, ...] | InvalidBatchInput:
|
||||
raw_lines: Final = tuple((number, raw) for number, raw in enumerate(content.splitlines(), start=1) if raw.strip())
|
||||
if not raw_lines:
|
||||
return InvalidBatchInput(None, "the input file has no requests")
|
||||
parsed: Final = tuple(_parse_line(number, raw, endpoint, reject_body) for number, raw in raw_lines)
|
||||
first_invalid: Final = next((item for item in parsed if isinstance(item, InvalidBatchInput)), None)
|
||||
if first_invalid is not None:
|
||||
return first_invalid
|
||||
lines: Final = tuple(item for item in parsed if isinstance(item, BatchInputLine))
|
||||
custom_ids: Final = sorted(line.custom_id for line in lines)
|
||||
duplicate: Final = next((first for first, second in pairwise(custom_ids) if first == second), None)
|
||||
if duplicate is not None:
|
||||
return InvalidBatchInput(None, f"custom_id {duplicate!r} is used more than once")
|
||||
return lines
|
||||
|
||||
|
||||
def batch_error(status_code: int, message: str) -> ProxyException:
|
||||
error_type: Final = "invalid_request_error" if status_code < 500 else ProxyErrorTypes.internal_server_error.value
|
||||
return ProxyException(message=message, type=error_type, param=None, code=status_code)
|
||||
|
||||
|
||||
def _validate_endpoint(endpoint: object) -> BatchEndpoint:
|
||||
try:
|
||||
return _BATCH_ENDPOINT_ADAPTER.validate_python(endpoint)
|
||||
except ValidationError:
|
||||
raise batch_error(400, f"endpoint {endpoint!r} is not supported for a LiteLLM-executed batch")
|
||||
|
||||
|
||||
def _status_code_of(error: Exception) -> int:
|
||||
status_code: Final[object] = getattr(error, "status_code", None)
|
||||
return status_code if isinstance(status_code, int) else 500
|
||||
|
||||
|
||||
def _error_body(error: Exception) -> _ErrorBody:
|
||||
body: Final[_ErrorBody] = {
|
||||
"error": {"message": str(error), "type": type(error).__name__, "param": None, "code": None}
|
||||
}
|
||||
return body
|
||||
|
||||
|
||||
def _line_response(outcome: RowOutcome | ExpiredRow) -> _ResultResponse | None:
|
||||
if isinstance(outcome, ExpiredRow):
|
||||
return None
|
||||
response: Final[_ResultResponse] = {
|
||||
"status_code": outcome.status_code,
|
||||
"request_id": f"req_{uuid_module.uuid4().hex[:24]}",
|
||||
"body": outcome.body,
|
||||
}
|
||||
return response
|
||||
|
||||
|
||||
def _line_error(outcome: RowOutcome | ExpiredRow) -> _LineError | None:
|
||||
if isinstance(outcome, RowOutcome):
|
||||
return None
|
||||
error: Final[_LineError] = {"code": "batch_expired", "message": _EXPIRED_MESSAGE}
|
||||
return error
|
||||
|
||||
|
||||
def _result_line(outcome: RowOutcome | ExpiredRow) -> _ResultLine:
|
||||
line: Final[_ResultLine] = {
|
||||
"id": f"batch_req_{uuid_module.uuid4().hex[:24]}",
|
||||
"custom_id": outcome.custom_id,
|
||||
"response": _line_response(outcome),
|
||||
"error": _line_error(outcome),
|
||||
}
|
||||
return line
|
||||
|
||||
|
||||
def _dump(response: object) -> Mapping[str, object]:
|
||||
if isinstance(response, BaseModel):
|
||||
return response.model_dump(mode="json")
|
||||
raise TypeError(f"Batch rows must return a single response object, got {type(response).__name__}")
|
||||
|
||||
|
||||
def _resolve_transition(current_status: str, requested: BatchStatus) -> BatchStatus:
|
||||
if current_status != "cancelling":
|
||||
return requested
|
||||
return _CANCELLING_TRANSITIONS.get(requested, requested)
|
||||
|
||||
|
||||
def executed_batch_runner_lost(status: str, updated_at: datetime) -> bool:
|
||||
if status in TERMINAL_BATCH_STATUSES:
|
||||
return False
|
||||
return (datetime.now(timezone.utc) - updated_at).total_seconds() > _STALE_AFTER_SECONDS
|
||||
|
||||
|
||||
class _StopWatch:
|
||||
def __init__(self, load_status: Callable[[], Awaitable[str | None]], interval_seconds: float) -> None:
|
||||
self._load_status = load_status
|
||||
self._interval_seconds = interval_seconds
|
||||
self._checked_at = float("-inf")
|
||||
self._stopped = False
|
||||
|
||||
async def stopped(self) -> bool:
|
||||
if self._stopped:
|
||||
return True
|
||||
now: Final = time.monotonic()
|
||||
if now - self._checked_at < self._interval_seconds:
|
||||
return False
|
||||
self._checked_at = now
|
||||
self._stopped = await self._load_status() in _STOP_STATUSES
|
||||
return self._stopped
|
||||
|
||||
|
||||
class LiteLLMExecutedBatchRunner:
|
||||
def __init__(
|
||||
self,
|
||||
llm_router: "Router",
|
||||
prisma_client: PrismaClient,
|
||||
managed_files: ManagedBatchStore,
|
||||
batches: ManagedBatchRepository,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
general_settings: Mapping[str, object],
|
||||
concurrency: int = LITELLM_EXECUTED_BATCH_CONCURRENCY,
|
||||
heartbeat_seconds: float = _HEARTBEAT_SECONDS,
|
||||
completion_window_seconds: float = _COMPLETION_WINDOW_SECONDS,
|
||||
storage_backend_factory: _StorageBackendFactory = get_storage_backend,
|
||||
upload_result_file: _ResultFileUploader = StorageBackendFileService.upload_file_to_storage_backend,
|
||||
) -> None:
|
||||
self.llm_router = llm_router
|
||||
self.prisma_client = prisma_client
|
||||
self.managed_files = managed_files
|
||||
self.batches = batches
|
||||
self.proxy_logging_obj = proxy_logging_obj
|
||||
self.general_settings = general_settings
|
||||
self.concurrency = concurrency
|
||||
self.heartbeat_seconds = heartbeat_seconds
|
||||
self.completion_window_seconds = completion_window_seconds
|
||||
self.storage_backend_factory = storage_backend_factory
|
||||
self.upload_result_file = upload_result_file
|
||||
|
||||
async def create(
|
||||
self,
|
||||
create_request: LiteLLMBatchCreateRequest,
|
||||
unified_input_file_id: str,
|
||||
model: str,
|
||||
provider: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request_tags: Sequence[str] | None,
|
||||
) -> LiteLLMBatch:
|
||||
endpoint: Final = _validate_endpoint(create_request.get("endpoint"))
|
||||
content: Final = await self._download_input(unified_input_file_id, user_api_key_dict)
|
||||
parsed: Final = parse_batch_input(content, endpoint, self._body_rejection(model))
|
||||
if isinstance(parsed, InvalidBatchInput):
|
||||
raise batch_error(400, f"Invalid batch input file: {parsed.describe()}")
|
||||
llm_batch_id: Final = f"{LITELLM_EXECUTED_BATCH_ID_PREFIX}{uuid_module.uuid4().hex}"
|
||||
model_id: Final = next(iter(self.llm_router.get_model_ids(model_name=model)), model)
|
||||
unified_batch_id: Final = self.managed_files.get_unified_batch_id(batch_id=llm_batch_id, model_id=model_id)
|
||||
now: Final = time.time()
|
||||
created_at: Final = int(now)
|
||||
batch: Final = LiteLLMBatch(
|
||||
id=unified_batch_id,
|
||||
object="batch",
|
||||
endpoint=endpoint,
|
||||
input_file_id=unified_input_file_id,
|
||||
completion_window="24h",
|
||||
status="validating",
|
||||
created_at=created_at,
|
||||
expires_at=created_at + int(self.completion_window_seconds),
|
||||
metadata=create_request.get("metadata"),
|
||||
model=model,
|
||||
request_counts=BatchRequestCounts(completed=0, failed=0, total=len(parsed)),
|
||||
)
|
||||
await self.managed_files.store_unified_object_id(
|
||||
unified_object_id=unified_batch_id,
|
||||
file_object=batch,
|
||||
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
model_object_id=llm_batch_id,
|
||||
file_purpose="batch",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_tags=request_tags,
|
||||
persist_attribution=True,
|
||||
batch_processed=True,
|
||||
)
|
||||
_record_batch_created(model, provider, user_api_key_dict)
|
||||
run: Final = _BatchRun(
|
||||
unified_batch_id=unified_batch_id,
|
||||
llm_batch_id=llm_batch_id,
|
||||
model=model,
|
||||
endpoint=endpoint,
|
||||
lines=parsed,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_tags=tuple(request_tags or ()),
|
||||
deadline=now + self.completion_window_seconds,
|
||||
)
|
||||
task: Final = asyncio.create_task(self._run(run))
|
||||
_RUNNING_BATCHES.add(task)
|
||||
task.add_done_callback(_RUNNING_BATCHES.discard)
|
||||
return batch
|
||||
|
||||
async def cancel(self, unified_batch_id: str, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch:
|
||||
current: Final = await self.batches.load_batch(unified_batch_id)
|
||||
if current is None:
|
||||
raise batch_error(404, f"Batch {unified_batch_id} not found")
|
||||
if current.status in TERMINAL_BATCH_STATUSES:
|
||||
raise batch_error(400, f"Cannot cancel a batch with status '{current.status}'")
|
||||
if current.status == "cancelling":
|
||||
return current
|
||||
cancelling: Final = current.model_copy(
|
||||
update=MappingProxyType({"status": "cancelling", "cancelling_at": int(time.time())})
|
||||
)
|
||||
unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status}
|
||||
if await self.batches.compare_and_set(cancelling, unchanged, user_api_key_dict.user_id):
|
||||
return cancelling
|
||||
return await self.cancel(unified_batch_id, user_api_key_dict)
|
||||
|
||||
async def fail_abandoned(self, batch: LiteLLMBatch, user_api_key_dict: UserAPIKeyAuth) -> LiteLLMBatch:
|
||||
error: Final = BatchError(message=_RUNNER_LOST_MESSAGE, code="runner_lost")
|
||||
errors: Final = Errors(data=[error], object="list") # mutable-ok: Errors.data is typed as a list
|
||||
failed: Final = batch.model_copy(
|
||||
update=MappingProxyType({"status": "failed", "failed_at": int(time.time()), "errors": errors})
|
||||
)
|
||||
untouched: Final[prisma_types.DateTimeFilter] = {
|
||||
"lt": datetime.now(timezone.utc) - timedelta(seconds=_STALE_AFTER_SECONDS)
|
||||
}
|
||||
still_abandoned: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {
|
||||
"status": batch.status,
|
||||
"updated_at": untouched,
|
||||
}
|
||||
if await self.batches.compare_and_set(failed, still_abandoned, user_api_key_dict.user_id):
|
||||
return failed
|
||||
return await self.batches.load_batch(batch.id) or batch
|
||||
|
||||
def _body_rejection(self, model: str) -> BodyRejection:
|
||||
def reject(body: Mapping[str, object]) -> str | None:
|
||||
try:
|
||||
is_request_body_safe(
|
||||
request_body=dict(body), # mutable-ok: is_request_body_safe takes a dict
|
||||
general_settings=dict(self.general_settings), # mutable-ok: is_request_body_safe takes a dict
|
||||
llm_router=self.llm_router,
|
||||
model=model,
|
||||
)
|
||||
except ValueError as e:
|
||||
return str(e)
|
||||
return None
|
||||
|
||||
return reject
|
||||
|
||||
async def _download_input(self, unified_input_file_id: str, user_api_key_dict: UserAPIKeyAuth) -> bytes:
|
||||
stored: Final = await self.managed_files.get_unified_file_id(
|
||||
unified_input_file_id, litellm_parent_otel_span=user_api_key_dict.parent_otel_span
|
||||
)
|
||||
if stored is None or not stored.storage_backend or not stored.storage_url:
|
||||
raise batch_error(
|
||||
400,
|
||||
f"LiteLLM does not hold the content of input file {unified_input_file_id}: "
|
||||
f"{LITELLM_EXECUTED_BATCH_UPLOAD_GUIDANCE}",
|
||||
)
|
||||
try:
|
||||
backend: Final = self.storage_backend_factory(stored.storage_backend, prisma_client=self.prisma_client)
|
||||
return await backend.download_file(stored.storage_url)
|
||||
except ValueError as e:
|
||||
raise batch_error(400, str(e))
|
||||
|
||||
async def _run(self, run: _BatchRun) -> None:
|
||||
heartbeat: Final = asyncio.create_task(self._heartbeat(run))
|
||||
try:
|
||||
await self._execute(run)
|
||||
except Exception as e: # noqa: BLE001 # whatever fails, the batch must end up marked failed
|
||||
verbose_proxy_logger.exception("LiteLLM-executed batch %s failed: %s", run.unified_batch_id, e)
|
||||
error: Final = BatchError(message=str(e), code="internal_error")
|
||||
errors: Final = Errors(data=[error], object="list") # mutable-ok: Errors.data is typed as a list
|
||||
try:
|
||||
await self._advance(run, "failed", MappingProxyType({"errors": errors}))
|
||||
except Exception as advance_error: # noqa: BLE001 # a failed status write is logged, never raised
|
||||
verbose_proxy_logger.exception(
|
||||
"LiteLLM-executed batch %s could not be marked failed: %s", run.unified_batch_id, advance_error
|
||||
)
|
||||
finally:
|
||||
heartbeat.cancel()
|
||||
|
||||
async def _heartbeat(self, run: _BatchRun) -> None:
|
||||
while True:
|
||||
await asyncio.sleep(self.heartbeat_seconds)
|
||||
try:
|
||||
await self._touch(run)
|
||||
except Exception as e: # noqa: BLE001 # a missed beat is logged and the next one retries
|
||||
verbose_proxy_logger.warning("LiteLLM-executed batch %s heartbeat failed: %s", run.unified_batch_id, e)
|
||||
|
||||
async def _touch(self, run: _BatchRun) -> None:
|
||||
await self.batches.touch(run.unified_batch_id, run.user_api_key_dict.user_id)
|
||||
|
||||
async def _execute(self, run: _BatchRun) -> None:
|
||||
await self._advance(run, "in_progress")
|
||||
watch: Final = _StopWatch(lambda: self.batches.load_status(run.unified_batch_id), _CANCEL_POLL_SECONDS)
|
||||
semaphore: Final = asyncio.Semaphore(self.concurrency)
|
||||
results: Final = await asyncio.gather(*(self._run_row(run, line, watch, semaphore) for line in run.lines))
|
||||
outcomes: Final = tuple(outcome for outcome in results if outcome is not None)
|
||||
if await self._advance(run, "finalizing") is None:
|
||||
return
|
||||
succeeded: Final = tuple(
|
||||
outcome for outcome in outcomes if isinstance(outcome, RowOutcome) and outcome.succeeded
|
||||
)
|
||||
failed: Final = tuple(
|
||||
outcome for outcome in outcomes if isinstance(outcome, ExpiredRow) or not outcome.succeeded
|
||||
)
|
||||
output_file_id: Final = await self._upload_results(run, "output", succeeded)
|
||||
error_file_id: Final = await self._upload_results(run, "error", failed)
|
||||
request_counts: Final = BatchRequestCounts(completed=len(succeeded), failed=len(failed), total=len(run.lines))
|
||||
final_status: Final[BatchStatus] = (
|
||||
"expired" if any(isinstance(outcome, ExpiredRow) for outcome in outcomes) else "completed"
|
||||
)
|
||||
await self._advance(
|
||||
run,
|
||||
final_status,
|
||||
MappingProxyType(
|
||||
{"output_file_id": output_file_id, "error_file_id": error_file_id, "request_counts": request_counts}
|
||||
),
|
||||
)
|
||||
|
||||
async def _run_row(
|
||||
self, run: _BatchRun, line: BatchInputLine, watch: _StopWatch, semaphore: asyncio.Semaphore
|
||||
) -> RowOutcome | ExpiredRow | None:
|
||||
async with semaphore:
|
||||
if await watch.stopped():
|
||||
return None
|
||||
remaining: Final = run.deadline - time.time()
|
||||
if remaining <= 0:
|
||||
return ExpiredRow(custom_id=line.custom_id)
|
||||
try:
|
||||
return await asyncio.wait_for(self._row_outcome(run, line), timeout=remaining)
|
||||
except asyncio.TimeoutError:
|
||||
return ExpiredRow(custom_id=line.custom_id)
|
||||
|
||||
async def _row_outcome(self, run: _BatchRun, line: BatchInputLine) -> RowOutcome:
|
||||
try:
|
||||
body: Final = await self._dispatch(run, line)
|
||||
except Exception as e: # noqa: BLE001 # a provider error becomes the row's error line, never a crashed batch
|
||||
return RowOutcome(
|
||||
custom_id=line.custom_id, status_code=_status_code_of(e), body=_error_body(e), succeeded=False
|
||||
)
|
||||
return RowOutcome(custom_id=line.custom_id, status_code=200, body=body, succeeded=True)
|
||||
|
||||
async def _dispatch(self, run: _BatchRun, line: BatchInputLine) -> Mapping[str, object]:
|
||||
params: Final = MappingProxyType(
|
||||
{**line.body, "model": run.model, "metadata": self._row_metadata(run), "disable_fallbacks": True}
|
||||
)
|
||||
return _dump(await self._router_call(run.endpoint)(**params))
|
||||
|
||||
def _router_call(self, endpoint: BatchEndpoint) -> _RouterCall:
|
||||
method: Final[object] = getattr(self.llm_router, _ROUTER_METHODS[endpoint], None)
|
||||
if not isinstance(method, _RouterCall):
|
||||
raise TypeError(f"the router has no callable for {endpoint}")
|
||||
return method
|
||||
|
||||
def _row_metadata(self, run: _BatchRun) -> dict[str, object]: # mutable-ok: router updates metadata in place
|
||||
return { # mutable-ok: the router updates request metadata in place
|
||||
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(run.user_api_key_dict),
|
||||
"user_api_key": run.user_api_key_dict.api_key,
|
||||
"user_api_end_user_max_budget": run.user_api_key_dict.end_user_max_budget,
|
||||
"tags": list(run.request_tags), # mutable-ok: litellm types request tags as a list
|
||||
"batch_id": run.unified_batch_id,
|
||||
}
|
||||
|
||||
async def _upload_results(
|
||||
self, run: _BatchRun, kind: Literal["output", "error"], outcomes: Sequence[RowOutcome | ExpiredRow]
|
||||
) -> str | None:
|
||||
if not outcomes:
|
||||
return None
|
||||
content: Final = "".join(f"{json.dumps(_result_line(outcome))}\n" for outcome in outcomes).encode()
|
||||
file_data: Final[ExtractedFileData] = {
|
||||
"filename": f"{run.llm_batch_id}_{kind}.jsonl",
|
||||
"content": content,
|
||||
"content_type": "application/jsonl",
|
||||
"headers": _NO_HEADERS,
|
||||
}
|
||||
file_object: Final = await self.upload_result_file(
|
||||
file_data=file_data,
|
||||
target_storage=LITELLM_DB_STORAGE_BACKEND_NAME,
|
||||
target_model_names=(run.model,),
|
||||
purpose="batch_output",
|
||||
proxy_logging_obj=self.proxy_logging_obj,
|
||||
user_api_key_dict=run.user_api_key_dict,
|
||||
prisma_client=self.prisma_client,
|
||||
)
|
||||
return file_object.id
|
||||
|
||||
async def _advance(
|
||||
self, run: _BatchRun, requested: BatchStatus, fields: Mapping[str, object] = _NO_FIELDS
|
||||
) -> BatchStatus | None:
|
||||
current: Final = await self.batches.load_batch(run.unified_batch_id)
|
||||
if current is None:
|
||||
raise RuntimeError(f"Batch {run.unified_batch_id} is no longer stored")
|
||||
if current.status in TERMINAL_BATCH_STATUSES:
|
||||
return None
|
||||
status: Final = _resolve_transition(current.status, requested)
|
||||
updated: Final = current.model_copy(
|
||||
update=MappingProxyType({**fields, "status": status, f"{status}_at": int(time.time())})
|
||||
)
|
||||
unchanged: Final[prisma_types.LiteLLM_ManagedObjectTableWhereInput] = {"status": current.status}
|
||||
if await self.batches.compare_and_set(updated, unchanged, run.user_api_key_dict.user_id):
|
||||
return status
|
||||
return await self._advance(run, requested, fields)
|
||||
|
||||
|
||||
def _record_batch_created(model: str, provider: str, user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
prometheus_logger: Final = PrometheusLogger.get_instance()
|
||||
if prometheus_logger is None:
|
||||
return
|
||||
prometheus_logger.record_managed_batch_created(
|
||||
model=model,
|
||||
api_provider=provider,
|
||||
user=user_api_key_dict.user_id or "",
|
||||
user_email=user_api_key_dict.user_email or "",
|
||||
api_key_alias=user_api_key_dict.key_alias or "",
|
||||
)
|
||||
|
|
@ -38,6 +38,7 @@ if TYPE_CHECKING:
|
|||
FILE_LIST_CONTINUATION_CHUNK_SIZE: Final = 500
|
||||
|
||||
BATCH_CREATE_HIDDEN_PARAM: Final = "batch_create"
|
||||
LITELLM_EXECUTED_BATCH_ID_PREFIX: Final = "litellm_batch_"
|
||||
|
||||
|
||||
def validate_file_list_limit(limit: int | None) -> None:
|
||||
|
|
@ -179,6 +180,11 @@ def get_batch_id_from_unified_batch_id(file_id: str) -> str:
|
|||
return re.split(r"[;,]", batch_id, maxsplit=1)[0]
|
||||
|
||||
|
||||
def is_litellm_executed_batch(decoded_unified_batch_id: str) -> bool:
|
||||
_, marker, batch_id = decoded_unified_batch_id.partition("llm_batch_id:")
|
||||
return bool(marker) and batch_id.startswith(LITELLM_EXECUTED_BATCH_ID_PREFIX)
|
||||
|
||||
|
||||
def encode_file_id_with_model(file_id: str, model: str, id_type: Literal["file", "batch"] = "file") -> str:
|
||||
"""
|
||||
Encode a file/batch ID with model routing information.
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@
|
|||
|
||||
import asyncio
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, BinaryIO, Final, TypedDict, cast, get_args
|
||||
|
||||
import httpx
|
||||
|
|
@ -32,10 +32,15 @@ from litellm.litellm_core_utils.cloud_storage_security import (
|
|||
is_managed_cloud_storage_uri,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
|
||||
from litellm.llms.base_llm.files.litellm_db_storage_backend import LITELLM_DB_STORAGE_BACKEND_NAME
|
||||
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
|
||||
from litellm.llms.base_llm.managed_resources.isolation import build_list_page
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.batches_endpoints.litellm_executed_batches import (
|
||||
litellm_executed_provider_of,
|
||||
resolve_litellm_executed_provider,
|
||||
)
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
|
|
@ -67,6 +72,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
_is_base64_encoded_unified_file_id,
|
||||
add_internal_model_credentials,
|
||||
apply_team_provider_credentials,
|
||||
authorize_model_for_key,
|
||||
encode_file_id_with_model,
|
||||
extract_file_creation_params,
|
||||
get_authorized_credentials_for_model,
|
||||
|
|
@ -86,7 +92,7 @@ from litellm.proxy.openai_files_endpoints.general_upload_validation import (
|
|||
coerce_optional_str_list_setting,
|
||||
raise_upload_validation_failure,
|
||||
)
|
||||
from litellm.proxy.utils import ProxyLogging, is_known_model
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, is_known_model
|
||||
from litellm.repositories.table_repositories import ManagedFileRepository
|
||||
from litellm.router import Router
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -99,6 +105,64 @@ from litellm.types.llms.openai import (
|
|||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
def _names_a_litellm_executed_provider(llm_router: Router, candidate: str, team_id: str | None) -> bool:
|
||||
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=candidate, team_id=team_id)
|
||||
return credentials is not None and litellm_executed_provider_of(credentials) is not None
|
||||
|
||||
|
||||
async def _litellm_executed_batch_input_model(
|
||||
llm_router: Router | None,
|
||||
purpose: OpenAIFilesPurpose,
|
||||
model: str | None,
|
||||
target_model_names_list: Sequence[str],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
explicit_storage: str | None,
|
||||
) -> str | None:
|
||||
if llm_router is None:
|
||||
return None
|
||||
candidates: Final = (model,) if model is not None else tuple(target_model_names_list)
|
||||
team_id: Final = user_api_key_dict.team_id
|
||||
await asyncio.gather(
|
||||
*(
|
||||
authorize_model_for_key(model_id=candidate, llm_router=llm_router, user_api_key_dict=user_api_key_dict)
|
||||
for candidate in candidates
|
||||
if _names_a_litellm_executed_provider(llm_router, candidate, team_id)
|
||||
)
|
||||
)
|
||||
if explicit_storage is not None:
|
||||
return None
|
||||
providers: Final = await asyncio.gather(
|
||||
*(resolve_litellm_executed_provider(llm_router, candidate, team_id) for candidate in candidates)
|
||||
)
|
||||
executed: Final = tuple(
|
||||
candidate for candidate, provider in zip(candidates, providers, strict=True) if provider is not None
|
||||
)
|
||||
if not executed:
|
||||
return None
|
||||
if purpose != "batch":
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"The server behind {', '.join(executed)} has no Files API, so LiteLLM keeps only batch input "
|
||||
f"files for it and runs the batch itself: upload with purpose=batch; got purpose={purpose}"
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="purpose",
|
||||
code=400,
|
||||
)
|
||||
if len(candidates) == 1:
|
||||
return executed[0]
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"LiteLLM runs batches for {', '.join(executed)} itself and keeps their input files, so a batch "
|
||||
f"input file can target only that one model; got target_model_names={', '.join(candidates)}"
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="target_model_names",
|
||||
code=400,
|
||||
)
|
||||
|
||||
|
||||
_MAX_BATCH_FILE_SIZE_MB_ADAPTER: Final = TypeAdapter(int | None)
|
||||
_LISTED_FILES_ADAPTER: Final = TypeAdapter(list[OpenAIFileObject])
|
||||
|
||||
|
|
@ -244,30 +308,41 @@ async def route_create_file(
|
|||
5. Else -> use custom_llm_provider with files_settings
|
||||
"""
|
||||
|
||||
# Handle custom storage backend
|
||||
if target_storage and target_storage != "default":
|
||||
explicit_storage: Final = target_storage if target_storage and target_storage != "default" else None
|
||||
if explicit_storage == LITELLM_DB_STORAGE_BACKEND_NAME:
|
||||
raise ProxyException(
|
||||
message=(
|
||||
f"target_storage={LITELLM_DB_STORAGE_BACKEND_NAME} is not a storage a caller can pick: LiteLLM "
|
||||
"chooses it on its own for the batch input files of a model whose batches it runs itself, so "
|
||||
"upload with purpose=batch and name that model instead of target_storage"
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="target_storage",
|
||||
code=400,
|
||||
)
|
||||
executed_model: Final = await _litellm_executed_batch_input_model(
|
||||
llm_router, purpose, model, target_model_names_list, user_api_key_dict, explicit_storage
|
||||
)
|
||||
storage: Final = explicit_storage or (LITELLM_DB_STORAGE_BACKEND_NAME if executed_model is not None else None)
|
||||
if storage is not None:
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
extract_file_data,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.storage_backend_service import (
|
||||
StorageBackendFileService,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
# Extract file data
|
||||
file_data: Final = extract_file_data(cast(Any, _create_file_request.get("file")))
|
||||
|
||||
# Use storage backend service to handle upload
|
||||
file_object: Final = await StorageBackendFileService.upload_file_to_storage_backend(
|
||||
file_data=file_data,
|
||||
target_storage=target_storage,
|
||||
target_model_names=target_model_names_list,
|
||||
return await StorageBackendFileService.upload_file_to_storage_backend(
|
||||
file_data=extract_file_data(cast(Any, _create_file_request.get("file"))),
|
||||
target_storage=storage,
|
||||
target_model_names=(executed_model,) if executed_model is not None else target_model_names_list,
|
||||
purpose=purpose,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
return file_object
|
||||
|
||||
# NEW: Handle model-based routing (no DB required)
|
||||
if model is not None:
|
||||
# Get credentials from model_list via router
|
||||
|
|
@ -848,7 +923,7 @@ async def get_file_content(
|
|||
|
||||
# Check if file is stored in a storage backend (check DB)
|
||||
if hasattr(managed_files_obj, "prisma_client") and getattr(managed_files_obj, "prisma_client", None):
|
||||
prisma_client: Final = getattr(managed_files_obj, "prisma_client")
|
||||
prisma_client: Final[PrismaClient] = getattr(managed_files_obj, "prisma_client")
|
||||
db_file: Final = await ManagedFileRepository(prisma_client).table.find_first(
|
||||
where={"unified_file_id": file_id}
|
||||
)
|
||||
|
|
@ -863,7 +938,7 @@ async def get_file_content(
|
|||
|
||||
try:
|
||||
# Get storage backend (uses same env vars as callback)
|
||||
storage_backend: Final = get_storage_backend(storage_backend_name)
|
||||
storage_backend: Final = get_storage_backend(storage_backend_name, prisma_client=prisma_client)
|
||||
file_content: Final = await storage_backend.download_file(storage_url)
|
||||
|
||||
# Return file content
|
||||
|
|
|
|||
|
|
@ -7,15 +7,16 @@ storage backends (e.g., Azure Blob Storage) and managing associated metadata.
|
|||
|
||||
import base64
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid as uuid_module
|
||||
from litellm.llms.base_llm.files.storage_backend import BaseFileStorageBackend
|
||||
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
|
||||
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.types.llms.openai import OpenAIFileObject, OpenAIFilesPurpose
|
||||
from litellm.types.utils import SpecialEnums
|
||||
|
||||
|
|
@ -35,21 +36,23 @@ class StorageBackendFileService:
|
|||
async def upload_file_to_storage_backend(
|
||||
file_data: Mapping[str, Any],
|
||||
target_storage: str,
|
||||
target_model_names: list[str],
|
||||
target_model_names: Sequence[str],
|
||||
purpose: OpenAIFilesPurpose,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient | None = None,
|
||||
) -> OpenAIFileObject:
|
||||
"""
|
||||
Upload a file to a storage backend and create a file object.
|
||||
|
||||
Args:
|
||||
file_data: File data dictionary from extract_file_data()
|
||||
target_storage: Storage backend name (e.g., "azure_storage")
|
||||
target_storage: Storage backend name (e.g., "azure_storage", "litellm_db")
|
||||
target_model_names: List of model names for managed files
|
||||
purpose: File purpose (e.g., "user_data", "batch")
|
||||
proxy_logging_obj: Proxy logging object for accessing hooks
|
||||
user_api_key_dict: User API key authentication data
|
||||
prisma_client: The proxy's database client, required by the "litellm_db" backend
|
||||
|
||||
Returns:
|
||||
OpenAIFileObject: Created file object with storage metadata
|
||||
|
|
@ -59,7 +62,7 @@ class StorageBackendFileService:
|
|||
"""
|
||||
# Get storage backend instance
|
||||
try:
|
||||
storage_backend: Final = get_storage_backend(target_storage)
|
||||
storage_backend: Final = get_storage_backend(target_storage, prisma_client=prisma_client)
|
||||
except ValueError as e:
|
||||
raise ProxyException(
|
||||
message=str(e),
|
||||
|
|
@ -103,8 +106,9 @@ class StorageBackendFileService:
|
|||
storage_url=storage_url,
|
||||
)
|
||||
|
||||
# Store in managed files if target_model_names provided
|
||||
if target_model_names:
|
||||
if not target_model_names:
|
||||
return file_object
|
||||
try:
|
||||
await StorageBackendFileService._store_in_managed_files(
|
||||
file_object=file_object,
|
||||
file_data=file_data,
|
||||
|
|
@ -114,9 +118,25 @@ class StorageBackendFileService:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
except Exception:
|
||||
await StorageBackendFileService._discard_orphaned_content(storage_backend, storage_url, target_storage)
|
||||
raise
|
||||
return file_object
|
||||
|
||||
@staticmethod
|
||||
async def _discard_orphaned_content(
|
||||
storage_backend: BaseFileStorageBackend, storage_url: str, target_storage: str
|
||||
) -> None:
|
||||
try:
|
||||
await storage_backend.delete_file(storage_url)
|
||||
except Exception as e: # noqa: BLE001 # the metadata failure is what surfaces; a failed cleanup is only logged
|
||||
verbose_proxy_logger.warning(
|
||||
"Could not delete orphaned content at %s on %s after its metadata write failed: %s",
|
||||
storage_url,
|
||||
target_storage,
|
||||
e,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _create_file_object_with_storage_metadata(
|
||||
file_content: bytes,
|
||||
|
|
@ -164,7 +184,7 @@ class StorageBackendFileService:
|
|||
@staticmethod
|
||||
def _create_unified_file_id(
|
||||
file_type: str,
|
||||
target_model_names: list[str],
|
||||
target_model_names: Sequence[str],
|
||||
file_id: str,
|
||||
) -> str:
|
||||
"""
|
||||
|
|
@ -194,7 +214,7 @@ class StorageBackendFileService:
|
|||
async def _store_in_managed_files(
|
||||
file_object: OpenAIFileObject,
|
||||
file_data: Mapping[str, Any],
|
||||
target_model_names: list[str],
|
||||
target_model_names: Sequence[str],
|
||||
target_storage: str,
|
||||
storage_url: str,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
48
litellm/repositories/managed_batch_repository.py
Normal file
48
litellm/repositories/managed_batch_repository.py
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.repositories.table_repositories import PrismaTableRepository
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
|
||||
|
||||
def _batch_of(blob: object) -> LiteLLMBatch:
|
||||
return LiteLLMBatch.model_validate_json(blob) if isinstance(blob, str) else LiteLLMBatch.model_validate(blob)
|
||||
|
||||
|
||||
class ManagedBatchRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedObjectTable"]):
|
||||
table_name = "litellm_managedobjecttable"
|
||||
|
||||
async def load_batch(self, unified_batch_id: str) -> LiteLLMBatch | None:
|
||||
row: Final = await self._find_row(unified_batch_id)
|
||||
return None if row is None or not row.file_object else _batch_of(row.file_object)
|
||||
|
||||
async def load_status(self, unified_batch_id: str) -> str | None:
|
||||
row: Final = await self._find_row(unified_batch_id)
|
||||
return row.status if row is not None else None
|
||||
|
||||
async def compare_and_set(
|
||||
self, batch: LiteLLMBatch, unchanged: Mapping[str, object], updated_by: str | None
|
||||
) -> bool:
|
||||
updated_rows: Final = await self.table.update_many(
|
||||
where={"unified_object_id": batch.id, **unchanged}, # mutable-ok: prisma filters are plain dicts
|
||||
data={ # mutable-ok: prisma payloads are plain dicts
|
||||
"file_object": batch.model_dump_json(),
|
||||
"status": batch.status,
|
||||
"updated_by": updated_by,
|
||||
},
|
||||
)
|
||||
return updated_rows > 0
|
||||
|
||||
async def touch(self, unified_batch_id: str, updated_by: str | None) -> None:
|
||||
await self.table.update_many(
|
||||
where={"unified_object_id": unified_batch_id}, # mutable-ok: prisma filters are plain dicts
|
||||
data={"updated_by": updated_by}, # mutable-ok: prisma payloads are plain dicts
|
||||
)
|
||||
|
||||
async def _find_row(self, unified_batch_id: str) -> "prisma_models.LiteLLM_ManagedObjectTable | None":
|
||||
return await self.table.find_first(
|
||||
where={"unified_object_id": unified_batch_id} # mutable-ok: prisma filters are plain dicts
|
||||
)
|
||||
32
litellm/repositories/managed_file_content_repository.py
Normal file
32
litellm/repositories/managed_file_content_repository.py
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.repositories.table_repositories import PrismaTableRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
|
||||
|
||||
class ManagedFileContentRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileContentTable"]):
|
||||
table_name = "litellm_managedfilecontenttable"
|
||||
|
||||
async def store(self, content: bytes) -> str:
|
||||
from prisma import Base64
|
||||
|
||||
row: Final = await self.table.create(
|
||||
data={"content": Base64.encode(content)} # mutable-ok: prisma payloads are plain dicts
|
||||
)
|
||||
return row.id
|
||||
|
||||
async def load(self, row_id: str) -> bytes | None:
|
||||
row: Final[prisma_models.LiteLLM_ManagedFileContentTable | None] = await self.table.find_unique(
|
||||
where={"id": row_id} # mutable-ok: prisma filters are plain dicts
|
||||
)
|
||||
return None if row is None else row.content.decode()
|
||||
|
||||
async def delete(self, row_id: str) -> None:
|
||||
from prisma.errors import RecordNotFoundError
|
||||
|
||||
try:
|
||||
await self.table.delete(where={"id": row_id}) # mutable-ok: prisma filters are plain dicts
|
||||
except RecordNotFoundError:
|
||||
return
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -4178,6 +4178,8 @@ FILE_CONTENT_STREAMING_PROVIDERS: Final[frozenset[str]] = frozenset(
|
|||
{*OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS, LlmProviders.VERTEX_AI.value}
|
||||
)
|
||||
|
||||
LITELLM_EXECUTED_BATCH_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.HOSTED_VLLM.value})
|
||||
|
||||
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai"]
|
||||
|
||||
LIST_BATCHES_SUPPORTED_PROVIDERS: Final[frozenset[str]] = frozenset(get_args(ListBatchesSupportedProvider))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1067,6 +1067,7 @@ async def test_afile_content_passes_trusted_model_credentials_to_router():
|
|||
managed_files = _make_managed_files_instance()
|
||||
unified_file_id = "unified-file-id"
|
||||
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
managed_files.get_unified_file_id = AsyncMock(return_value=None)
|
||||
managed_files.get_model_file_id_mapping = AsyncMock(
|
||||
return_value={unified_file_id: {"model-123": s3_uri}}
|
||||
)
|
||||
|
|
@ -1238,6 +1239,7 @@ async def test_afile_content_bedrock_unified_id_end_to_end(monkeypatch):
|
|||
managed_files = _make_managed_files_instance()
|
||||
unified_file_id = "unified-file-id"
|
||||
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
managed_files.get_unified_file_id = AsyncMock(return_value=None)
|
||||
managed_files.get_model_file_id_mapping = AsyncMock(
|
||||
return_value={unified_file_id: {"model-123": s3_uri}}
|
||||
)
|
||||
|
|
@ -1268,6 +1270,7 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri():
|
|||
managed_files = _make_managed_files_instance()
|
||||
unified_file_id = "litellm_proxy_unified_id_abc"
|
||||
s3_uri = "s3://my-bucket/litellm-batch-outputs/job-123/input.jsonl.out"
|
||||
managed_files.get_unified_file_id = AsyncMock(return_value=None)
|
||||
managed_files.get_model_file_id_mapping = AsyncMock(
|
||||
return_value={unified_file_id: {"model-123": s3_uri}}
|
||||
)
|
||||
|
|
@ -1732,6 +1735,40 @@ async def test_batch_retrieve_hook_does_not_claim_attribution():
|
|||
assert managed_files.store_unified_object_id.await_args.kwargs["persist_attribution"] is False
|
||||
|
||||
|
||||
def _unified_batch_id(llm_batch_id: str) -> str:
|
||||
decoded = f"litellm_proxy;model_id:my-vllm;llm_batch_id:{llm_batch_id}"
|
||||
return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"llm_batch_id, stores",
|
||||
[("litellm_batch_abc", False), ("batch_abc", True)],
|
||||
ids=["litellm-executed batch is left alone", "provider batch is still stored"],
|
||||
)
|
||||
async def test_post_call_hook_leaves_litellm_executed_batches_untouched(llm_batch_id: str, stores: bool):
|
||||
managed_files = _make_managed_files_instance()
|
||||
response = _make_batch_response(status="in_progress", output_file_id=None)
|
||||
response.id = _unified_batch_id(llm_batch_id)
|
||||
response._hidden_params = {
|
||||
"unified_batch_id": response.id,
|
||||
"model_id": "my-vllm",
|
||||
"model_name": "hosted_vllm/qwen",
|
||||
}
|
||||
original_id = response.id
|
||||
|
||||
returned = await managed_files.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-the-poller", user_id="bob", parent_otel_span=None),
|
||||
response=response,
|
||||
)
|
||||
|
||||
assert returned is response
|
||||
assert managed_files.store_unified_object_id.await_count == (1 if stores else 0)
|
||||
if not stores:
|
||||
assert response.id == original_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_delete_passes_trusted_model_credentials_to_router():
|
||||
"""
|
||||
|
|
@ -1743,6 +1780,7 @@ async def test_afile_delete_passes_trusted_model_credentials_to_router():
|
|||
managed_files = _make_managed_files_instance()
|
||||
unified_file_id = "unified-file-id"
|
||||
s3_uri = "s3://my-bucket/litellm-bedrock-files/job-123/input.jsonl"
|
||||
managed_files.get_unified_file_id = AsyncMock(return_value=None)
|
||||
managed_files.get_model_file_id_mapping = AsyncMock(return_value={unified_file_id: {"model-123": s3_uri}})
|
||||
managed_files.delete_unified_file_id = AsyncMock(return_value=_make_file_object(unified_file_id))
|
||||
|
||||
|
|
@ -1809,6 +1847,7 @@ async def test_afile_delete_bedrock_unified_id_end_to_end(monkeypatch):
|
|||
managed_files = _make_managed_files_instance()
|
||||
unified_file_id = "unified-file-id"
|
||||
s3_uri = "s3://my-bucket/litellm-bedrock-files/job-123/input.jsonl"
|
||||
managed_files.get_unified_file_id = AsyncMock(return_value=None)
|
||||
managed_files.get_model_file_id_mapping = AsyncMock(return_value={unified_file_id: {"model-123": s3_uri}})
|
||||
managed_files.delete_unified_file_id = AsyncMock(return_value=_make_file_object(unified_file_id))
|
||||
|
||||
|
|
@ -1827,3 +1866,147 @@ async def test_afile_delete_bedrock_unified_id_end_to_end(monkeypatch):
|
|||
assert response.id == unified_file_id
|
||||
assert response.model_dump() == {"id": unified_file_id, "object": "file", "deleted": True}
|
||||
managed_files.delete_unified_file_id.assert_awaited_once_with(unified_file_id, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_delete_storage_backed_row_deletes_stored_content_not_provider_files():
|
||||
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
|
||||
from openai.types import FileDeleted
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.models.managed_files import LiteLLM_ManagedFileTable
|
||||
|
||||
storage_url = "litellm_db://content-row-1"
|
||||
unified_file_id = _managed_deletion_file_id(storage_url)
|
||||
row = LiteLLM_ManagedFileTable(
|
||||
unified_file_id=unified_file_id,
|
||||
model_mappings={"vllm-batch": storage_url},
|
||||
flat_model_file_ids=[storage_url],
|
||||
file_object=_make_file_object(unified_file_id),
|
||||
storage_backend="litellm_db",
|
||||
storage_url=storage_url,
|
||||
)
|
||||
file_table = MagicMock(find_first=AsyncMock(return_value=row), delete=AsyncMock())
|
||||
content_table = MagicMock(delete=AsyncMock())
|
||||
managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
internal_usage_cache=DualCache(),
|
||||
prisma_client=MagicMock(
|
||||
db=MagicMock(litellm_managedfiletable=file_table, litellm_managedfilecontenttable=content_table)
|
||||
),
|
||||
)
|
||||
router = MagicMock(
|
||||
get_deployment_credentials_with_provider=MagicMock(return_value=None),
|
||||
afile_delete=AsyncMock(),
|
||||
)
|
||||
|
||||
response = await managed_files.afile_delete(
|
||||
file_id=unified_file_id,
|
||||
litellm_parent_otel_span=None,
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
content_table.delete.assert_awaited_once_with(where={"id": "content-row-1"})
|
||||
router.afile_delete.assert_not_awaited()
|
||||
file_table.delete.assert_awaited_once_with(where={"unified_file_id": unified_file_id})
|
||||
assert response == FileDeleted(id=unified_file_id, object="file", deleted=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_storage_backed_row_returns_stored_bytes_not_provider_content():
|
||||
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
|
||||
from prisma import Base64
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.models.managed_files import LiteLLM_ManagedFileTable
|
||||
|
||||
storage_url = "litellm_db://content-row-1"
|
||||
unified_file_id = _managed_deletion_file_id(storage_url)
|
||||
stored_bytes = b'{"custom_id": "line-1", "method": "POST", "url": "/v1/chat/completions", "body": {}}\n'
|
||||
row = LiteLLM_ManagedFileTable(
|
||||
unified_file_id=unified_file_id,
|
||||
model_mappings={"vllm-batch": storage_url},
|
||||
flat_model_file_ids=[storage_url],
|
||||
file_object=_make_file_object(unified_file_id),
|
||||
storage_backend="litellm_db",
|
||||
storage_url=storage_url,
|
||||
)
|
||||
file_table = MagicMock(find_first=AsyncMock(return_value=row))
|
||||
content_table = MagicMock(find_unique=AsyncMock(return_value=MagicMock(content=Base64.encode(stored_bytes))))
|
||||
managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
internal_usage_cache=DualCache(),
|
||||
prisma_client=MagicMock(
|
||||
db=MagicMock(litellm_managedfiletable=file_table, litellm_managedfilecontenttable=content_table)
|
||||
),
|
||||
)
|
||||
router = MagicMock(
|
||||
get_deployment_credentials_with_provider=MagicMock(return_value=None),
|
||||
afile_content=AsyncMock(),
|
||||
)
|
||||
|
||||
response = await managed_files.afile_content(
|
||||
file_id=unified_file_id,
|
||||
litellm_parent_otel_span=None,
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
assert response.content == stored_bytes
|
||||
content_table.find_unique.assert_awaited_once_with(where={"id": "content-row-1"})
|
||||
router.afile_content.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_unified_object_id_batch_processed_is_written_only_when_asked():
|
||||
managed_files, mock_prisma = _make_object_store_instance()
|
||||
upsert = mock_prisma.db.litellm_managedobjecttable.upsert
|
||||
creator = UserAPIKeyAuth(api_key="sk-creator", user_id="alice", team_id="team-alpha", parent_otel_span=None)
|
||||
|
||||
await managed_files.store_unified_object_id(
|
||||
unified_object_id="uoi-processed",
|
||||
file_object=_make_batch_response(status="completed"),
|
||||
litellm_parent_otel_span=None,
|
||||
model_object_id="batch-processed",
|
||||
file_purpose="batch",
|
||||
user_api_key_dict=creator,
|
||||
batch_processed=True,
|
||||
)
|
||||
await managed_files.store_unified_object_id(
|
||||
unified_object_id="uoi-default",
|
||||
file_object=_make_batch_response(status="completed"),
|
||||
litellm_parent_otel_span=None,
|
||||
model_object_id="batch-default",
|
||||
file_purpose="batch",
|
||||
user_api_key_dict=creator,
|
||||
)
|
||||
|
||||
processed_create, default_create = (call.kwargs["data"]["create"] for call in upsert.await_args_list)
|
||||
assert processed_create["batch_processed"] is True
|
||||
assert default_create["batch_processed"] is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_unified_file_id_caches_the_storage_location_the_db_row_gets():
|
||||
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
|
||||
|
||||
from litellm.caching import DualCache
|
||||
|
||||
file_table = MagicMock(upsert=AsyncMock(), find_first=AsyncMock(side_effect=AssertionError("cache miss")))
|
||||
managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
internal_usage_cache=DualCache(),
|
||||
prisma_client=MagicMock(db=MagicMock(litellm_managedfiletable=file_table)),
|
||||
)
|
||||
stored = _make_file_object("file-kept").model_copy(update={"purpose": "batch"})
|
||||
stored._hidden_params = {"storage_backend": "litellm_db", "storage_url": "litellm_db://content-row-1"}
|
||||
|
||||
await managed_files.store_unified_file_id(
|
||||
file_id="unified-kept",
|
||||
file_object=stored,
|
||||
litellm_parent_otel_span=None,
|
||||
model_mappings={"vllm-batch": "litellm_db://content-row-1"},
|
||||
user_api_key_dict=_make_user_api_key_dict(),
|
||||
)
|
||||
cached = await managed_files.get_unified_file_id("unified-kept")
|
||||
|
||||
assert cached is not None
|
||||
assert (cached.storage_backend, cached.storage_url) == ("litellm_db", "litellm_db://content-row-1")
|
||||
create_data = file_table.upsert.await_args.kwargs["data"]["create"]
|
||||
assert (create_data["storage_backend"], create_data["storage_url"]) == ("litellm_db", "litellm_db://content-row-1")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -0,0 +1,34 @@
|
|||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.base_llm.files.litellm_db_storage_backend import (
|
||||
LITELLM_DB_STORAGE_BACKEND_NAME,
|
||||
LITELLM_DB_STORAGE_URL_PREFIX,
|
||||
LiteLLMDbStorageBackend,
|
||||
)
|
||||
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_db_backend_stores_through_the_given_prisma_client():
|
||||
table = MagicMock(create=AsyncMock(return_value=SimpleNamespace(id="row-1")))
|
||||
prisma_client = MagicMock(db=MagicMock(litellm_managedfilecontenttable=table))
|
||||
|
||||
backend = get_storage_backend(LITELLM_DB_STORAGE_BACKEND_NAME, prisma_client=prisma_client)
|
||||
|
||||
assert isinstance(backend, LiteLLMDbStorageBackend)
|
||||
stored_at = await backend.upload_file(file_content=b"line\n", filename="input.jsonl", content_type="text/plain")
|
||||
assert stored_at == f"{LITELLM_DB_STORAGE_URL_PREFIX}row-1"
|
||||
table.create.assert_awaited_once()
|
||||
|
||||
|
||||
def test_litellm_db_backend_without_a_database_is_rejected():
|
||||
with pytest.raises(ValueError, match="database-connected proxy"):
|
||||
get_storage_backend(LITELLM_DB_STORAGE_BACKEND_NAME)
|
||||
|
||||
|
||||
def test_unknown_backend_is_still_rejected():
|
||||
with pytest.raises(ValueError, match="Unsupported storage backend type: nope"):
|
||||
get_storage_backend("nope", prisma_client=MagicMock())
|
||||
|
|
@ -34,10 +34,13 @@ import json
|
|||
import logging
|
||||
from contextlib import ExitStack
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
|
||||
|
||||
import litellm
|
||||
|
|
@ -52,7 +55,7 @@ from litellm.proxy.utils import ProxyLogging
|
|||
from litellm.router import Router
|
||||
from litellm.types.llms.openai import BatchJobStatus
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
from litellm.types.utils import CredentialItem, LiteLLMBatch
|
||||
from litellm.types.utils import CredentialItem, LiteLLMBatch, SpecialEnums
|
||||
|
||||
from fastapi import Request, Response
|
||||
|
||||
|
|
@ -74,6 +77,12 @@ CREDS: Dict[str, Dict[str, str]] = {
|
|||
"api_base": "https://vertex.test",
|
||||
"model": "vertex_ai/gemini-2.0",
|
||||
},
|
||||
"my-vllm": {
|
||||
"custom_llm_provider": "hosted_vllm",
|
||||
"api_key": "sk-vllm",
|
||||
"api_base": "http://vllm.test/v1",
|
||||
"model": "hosted_vllm/qwen",
|
||||
},
|
||||
}
|
||||
|
||||
# A real model-encoded file id: decodes to "azure/gpt-4o", strips to "file-original123".
|
||||
|
|
@ -147,6 +156,7 @@ class Harness:
|
|||
router: MagicMock
|
||||
logging: MagicMock
|
||||
creds_resolver: MagicMock
|
||||
upstream_files_route: respx.Route
|
||||
|
||||
@property
|
||||
def router_acreate(self) -> AsyncMock:
|
||||
|
|
@ -162,13 +172,14 @@ class Harness:
|
|||
return dict(self.router_acreate.call_args.kwargs)
|
||||
|
||||
|
||||
def _creds_lookup(*, model_id: str) -> Dict[str, str]:
|
||||
# KeyError on an unknown/hardcoded model_id - the bug cannot hide.
|
||||
return dict(CREDS[model_id])
|
||||
def _creds_lookup(*, model_id: str, team_id: str | None = None) -> dict[str, str] | None:
|
||||
# An unknown/hardcoded model_id resolves to None exactly like the real router,
|
||||
# which the endpoint turns into a 400 and a missing dispatch - the bug cannot hide.
|
||||
return dict(CREDS[model_id]) if model_id in CREDS else None
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def harness():
|
||||
def harness(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Seam harness. Patches only true I/O boundaries; pure encode/decode/merge
|
||||
helpers run for real. Object mocks are spec'd so unknown method calls raise."""
|
||||
body_holder: Dict[str, Any] = {}
|
||||
|
|
@ -192,6 +203,7 @@ def harness():
|
|||
provider_from_headers = MagicMock(return_value=None)
|
||||
is_known_model = MagicMock(return_value=False)
|
||||
litellm_acreate = AsyncMock(return_value=make_batch())
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch.object(endpoints, "_read_request_body", read_body))
|
||||
|
|
@ -213,6 +225,10 @@ def harness():
|
|||
stack.enter_context(patch.object(endpoints, "is_known_model", is_known_model))
|
||||
stack.enter_context(patch.object(litellm, "acreate_batch", litellm_acreate))
|
||||
stack.enter_context(patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False))
|
||||
upstream = stack.enter_context(respx.mock(assert_all_called=False))
|
||||
upstream_files_route = upstream.get(f"{CREDS['my-vllm']['api_base']}/files").mock(
|
||||
return_value=httpx.Response(404, json={"detail": "Not Found"})
|
||||
)
|
||||
stack.enter_context(patch.object(proxy_server, "llm_router", router))
|
||||
stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging))
|
||||
stack.enter_context(patch.object(proxy_server, "general_settings", {}))
|
||||
|
|
@ -231,6 +247,7 @@ def harness():
|
|||
router=router,
|
||||
logging=logging,
|
||||
creds_resolver=router.get_deployment_credentials_with_provider,
|
||||
upstream_files_route=upstream_files_route,
|
||||
)
|
||||
yield h
|
||||
|
||||
|
|
@ -255,6 +272,25 @@ async def call_create(
|
|||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def executed_runner():
|
||||
runner = MagicMock(spec=endpoints.LiteLLMExecutedBatchRunner)
|
||||
runner.create = AsyncMock(return_value=make_batch(id="litellm-executed-batch"))
|
||||
runner.cancel = AsyncMock(return_value=make_batch(id="litellm-executed-batch", status="cancelling"))
|
||||
factory = MagicMock(return_value=runner)
|
||||
with patch.object( # test-quality-ok: the route builds its runner from proxy_server globals; the factory is the only seam
|
||||
endpoints, "_litellm_executed_batch_runner", factory
|
||||
):
|
||||
yield runner, factory
|
||||
|
||||
|
||||
def _managed_input_file_id(model: str) -> str:
|
||||
unified = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
|
||||
"application/jsonl", "managed-id", model, "file-id", "file-model-id"
|
||||
)
|
||||
return base64.urlsafe_b64encode(unified.encode()).decode().rstrip("=")
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# SCENARIO 1 - input_file_id encoded with model. The full showcase: every
|
||||
# assertion type from the design lives here.
|
||||
|
|
@ -766,6 +802,137 @@ async def test_create__unified_file_id_legacy_row_without_storage_url_dispatches
|
|||
assert harness.router_kwargs()["input_file_id"] == "litellm_proxy_unified_id"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# LiteLLM-executed batches: a unified file targeting a provider whose API has
|
||||
# no /v1/batches (hosted_vllm) runs inside LiteLLM instead of being forwarded.
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_executed_provider_runs_inside_litellm(harness, executed_runner):
|
||||
runner, factory = executed_runner
|
||||
caller = UserAPIKeyAuth(api_key="sk-test", team_id="team-vllm")
|
||||
input_file_id = _managed_input_file_id("my-vllm")
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": input_file_id,
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
"litellm_metadata": {"tags": ["batch-tag"]},
|
||||
},
|
||||
)
|
||||
resp = await call_create(harness, user=caller)
|
||||
|
||||
harness.router_acreate.assert_not_called()
|
||||
harness.litellm_acreate.assert_not_called()
|
||||
harness.creds_resolver.assert_called_once_with(model_id="my-vllm", team_id="team-vllm")
|
||||
factory.assert_called_once_with(harness.router, harness.logging)
|
||||
runner.create.assert_awaited_once()
|
||||
create_kwargs = runner.create.call_args.kwargs
|
||||
assert create_kwargs["unified_input_file_id"] == input_file_id
|
||||
assert create_kwargs["model"] == "my-vllm"
|
||||
assert create_kwargs["provider"] == "hosted_vllm"
|
||||
assert create_kwargs["request_tags"] == ("batch-tag",)
|
||||
assert create_kwargs["user_api_key_dict"] is caller
|
||||
assert create_kwargs["create_request"]["model"] == "my-vllm"
|
||||
assert resp.id == "litellm-executed-batch"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_executed_provider_without_database_400(harness):
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": _managed_input_file_id("my-vllm"),
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_create(harness)
|
||||
|
||||
assert exc.value.code == "400"
|
||||
assert "need a database" in exc.value.message
|
||||
harness.router_acreate.assert_not_called()
|
||||
harness.litellm_acreate.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_executed_provider_with_its_own_files_api_goes_to_the_provider(harness, executed_runner):
|
||||
runner, factory = executed_runner
|
||||
harness.upstream_files_route.mock(return_value=httpx.Response(200, json={"object": "list", "data": []}))
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": _managed_input_file_id("my-vllm"),
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
await call_create(harness)
|
||||
|
||||
factory.assert_not_called()
|
||||
runner.create.assert_not_called()
|
||||
assert harness.router_kwargs()["model"] == "my-vllm"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_provider_model_never_touches_executed_runner(harness, executed_runner):
|
||||
runner, factory = executed_runner
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": _managed_input_file_id("azure/gpt-4o"),
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
await call_create(harness)
|
||||
|
||||
factory.assert_not_called()
|
||||
runner.create.assert_not_called()
|
||||
harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o", team_id=None)
|
||||
assert harness.router_kwargs()["model"] == "azure/gpt-4o"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("via", ["body", "header"])
|
||||
async def test_create__raw_file_with_executed_model_400_with_upload_guidance(harness, via):
|
||||
body = {"input_file_id": "file-plain", "endpoint": "/v1/chat/completions", "completion_window": "24h"}
|
||||
set_body(harness, {**body, "model": "my-vllm"} if via == "body" else body)
|
||||
headers = {"x-litellm-model": "my-vllm"} if via == "header" else None
|
||||
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_create(harness, headers=headers)
|
||||
|
||||
assert exc.value.code == "400"
|
||||
assert "POST /v1/files" in exc.value.message
|
||||
assert "x-litellm-model" in exc.value.message
|
||||
harness.litellm_acreate.assert_not_called()
|
||||
harness.router_acreate.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"upstream_answer",
|
||||
[httpx.Response(200, json={"object": "list", "data": []}), httpx.Response(405), httpx.ConnectError("refused")],
|
||||
ids=["lists files", "files route without list", "unreachable"],
|
||||
)
|
||||
async def test_create__raw_file_with_executed_model_is_forwarded_unless_the_server_lacks_a_files_api(
|
||||
harness, upstream_answer
|
||||
):
|
||||
harness.upstream_files_route.mock(side_effect=[upstream_answer])
|
||||
set_body(harness, {"input_file_id": "file-plain", "endpoint": "/v1/chat/completions", "completion_window": "24h"})
|
||||
|
||||
await call_create(harness, headers={"x-litellm-model": "my-vllm"})
|
||||
|
||||
forwarded = harness.acreate_kwargs()
|
||||
assert forwarded["input_file_id"] == "file-plain"
|
||||
assert forwarded["custom_llm_provider"] == "hosted_vllm"
|
||||
assert forwarded["api_base"] == CREDS["my-vllm"]["api_base"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__model_encoded_beats_unified(harness):
|
||||
"""Precedence row: a file id that is BOTH model-encoded and (pretend) unified
|
||||
|
|
@ -1146,6 +1313,11 @@ AZURE_BATCH_ID = encode_file_id_with_model("batch_orig123", "azure/gpt-4o", id_t
|
|||
# returns). model_id / llm_batch_id are parsed out of this by the real helpers.
|
||||
UNIFIED_BATCH_ID = "litellm_proxy;model_id:gpt-4o-mini;llm_batch_id:batch-raw-xyz"
|
||||
|
||||
# A decoded unified id of a batch LiteLLM runs itself: the llm_batch_id carries
|
||||
# the litellm_batch_ prefix, so no provider holds a batch to sync with.
|
||||
EXECUTED_BATCH_ID = "litellm_proxy;model_id:my-vllm;llm_batch_id:litellm_batch_abc"
|
||||
EXECUTED_BATCH_B64 = base64.urlsafe_b64encode(EXECUTED_BATCH_ID.encode()).decode().rstrip("=")
|
||||
|
||||
|
||||
@dataclass
|
||||
class RetrieveHarness:
|
||||
|
|
@ -1580,6 +1752,69 @@ async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harn
|
|||
assert retrieve_harness.update_batch_in_db.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("status", ["validating", "in_progress", "finalizing", "cancelling"])
|
||||
async def test_retrieve__executed_batch_served_from_db_in_every_status(retrieve_harness, status):
|
||||
db_response = make_batch(id="litellm-executed-batch", status=status)
|
||||
db_batch_object = MagicMock()
|
||||
db_batch_object.updated_at = datetime.now(timezone.utc)
|
||||
retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response)
|
||||
|
||||
resp = await call_retrieve(retrieve_harness, EXECUTED_BATCH_B64)
|
||||
|
||||
assert resp is db_response
|
||||
retrieve_harness.litellm_aretrieve.assert_not_called()
|
||||
retrieve_harness.router_aretrieve.assert_not_called()
|
||||
retrieve_harness.update_batch_in_db.assert_not_called()
|
||||
retrieve_harness.ensure_managed_files.assert_called_once()
|
||||
assert retrieve_harness.ensure_managed_files.call_args.kwargs["unified_batch_id"] == EXECUTED_BATCH_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve__executed_batch_abandoned_by_its_runner_is_served_failed(retrieve_harness, executed_runner):
|
||||
runner, _ = executed_runner
|
||||
failed = make_batch(id="litellm-executed-batch", status="failed")
|
||||
runner.fail_abandoned = AsyncMock(return_value=failed)
|
||||
db_response = make_batch(id="litellm-executed-batch", status="in_progress")
|
||||
db_batch_object = MagicMock()
|
||||
db_batch_object.updated_at = datetime.now(timezone.utc) - timedelta(minutes=10)
|
||||
retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response)
|
||||
user = UserAPIKeyAuth(api_key="sk-test", user_id="user-1")
|
||||
|
||||
resp = await call_retrieve(retrieve_harness, EXECUTED_BATCH_B64, user=user)
|
||||
|
||||
assert resp is failed
|
||||
runner.fail_abandoned.assert_awaited_once_with(db_response, user)
|
||||
retrieve_harness.litellm_aretrieve.assert_not_called()
|
||||
retrieve_harness.router_aretrieve.assert_not_called()
|
||||
retrieve_harness.ensure_managed_files.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve__executed_batch_with_a_fresh_heartbeat_is_left_running(retrieve_harness, executed_runner):
|
||||
runner, _ = executed_runner
|
||||
runner.fail_abandoned = AsyncMock()
|
||||
db_response = make_batch(id="litellm-executed-batch", status="in_progress")
|
||||
db_batch_object = MagicMock()
|
||||
db_batch_object.updated_at = datetime.now(timezone.utc) - timedelta(seconds=30)
|
||||
retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response)
|
||||
|
||||
resp = await call_retrieve(retrieve_harness, EXECUTED_BATCH_B64)
|
||||
|
||||
assert resp is db_response
|
||||
runner.fail_abandoned.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve__executed_batch_without_db_row_404(retrieve_harness):
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_retrieve(retrieve_harness, EXECUTED_BATCH_B64)
|
||||
|
||||
assert exc.value.code == "404"
|
||||
retrieve_harness.litellm_aretrieve.assert_not_called()
|
||||
retrieve_harness.router_aretrieve.assert_not_called()
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Cross-cutting: enrichment route_type and failure-hook on provider error.
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
|
@ -2299,6 +2534,35 @@ async def test_cancel__unified_no_router_500(cancel_harness):
|
|||
assert exc.value.code == "500"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel__executed_batch_routes_to_runner(cancel_harness, executed_runner):
|
||||
runner, factory = executed_runner
|
||||
caller = UserAPIKeyAuth(api_key="sk-test", user_id="user-cancel-2")
|
||||
resp = await call_cancel(cancel_harness, EXECUTED_BATCH_B64, user=caller)
|
||||
|
||||
runner.cancel.assert_awaited_once_with(EXECUTED_BATCH_B64, caller)
|
||||
factory.assert_called_once_with(cancel_harness.router, cancel_harness.logging)
|
||||
cancel_harness.router_acancel.assert_not_called()
|
||||
cancel_harness.litellm_acancel.assert_not_called()
|
||||
cancel_harness.creds_resolver.assert_not_called()
|
||||
assert resp is runner.cancel.return_value
|
||||
assert cancel_harness.update_batch_in_db.call_args.kwargs["operation"] == "cancel"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel__executed_batch_no_router_500(cancel_harness, executed_runner):
|
||||
runner, factory = executed_runner
|
||||
with patch.object( # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
proxy_server, "llm_router", None
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_cancel(cancel_harness, EXECUTED_BATCH_B64)
|
||||
|
||||
assert exc.value.code == "500"
|
||||
factory.assert_not_called()
|
||||
runner.cancel.assert_not_called()
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# SCENARIO 3 - fallback to custom_llm_provider. Rebuilds a CancelBatchRequest
|
||||
# and forwards only {custom_llm_provider, batch_id}.
|
||||
|
|
@ -2956,3 +3220,16 @@ async def test_cancel__unified_batch_id_rejects_key_without_model_grant(cancel_h
|
|||
|
||||
assert exc_info.value.code == "403"
|
||||
cancel_harness.router_acancel.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel__executed_batch_rejects_key_without_model_grant(cancel_harness, executed_runner):
|
||||
runner, factory = executed_runner
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await call_cancel(cancel_harness, EXECUTED_BATCH_B64, user=_key_restricted_to("vertex-model"))
|
||||
|
||||
assert exc_info.value.code == "403"
|
||||
factory.assert_not_called()
|
||||
runner.cancel.assert_not_called()
|
||||
cancel_harness.router_acancel.assert_not_called()
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -4,10 +4,10 @@ from unittest.mock import AsyncMock, MagicMock
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
apply_unified_file_ids,
|
||||
get_credentials_for_model,
|
||||
is_litellm_executed_batch,
|
||||
map_raw_file_ids_to_unified,
|
||||
)
|
||||
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
|
||||
|
|
@ -500,3 +500,17 @@ class TestCompletedBatchSafeToRetire:
|
|||
|
||||
def test_no_output_and_unknown_counts_is_not_safe(self):
|
||||
assert _completed_batch_safe_to_retire(_completed_batch_for_retire(None)) is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"decoded_unified_batch_id, executed",
|
||||
[
|
||||
("litellm_proxy;model_id:my-vllm;llm_batch_id:litellm_batch_0123abcd", True),
|
||||
("litellm_proxy;model_id:my-vllm;llm_batch_id:batch_0123abcd", False),
|
||||
("litellm_proxy;model_id:my-vllm;generic_response_id:resp_0123abcd", False),
|
||||
("litellm_proxy;model_id:my-vllm;llm_output_file_id:file-0123abcd", False),
|
||||
("batch_0123abcd", False),
|
||||
],
|
||||
)
|
||||
def test_is_litellm_executed_batch_reads_the_llm_batch_id_prefix(decoded_unified_batch_id: str, executed: bool):
|
||||
assert is_litellm_executed_batch(decoded_unified_batch_id) is executed
|
||||
|
|
|
|||
|
|
@ -609,6 +609,246 @@ def test_target_storage_with_target_models(
|
|||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
BATCH_JSONL_LINE = (
|
||||
b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions", '
|
||||
b'"body": {"model": "my-vllm", "messages": [{"role": "user", "content": "hi"}]}}\n'
|
||||
)
|
||||
|
||||
|
||||
def _router_with_executed_batch_model() -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "my-vllm",
|
||||
"litellm_params": {
|
||||
"model": "hosted_vllm/qwen",
|
||||
"api_key": "sk-vllm",
|
||||
"api_base": "http://vllm.test/v1",
|
||||
},
|
||||
"model_info": {"id": "my-vllm-id"},
|
||||
},
|
||||
{
|
||||
"model_name": "gemini-2.0-flash",
|
||||
"litellm_params": {"model": "gemini/gemini-2.0-flash"},
|
||||
"model_info": {"id": "gemini-2.0-flash-id"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def batch_upload_seams(mocker: MockerFixture, monkeypatch):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
llm_router = _router_with_executed_batch_model()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
|
||||
setup_proxy_logging_object(monkeypatch, llm_router)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"
|
||||
)
|
||||
uploaded = OpenAIFileObject(
|
||||
id="file-kept",
|
||||
object="file",
|
||||
purpose="batch",
|
||||
created_at=0,
|
||||
bytes=len(BATCH_JSONL_LINE),
|
||||
filename="batch.jsonl",
|
||||
status="uploaded",
|
||||
)
|
||||
stored = mocker.patch( # test-quality-ok: the route calls the storage service directly with no injection seam
|
||||
"litellm.proxy.openai_files_endpoints.storage_backend_service.StorageBackendFileService.upload_file_to_storage_backend",
|
||||
new=mocker.AsyncMock(return_value=uploaded),
|
||||
)
|
||||
provider_upload = mocker.patch( # test-quality-ok: the route calls litellm.acreate_file directly with no injection seam
|
||||
"litellm.acreate_file", new=mocker.AsyncMock(return_value=uploaded)
|
||||
)
|
||||
try:
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
upstream_files_route = upstream.get("http://vllm.test/v1/files").mock(
|
||||
return_value=httpx.Response(404, json={"detail": "Not Found"})
|
||||
)
|
||||
yield stored, provider_upload, upstream_files_route
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def _upload_batch_file(headers: dict[str, str], form: dict[str, str]):
|
||||
return client.post(
|
||||
"/v1/files",
|
||||
files={"file": ("batch.jsonl", BATCH_JSONL_LINE, "application/jsonl")},
|
||||
data={"purpose": "batch", **form},
|
||||
headers={"Authorization": "Bearer test-key", **headers},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"headers, form",
|
||||
[({"x-litellm-model": "my-vllm"}, {}), ({}, {"target_model_names": "my-vllm"})],
|
||||
ids=["x-litellm-model header", "target_model_names form field"],
|
||||
)
|
||||
def test_batch_upload_for_a_litellm_executed_model_is_kept_by_litellm(
|
||||
batch_upload_seams, headers: dict[str, str], form: dict[str, str]
|
||||
):
|
||||
stored, provider_upload, _ = batch_upload_seams
|
||||
|
||||
response = _upload_batch_file(headers, form)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
provider_upload.assert_not_awaited()
|
||||
stored.assert_awaited_once()
|
||||
kwargs = stored.call_args.kwargs
|
||||
assert kwargs["target_storage"] == "litellm_db"
|
||||
assert tuple(kwargs["target_model_names"]) == ("my-vllm",)
|
||||
assert kwargs["purpose"] == "batch"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"headers, form",
|
||||
[({"x-litellm-model": "my-vllm"}, {}), ({}, {"target_model_names": "my-vllm"})],
|
||||
ids=["x-litellm-model header", "target_model_names form field"],
|
||||
)
|
||||
def test_batch_upload_for_a_litellm_executed_model_the_key_cannot_call_is_refused_before_the_server_is_probed(
|
||||
batch_upload_seams, headers: dict[str, str], form: dict[str, str]
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
stored, provider_upload, upstream_files_route = batch_upload_seams
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="restricted-user", models=["gemini-2.0-flash"]
|
||||
)
|
||||
|
||||
response = _upload_batch_file(headers, form)
|
||||
|
||||
assert response.status_code == 403, response.text
|
||||
assert "my-vllm" in response.text
|
||||
assert upstream_files_route.call_count == 0
|
||||
stored.assert_not_awaited()
|
||||
provider_upload.assert_not_awaited()
|
||||
|
||||
|
||||
def test_batch_upload_naming_an_executed_and_a_provider_model_is_rejected(batch_upload_seams):
|
||||
stored, provider_upload, _ = batch_upload_seams
|
||||
|
||||
response = _upload_batch_file({}, {"target_model_names": "my-vllm,gemini-2.0-flash"})
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
assert "my-vllm" in response.text
|
||||
assert "target_model_names" in response.text
|
||||
stored.assert_not_awaited()
|
||||
provider_upload.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("purpose", ["assistants", "user_data"])
|
||||
def test_non_batch_upload_for_a_litellm_executed_model_is_rejected_with_the_purpose_to_use(
|
||||
batch_upload_seams, purpose: str
|
||||
):
|
||||
stored, provider_upload, _ = batch_upload_seams
|
||||
|
||||
response = _upload_batch_file({"x-litellm-model": "my-vllm"}, {"purpose": purpose})
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
error = response.json()["error"]
|
||||
assert error["type"] == "invalid_request_error"
|
||||
assert error["param"] == "purpose"
|
||||
assert "purpose=batch" in error["message"]
|
||||
assert f"purpose={purpose}" in error["message"]
|
||||
stored.assert_not_awaited()
|
||||
provider_upload.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("purpose", ["batch", "assistants"])
|
||||
@pytest.mark.parametrize(
|
||||
"upstream_answer",
|
||||
[httpx.Response(200, json={"object": "list", "data": []}), httpx.Response(405), httpx.ConnectError("refused")],
|
||||
ids=["lists files", "files route without list", "unreachable"],
|
||||
)
|
||||
def test_upload_for_a_litellm_executed_model_goes_to_the_provider_unless_the_server_lacks_a_files_api(
|
||||
batch_upload_seams, upstream_answer: httpx.Response | httpx.ConnectError, purpose: str
|
||||
):
|
||||
stored, provider_upload, upstream_files_route = batch_upload_seams
|
||||
upstream_files_route.mock(side_effect=[upstream_answer])
|
||||
|
||||
response = _upload_batch_file({"x-litellm-model": "my-vllm"}, {"purpose": purpose})
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
stored.assert_not_awaited()
|
||||
provider_upload.assert_awaited_once()
|
||||
assert provider_upload.call_args.kwargs["custom_llm_provider"] == "hosted_vllm"
|
||||
assert provider_upload.call_args.kwargs["api_base"] == "http://vllm.test/v1"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"form",
|
||||
[{}, {"target_model_names": "my-vllm"}, {"target_model_names": "gemini-2.0-flash"}],
|
||||
ids=["no model", "litellm-executed model", "provider model"],
|
||||
)
|
||||
def test_upload_naming_litellm_db_as_target_storage_is_rejected(batch_upload_seams, form: dict[str, str]):
|
||||
stored, provider_upload, upstream_files_route = batch_upload_seams
|
||||
|
||||
response = _upload_batch_file({}, {**form, "target_storage": "litellm_db"})
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
error = response.json()["error"]
|
||||
assert error["type"] == "invalid_request_error"
|
||||
assert error["param"] == "target_storage"
|
||||
assert "litellm_db" in error["message"]
|
||||
assert upstream_files_route.call_count == 0
|
||||
stored.assert_not_awaited()
|
||||
provider_upload.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("purpose", ["user_data", "batch"])
|
||||
def test_upload_with_an_explicit_target_storage_goes_where_the_caller_said_without_probing_the_server(
|
||||
batch_upload_seams, purpose: str
|
||||
):
|
||||
stored, provider_upload, upstream_files_route = batch_upload_seams
|
||||
|
||||
response = _upload_batch_file(
|
||||
{}, {"purpose": purpose, "target_model_names": "my-vllm", "target_storage": "azure_storage"}
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert upstream_files_route.call_count == 0
|
||||
provider_upload.assert_not_awaited()
|
||||
stored.assert_awaited_once()
|
||||
kwargs = stored.call_args.kwargs
|
||||
assert kwargs["target_storage"] == "azure_storage"
|
||||
assert tuple(kwargs["target_model_names"]) == ("my-vllm",)
|
||||
assert kwargs["purpose"] == purpose
|
||||
|
||||
|
||||
def test_upload_with_an_explicit_target_storage_still_refuses_a_key_without_the_executed_model(batch_upload_seams):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
stored, provider_upload, upstream_files_route = batch_upload_seams
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="restricted-user", models=["gemini-2.0-flash"]
|
||||
)
|
||||
|
||||
response = _upload_batch_file({}, {"target_model_names": "my-vllm", "target_storage": "azure_storage"})
|
||||
|
||||
assert response.status_code == 403, response.text
|
||||
assert "my-vllm" in response.text
|
||||
assert upstream_files_route.call_count == 0
|
||||
stored.assert_not_awaited()
|
||||
provider_upload.assert_not_awaited()
|
||||
|
||||
|
||||
def test_batch_upload_for_a_provider_model_still_goes_to_the_provider(batch_upload_seams):
|
||||
stored, provider_upload, _ = batch_upload_seams
|
||||
|
||||
response = _upload_batch_file({"x-litellm-model": "gemini-2.0-flash"}, {})
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
stored.assert_not_awaited()
|
||||
provider_upload.assert_awaited_once()
|
||||
assert provider_upload.call_args.kwargs["custom_llm_provider"] == "gemini"
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="mock respx fails on ci/cd - unclear why")
|
||||
def test_create_file_and_call_chat_completion_e2e(
|
||||
mocker: MockerFixture, monkeypatch, llm_router: Router
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
|
||||
|
|
@ -6,16 +8,24 @@ from litellm.proxy.openai_files_endpoints import storage_backend_service
|
|||
from litellm.proxy.openai_files_endpoints.storage_backend_service import (
|
||||
StorageBackendFileService,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
class _RecordingStorageBackend:
|
||||
def __init__(self):
|
||||
def __init__(self, delete_error: Exception | None = None):
|
||||
self.upload_calls = []
|
||||
self.delete_calls: list[str] = []
|
||||
self.delete_error = delete_error
|
||||
|
||||
async def upload_file(self, **kwargs):
|
||||
self.upload_calls.append(kwargs)
|
||||
return "https://storage.example/blob-1"
|
||||
|
||||
async def delete_file(self, storage_url: str) -> None:
|
||||
self.delete_calls.append(storage_url)
|
||||
if self.delete_error is not None:
|
||||
raise self.delete_error
|
||||
|
||||
|
||||
class _FakeManagedFilesHook(BaseFileEndpoints):
|
||||
def __init__(self):
|
||||
|
|
@ -42,6 +52,11 @@ class _FakeManagedFilesHook(BaseFileEndpoints):
|
|||
self.stored.append(kwargs)
|
||||
|
||||
|
||||
class _FailingManagedFilesHook(_FakeManagedFilesHook):
|
||||
async def store_unified_file_id(self, **kwargs):
|
||||
raise RuntimeError("db down")
|
||||
|
||||
|
||||
class _FakeProxyLogging:
|
||||
def __init__(self, hook):
|
||||
self._hook = hook
|
||||
|
|
@ -57,7 +72,7 @@ def _file_data():
|
|||
@pytest.mark.asyncio
|
||||
async def test_upload_with_target_model_names_but_no_hook_raises_before_uploading(monkeypatch):
|
||||
backend = _RecordingStorageBackend()
|
||||
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend)
|
||||
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name, prisma_client=None: backend)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await StorageBackendFileService.upload_file_to_storage_backend(
|
||||
|
|
@ -80,7 +95,7 @@ async def test_upload_with_target_model_names_but_no_hook_raises_before_uploadin
|
|||
@pytest.mark.asyncio
|
||||
async def test_upload_without_target_model_names_skips_hook_requirement(monkeypatch):
|
||||
backend = _RecordingStorageBackend()
|
||||
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend)
|
||||
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name, prisma_client=None: backend)
|
||||
|
||||
file_object = await StorageBackendFileService.upload_file_to_storage_backend(
|
||||
file_data=_file_data(),
|
||||
|
|
@ -101,7 +116,7 @@ async def test_upload_without_target_model_names_skips_hook_requirement(monkeypa
|
|||
@pytest.mark.asyncio
|
||||
async def test_upload_with_target_model_names_and_hook_stores_unified_id(monkeypatch):
|
||||
backend = _RecordingStorageBackend()
|
||||
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend)
|
||||
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name, prisma_client=None: backend)
|
||||
hook = _FakeManagedFilesHook()
|
||||
|
||||
file_object = await StorageBackendFileService.upload_file_to_storage_backend(
|
||||
|
|
@ -125,3 +140,50 @@ async def test_upload_with_target_model_names_and_hook_stores_unified_id(monkeyp
|
|||
"stored_id_matches_response": True,
|
||||
"model_mappings": {"gpt-x": "https://storage.example/blob-1"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_hands_the_prisma_client_to_the_storage_backend_factory(monkeypatch: pytest.MonkeyPatch):
|
||||
backend = _RecordingStorageBackend()
|
||||
factory_calls: list[tuple[str, PrismaClient | None]] = []
|
||||
|
||||
def _factory(name: str, prisma_client: PrismaClient | None = None) -> _RecordingStorageBackend:
|
||||
factory_calls.append((name, prisma_client))
|
||||
return backend
|
||||
|
||||
monkeypatch.setattr(storage_backend_service, "get_storage_backend", _factory)
|
||||
prisma_client = MagicMock()
|
||||
|
||||
await StorageBackendFileService.upload_file_to_storage_backend(
|
||||
file_data=_file_data(),
|
||||
target_storage="litellm_db",
|
||||
target_model_names=[],
|
||||
purpose="batch",
|
||||
proxy_logging_obj=_FakeProxyLogging(hook=None),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
assert factory_calls == [("litellm_db", prisma_client)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("delete_error", [None, OSError("blob locked")], ids=["delete succeeds", "delete fails"])
|
||||
async def test_upload_deletes_the_uploaded_content_when_the_metadata_write_fails(
|
||||
monkeypatch: pytest.MonkeyPatch, delete_error: Exception | None
|
||||
):
|
||||
backend = _RecordingStorageBackend(delete_error=delete_error)
|
||||
monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name, prisma_client=None: backend)
|
||||
|
||||
with pytest.raises(RuntimeError, match="db down"):
|
||||
await StorageBackendFileService.upload_file_to_storage_backend(
|
||||
file_data=_file_data(),
|
||||
target_storage="azure_storage",
|
||||
target_model_names=["gpt-x"],
|
||||
purpose="batch",
|
||||
proxy_logging_obj=_FakeProxyLogging(hook=_FailingManagedFilesHook()),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
)
|
||||
|
||||
assert len(backend.upload_calls) == 1
|
||||
assert backend.delete_calls == ["https://storage.example/blob-1"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue