mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
f3633c5af1
commit
bda8c0a1c3
6 changed files with 414 additions and 191 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"',
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue