added streaming support for base_llm_provider handler and azure.

TODO: Add support for mock GCS / S3 services by passing in API base. Add more tests
This commit is contained in:
harish876 2026-04-07 17:56:10 +00:00
parent f3633c5af1
commit bda8c0a1c3
6 changed files with 414 additions and 191 deletions

View file

@ -1041,10 +1041,12 @@ def file_content_streaming(
"""
Prototype API: Returns a byte iterator for file contents.
Currently supports OpenAI provider only.
Supports OpenAI-compatible providers and Azure.
"""
try:
optional_params = GenericLiteLLMParams(**kwargs)
litellm_params_dict = get_litellm_params(**kwargs)
client = kwargs.get("client")
try:
if model is not None:
@ -1054,26 +1056,11 @@ def file_content_streaming(
except Exception:
pass
resolved_provider = cast(Optional[str], custom_llm_provider) or "openai"
if resolved_provider != "openai":
raise litellm.exceptions.BadRequestError(
message="LiteLLM doesn't support {} for 'file_content_v2'. Supported providers are 'openai'.".format(
resolved_provider
),
model="n/a",
llm_provider=resolved_provider,
response=httpx.Response(
status_code=400,
content="Unsupported provider",
request=httpx.Request(method="file_content_v2", url="https://github.com/BerriAI/litellm"), # type: ignore
),
)
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
if (
timeout is not None
and isinstance(timeout, httpx.Timeout)
and supports_httpx_timeout(resolved_provider) is False
and supports_httpx_timeout(cast(str, custom_llm_provider)) is False
):
timeout = timeout.read or 600
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
@ -1081,16 +1068,66 @@ def file_content_streaming(
elif timeout is None:
timeout = 600.0
openai_creds = get_openai_credentials(
api_base=optional_params.api_base,
api_key=optional_params.api_key,
organization=optional_params.organization,
)
_is_async = kwargs.pop("afile_content_streaming", False) is True
response = cast(Union[Iterator[bytes], AsyncIterator[bytes]], iter(()))
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
provider_config = ProviderConfigManager.get_provider_files_config(
model="",
provider=LlmProviders(custom_llm_provider),
)
if provider_config is not None:
litellm_params_dict["api_key"] = optional_params.api_key
litellm_params_dict["api_base"] = optional_params.api_base
logging_obj = cast(
Optional[LiteLLMLoggingObj], kwargs.get("litellm_logging_obj")
)
if logging_obj is None:
logging_obj = LiteLLMLoggingObj(
model="",
messages=[],
stream=True,
call_type=(
"afile_content_streaming"
if _is_async
else "file_content_streaming"
),
start_time=time.time(),
litellm_call_id=kwargs.get(
"litellm_call_id", str(uuid_module.uuid4())
),
function_id=str(kwargs.get("id") or ""),
)
response = cast(
Union[Iterator[bytes], AsyncIterator[bytes]],
base_llm_http_handler.retrieve_file_content_streaming(
file_content_request=FileContentRequest(
file_id=file_id,
extra_headers=extra_headers,
extra_body=extra_body,
),
provider_config=provider_config,
litellm_params=litellm_params_dict,
headers=extra_headers or {},
logging_obj=logging_obj,
_is_async=_is_async,
client=(
client
if client is not None
and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
else None
),
timeout=timeout,
chunk_size=chunk_size,
),
)
elif custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
openai_creds = get_openai_credentials(
api_base=optional_params.api_base,
api_key=optional_params.api_key,
organization=optional_params.organization,
)
response = openai_files_instance.file_content_streaming(
_is_async=_is_async,
file_content_request=FileContentRequest(
@ -1105,6 +1142,28 @@ def file_content_streaming(
organization=openai_creds.organization,
chunk_size=chunk_size,
)
elif custom_llm_provider == "azure":
azure_creds = get_azure_credentials(
api_base=optional_params.api_base,
api_key=optional_params.api_key,
api_version=optional_params.api_version,
)
response = azure_files_instance.file_content_streaming(
_is_async=_is_async,
file_content_request=FileContentRequest(
file_id=file_id,
extra_headers=extra_headers,
extra_body=extra_body,
),
api_base=azure_creds.api_base,
api_key=azure_creds.api_key,
timeout=timeout,
max_retries=optional_params.max_retries,
api_version=azure_creds.api_version,
chunk_size=chunk_size,
client=client,
litellm_params=litellm_params_dict,
)
else:
raise litellm.exceptions.BadRequestError(
message="LiteLLM doesn't support {} for 'file_content'. Supported providers are 'openai', 'azure', 'vertex_ai', 'bedrock', 'manus', 'anthropic'.".format(
@ -1118,7 +1177,7 @@ def file_content_streaming(
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
),
)
return response
except Exception as e:
raise e

View file

@ -1,4 +1,5 @@
from typing import Any, Coroutine, Optional, Union, cast
from typing import Any, AsyncIterator, Coroutine, Iterator, Optional, Union, cast
import httpx
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
@ -143,6 +144,67 @@ class AzureOpenAIFilesAPI(BaseAzureLLM):
return HttpxBinaryResponseContent(response=response.response)
async def afile_content_streaming(
self,
file_content_request: FileContentRequest,
openai_client: Union[AsyncAzureOpenAI, AsyncOpenAI],
chunk_size: int = 1024 * 1024,
) -> AsyncIterator[bytes]:
async with openai_client.files.with_streaming_response.content(
**file_content_request
) as response:
async for chunk in response.iter_bytes(chunk_size=chunk_size):
yield chunk
def file_content_streaming(
self,
_is_async: bool,
file_content_request: FileContentRequest,
api_base: Optional[str],
api_key: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
api_version: Optional[str] = None,
chunk_size: int = 1024 * 1024,
client: Optional[
Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]
] = None,
litellm_params: Optional[dict] = None,
) -> Union[Iterator[bytes], AsyncIterator[bytes]]:
openai_client: Optional[
Union[AzureOpenAI, AsyncAzureOpenAI, OpenAI, AsyncOpenAI]
] = self.get_azure_openai_client(
litellm_params=litellm_params or {},
api_key=api_key,
api_base=api_base,
api_version=api_version,
client=client,
_is_async=_is_async,
)
if openai_client is None:
raise ValueError(
"AzureOpenAI client is not initialized. Make sure api_key is passed or OPENAI_API_KEY is set in the environment."
)
if _is_async is True:
if not isinstance(openai_client, (AsyncAzureOpenAI, AsyncOpenAI)):
raise ValueError(
"AzureOpenAI client is not an instance of AsyncAzureOpenAI. Make sure you passed an AsyncAzureOpenAI client."
)
return self.afile_content_streaming(
file_content_request=file_content_request,
openai_client=openai_client,
chunk_size=chunk_size,
)
def _stream() -> Iterator[bytes]:
with cast(Union[AzureOpenAI, OpenAI], openai_client).files.with_streaming_response.content(
**file_content_request
) as response:
yield from response.iter_bytes(chunk_size=chunk_size)
return _stream()
async def aretrieve_file(
self,
file_id: str,

View file

@ -6,6 +6,7 @@ from typing import (
AsyncIterator,
Coroutine,
Dict,
Iterator,
List,
Literal,
Optional,
@ -4373,6 +4374,164 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params,
)
def retrieve_file_content_streaming(
self,
file_content_request: "FileContentRequest",
provider_config: BaseFilesConfig,
litellm_params: dict,
headers: dict,
logging_obj: LiteLLMLoggingObj,
_is_async: bool = False,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
chunk_size: int = 1024 * 1024,
) -> Union[
Iterator[bytes],
Coroutine[Any, Any, AsyncIterator[bytes]],
]:
"""
Retrieve file content by ID as a streaming iterator.
"""
if _is_async:
return self.async_retrieve_file_content_streaming(
file_content_request=file_content_request,
provider_config=provider_config,
litellm_params=litellm_params,
headers=headers,
logging_obj=logging_obj,
client=client,
timeout=timeout,
chunk_size=chunk_size,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_handler = _get_httpx_client()
else:
sync_httpx_handler = client
# Get URL and params from provider config
url, params = provider_config.transform_file_content_request(
file_content_request=file_content_request,
optional_params={},
litellm_params=litellm_params,
)
# Validate environment and get headers
headers = provider_config.validate_environment(
api_key=litellm_params.get("api_key"),
headers=headers,
model="",
messages=[],
optional_params={},
litellm_params=litellm_params,
)
logging_obj.pre_call(
input="",
api_key="",
additional_args={
"api_base": url,
"headers": headers,
"file_id": file_content_request.get("file_id"),
"stream": True,
},
)
try:
request = sync_httpx_handler.client.build_request(
"GET",
url,
headers=headers,
params=params,
timeout=timeout,
)
response = sync_httpx_handler.client.send(request, stream=True)
response.raise_for_status()
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
def _sync_stream_iterator() -> Iterator[bytes]:
try:
for chunk in response.iter_bytes(chunk_size=chunk_size):
if chunk:
yield chunk
finally:
response.close()
return _sync_stream_iterator()
async def async_retrieve_file_content_streaming(
self,
file_content_request: "FileContentRequest",
provider_config: BaseFilesConfig,
litellm_params: dict,
headers: dict,
logging_obj: LiteLLMLoggingObj,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
chunk_size: int = 1024 * 1024,
) -> AsyncIterator[bytes]:
"""
Async retrieve file content by ID as a streaming iterator.
"""
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_handler = get_async_httpx_client(
llm_provider=provider_config.custom_llm_provider
)
else:
async_httpx_handler = client
# Get URL and params from provider config
url, params = provider_config.transform_file_content_request(
file_content_request=file_content_request,
optional_params={},
litellm_params=litellm_params,
)
# Validate environment and get headers
headers = provider_config.validate_environment(
api_key=litellm_params.get("api_key"),
headers=headers,
model="",
messages=[],
optional_params={},
litellm_params=litellm_params,
)
logging_obj.pre_call(
input="",
api_key="",
additional_args={
"api_base": url,
"headers": headers,
"file_id": file_content_request.get("file_id"),
"stream": True,
},
)
try:
request = async_httpx_handler.client.build_request(
"GET",
url,
headers=headers,
params=params,
timeout=timeout,
)
response = await async_httpx_handler.client.send(request, stream=True)
response.raise_for_status()
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
async def _async_stream_iterator() -> AsyncIterator[bytes]:
try:
async for chunk in response.aiter_bytes(chunk_size=chunk_size):
if chunk:
yield chunk
finally:
await response.aclose()
return _async_stream_iterator()
async def async_retrieve_file_content(
self,
file_content_request: "FileContentRequest",
@ -4691,18 +4850,30 @@ class BaseLLMHTTPHandler:
BaseEvalsAPIConfig,
],
):
def _safe_response_text(response: httpx.Response) -> str:
try:
return response.text
except RuntimeError:
# Streamed responses require an explicit read before accessing .text.
try:
return response.read().decode("utf-8", errors="replace")
except Exception:
return ""
status_code = getattr(e, "status_code", 500)
error_headers = getattr(e, "headers", None)
if isinstance(e, httpx.HTTPStatusError):
error_text = e.response.text
error_text = _safe_response_text(e.response) or str(e)
status_code = e.response.status_code
else:
error_text = getattr(e, "text", str(e))
error_response = getattr(e, "response", None)
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
if error_response and hasattr(error_response, "text"):
error_text = getattr(error_response, "text", error_text)
if isinstance(error_response, httpx.Response):
safe_text = _safe_response_text(error_response)
if safe_text:
error_text = safe_text
if error_headers:
error_headers = dict(error_headers)
else:

View file

@ -22,7 +22,6 @@ from fastapi import (
status,
)
from fastapi.responses import StreamingResponse
from openai import OpenAI
import litellm
from litellm import CreateFileRequest, get_secret_str
@ -57,7 +56,6 @@ from .common_utils import (
handle_model_based_routing,
prepare_data_with_credentials,
)
import os
from .storage_backend_service import StorageBackendFileService
router = APIRouter()
@ -116,11 +114,6 @@ def get_model_from_json_obj(json_object: dict) -> Optional[str]:
return model
def _stream_openai_file_content(file_id: str, client: OpenAI):
with client.files.with_streaming_response.content(file_id) as response:
yield from response.iter_bytes(chunk_size=1024 * 1024)
async def _deprecated_loadbalanced_create_file(
llm_router: Optional[Router],
router_model: str,
@ -843,7 +836,7 @@ async def get_file_content( # noqa: PLR0915
dependencies=[Depends(user_api_key_auth)],
tags=["files"],
)
async def get_file_content_v2(
async def get_file_content_streaming(
request: Request,
fastapi_response: Response,
file_id: str,
@ -859,7 +852,6 @@ async def get_file_content_v2(
data: Dict = {"file_id": file_id}
try:
# Include original request and headers in the data (same as v1 flow)
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
(
data,
@ -871,7 +863,7 @@ async def get_file_content_v2(
version=version,
proxy_logging_obj=proxy_logging_obj,
proxy_config=proxy_config,
route_type="afile_content",
route_type="afile_content_streaming",
)
custom_llm_provider = (
@ -881,15 +873,12 @@ async def get_file_content_v2(
or await get_custom_llm_provider_from_request_body(request=request)
or "openai"
)
if custom_llm_provider != "openai":
raise HTTPException(
status_code=400,
detail={"error": "v2 files content currently only supports openai provider"},
)
client = OpenAI(
base_url=os.getenv("OPENAI_BASE_URL"),
api_key=os.getenv("OPENAI_API_KEY"),
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,
)
asyncio.create_task(
@ -910,7 +899,7 @@ async def get_file_content_v2(
)
return StreamingResponse(
_stream_openai_file_content(file_id, client),
content=stream_iterator,
media_type="application/octet-stream",
headers={
"content-disposition": f'attachment; filename="{file_id}.bin"',

View file

@ -0,0 +1,84 @@
from unittest.mock import MagicMock
import pytest
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from litellm.proxy._types import ProxyException
from litellm.proxy.openai_files_endpoints.files_endpoints import get_file_content_streaming
async def _fake_stream_iterator():
yield b"hello "
yield b"world"
class _FakeProxyLogging:
async def update_request_status(self, litellm_call_id: str, status: str):
return None
async def post_call_failure_hook(
self, user_api_key_dict, original_exception: Exception, request_data
):
return None
async def _fake_common_processing_pre_call_logic(self, **kwargs):
return self.data, MagicMock()
@pytest.fixture
def request_obj() -> Request:
scope = {
"type": "http",
"method": "GET",
"path": "/v2/files/file-123/content",
"headers": [],
"query_string": b"",
}
return Request(scope)
@pytest.mark.asyncio
async def test_get_file_content_streaming_returns_streaming_response(monkeypatch, request_obj):
async def _mock_afile_content_streaming(**kwargs):
return _fake_stream_iterator()
monkeypatch.setattr("litellm.afile_content_streaming", _mock_afile_content_streaming)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.ProxyBaseLLMRequestProcessing.common_processing_pre_call_logic",
_fake_common_processing_pre_call_logic,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj",
_FakeProxyLogging(),
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{},
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_config",
None,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.version",
"test-version",
)
response = await get_file_content_streaming(
request=request_obj,
fastapi_response=Response(),
file_id="file-123",
provider="openai",
user_api_key_dict=MagicMock(),
)
assert isinstance(response, StreamingResponse)
assert response.media_type == "application/octet-stream"
chunks = []
async for chunk in response.body_iterator:
chunks.append(chunk)
assert chunks == [b"hello ", b"world"]

View file

@ -1,142 +0,0 @@
from unittest.mock import MagicMock
import pytest
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from litellm.proxy.openai_files_endpoints.files_endpoints import get_file_content_v2
from litellm.proxy._types import ProxyException
class _FakeStreamingResponse:
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def iter_bytes(self, chunk_size=1024 * 1024):
yield b"hello "
yield b"world"
class _FakeFilesClient:
class _WithStreamingResponse:
def content(self, file_id):
assert file_id == "file-123"
return _FakeStreamingResponse()
def __init__(self):
self.with_streaming_response = self._WithStreamingResponse()
class _FakeOpenAIClient:
def __init__(self, *args, **kwargs):
self.files = _FakeFilesClient()
class _FakeProxyLogging:
async def update_request_status(self, litellm_call_id: str, status: str):
return None
async def post_call_failure_hook(
self, user_api_key_dict, original_exception: Exception, request_data
):
return None
async def _fake_common_processing_pre_call_logic(self, **kwargs):
return self.data, MagicMock()
@pytest.fixture
def request_obj() -> Request:
scope = {
"type": "http",
"method": "GET",
"path": "/v2/files/file-123/content",
"headers": [],
"query_string": b"",
}
return Request(scope)
@pytest.mark.asyncio
async def test_get_file_content_v2_returns_streaming_response(monkeypatch, request_obj):
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.OpenAI",
_FakeOpenAIClient,
)
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.ProxyBaseLLMRequestProcessing.common_processing_pre_call_logic",
_fake_common_processing_pre_call_logic,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj",
_FakeProxyLogging(),
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{},
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_config",
None,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.version",
"test-version",
)
response = await get_file_content_v2(
request=request_obj,
fastapi_response=Response(),
file_id="file-123",
provider="openai",
user_api_key_dict=MagicMock(),
)
assert isinstance(response, StreamingResponse)
assert response.media_type == "application/octet-stream"
chunks = []
async for chunk in response.body_iterator:
chunks.append(chunk)
assert chunks == [b"hello ", b"world"]
@pytest.mark.asyncio
async def test_get_file_content_v2_rejects_non_openai_provider(monkeypatch, request_obj):
monkeypatch.setattr(
"litellm.proxy.openai_files_endpoints.files_endpoints.ProxyBaseLLMRequestProcessing.common_processing_pre_call_logic",
_fake_common_processing_pre_call_logic,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj",
_FakeProxyLogging(),
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{},
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_config",
None,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.version",
"test-version",
)
with pytest.raises(ProxyException) as exc_info:
await get_file_content_v2(
request=request_obj,
fastapi_response=Response(),
file_id="file-123",
provider="anthropic",
user_api_key_dict=MagicMock(),
)
assert exc_info.value.code == str(400)
assert "only supports openai provider" in exc_info.value.message