backwards compatibility resolution and code fixes suggestions by greptile

This commit is contained in:
harish876 2026-04-08 03:20:50 +00:00
parent 8aef24c871
commit 6ab02e6e51
5 changed files with 65 additions and 31 deletions

View file

@ -13,7 +13,6 @@ from functools import partial
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[
@ -998,7 +997,7 @@ async def afile_content_streaming(
Async wrapper for file_content_streaming.
"""
try:
loop = asyncio.get_event_loop()
loop = asyncio.get_running_loop()
kwargs["afile_content_streaming"] = True
model = kwargs.pop("model", None)

View file

@ -349,14 +349,25 @@ class AzureBlobStorageBackend(BaseFileStorageBackend, AzureBlobStorageLogger):
file_client = file_system_client.get_file_client(file_path)
download_response = await file_client.download_file()
if not hasattr(download_response, "chunks"):
chunks_method = getattr(download_response, "chunks", None)
if chunks_method is None or not callable(chunks_method):
raise RuntimeError(
"Azure SDK download response does not support chunk streaming"
)
async for chunk in download_response.chunks():
if chunk:
yield chunk
chunks_iterable = chunks_method()
if hasattr(chunks_iterable, "__aiter__"):
async for chunk in chunks_iterable: #type: ignore
if chunk:
yield chunk
elif hasattr(chunks_iterable, "__iter__"):
for chunk in chunks_iterable: #type: ignore
if chunk:
yield chunk
else:
raise RuntimeError(
"Azure SDK download response chunks() did not return an iterable"
)
async def _download_file_with_azure_ad(self, file_path: str) -> bytes:
"""Download file using REST API with Azure AD token."""

View file

@ -568,12 +568,17 @@ async def create_file( # noqa: PLR0915
@router.get(
"/{provider}/v0/files/{file_id:path}/content",
"/{provider}/v1/files/{file_id:path}/content",
dependencies=[Depends(user_api_key_auth)],
tags=["files"],
)
@router.get(
"/v0/files/{file_id:path}/content",
"/v1/files/{file_id:path}/content",
dependencies=[Depends(user_api_key_auth)],
tags=["files"],
)
@router.get(
"/files/{file_id:path}/content",
dependencies=[Depends(user_api_key_auth)],
tags=["files"],
)
@ -821,17 +826,12 @@ async def get_file_content( # noqa: PLR0915
@router.get(
"/{provider}/v1/files/{file_id:path}/content",
"/{provider}/v2/files/{file_id:path}/content",
dependencies=[Depends(user_api_key_auth)],
tags=["files"],
)
@router.get(
"/v1/files/{file_id:path}/content",
dependencies=[Depends(user_api_key_auth)],
tags=["files"],
)
@router.get(
"/files/{file_id:path}/content",
"/v2/files/{file_id:path}/content",
dependencies=[Depends(user_api_key_auth)],
tags=["files"],
)
@ -995,11 +995,35 @@ async def get_file_content_streaming(
**data,
)
asyncio.create_task(
proxy_logging_obj.update_request_status(
litellm_call_id=data.get("litellm_call_id", ""), status="success"
)
)
async def _stream_with_logging():
try:
# Handle both async and sync iterators
if hasattr(stream_iterator, '__aiter__'):
async for chunk in stream_iterator: # type: ignore
yield chunk
else:
for chunk in stream_iterator:# type: ignore
yield chunk
asyncio.create_task(
proxy_logging_obj.update_request_status(
litellm_call_id=data.get("litellm_call_id", ""), status="success"
)
)
except Exception as e:
verbose_proxy_logger.exception(
"File streaming failed mid-transfer - {}".format(str(e))
)
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
request_data=data,
)
asyncio.create_task(
proxy_logging_obj.update_request_status(
litellm_call_id=data.get("litellm_call_id", ""), status="fail"
)
)
raise
fastapi_response.headers.update(
ProxyBaseLLMRequestProcessing.get_custom_headers(
@ -1013,7 +1037,7 @@ async def get_file_content_streaming(
)
return StreamingResponse(
content=stream_iterator,
content=_stream_with_logging(),
media_type="application/octet-stream",
headers={
"content-disposition": f'attachment; filename="{file_id}.bin"',

View file

@ -71,7 +71,7 @@ def request_obj() -> Request:
scope = {
"type": "http",
"method": "GET",
"path": "/v1/files/file-123/content",
"path": "/v2/files/file-123/content",
"headers": [],
"query_string": b"",
}

View file

@ -1658,7 +1658,7 @@ def test_get_file_content_streaming_uses_storage_backend(
try:
response = client.get(
f"/v1/files/{encoded_file_id}/content",
f"/v2/files/{encoded_file_id}/content",
headers={"Authorization": "Bearer test-key"},
)
@ -1696,7 +1696,7 @@ def test_get_file_content_streaming_uses_managed_files_hook(
try:
response = client.get(
f"/v1/files/{encoded_file_id}/content",
f"/v2/files/{encoded_file_id}/content",
headers={"Authorization": "Bearer test-key"},
)
@ -1732,7 +1732,7 @@ def test_get_file_content_streaming_uses_litellm_for_managed_model_branch(
try:
response = client.get(
f"/v1/files/{encoded_file_id}/content",
f"/v2/files/{encoded_file_id}/content",
headers={"Authorization": "Bearer test-key"},
)
@ -1783,7 +1783,7 @@ def test_get_file_content_streaming_routes_model_request(
try:
response = client.get(
"/v1/files/file-plain-123/content",
"/v2/files/file-plain-123/content",
headers={"Authorization": "Bearer test-key"},
)
@ -1820,7 +1820,7 @@ def test_get_file_content_streaming_falls_back_to_provider(
try:
response = client.get(
"/v1/files/file-plain-456/content",
"/v2/files/file-plain-456/content",
headers={"Authorization": "Bearer test-key"},
)
@ -1846,7 +1846,7 @@ def test_get_file_content_streaming_errors_when_managed_files_hook_missing(
try:
response = client.get(
f"/v1/files/{encoded_file_id}/content",
f"/v2/files/{encoded_file_id}/content",
headers={"Authorization": "Bearer test-key"},
)
@ -1868,7 +1868,7 @@ def test_get_file_content_streaming_errors_when_managed_files_hook_is_wrong_type
try:
response = client.get(
f"/v1/files/{encoded_file_id}/content",
f"/v2/files/{encoded_file_id}/content",
headers={"Authorization": "Bearer test-key"},
)
@ -1892,7 +1892,7 @@ def test_get_file_content_streaming_errors_when_router_is_missing(
try:
response = client.get(
f"/v1/files/{encoded_file_id}/content",
f"/v2/files/{encoded_file_id}/content",
headers={"Authorization": "Bearer test-key"},
)
@ -1933,7 +1933,7 @@ def test_get_file_content_streaming_errors_when_storage_backend_is_invalid(
try:
response = client.get(
f"/v1/files/{encoded_file_id}/content",
f"/v2/files/{encoded_file_id}/content",
headers={"Authorization": "Bearer test-key"},
)