From bda8c0a1c36946d71c2fe3fed60cf3a8e72c545b Mon Sep 17 00:00:00 2001 From: harish876 Date: Tue, 7 Apr 2026 17:56:10 +0000 Subject: [PATCH] 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 --- litellm/files/main.py | 111 ++++++++--- litellm/llms/azure/files/handler.py | 64 ++++++- litellm/llms/custom_httpx/llm_http_handler.py | 177 +++++++++++++++++- .../openai_files_endpoints/files_endpoints.py | 27 +-- .../test_openai_files_streaming_endpoint.py | 84 +++++++++ .../test_openai_files_v2_endpoint.py | 142 -------------- 6 files changed, 414 insertions(+), 191 deletions(-) create mode 100644 tests/proxy_unit_tests/test_openai_files_streaming_endpoint.py delete mode 100644 tests/proxy_unit_tests/test_openai_files_v2_endpoint.py diff --git a/litellm/files/main.py b/litellm/files/main.py index 3d1c1f15c94..3181cb5b46b 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -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 diff --git a/litellm/llms/azure/files/handler.py b/litellm/llms/azure/files/handler.py index 72cbcba8a9a..5a2464c72fc 100644 --- a/litellm/llms/azure/files/handler.py +++ b/litellm/llms/azure/files/handler.py @@ -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, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 4c9abaad908..2d30fa55c96 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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: diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 9ff9fd82ad7..314205d0b66 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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"', diff --git a/tests/proxy_unit_tests/test_openai_files_streaming_endpoint.py b/tests/proxy_unit_tests/test_openai_files_streaming_endpoint.py new file mode 100644 index 00000000000..8f58dbbe809 --- /dev/null +++ b/tests/proxy_unit_tests/test_openai_files_streaming_endpoint.py @@ -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"] \ No newline at end of file diff --git a/tests/proxy_unit_tests/test_openai_files_v2_endpoint.py b/tests/proxy_unit_tests/test_openai_files_v2_endpoint.py deleted file mode 100644 index 17a3573d6e7..00000000000 --- a/tests/proxy_unit_tests/test_openai_files_v2_endpoint.py +++ /dev/null @@ -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