mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Azure storage backend handlers to streaming
This commit is contained in:
parent
d18190ed16
commit
06ddc2c0db
4 changed files with 362 additions and 21 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue