Azure storage backend handlers to streaming

This commit is contained in:
harish876 2026-04-07 23:08:05 +00:00
parent d18190ed16
commit 06ddc2c0db
4 changed files with 362 additions and 21 deletions

View file

@ -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

View file

@ -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.

View file

@ -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(

View file

@ -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"]