diff --git a/litellm/llms/base_llm/files/azure_blob_storage_backend.py b/litellm/llms/base_llm/files/azure_blob_storage_backend.py index a2155df4047..fabe42308bc 100644 --- a/litellm/llms/base_llm/files/azure_blob_storage_backend.py +++ b/litellm/llms/base_llm/files/azure_blob_storage_backend.py @@ -7,7 +7,7 @@ to reuse all authentication and Azure Storage operations. """ import time -from typing import Optional +from typing import AsyncIterator, Optional from urllib.parse import quote from litellm._logging import verbose_logger @@ -252,19 +252,7 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): 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 + file_path = self._extract_file_path_from_storage_url(storage_url) if self.azure_storage_account_key: # Use Azure SDK (reuse logger's service client) @@ -279,6 +267,54 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): ) raise + async def download_file_streaming( + self, storage_url: str, chunk_size: int = 1024 * 1024 + ) -> AsyncIterator[bytes]: + """ + Stream-download a file from Azure Blob Storage. + + Args: + storage_url: Blob URL in format: https://{account}.blob.core.windows.net/{container}/{path} + chunk_size: Chunk size in bytes for streamed reads + + Yields: + bytes: File content chunks + """ + try: + file_path = self._extract_file_path_from_storage_url(storage_url) + + if self.azure_storage_account_key: + async for chunk in self._download_file_with_account_key_streaming( + file_path=file_path, + chunk_size=chunk_size, + ): + yield chunk + else: + async for chunk in self._download_file_with_azure_ad_streaming( + file_path=file_path, + chunk_size=chunk_size, + ): + yield chunk + except Exception as e: + verbose_logger.exception( + f"Error streaming file from Azure Blob Storage: {str(e)}" + ) + raise + + def _extract_file_path_from_storage_url(self, storage_url: str) -> str: + # Parse blob URL to extract path + # URL format: https://{account}.blob.core.windows.net/{container}/{path} + """Extract file path from Azure blob URL.""" + if ".blob.core.windows.net/" not in storage_url: + raise ValueError(f"Invalid Azure Blob Storage URL: {storage_url}") + + 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}") + + return path_parts[1] + 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 @@ -297,6 +333,31 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): file_content = await download_response.readall() return file_content + async def _download_file_with_account_key_streaming( + self, file_path: str, chunk_size: int = 1024 * 1024 + ) -> AsyncIterator[bytes]: + """Stream-download file using Azure SDK with account key.""" + service_client = await self.get_service_client() + file_system_client = service_client.get_file_system_client( + file_system=self.azure_storage_file_system + ) + 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_response = await file_client.download_file() + + if not hasattr(download_response, "chunks"): + raise RuntimeError( + "Azure SDK download response does not support chunk streaming" + ) + + async for chunk in download_response.chunks(): + if chunk: + yield chunk + 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 @@ -323,3 +384,31 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger): response = await async_client.get(blob_url, headers=headers) response.raise_for_status() return response.content + + async def _download_file_with_azure_ad_streaming( + self, file_path: str, chunk_size: int = 1024 * 1024 + ) -> AsyncIterator[bytes]: + """Stream-download file using REST API with Azure AD token.""" + 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 + ) + + 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}", + } + + async with async_client.client.stream("GET", blob_url, headers=headers) as response: + response.raise_for_status() + async for chunk in response.aiter_bytes(chunk_size=chunk_size): + if chunk: + yield chunk diff --git a/litellm/llms/base_llm/files/storage_backend.py b/litellm/llms/base_llm/files/storage_backend.py index 31e68a7002a..59b6201ab11 100644 --- a/litellm/llms/base_llm/files/storage_backend.py +++ b/litellm/llms/base_llm/files/storage_backend.py @@ -6,7 +6,7 @@ This module defines the abstract base class that all file storage backends """ from abc import ABC, abstractmethod -from typing import Optional +from typing import AsyncIterator, Optional class BaseFileStorageBackend(ABC): @@ -60,6 +60,29 @@ class BaseFileStorageBackend(ABC): """ pass + async def download_file_streaming( + self, storage_url: str, chunk_size: int = 1024 * 1024 + ) -> AsyncIterator[bytes]: + """ + Stream-download a file from the storage backend. + + Backends can override this for true streaming. The default implementation + falls back to a buffered download and yields a single chunk. + + Args: + storage_url: The storage URL returned from upload_file + chunk_size: Preferred chunk size in bytes + + Yields: + bytes: File content chunks + + Raises: + Exception: If download fails + """ + file_content = await self.download_file(storage_url) + if file_content: + yield file_content + async def delete_file(self, storage_url: str) -> None: """ Delete a file from the storage backend. diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 314205d0b66..e6b92dfcf47 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -845,6 +845,7 @@ async def get_file_content_streaming( ): from litellm.proxy.proxy_server import ( general_settings, + llm_router, proxy_config, proxy_logging_obj, version, @@ -874,12 +875,140 @@ async def get_file_content_streaming( or "openai" ) - data.pop("file_id", None) - stream_iterator = await litellm.afile_content_streaming( - file_id=file_id, - custom_llm_provider=custom_llm_provider, # type: ignore - **data, - ) + ## check if file_id is a litellm managed file (same top-level flow as v1) + is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + if is_base64_unified_file_id: + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") + if managed_files_obj is None: + raise ProxyException( + message="Managed files hook not found", + type="None", + param="None", + code=500, + ) + if llm_router is None: + raise ProxyException( + message="LLM Router not found", + type="None", + param="None", + code=500, + ) + if not isinstance(managed_files_obj, BaseFileEndpoints): + raise ProxyException( + message="Managed files hook is not a BaseFileEndpoints", + type="None", + param="None", + code=500, + ) + + if hasattr(managed_files_obj, "prisma_client") and getattr( + managed_files_obj, "prisma_client", None + ): + prisma_client = getattr(managed_files_obj, "prisma_client") + db_file = await 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: + 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: + storage_backend = get_storage_backend(storage_backend_name) + + return StreamingResponse( + content=storage_backend.download_file_streaming(storage_url), + media_type="application/octet-stream", + headers={ + "content-disposition": f'attachment; filename="{file_id}.bin"', + }, + ) + 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: + # TODO: Add Streaming version here + # response = await llm_router.afile_content( + # **{ + # "model": model, + # "file_id": file_id, + # **data, + # } + # ) # type: ignore + raise ProxyException( + message="Managed files streaming path is pending implementation", + type="None", + param="file_id", + code=501, + ) + + else: + # TODO: Add Streaming version here + # response = await managed_files_obj.afile_content( + # **{ + # "file_id": file_id, + # "litellm_parent_otel_span": user_api_key_dict.parent_otel_span, + # "llm_router": llm_router, + # **data, + # } + # ) + raise ProxyException( + message="Managed files streaming path is pending implementation", + type="None", + param="file_id", + code=501, + ) + else: + # Check for model-based credential routing + ( + should_route, + model_used, + original_file_id, + credentials, + ) = handle_model_based_routing( + file_id=file_id, + request=request, + llm_router=llm_router, + data=data, + check_file_id_encoding=True, + ) + + if should_route: + # Use model-based routing with credentials from config + prepare_data_with_credentials( + data=data, + credentials=credentials, # type: ignore + file_id=original_file_id, # Use decoded file ID if from encoded ID + ) + + stream_iterator = await litellm.afile_content_streaming( + custom_llm_provider=credentials["custom_llm_provider"], # type: ignore + **data, + ) # type: ignore + + verbose_proxy_logger.debug( + f"Retrieved file content stream using model: {model_used}" + + ( + f", file_id: {file_id} -> {original_file_id}" + if original_file_id + else "" + ) + ) + else: + data.pop("file_id", None) + stream_iterator = await litellm.afile_content_streaming( + custom_llm_provider=custom_llm_provider, # type: ignore + file_id=file_id, + **data, + ) asyncio.create_task( proxy_logging_obj.update_request_status( diff --git a/tests/test_litellm/integrations/azure_storage/test_azure_storage.py b/tests/test_litellm/integrations/azure_storage/test_azure_storage.py index d7bb2a900d9..dd5bb8a743e 100644 --- a/tests/test_litellm/integrations/azure_storage/test_azure_storage.py +++ b/tests/test_litellm/integrations/azure_storage/test_azure_storage.py @@ -9,6 +9,9 @@ sys.path.insert( ) # Adds the parent directory to the system path from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLogger +from litellm.llms.base_llm.files.azure_blob_storage_backend import ( + AzureBlobStorageBackend, +) from litellm.types.utils import StandardLoggingPayload @@ -94,3 +97,100 @@ async def test_async_upload_payload_to_azure_blob_storage(mock_env_vars): # Verify raise_for_status was called on all responses assert mock_response.raise_for_status.call_count == 3 + + +async def _collect_async_chunks(async_iterable): + chunks = [] + async for chunk in async_iterable: + chunks.append(chunk) + return chunks + + +@pytest.mark.asyncio +async def test_download_file_account_key_flow(mock_env_vars, monkeypatch): + monkeypatch.setenv("AZURE_STORAGE_ACCOUNT_KEY", "test-account-key") + + backend = AzureBlobStorageBackend() + + service_client = MagicMock() + file_system_client = MagicMock() + file_client = MagicMock() + download_response = MagicMock() + + backend.get_service_client = AsyncMock(return_value=service_client) # type: ignore + service_client.get_file_system_client.return_value = file_system_client + file_system_client.exists = AsyncMock(return_value=True) + file_system_client.get_file_client.return_value = file_client + file_client.download_file = AsyncMock(return_value=download_response) + download_response.readall = AsyncMock(return_value=b"account-key-content") + + storage_url = "https://test-account.blob.core.windows.net/test-container/path/to/file.txt" + result = await backend.download_file(storage_url) + + assert result == b"account-key-content" + service_client.get_file_system_client.assert_called_once_with( + file_system="test-container" + ) + file_system_client.get_file_client.assert_called_once_with("path/to/file.txt") + + +@pytest.mark.asyncio +async def test_download_file_azure_ad_flow(mock_env_vars, monkeypatch): + monkeypatch.delenv("AZURE_STORAGE_ACCOUNT_KEY", raising=False) + + backend = AzureBlobStorageBackend() + backend.azure_auth_token = "mock-azure-ad-token" + backend.set_valid_azure_ad_token = AsyncMock() # type: ignore + + mock_async_httpx_client = MagicMock() + mock_response = MagicMock() + mock_response.raise_for_status = MagicMock() + mock_response.content = b"azure-ad-content" + mock_async_httpx_client.get = AsyncMock(return_value=mock_response) + + with patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_async_httpx_client, + ): + storage_url = "https://test-account.blob.core.windows.net/test-container/path/to/file.txt" + result = await backend.download_file(storage_url) + + assert result == b"azure-ad-content" + backend.set_valid_azure_ad_token.assert_awaited_once() + mock_async_httpx_client.get.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_download_file_streaming_account_key_flow(mock_env_vars, monkeypatch): + monkeypatch.setenv("AZURE_STORAGE_ACCOUNT_KEY", "test-account-key") + + backend = AzureBlobStorageBackend() + + async def _mock_stream(*args, **kwargs): + yield b"chunk-1" + yield b"chunk-2" + + backend._download_file_with_account_key_streaming = _mock_stream # type: ignore + + storage_url = "https://test-account.blob.core.windows.net/test-container/path/to/file.txt" + chunks = await _collect_async_chunks(backend.download_file_streaming(storage_url)) + + assert chunks == [b"chunk-1", b"chunk-2"] + + +@pytest.mark.asyncio +async def test_download_file_streaming_azure_ad_flow(mock_env_vars, monkeypatch): + monkeypatch.delenv("AZURE_STORAGE_ACCOUNT_KEY", raising=False) + + backend = AzureBlobStorageBackend() + + async def _mock_stream(*args, **kwargs): + yield b"ad-chunk-1" + yield b"ad-chunk-2" + + backend._download_file_with_azure_ad_streaming = _mock_stream # type: ignore + + storage_url = "https://test-account.blob.core.windows.net/test-container/path/to/file.txt" + chunks = await _collect_async_chunks(backend.download_file_streaming(storage_url)) + + assert chunks == [b"ad-chunk-1", b"ad-chunk-2"]