polished streaming implementation. TODO: Add support for azure, bedrock and managed files and base_llm_provider path

This commit is contained in:
harish876 2026-04-06 18:57:59 +00:00
parent 6b63082ddf
commit f3633c5af1
3 changed files with 199 additions and 1 deletions

View file

@ -10,9 +10,10 @@ import contextvars
import time
import uuid as uuid_module
from functools import partial
from typing import Any, Coroutine, Dict, Literal, Optional, Union, cast
from typing import Any, AsyncIterator, Coroutine, Dict, Iterator, Literal, Optional, Union, cast
import httpx
from openai import AsyncOpenAI, OpenAI
# Type aliases for provider parameters
FileCreateProvider = Literal[
@ -982,3 +983,142 @@ def file_content(
return response
except Exception as e:
raise e
@client
async def afile_content_streaming(
file_id: str,
custom_llm_provider: FileContentProvider = "openai",
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
chunk_size: int = 1024 * 1024,
**kwargs,
) -> Union[Iterator[bytes], AsyncIterator[bytes]]:
"""
Async wrapper for file_content_streaming.
"""
try:
loop = asyncio.get_event_loop()
kwargs["afile_content_streaming"] = True
model = kwargs.pop("model", None)
# Use a partial function to pass your keyword arguments
func = partial(
file_content_streaming,
file_id,
model,
custom_llm_provider,
extra_headers,
extra_body,
chunk_size,
**kwargs,
)
# Add the context to the function
ctx = contextvars.copy_context()
func_with_context = partial(ctx.run, func)
init_response = await loop.run_in_executor(None, func_with_context)
if asyncio.iscoroutine(init_response):
response = await init_response
else:
response = init_response # type: ignore
return response
except Exception as e:
raise e
@client
def file_content_streaming(
file_id: str,
model: Optional[str] = None,
custom_llm_provider: Optional[Union[FileContentProvider, str]] = None,
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
chunk_size: int = 1024 * 1024,
**kwargs,
) -> Union[Iterator[bytes], AsyncIterator[bytes]]:
"""
Prototype API: Returns a byte iterator for file contents.
Currently supports OpenAI provider only.
"""
try:
optional_params = GenericLiteLLMParams(**kwargs)
try:
if model is not None:
_, custom_llm_provider, _, _ = get_llm_provider(
model, custom_llm_provider
)
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
):
timeout = timeout.read or 600
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
timeout = float(timeout) # type: ignore
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:
response = openai_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=openai_creds.api_base,
api_key=openai_creds.api_key,
timeout=timeout,
max_retries=optional_params.max_retries,
organization=openai_creds.organization,
chunk_size=chunk_size,
)
else:
raise litellm.exceptions.BadRequestError(
message="LiteLLM doesn't support {} for 'file_content'. Supported providers are 'openai', 'azure', 'vertex_ai', 'bedrock', 'manus', 'anthropic'.".format(
custom_llm_provider
),
model="n/a",
llm_provider=custom_llm_provider,
response=httpx.Response(
status_code=400,
content="Unsupported provider",
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
),
)
return response
except Exception as e:
raise e

View file

@ -1751,6 +1751,63 @@ class OpenAIFilesAPI(BaseLLM):
return HttpxBinaryResponseContent(response=response.response)
async def afile_content_streaming(
self,
file_content_request: FileContentRequest,
openai_client: 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: str,
api_key: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
organization: Optional[str],
chunk_size: int = 1024 * 1024,
client: Optional[Union[OpenAI, AsyncOpenAI]] = None,
) -> Union[Iterator[bytes], AsyncIterator[bytes]]:
openai_client: Optional[Union[OpenAI, AsyncOpenAI]] = self.get_openai_client(
api_key=api_key,
api_base=api_base,
timeout=timeout,
max_retries=max_retries,
organization=organization,
client=client,
_is_async=_is_async,
)
if openai_client is None:
raise ValueError(
"OpenAI 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, AsyncOpenAI):
raise ValueError(
"OpenAI client is not an instance of AsyncOpenAI. Make sure you passed an AsyncOpenAI client."
)
return self.afile_content_streaming( # type: ignore
file_content_request=file_content_request,
openai_client=openai_client,
chunk_size=chunk_size,
)
def _stream() -> Iterator[bytes]:
with cast(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

@ -595,6 +595,7 @@ class ProxyBaseLLMRequestProcessing:
"alist_batches",
"acancel_batch",
"afile_content",
"afile_content_streaming",
"afile_retrieve",
"afile_delete",
"atext_completion",