From 7b1cef86a78848c2f3dc75da6db5060218dd5919 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 11 Dec 2025 15:08:17 +0530 Subject: [PATCH] Add support for target_storage param --- .../proxy/hooks/managed_files.py | 159 ++++++++- .../migration.sql | 4 + .../files/azure_blob_storage_backend.py | 312 ++++++++++++++++++ .../llms/base_llm/files/storage_backend.py | 79 +++++ .../base_llm/files/storage_backend_factory.py | 41 +++ litellm/proxy/_types.py | 2 + .../openai_files_endpoints/common_utils.py | 296 +++++++++++++++++ .../openai_files_endpoints/files_endpoints.py | 89 ++++- .../storage_backend_service.py | 244 ++++++++++++++ schema.prisma | 2 + .../test_files_endpoint.py | 92 ++++++ 11 files changed, 1300 insertions(+), 20 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20251210125210_add_storage_backend_to_managed_files/migration.sql create mode 100644 litellm/llms/base_llm/files/azure_blob_storage_backend.py create mode 100644 litellm/llms/base_llm/files/storage_backend.py create mode 100644 litellm/llms/base_llm/files/storage_backend_factory.py create mode 100644 litellm/proxy/openai_files_endpoints/storage_backend_service.py diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 608bb49588..6620db5ffa 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -22,7 +22,6 @@ from litellm.proxy._types import ( ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, - convert_b64_uid_to_unified_uid, get_batch_id_from_unified_batch_id, get_model_id_from_unified_batch_id, ) @@ -42,6 +41,10 @@ from litellm.types.utils import ( LLMResponseTypes, SpecialEnums, ) +from litellm.proxy.openai_files_endpoints.common_utils import ( + get_content_type_from_file_object, + normalize_mime_type_for_provider, +) if TYPE_CHECKING: from litellm.types.llms.openai import HttpxBinaryResponseContent @@ -108,6 +111,17 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if file_object is not None: db_data["file_object"] = file_object.model_dump_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"] + if "storage_url" in hidden_params: + db_data["storage_url"] = hidden_params["storage_url"] + + verbose_logger.debug( + f"Storage metadata: storage_backend={db_data.get('storage_backend')}, " + f"storage_url={db_data.get('storage_url')}" + ) result = await self.prisma_client.db.litellm_managedfiletable.create( data=db_data @@ -268,7 +282,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) return False - async def async_pre_call_hook( + async def async_pre_call_hook( # noqa: PLR0915 self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, @@ -287,15 +301,31 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): await self.check_managed_file_id_access(data, user_api_key_dict) ### HANDLE TRANSFORMATIONS ### - if call_type == CallTypes.completion.value: + # Check both completion and acompletion call types + is_completion_call = ( + call_type == CallTypes.completion.value + or call_type == CallTypes.acompletion.value + ) + + if is_completion_call: messages = data.get("messages") + model = data.get("model", "") if messages: file_ids = self.get_file_ids_from_messages(messages) if file_ids: + # Check if any files are stored in storage backends and need base64 conversion + # This is needed for Vertex AI/Gemini which requires base64 content + is_vertex_ai = model and ("vertex_ai" in model or "gemini" in model.lower()) + if is_vertex_ai: + await self._convert_storage_files_to_base64( + messages=messages, + file_ids=file_ids, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + model_file_id_mapping = await self.get_model_file_id_mapping( file_ids, user_api_key_dict.parent_otel_span ) - data["model_file_id_mapping"] = model_file_id_mapping elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value: # Handle managed files in responses API input @@ -865,3 +895,124 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) else: raise Exception(f"LiteLLM Managed File object with id={file_id} not found") + + async def _convert_storage_files_to_base64( + self, + messages: List[AllMessageValues], + file_ids: List[str], + litellm_parent_otel_span: Optional[Span], + ) -> None: + """ + Convert files stored in storage backends to base64 format for Vertex AI/Gemini. + + This method checks if any managed files are stored in storage backends, + downloads them, and converts them to base64 format in the messages. + """ + # Check each file_id to see if it's stored in a storage backend + for file_id in file_ids: + # Check if this is a base64 encoded unified file ID + decoded_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + + if not decoded_unified_file_id: + continue + + # Check database for storage backend info + # IMPORTANT: The database stores the base64 encoded unified_file_id (not the decoded version) + # So we query with the original file_id (which is base64 encoded) + db_file = await self.prisma_client.db.litellm_managedfiletable.find_first( + where={"unified_file_id": file_id} + ) + + if not db_file or not db_file.storage_backend or not db_file.storage_url: + continue + + # 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) + except ValueError as e: + verbose_logger.warning( + f"Storage backend '{storage_backend_name}' error for file {file_id}: {str(e)}" + ) + continue + + file_content = await storage_backend.download_file(storage_url) + + # Determine content type from file object + content_type = self._get_content_type_from_file_object(db_file.file_object) + + # Convert to base64 + base64_data = base64.b64encode(file_content).decode("utf-8") + base64_data_uri = f"data:{content_type};base64,{base64_data}" + + # Update messages to use base64 instead of file_id + self._update_messages_with_base64_data(messages, file_id, base64_data_uri, content_type) + except Exception as e: + verbose_logger.exception( + f"Error converting file {file_id} from storage backend to base64: {str(e)}" + ) + # Continue with other files even if one fails + continue + + def _get_content_type_from_file_object(self, file_object: Optional[Any]) -> str: + """ + Determine content type from file object. + + Uses the MIME type utility for consistent detection and normalization. + + Args: + file_object: The file object from the database (can be dict, JSON string, or None) + + Returns: + str: MIME type (defaults to "application/octet-stream" if cannot be determined) + """ + # Use utility function for detection + content_type = get_content_type_from_file_object(file_object) + + # Normalize for Gemini/Vertex AI (requires image/jpeg, not image/jpg) + content_type = normalize_mime_type_for_provider(content_type, provider="gemini") + + return content_type + + def _update_messages_with_base64_data( + self, + messages: List[AllMessageValues], + file_id: str, + base64_data_uri: str, + content_type: str, + ) -> None: + """ + Update messages to replace file_id with base64 data URI. + + Args: + messages: List of messages to update + file_id: The file ID to replace + base64_data_uri: The base64 data URI to use as replacement + content_type: The MIME type of the file (e.g., "image/jpeg", "application/pdf") + """ + for message in messages: + if message.get("role") == "user": + content = message.get("content") + if content and isinstance(content, list): + for element in content: + if element.get("type") == "file": + file_element = cast(ChatCompletionFileObject, element) + file_element_file = file_element.get("file", {}) + + if file_element_file.get("file_id") == file_id: + # Replace file_id with base64 data + file_element_file["file_data"] = base64_data_uri + # Set format to help Gemini determine mime type + file_element_file["format"] = content_type + # Remove file_id to ensure only file_data is used + file_element_file.pop("file_id", None) + + verbose_logger.debug( + f"Converted file {file_id} from storage backend to base64 with format {content_type}" + ) diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20251210125210_add_storage_backend_to_managed_files/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20251210125210_add_storage_backend_to_managed_files/migration.sql new file mode 100644 index 0000000000..26f8d31d27 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20251210125210_add_storage_backend_to_managed_files/migration.sql @@ -0,0 +1,4 @@ +-- AlterTable +ALTER TABLE "LiteLLM_ManagedFileTable" ADD COLUMN IF NOT EXISTS "storage_backend" TEXT; +ALTER TABLE "LiteLLM_ManagedFileTable" ADD COLUMN IF NOT EXISTS "storage_url" TEXT; + diff --git a/litellm/llms/base_llm/files/azure_blob_storage_backend.py b/litellm/llms/base_llm/files/azure_blob_storage_backend.py new file mode 100644 index 0000000000..db3aa50d89 --- /dev/null +++ b/litellm/llms/base_llm/files/azure_blob_storage_backend.py @@ -0,0 +1,312 @@ +""" +Azure Blob Storage backend implementation for file storage. + +This module implements the Azure Blob Storage backend for storing files +in Azure Data Lake Storage Gen2. It inherits from AzureBlobStorageLogger +to reuse all authentication and Azure Storage operations. +""" + +import time +from typing import Optional +from urllib.parse import quote + +from litellm._logging import verbose_logger +from litellm._uuid import uuid + +from .storage_backend import BaseFileStorageBackend +from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLogger + + +class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): + """ + Azure Blob Storage backend implementation. + + Inherits from AzureBlobStorageLogger to reuse: + - Authentication (account key and Azure AD) + - Service client management + - Token management + - All Azure Storage helper methods + + Reads configuration from the same environment variables as AzureBlobStorageLogger. + """ + + def __init__(self, **kwargs): + """ + Initialize Azure Blob Storage backend. + + Inherits all functionality from AzureBlobStorageLogger which handles: + - Reading environment variables + - Authentication (account key and Azure AD) + - Service client management + - Token management + + Environment variables (same as AzureBlobStorageLogger): + - AZURE_STORAGE_ACCOUNT_NAME (required) + - AZURE_STORAGE_FILE_SYSTEM (required) + - AZURE_STORAGE_ACCOUNT_KEY (optional, if using account key auth) + - AZURE_STORAGE_TENANT_ID (optional, if using Azure AD) + - AZURE_STORAGE_CLIENT_ID (optional, if using Azure AD) + - AZURE_STORAGE_CLIENT_SECRET (optional, if using Azure AD) + + Note: We skip periodic_flush since we're not using this as a logger. + """ + # Initialize AzureBlobStorageLogger (handles all auth and config) + AzureBlobStorageLogger.__init__(self, **kwargs) + + # Disable logging functionality - we're only using this for file storage + # The periodic_flush task will be created but will do nothing since we override it + + async def periodic_flush(self): + """ + Override to do nothing - we're not using this as a logger. + This prevents the periodic flush task from doing any work. + """ + # Do nothing - this class is used for file storage, not logging + return + + async def async_log_success_event(self, *args, **kwargs): + """ + Override to do nothing - we're not using this as a logger. + """ + # Do nothing - this class is used for file storage, not logging + pass + + async def async_log_failure_event(self, *args, **kwargs): + """ + Override to do nothing - we're not using this as a logger. + """ + # Do nothing - this class is used for file storage, not logging + pass + + def _generate_file_name( + self, original_filename: str, file_naming_strategy: str + ) -> str: + """Generate file name based on naming strategy.""" + if file_naming_strategy == "original_filename": + # Use original filename, but sanitize it + return quote(original_filename, safe="") + elif file_naming_strategy == "timestamp": + # Use timestamp + extension = original_filename.split(".")[-1] if "." in original_filename else "" + timestamp = int(time.time() * 1000) # milliseconds + return f"{timestamp}.{extension}" if extension else str(timestamp) + else: # default to "uuid" + # Use UUID + extension = original_filename.split(".")[-1] if "." in original_filename else "" + file_uuid = str(uuid.uuid4()) + return f"{file_uuid}.{extension}" if extension else file_uuid + + async def upload_file( + self, + file_content: bytes, + filename: str, + content_type: str, + path_prefix: Optional[str] = None, + file_naming_strategy: str = "uuid", + ) -> str: + """ + Upload a file to Azure Blob Storage. + + Returns the blob URL in format: https://{account}.blob.core.windows.net/{container}/{path} + """ + try: + # Generate file name + file_name = self._generate_file_name(filename, file_naming_strategy) + + # Build full path + if path_prefix: + # Remove leading/trailing slashes and normalize + prefix = path_prefix.strip("/") + full_path = f"{prefix}/{file_name}" + else: + full_path = file_name + + if self.azure_storage_account_key: + # Use Azure SDK with account key (reuse logger's method) + storage_url = await self._upload_file_with_account_key( + file_content=file_content, + full_path=full_path, + ) + else: + # Use REST API with Azure AD token (reuse logger's methods) + storage_url = await self._upload_file_with_azure_ad( + file_content=file_content, + full_path=full_path, + ) + + verbose_logger.debug( + f"Successfully uploaded file to Azure Blob Storage: {storage_url}" + ) + return storage_url + + except Exception as e: + verbose_logger.exception(f"Error uploading file to Azure Blob Storage: {str(e)}") + raise + + async def _upload_file_with_account_key( + self, file_content: bytes, full_path: str + ) -> str: + """Upload file using Azure SDK with account key authentication.""" + # Reuse the logger's service client method + service_client = await self.get_service_client() + file_system_client = service_client.get_file_system_client( + file_system=self.azure_storage_file_system + ) + + # Create filesystem (container) if it doesn't exist + if not await file_system_client.exists(): + await file_system_client.create_file_system() + verbose_logger.debug(f"Created filesystem: {self.azure_storage_file_system}") + + # Extract directory and filename (similar to logger's pattern) + path_parts = full_path.split("/") + if len(path_parts) > 1: + directory_path = "/".join(path_parts[:-1]) + file_name = path_parts[-1] + + # Create directory if needed (like logger does) + directory_client = file_system_client.get_directory_client(directory_path) + if not await directory_client.exists(): + await directory_client.create_directory() + verbose_logger.debug(f"Created directory: {directory_path}") + + # Get file client from directory (same pattern as logger) + file_client = directory_client.get_file_client(file_name) + else: + # No directory, create file directly in root + file_client = file_system_client.get_file_client(full_path) + + # Create, append, and flush (same pattern as logger's upload_to_azure_data_lake_with_azure_account_key) + await file_client.create_file() + await file_client.append_data(data=file_content, offset=0, length=len(file_content)) + await file_client.flush_data(position=len(file_content), offset=0) + + # Return blob URL (not DFS URL) + blob_url = f"https://{self.azure_storage_account_name}.blob.core.windows.net/{self.azure_storage_file_system}/{full_path}" + return blob_url + + async def _upload_file_with_azure_ad( + self, file_content: bytes, full_path: str + ) -> str: + """Upload file using REST API with Azure AD authentication.""" + # Reuse the logger's token management + await self.set_valid_azure_ad_token() + + from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, + ) + + async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.LoggingCallback + ) + + # Use DFS endpoint for upload + base_url = f"https://{self.azure_storage_account_name}.dfs.core.windows.net/{self.azure_storage_file_system}/{full_path}" + + # Execute 3-step upload process: create, append, flush + # Reuse the logger's helper methods + await self._create_file(async_client, base_url) + # Append data - logger's _append_data expects string, so we create our own for bytes + await self._append_data_bytes(async_client, base_url, file_content) + await self._flush_data(async_client, base_url, len(file_content)) + + # Return blob URL (not DFS URL) + blob_url = f"https://{self.azure_storage_account_name}.blob.core.windows.net/{self.azure_storage_file_system}/{full_path}" + return blob_url + + async def _append_data_bytes( + self, client, base_url: str, file_content: bytes + ): + """Append binary data to file using REST API.""" + from litellm.constants import AZURE_STORAGE_MSFT_VERSION + + headers = { + "x-ms-version": AZURE_STORAGE_MSFT_VERSION, + "Content-Type": "application/octet-stream", + "Authorization": f"Bearer {self.azure_auth_token}", + } + response = await client.patch( + f"{base_url}?action=append&position=0", + headers=headers, + content=file_content, + ) + response.raise_for_status() + + async def download_file(self, storage_url: str) -> bytes: + """ + Download a file from Azure Blob Storage. + + Args: + storage_url: Blob URL in format: https://{account}.blob.core.windows.net/{container}/{path} + + Returns: + bytes: File content + """ + try: + # Parse blob URL to extract path + # URL format: https://{account}.blob.core.windows.net/{container}/{path} + if ".blob.core.windows.net/" not in storage_url: + raise ValueError(f"Invalid Azure Blob Storage URL: {storage_url}") + + # Extract path after container name + container_and_path = storage_url.split(".blob.core.windows.net/", 1)[1] + path_parts = container_and_path.split("/", 1) + if len(path_parts) < 2: + raise ValueError(f"Invalid Azure Blob Storage URL format: {storage_url}") + file_path = path_parts[1] # Path after container name + + if self.azure_storage_account_key: + # Use Azure SDK (reuse logger's service client) + return await self._download_file_with_account_key(file_path) + else: + # Use REST API (reuse logger's token management) + return await self._download_file_with_azure_ad(file_path) + + except Exception as e: + verbose_logger.exception(f"Error downloading file from Azure Blob Storage: {str(e)}") + raise + + async def _download_file_with_account_key(self, file_path: str) -> bytes: + """Download file using Azure SDK with account key.""" + # Reuse the logger's service client method + service_client = await self.get_service_client() + file_system_client = service_client.get_file_system_client( + file_system=self.azure_storage_file_system + ) + # Ensure filesystem exists (should already exist, but check for safety) + if not await file_system_client.exists(): + raise ValueError(f"Filesystem {self.azure_storage_file_system} does not exist") + file_client = file_system_client.get_file_client(file_path) + # Download file + download_response = await file_client.download_file() + file_content = await download_response.readall() + return file_content + + async def _download_file_with_azure_ad(self, file_path: str) -> bytes: + """Download file using REST API with Azure AD token.""" + # Reuse the logger's token management + await self.set_valid_azure_ad_token() + + from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, + ) + from litellm.constants import AZURE_STORAGE_MSFT_VERSION + + async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.LoggingCallback + ) + + # Use blob endpoint for download (simpler than DFS) + blob_url = f"https://{self.azure_storage_account_name}.blob.core.windows.net/{self.azure_storage_file_system}/{file_path}" + + headers = { + "x-ms-version": AZURE_STORAGE_MSFT_VERSION, + "Authorization": f"Bearer {self.azure_auth_token}", + } + + response = await async_client.get(blob_url, headers=headers) + response.raise_for_status() + return response.content + diff --git a/litellm/llms/base_llm/files/storage_backend.py b/litellm/llms/base_llm/files/storage_backend.py new file mode 100644 index 0000000000..d957045295 --- /dev/null +++ b/litellm/llms/base_llm/files/storage_backend.py @@ -0,0 +1,79 @@ +""" +Base storage backend interface for file storage backends. + +This module defines the abstract base class that all file storage backends +(e.g., Azure Blob Storage, S3, GCS) must implement. +""" + +from abc import ABC, abstractmethod +from typing import Optional + + +class BaseFileStorageBackend(ABC): + """ + Abstract base class for file storage backends. + + All storage backends (Azure Blob Storage, S3, GCS, etc.) must implement + these methods to provide a consistent interface for file operations. + """ + + @abstractmethod + async def upload_file( + self, + file_content: bytes, + filename: str, + content_type: str, + path_prefix: Optional[str] = None, + file_naming_strategy: str = "uuid", + ) -> str: + """ + Upload a file to the storage backend. + + Args: + file_content: The file content as bytes + filename: Original filename (may be used for naming strategy) + content_type: MIME type of the file + path_prefix: Optional path prefix for organizing files + file_naming_strategy: Strategy for naming files ("uuid", "timestamp", "original_filename") + + Returns: + str: The storage URL where the file can be accessed/downloaded + + Raises: + Exception: If upload fails + """ + pass + + @abstractmethod + async def download_file(self, storage_url: str) -> bytes: + """ + Download a file from the storage backend. + + Args: + storage_url: The storage URL returned from upload_file + + Returns: + bytes: The file content + + Raises: + Exception: If download fails + """ + pass + + async def delete_file(self, storage_url: str) -> None: + """ + Delete a file from the storage backend. + + This is optional and can be overridden by backends that support deletion. + Default implementation does nothing. + + Args: + storage_url: The storage URL of the file to delete + + Raises: + Exception: If deletion fails + """ + # Default implementation: no-op + # Backends can override if they support deletion + pass + diff --git a/litellm/llms/base_llm/files/storage_backend_factory.py b/litellm/llms/base_llm/files/storage_backend_factory.py new file mode 100644 index 0000000000..1685f3fbd2 --- /dev/null +++ b/litellm/llms/base_llm/files/storage_backend_factory.py @@ -0,0 +1,41 @@ +""" +Factory for creating storage backend instances. + +This module provides a factory function to instantiate the correct storage backend +based on the backend type. Backends use the same configuration as their corresponding +callbacks (e.g., azure_storage uses the same env vars as AzureBlobStorageLogger). +""" + +from litellm._logging import verbose_logger + +from .azure_blob_storage_backend import AzureBlobStorageBackend +from .storage_backend import BaseFileStorageBackend + + +def get_storage_backend(backend_type: str) -> 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. + + Args: + backend_type: Backend type identifier (e.g., "azure_storage") + + Returns: + BaseFileStorageBackend: Instance of the appropriate storage backend + + Raises: + ValueError: If backend_type is not supported + """ + verbose_logger.debug(f"Creating storage backend: type={backend_type}") + + if backend_type == "azure_storage": + return AzureBlobStorageBackend() + else: + raise ValueError( + f"Unsupported storage backend type: {backend_type}. " + f"Supported types: azure_storage" + ) + diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 083ac07340..b7f128a247 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3680,6 +3680,8 @@ class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase): flat_model_file_ids: List[str] created_by: Optional[str] updated_by: Optional[str] + storage_backend: Optional[str] = None + storage_url: Optional[str] = None class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index c5b58e06d4..d51336ef0b 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -1,7 +1,12 @@ import base64 +import mimetypes import re +from dataclasses import dataclass, field from typing import List, Literal, Optional, Union +from fastapi import Request + +from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.types.utils import SpecialEnums @@ -339,3 +344,294 @@ def handle_model_based_routing( # No model-based routing needed return False, None, None, None + + +# ============================================================================ +# MIME TYPE DETECTION AND NORMALIZATION +# ============================================================================ + + +# Gemini-supported image MIME types +GEMINI_SUPPORTED_IMAGE_TYPES = { + "image/png", + "image/jpeg", + "image/webp", +} + +# Gemini-supported video MIME types +GEMINI_SUPPORTED_VIDEO_TYPES = { + "video/3gpp", + "video/wmv", + "video/webm", + "video/mp4", + "video/mpg", + "video/mpegps", + "video/mpeg", + "video/quicktime", + "video/x-flv", +} + +# Gemini-supported audio MIME types +GEMINI_SUPPORTED_AUDIO_TYPES = { + "audio/webm", + "audio/wav", + "audio/pcm", + "audio/opus", + "audio/mp4", + "audio/mpga", + "audio/mpeg", + "audio/m4a", + "audio/mp3", + "audio/flac", + "audio/aac", +} + +# Gemini-supported document MIME types +GEMINI_SUPPORTED_DOCUMENT_TYPES = { + "text/plain", + "application/pdf", +} + +# Mapping of common file extensions to MIME types +# This extends Python's mimetypes with custom mappings +EXTENSION_TO_MIME_TYPE = { + ".jpg": "image/jpeg", # Normalize jpg to jpeg + ".jpeg": "image/jpeg", + ".png": "image/png", + ".webp": "image/webp", + ".pdf": "application/pdf", + ".mp3": "audio/mpeg", + ".wav": "audio/wav", + ".m4a": "audio/mp4", +} + + +def detect_content_type_from_filename(filename: str) -> str: + """ + Detect content type from filename using extension. + + Uses Python's mimetypes module with custom overrides for common cases. + Normalizes jpg to jpeg for consistency. + """ + if not filename: + return "application/octet-stream" + + # Try custom mapping first + filename_lower = filename.lower() + for ext, mime_type in EXTENSION_TO_MIME_TYPE.items(): + if filename_lower.endswith(ext): + return mime_type + + # Fall back to Python's mimetypes + mime_type_guess, _ = mimetypes.guess_type(filename) + if mime_type_guess is not None: + return mime_type_guess + + return "application/octet-stream" + + +def normalize_mime_type_for_provider( + mime_type: str, provider: Optional[str] = None +) -> str: + """ + Normalize MIME type for specific provider requirements. + + Currently handles: + - Gemini: Normalizes image/jpg to image/jpeg + + Args: + mime_type: Original MIME type + provider: Provider name (e.g., "gemini", "vertex_ai") + + Returns: + str: Normalized MIME type + """ + normalized = mime_type.lower().strip() + + # Gemini/Vertex AI requires image/jpeg, not image/jpg + if provider and ("gemini" in provider.lower() or "vertex_ai" in provider.lower()): + if normalized == "image/jpg": + normalized = "image/jpeg" + + # General normalization: always normalize jpg to jpeg + if normalized == "image/jpg": + normalized = "image/jpeg" + + return normalized + + +def is_gemini_supported_mime_type(mime_type: str) -> bool: + """ + Check if a MIME type is supported by Gemini multimodal models. + + Supported categories: + - Images: image/png, image/jpeg, image/webp + - Video: 3gpp, wmv, webm, mp4, mpg, mpegps, mpeg, quicktime, x-flv + - Audio: webm, wav, pcm, opus, mp4, mpga, mpeg, m4a, mp3, flac, aac + - Documents: text/plain, application/pdf + + Args: + mime_type: MIME type to check + + Returns: + bool: True if supported, False otherwise + """ + normalized = normalize_mime_type_for_provider(mime_type, provider="gemini") + return normalized in ( + GEMINI_SUPPORTED_IMAGE_TYPES + | GEMINI_SUPPORTED_VIDEO_TYPES + | GEMINI_SUPPORTED_AUDIO_TYPES + | GEMINI_SUPPORTED_DOCUMENT_TYPES + ) + + +def get_content_type_from_file_object(file_object: Optional[dict]) -> str: + """ + Determine content type from file object (from database or API response). + + Extracts filename from file object and uses detect_content_type_from_filename. + Falls back to default if file object is invalid or filename not found. + + Args: + file_object: File object dictionary (can be None) + + Returns: + str: MIME type (defaults to "application/octet-stream" if cannot be determined) + """ + if not file_object: + return "application/octet-stream" + + # Handle JSON string + if isinstance(file_object, str): + import json + try: + file_object = json.loads(file_object) + except json.JSONDecodeError: + return "application/octet-stream" + + if not isinstance(file_object, dict): + return "application/octet-stream" + + # Try to get filename + filename = file_object.get("filename", "") + if filename: + return detect_content_type_from_filename(filename) + + return "application/octet-stream" + + +# ============================================================================ +# REQUEST PARAMETER EXTRACTION +# ============================================================================ + + +@dataclass +class FileCreationParams: + """ + Structured parameters extracted from file creation requests. + + Attributes: + target_storage: Storage backend name (e.g., "azure_storage", "default") + target_model_names: List of model names for managed files + model: Model parameter for multi-account routing + """ + + target_storage: str = "default" + target_model_names: List[str] = field(default_factory=list) + model: Optional[str] = None + + def __post_init__(self): + """Normalize and validate parameters after initialization.""" + if self.target_model_names is None: + self.target_model_names = [] + + # Normalize target_storage + if not self.target_storage: + self.target_storage = "default" + + # Strip whitespace from model names + self.target_model_names = [name.strip() for name in self.target_model_names if name.strip()] + + +async def extract_file_creation_params( + request: Request, + request_body: Optional[dict] = None, + target_model_names_form: Optional[str] = None, + target_storage_form: Optional[str] = None, +) -> FileCreationParams: + """ + Extract file creation parameters from request. + + Args: + request: FastAPI request object + request_body: Optional pre-parsed request body + target_model_names_form: target_model_names from form field (comma-separated string) + target_storage_form: target_storage from form field (defaults to "default") + + Returns: + FileCreationParams: Structured parameters extracted from the request + """ + if request_body is None: + request_body = await _read_request_body(request=request) or {} + + # Extract target_storage (simplified - just use form parameter) + target_storage = _extract_target_storage_simple(target_storage_form) + + # Extract target_model_names (simplified - just use form parameter) + target_model_names = _extract_target_model_names_simple(target_model_names_form) + + # Extract model parameter + model = _extract_model_param(request, request_body) + + return FileCreationParams( + target_storage=target_storage, + target_model_names=target_model_names, + model=model, + ) + + +def _extract_target_storage_simple(target_storage_form: Optional[str] = None) -> str: + """ + Extract target_storage parameter from form field. + + Args: + target_storage_form: target_storage from form field + + Returns: + str: Target storage backend name, or "default" + """ + if target_storage_form: + return target_storage_form.strip() + return "default" + + +def _extract_target_model_names_simple(target_model_names_form: Optional[str] = None) -> List[str]: + """ + Extract target_model_names parameter from form field. + """ + if not target_model_names_form: + return [] + + # Parse comma-separated string into list + if isinstance(target_model_names_form, str): + return [name.strip() for name in target_model_names_form.split(",") if name.strip()] + elif isinstance(target_model_names_form, list): + return [str(name).strip() for name in target_model_names_form if name] + + return [] + + +def _extract_model_param(request: Request, request_body: dict) -> Optional[str]: + """ + Extract model parameter from request. + + Priority: + 1. request_body.model + 2. Query parameter (?model=) + 3. Header (x-litellm-model) + """ + return ( + request_body.get("model") + or request.query_params.get("model") + or request.headers.get("x-litellm-model") + ) diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 3f08a4ec36..9738dd5738 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -7,7 +7,7 @@ import asyncio import traceback -from typing import Optional, cast, get_args +from typing import Any, Optional, cast, get_args import httpx from fastapi import ( @@ -46,10 +46,12 @@ from litellm.types.llms.openai import ( from .common_utils import ( _is_base64_encoded_unified_file_id, encode_file_id_with_model, + extract_file_creation_params, get_credentials_for_model, handle_model_based_routing, prepare_data_with_credentials, ) +from .storage_backend_service import StorageBackendFileService router = APIRouter() @@ -135,17 +137,38 @@ async def route_create_file( router_model: Optional[str], custom_llm_provider: str, model: Optional[str] = None, + target_storage: Optional[str] = "default", ) -> OpenAIFileObject: """ Route file creation request to the appropriate provider. Priority: - 1. If model parameter provided -> use model credentials and encode ID - 2. If enable_loadbalancing_on_batch_endpoints -> deprecated loadbalancing - 3. If target_model_names_list -> managed files (requires DB) - 4. Else -> use custom_llm_provider with files_settings + 1. If target_storage is specified and not "default" -> use storage backend + 2. If model parameter provided -> use model credentials and encode ID + 3. If enable_loadbalancing_on_batch_endpoints -> deprecated loadbalancing + 4. If target_model_names_list -> managed files (requires DB) + 5. Else -> use custom_llm_provider with files_settings """ + # Handle custom storage backend + if target_storage and target_storage != "default": + from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data + + # Extract file data + file_data = extract_file_data(cast(Any, _create_file_request.get("file"))) + + # Use storage backend service to handle upload + file_object = await StorageBackendFileService.upload_file_to_storage_backend( + file_data=file_data, + target_storage=target_storage, + target_model_names=target_model_names_list, + purpose=purpose, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=user_api_key_dict, + ) + + return file_object + # NEW: Handle model-based routing (no DB required) if model is not None: # Get credentials from model_list via router @@ -254,6 +277,7 @@ async def create_file( fastapi_response: Response, purpose: str = Form(...), target_model_names: str = Form(default=""), + target_storage: str = Form(default="default"), provider: Optional[str] = None, custom_llm_provider: str = Form(default="openai"), file: UploadFile = File(...), @@ -297,18 +321,18 @@ async def create_file( or "openai" ) - # NEW: Extract model parameter for multi-account routing + # Extract file creation parameters using utility function request_body = await _read_request_body(request=request) or {} - model_param = ( - request_body.get("model") - or request.query_params.get("model") - or request.headers.get("x-litellm-model") + file_params = await extract_file_creation_params( + request=request, + request_body=request_body, + target_model_names_form=target_model_names, + target_storage_form=target_storage, ) - - target_model_names_list = ( - target_model_names.split(",") if target_model_names else [] - ) - target_model_names_list = [model.strip() for model in target_model_names_list] + + target_storage = file_params.target_storage + target_model_names_list = file_params.target_model_names + model_param = file_params.model # Prepare the data for forwarding # Replace with: @@ -368,6 +392,7 @@ async def create_file( router_model=router_model, custom_llm_provider=custom_llm_provider, model=model_param, + target_storage=target_storage, ) if response is None: @@ -447,7 +472,7 @@ async def create_file( dependencies=[Depends(user_api_key_auth)], tags=["files"], ) -async def get_file_content( +async def get_file_content( # noqa: PLR0915 request: Request, fastapi_response: Response, file_id: str, @@ -525,6 +550,38 @@ async def get_file_content( param="None", code=500, ) + + # Check if file is stored in a storage backend (check DB) + if hasattr(managed_files_obj, "prisma_client") and managed_files_obj.prisma_client: + db_file = await managed_files_obj.prisma_client.db.litellm_managedfiletable.find_first( + where={"unified_file_id": file_id} + ) + if db_file and db_file.storage_backend and db_file.storage_url: + # File is stored in a storage backend, download it + 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 + + try: + # Get storage backend (uses same env vars as callback) + storage_backend = get_storage_backend(storage_backend_name) + file_content = await storage_backend.download_file(storage_url) + + # Return file content + from fastapi.responses import Response as FastAPIResponse + return FastAPIResponse( + content=file_content, + media_type="application/octet-stream", + ) + except ValueError as e: + raise ProxyException( + message=f"Storage backend error: {str(e)}", + type="invalid_request_error", + param="file_id", + code=400, + ) + model = cast(Optional[str], data.get("model")) if model: response = await llm_router.afile_content( diff --git a/litellm/proxy/openai_files_endpoints/storage_backend_service.py b/litellm/proxy/openai_files_endpoints/storage_backend_service.py new file mode 100644 index 0000000000..727a7876a5 --- /dev/null +++ b/litellm/proxy/openai_files_endpoints/storage_backend_service.py @@ -0,0 +1,244 @@ +""" +Storage backend service for file upload operations. + +This module provides a service class for handling file uploads to custom +storage backends (e.g., Azure Blob Storage) and managing associated metadata. +""" + +import base64 +import time +from typing import Any, List, Mapping, cast + +from litellm._logging import verbose_proxy_logger +from litellm._uuid import uuid as uuid_module +from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend +from litellm.llms.base_llm.files.transformation import BaseFileEndpoints +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.utils import ProxyLogging +from litellm.types.llms.openai import OpenAIFileObject +from litellm.types.utils import SpecialEnums + + +class StorageBackendFileService: + """ + Service for handling file uploads to storage backends. + + This service encapsulates the logic for: + - Uploading files to storage backends + - Creating file objects with storage metadata + - Generating unified file IDs for managed files + - Storing files in the managed files system + """ + + @staticmethod + async def upload_file_to_storage_backend( + file_data: Mapping[str, Any], + target_storage: str, + target_model_names: List[str], + purpose: str, + proxy_logging_obj: ProxyLogging, + user_api_key_dict: UserAPIKeyAuth, + ) -> 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_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 + + Returns: + OpenAIFileObject: Created file object with storage metadata + + Raises: + ProxyException: If storage backend is invalid or upload fails + """ + # Get storage backend instance + try: + storage_backend = get_storage_backend(target_storage) + except ValueError as e: + raise ProxyException( + message=str(e), + type="invalid_request_error", + param="target_storage", + code=400, + ) + + # Extract file information + file_content = file_data["content"] + filename = file_data.get("filename", "file") + content_type = file_data.get("content_type", "application/octet-stream") + + # Upload to storage backend + storage_url = await storage_backend.upload_file( + file_content=file_content, + filename=filename, + content_type=content_type, + path_prefix="", + file_naming_strategy="uuid", + ) + + verbose_proxy_logger.debug( + f"Storage backend upload complete: backend={target_storage}, url={storage_url}" + ) + + # Create file object with storage metadata + file_object = StorageBackendFileService._create_file_object_with_storage_metadata( + file_content=file_content, + filename=filename, + purpose=purpose, + target_storage=target_storage, + storage_url=storage_url, + ) + + # Store in managed files if target_model_names provided + if target_model_names: + await StorageBackendFileService._store_in_managed_files( + file_object=file_object, + file_data=file_data, + target_model_names=target_model_names, + target_storage=target_storage, + storage_url=storage_url, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=user_api_key_dict, + ) + + return file_object + + @staticmethod + def _create_file_object_with_storage_metadata( + file_content: bytes, + filename: str, + purpose: str, + target_storage: str, + storage_url: str, + ) -> OpenAIFileObject: + """ + Create an OpenAIFileObject with storage backend metadata. + + Args: + file_content: File content bytes + filename: Original filename + purpose: File purpose + target_storage: Storage backend name + storage_url: URL where file is stored + + Returns: + OpenAIFileObject: File object with storage metadata in _hidden_params + """ + file_id = f"file-{uuid_module.uuid4().hex[:24]}" + file_object = OpenAIFileObject( + id=file_id, + object="file", + purpose=purpose, + created_at=int(time.time()), + bytes=len(file_content), + filename=filename, + status="uploaded", + ) + + # Store storage metadata in hidden params + if not hasattr(file_object, "_hidden_params") or file_object._hidden_params is None: + file_object._hidden_params = {} + file_object._hidden_params.update({ + "storage_backend": target_storage, + "storage_url": storage_url, + }) + + return file_object + + @staticmethod + def _create_unified_file_id( + file_type: str, + target_model_names: List[str], + file_id: str, + ) -> str: + """ + Create a base64-encoded unified file ID for managed files. + + Args: + file_type: MIME type of the file + target_model_names: List of model names + file_id: Original file ID + + Returns: + str: Base64-encoded unified file ID + """ + unified_file_id_str = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( + file_type, + str(uuid_module.uuid4()), + ",".join(target_model_names), + file_id, + None, + ) + + base64_unified_file_id = ( + base64.urlsafe_b64encode(unified_file_id_str.encode()).decode().rstrip("=") + ) + + return base64_unified_file_id + + @staticmethod + async def _store_in_managed_files( + file_object: OpenAIFileObject, + file_data: Mapping[str, Any], + target_model_names: List[str], + target_storage: str, + storage_url: str, + proxy_logging_obj: ProxyLogging, + user_api_key_dict: UserAPIKeyAuth, + ) -> None: + """ + Store file in managed files system with unified file ID. + + Args: + file_object: File object to store + file_data: File data dictionary + target_model_names: List of model names + target_storage: Storage backend name + storage_url: URL where file is stored + proxy_logging_obj: Proxy logging object + user_api_key_dict: User API key authentication data + """ + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") + if not managed_files_obj or not isinstance(managed_files_obj, BaseFileEndpoints): + verbose_proxy_logger.warning( + "Managed files hook not available, skipping managed files storage" + ) + return + managed_files_obj = cast(Any, managed_files_obj) + + # Create model mappings using storage URL + model_mappings = { + model_name: storage_url + for model_name in target_model_names + } + + # Create unified file ID + file_type = file_data.get("content_type", "application/octet-stream") + base64_unified_file_id = StorageBackendFileService._create_unified_file_id( + file_type=file_type, + target_model_names=target_model_names, + file_id=file_object.id, + ) + + # Update file object ID to unified ID + file_object.id = base64_unified_file_id + + verbose_proxy_logger.debug( + f"Storing file in managed files: unified_id={base64_unified_file_id}, " + f"storage_backend={target_storage}, storage_url={storage_url}" + ) + + # Store in managed files + await managed_files_obj.store_unified_file_id( + file_id=base64_unified_file_id, + file_object=file_object, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + model_mappings=model_mappings, + user_api_key_dict=user_api_key_dict, + ) + diff --git a/schema.prisma b/schema.prisma index e227c41f93..1aecd4c914 100644 --- a/schema.prisma +++ b/schema.prisma @@ -573,6 +573,8 @@ model LiteLLM_ManagedFileTable { file_object Json? // Stores the OpenAIFileObject model_mappings Json flat_model_file_ids String[] @default([]) // Flat list of model file id's - for faster querying of model id -> unified file id + storage_backend String? // Storage backend name (e.g., "azure_storage", "gcs", "default") + storage_url String? // The actual storage URL where the file is stored created_at DateTime @default(now()) created_by String? updated_at DateTime @updatedAt diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 521faae3ca..754b694148 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -18,6 +18,7 @@ from litellm.proxy._types import LiteLLM_UserTableFiltered, UserAPIKeyAuth from litellm.proxy.hooks import get_proxy_hook from litellm.proxy.management_endpoints.internal_user_endpoints import ui_view_users from litellm.proxy.proxy_server import app +from litellm.types.llms.openai import OpenAIFileObject client = TestClient(app) from litellm.caching.caching import DualCache @@ -225,6 +226,97 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router: assert openai_call_found, "OpenAI call not found with expected parameters" +def test_target_storage_invokes_storage_backend( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + """ + Ensure target_storage is parsed and invokes the storage backend service. + """ + setup_proxy_logging_object(monkeypatch, llm_router) + + async_mock = mocker.AsyncMock( + return_value=OpenAIFileObject( + id="file-test", + object="file", + purpose="user_data", + created_at=0, + bytes=3, + filename="abc.txt", + status="uploaded", + ) + ) + mocker.patch( + "litellm.proxy.openai_files_endpoints.files_endpoints.StorageBackendFileService.upload_file_to_storage_backend", + new=async_mock, + ) + + test_file_content = b"abc" + test_file = ("abc.txt", test_file_content, "text/plain") + + response = client.post( + "/v1/files", + files={"file": test_file}, + data={ + "purpose": "user_data", + "target_storage": "azure_storage", + }, + headers={"Authorization": "Bearer test-key"}, + ) + + assert response.status_code == 200 + async_mock.assert_awaited_once() + called_kwargs = async_mock.call_args.kwargs + assert called_kwargs["target_storage"] == "azure_storage" + assert called_kwargs["target_model_names"] == [] + assert called_kwargs["purpose"] == "user_data" + + +def test_target_storage_with_target_models( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + """ + Ensure target_storage and target_model_names are parsed and passed through. + """ + setup_proxy_logging_object(monkeypatch, llm_router) + + async_mock = mocker.AsyncMock( + return_value=OpenAIFileObject( + id="file-test", + object="file", + purpose="user_data", + created_at=0, + bytes=3, + filename="abc.txt", + status="uploaded", + ) + ) + mocker.patch( + "litellm.proxy.openai_files_endpoints.files_endpoints.StorageBackendFileService.upload_file_to_storage_backend", + new=async_mock, + ) + + test_file_content = b"abc" + test_file = ("abc.txt", test_file_content, "text/plain") + + response = client.post( + "/v1/files", + files={"file": test_file}, + data={ + "purpose": "user_data", + "target_storage": "azure_storage", + "target_model_names": "gemini-2.0-flash", + }, + headers={"Authorization": "Bearer test-key"}, + ) + + assert response.status_code == 200 + async_mock.assert_awaited_once() + called_kwargs = async_mock.call_args.kwargs + assert called_kwargs["target_storage"] == "azure_storage" + assert called_kwargs["target_model_names"] == ["gemini-2.0-flash"] + assert called_kwargs["purpose"] == "user_data" + + @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