mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
add tests for streaming file contents. Deprecate buffered version
This commit is contained in:
parent
06ddc2c0db
commit
95af7a55f6
6 changed files with 595 additions and 47 deletions
|
|
@ -4,7 +4,7 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, Iterator, List, Literal, Optional, Union, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -1472,6 +1472,47 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
else:
|
||||
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
||||
|
||||
async def afile_content_streaming(
|
||||
self,
|
||||
file_id: str,
|
||||
litellm_parent_otel_span: Optional[Span],
|
||||
llm_router: Router,
|
||||
chunk_size: int = 1024 * 1024,
|
||||
**data: Dict,
|
||||
) -> Union[Iterator[bytes], AsyncIterator[bytes]]:
|
||||
"""
|
||||
Stream the content of a file from the first model that has it.
|
||||
"""
|
||||
data.pop("llm_router", None)
|
||||
data.pop("litellm_parent_otel_span", None)
|
||||
data.pop("model", None)
|
||||
|
||||
model_file_id_mapping = data.pop("model_file_id_mapping", None)
|
||||
model_file_id_mapping = (
|
||||
model_file_id_mapping
|
||||
or await self.get_model_file_id_mapping([file_id], litellm_parent_otel_span)
|
||||
)
|
||||
|
||||
specific_model_file_id_mapping = model_file_id_mapping.get(file_id)
|
||||
|
||||
if specific_model_file_id_mapping:
|
||||
exception_dict = {}
|
||||
for model_id, provider_file_id in specific_model_file_id_mapping.items():
|
||||
try:
|
||||
return await litellm.afile_content_streaming(
|
||||
model=model_id,
|
||||
file_id=provider_file_id,
|
||||
chunk_size=chunk_size,
|
||||
**data, #type: ignore
|
||||
)
|
||||
except Exception as e:
|
||||
exception_dict[model_id] = str(e)
|
||||
raise Exception(
|
||||
f"LiteLLM Managed File object with id={file_id} not found. Checked model id's: {specific_model_file_id_mapping.keys()}. Errors: {exception_dict}"
|
||||
)
|
||||
else:
|
||||
raise Exception(f"LiteLLM Managed File object with id={file_id} not found")
|
||||
|
||||
async def _convert_storage_files_to_base64(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, Iterator, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
|
|
@ -259,3 +259,13 @@ class BaseFileEndpoints(ABC):
|
|||
**data: Dict,
|
||||
) -> "HttpxBinaryResponseContent":
|
||||
pass
|
||||
|
||||
async def afile_content_streaming(
|
||||
self,
|
||||
file_id: str,
|
||||
litellm_parent_otel_span: Optional[Span],
|
||||
llm_router: Router,
|
||||
chunk_size: int = 1024 * 1024,
|
||||
**data: Dict,
|
||||
) -> Union[Iterator[bytes], AsyncIterator[bytes]]:
|
||||
raise NotImplementedError()
|
||||
|
|
|
|||
|
|
@ -568,20 +568,20 @@ async def create_file( # noqa: PLR0915
|
|||
|
||||
|
||||
@router.get(
|
||||
"/{provider}/v1/files/{file_id:path}/content",
|
||||
"/{provider}/v0/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",
|
||||
"/v0/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"],
|
||||
# )
|
||||
async def get_file_content( # noqa: PLR0915
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
|
|
@ -825,14 +825,13 @@ async def get_file_content( # noqa: PLR0915
|
|||
)
|
||||
|
||||
|
||||
# NOTE: Rough Prototype to check memory usage
|
||||
@router.get(
|
||||
"/{provider}/v2/files/{file_id:path}/content",
|
||||
"/{provider}/v1/files/{file_id:path}/content",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["files"],
|
||||
)
|
||||
@router.get(
|
||||
"/v2/files/{file_id:path}/content",
|
||||
"/v1/files/{file_id:path}/content",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
tags=["files"],
|
||||
)
|
||||
|
|
@ -856,7 +855,7 @@ async def get_file_content_streaming(
|
|||
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
(
|
||||
data,
|
||||
litellm_logging_obj,
|
||||
_,
|
||||
) = await base_llm_response_processor.common_processing_pre_call_logic(
|
||||
request=request,
|
||||
general_settings=general_settings,
|
||||
|
|
@ -935,36 +934,21 @@ async def get_file_content_streaming(
|
|||
|
||||
model = cast(Optional[str], data.get("model"))
|
||||
if model:
|
||||
# TODO: Add Streaming version here
|
||||
# response = await llm_router.afile_content(
|
||||
# **{
|
||||
# "model": model,
|
||||
# "file_id": file_id,
|
||||
# **data,
|
||||
# }
|
||||
# ) # type: ignore
|
||||
raise ProxyException(
|
||||
message="Managed files streaming path is pending implementation",
|
||||
type="None",
|
||||
param="file_id",
|
||||
code=501,
|
||||
)
|
||||
|
||||
stream_iterator = await litellm.afile_content_streaming(
|
||||
**{
|
||||
"model": model,
|
||||
"file_id": file_id,
|
||||
**data,
|
||||
}
|
||||
) # type: ignore
|
||||
else:
|
||||
# TODO: Add Streaming version here
|
||||
# response = await managed_files_obj.afile_content(
|
||||
# **{
|
||||
# "file_id": file_id,
|
||||
# "litellm_parent_otel_span": user_api_key_dict.parent_otel_span,
|
||||
# "llm_router": llm_router,
|
||||
# **data,
|
||||
# }
|
||||
# )
|
||||
raise ProxyException(
|
||||
message="Managed files streaming path is pending implementation",
|
||||
type="None",
|
||||
param="file_id",
|
||||
code=501,
|
||||
stream_iterator = await managed_files_obj.afile_content_streaming(
|
||||
**{
|
||||
"file_id": file_id,
|
||||
"litellm_parent_otel_span": user_api_key_dict.parent_otel_span,
|
||||
"llm_router": llm_router,
|
||||
**data,
|
||||
}
|
||||
)
|
||||
else:
|
||||
# Check for model-based credential routing
|
||||
|
|
@ -989,6 +973,7 @@ async def get_file_content_streaming(
|
|||
file_id=original_file_id, # Use decoded file ID if from encoded ID
|
||||
)
|
||||
|
||||
data.pop("custom_llm_provider", None)
|
||||
stream_iterator = await litellm.afile_content_streaming(
|
||||
custom_llm_provider=credentials["custom_llm_provider"], # type: ignore
|
||||
**data,
|
||||
|
|
@ -1035,6 +1020,8 @@ async def get_file_content_streaming(
|
|||
},
|
||||
)
|
||||
except Exception as e:
|
||||
if isinstance(e, ProxyException):
|
||||
raise e
|
||||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
|
|
|
|||
|
|
@ -38,6 +38,11 @@ def test_get_file_ids_from_messages():
|
|||
]
|
||||
|
||||
|
||||
async def _fake_stream_iterator():
|
||||
yield b"stream-1"
|
||||
yield b"stream-2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_hook_batch_retrieve():
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -167,6 +172,40 @@ async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_uncha
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_afile_content_streaming_delegates_to_litellm(monkeypatch):
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
DualCache(), prisma_client=MagicMock()
|
||||
)
|
||||
proxy_managed_files.get_model_file_id_mapping = AsyncMock(
|
||||
return_value={"managed-file-abc": {"gpt-4o": "provider-file-xyz"}}
|
||||
)
|
||||
|
||||
mock_afile_content_streaming = AsyncMock(return_value=_fake_stream_iterator())
|
||||
monkeypatch.setattr("litellm.afile_content_streaming", mock_afile_content_streaming)
|
||||
|
||||
response = await proxy_managed_files.afile_content_streaming(
|
||||
file_id="managed-file-abc",
|
||||
litellm_parent_otel_span=MagicMock(),
|
||||
llm_router=MagicMock(),
|
||||
chunk_size=4096,
|
||||
request_timeout=30,
|
||||
model="ignored",
|
||||
)
|
||||
|
||||
chunks = []
|
||||
async for chunk in response:
|
||||
chunks.append(chunk)
|
||||
|
||||
assert chunks == [b"stream-1", b"stream-2"]
|
||||
mock_afile_content_streaming.assert_awaited_once()
|
||||
call_kwargs = mock_afile_content_streaming.call_args.kwargs
|
||||
assert call_kwargs["model"] == "gpt-4o"
|
||||
assert call_kwargs["file_id"] == "provider-file-xyz"
|
||||
assert call_kwargs["chunk_size"] == 4096
|
||||
assert call_kwargs["request_timeout"] == 30
|
||||
|
||||
|
||||
# def test_list_managed_files():
|
||||
# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache())
|
||||
|
||||
|
|
|
|||
|
|
@ -1,11 +1,11 @@
|
|||
from unittest.mock import MagicMock
|
||||
from unittest.mock import AsyncMock, 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
|
||||
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
|
||||
from litellm.proxy.openai_files_endpoints import files_endpoints
|
||||
|
||||
|
||||
async def _fake_stream_iterator():
|
||||
|
|
@ -23,6 +23,45 @@ class _FakeProxyLogging:
|
|||
return None
|
||||
|
||||
|
||||
class _FakeManagedFiles(BaseFileEndpoints):
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
async def acreate_file(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
async def afile_retrieve(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
async def afile_list(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
async def afile_delete(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
async def afile_content(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
async def afile_content_streaming(
|
||||
self,
|
||||
file_id: str,
|
||||
litellm_parent_otel_span,
|
||||
llm_router,
|
||||
chunk_size: int = 1024 * 1024,
|
||||
**data,
|
||||
):
|
||||
self.calls.append(
|
||||
{
|
||||
"file_id": file_id,
|
||||
"litellm_parent_otel_span": litellm_parent_otel_span,
|
||||
"llm_router": llm_router,
|
||||
"chunk_size": chunk_size,
|
||||
"data": data,
|
||||
}
|
||||
)
|
||||
return _fake_stream_iterator()
|
||||
|
||||
|
||||
async def _fake_common_processing_pre_call_logic(self, **kwargs):
|
||||
return self.data, MagicMock()
|
||||
|
||||
|
|
@ -66,7 +105,7 @@ async def test_get_file_content_streaming_returns_streaming_response(monkeypatch
|
|||
"test-version",
|
||||
)
|
||||
|
||||
response = await get_file_content_streaming(
|
||||
response = await files_endpoints.get_file_content_streaming(
|
||||
request=request_obj,
|
||||
fastapi_response=Response(),
|
||||
file_id="file-123",
|
||||
|
|
@ -81,4 +120,47 @@ async def test_get_file_content_streaming_returns_streaming_response(monkeypatch
|
|||
async for chunk in response.body_iterator:
|
||||
chunks.append(chunk)
|
||||
|
||||
assert chunks == [b"hello ", b"world"]
|
||||
assert chunks == [b"hello ", b"world"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_file_content_streaming_uses_managed_files_hook(monkeypatch, request_obj):
|
||||
fake_managed_files = _FakeManagedFiles()
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.get_proxy_hook.return_value = fake_managed_files
|
||||
proxy_logging.update_request_status = AsyncMock(return_value=None)
|
||||
proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
|
||||
async def _fake_common_processing_pre_call_logic(self, **kwargs):
|
||||
return self.data, MagicMock()
|
||||
|
||||
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", proxy_logging)
|
||||
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")
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock())
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.openai_files_endpoints.files_endpoints._is_base64_encoded_unified_file_id",
|
||||
lambda file_id: file_id,
|
||||
)
|
||||
|
||||
response = await files_endpoints.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)
|
||||
|
||||
chunks = []
|
||||
async for chunk in response.body_iterator:
|
||||
chunks.append(chunk)
|
||||
|
||||
assert chunks == [b"hello ", b"world"]
|
||||
assert fake_managed_files.calls[0]["file_id"].startswith("file-123")
|
||||
|
|
@ -1,7 +1,9 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import ANY
|
||||
from typing import Any, Optional, cast
|
||||
|
||||
import pytest
|
||||
import respx
|
||||
|
|
@ -15,9 +17,14 @@ sys.path.insert(
|
|||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.proxy._types import LiteLLM_UserTableFiltered, UserAPIKeyAuth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.enterprise.litellm_enterprise.proxy.hooks.managed_files import (
|
||||
_PROXY_LiteLLMManagedFiles,
|
||||
)
|
||||
from litellm.proxy.hooks import get_proxy_hook
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import ui_view_users
|
||||
from litellm.proxy.proxy_server import app
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import encode_file_id_with_model
|
||||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
|
||||
client = TestClient(app)
|
||||
|
|
@ -77,6 +84,67 @@ def setup_proxy_logging_object(monkeypatch, llm_router: Router) -> ProxyLogging:
|
|||
return proxy_logging_object
|
||||
|
||||
|
||||
async def _stream_bytes(*chunks: bytes):
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
|
||||
def _build_managed_files_hook(prisma_client=None) -> _PROXY_LiteLLMManagedFiles:
|
||||
return _PROXY_LiteLLMManagedFiles(
|
||||
internal_usage_cache=cast(Any, DualCache(default_in_memory_ttl=1)),
|
||||
prisma_client=cast(Any, prisma_client),
|
||||
)
|
||||
|
||||
|
||||
def _setup_streaming_proxy(
|
||||
monkeypatch, llm_router: Router, managed_files_obj
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
|
||||
|
||||
proxy_logging_obj = ProxyLogging(
|
||||
user_api_key_cache=DualCache(default_in_memory_ttl=1)
|
||||
)
|
||||
proxy_logging_obj.proxy_hook_mapping["managed_files"] = managed_files_obj
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
|
||||
)
|
||||
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id="test-user",
|
||||
)
|
||||
|
||||
return proxy_logging_obj
|
||||
|
||||
|
||||
def _patch_managed_file_detection(monkeypatch, is_managed: bool) -> None:
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.openai_files_endpoints.files_endpoints._is_base64_encoded_unified_file_id",
|
||||
lambda file_id: file_id if is_managed else False,
|
||||
)
|
||||
|
||||
|
||||
def _patch_common_processing(
|
||||
monkeypatch, model: Optional[str] = None
|
||||
) -> None:
|
||||
async def fake_common_processing(self, **kwargs):
|
||||
data = {**self.data}
|
||||
if model is not None:
|
||||
data["model"] = model
|
||||
return data, None
|
||||
|
||||
monkeypatch.setattr(
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
"common_processing_pre_call_logic",
|
||||
fake_common_processing,
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_purpose(mocker: MockerFixture, monkeypatch, llm_router: Router):
|
||||
"""
|
||||
Asserts 'create_file' is called with the correct arguments
|
||||
|
|
@ -1552,3 +1620,324 @@ def test_file_invalid_anchor_returns_500(
|
|||
)
|
||||
assert response.status_code == 500
|
||||
assert "created_at" in response.json()["error"]["message"]
|
||||
|
||||
|
||||
def test_get_file_content_streaming_uses_storage_backend(
|
||||
mocker: MockerFixture, monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
encoded_file_id = encode_file_id_with_model("file-abc123", "gpt-3.5-turbo")
|
||||
hook = _build_managed_files_hook()
|
||||
proxy_logging_obj = _setup_streaming_proxy(monkeypatch, llm_router, hook)
|
||||
_patch_managed_file_detection(monkeypatch, True)
|
||||
_patch_common_processing(monkeypatch)
|
||||
|
||||
find_first = mocker.AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
storage_backend="azure_storage",
|
||||
storage_url="https://storage.example/file.bin",
|
||||
)
|
||||
)
|
||||
hook.prisma_client = SimpleNamespace(
|
||||
db=SimpleNamespace(
|
||||
litellm_managedfiletable=SimpleNamespace(find_first=find_first)
|
||||
)
|
||||
)
|
||||
|
||||
class DummyStorageBackend:
|
||||
async def download_file_streaming(self, storage_url, chunk_size=1024 * 1024):
|
||||
assert storage_url == "https://storage.example/file.bin"
|
||||
yield b"storage-"
|
||||
yield b"stream"
|
||||
|
||||
mocker.patch(
|
||||
"litellm.llms.base_llm.files.storage_backend_factory.get_storage_backend",
|
||||
return_value=DummyStorageBackend(),
|
||||
)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
f"/v2/files/{encoded_file_id}/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.content == b"storage-stream"
|
||||
assert (
|
||||
response.headers["content-disposition"]
|
||||
== f'attachment; filename="{encoded_file_id}.bin"'
|
||||
)
|
||||
find_first.assert_awaited_once_with(where={"unified_file_id": encoded_file_id})
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_get_file_content_streaming_uses_managed_files_hook(
|
||||
mocker: MockerFixture, monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
encoded_file_id = encode_file_id_with_model("file-abc123", "gpt-3.5-turbo")
|
||||
hook = _build_managed_files_hook()
|
||||
_setup_streaming_proxy(monkeypatch, llm_router, hook)
|
||||
_patch_managed_file_detection(monkeypatch, True)
|
||||
_patch_common_processing(monkeypatch)
|
||||
|
||||
hook.get_model_file_id_mapping = mocker.AsyncMock(
|
||||
return_value={
|
||||
encoded_file_id: {"azure-gpt-3-5-turbo": "file-provider-123"}
|
||||
}
|
||||
)
|
||||
stream_mock = mocker.patch(
|
||||
"litellm.afile_content_streaming",
|
||||
new=mocker.AsyncMock(return_value=_stream_bytes(b"managed-", b"hook")),
|
||||
)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
f"/v2/files/{encoded_file_id}/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.content == b"managed-hook"
|
||||
hook.get_model_file_id_mapping.assert_awaited_once_with(
|
||||
[encoded_file_id], None
|
||||
)
|
||||
stream_mock.assert_awaited_once()
|
||||
called_kwargs = stream_mock.await_args.kwargs
|
||||
assert called_kwargs["model"] == "azure-gpt-3-5-turbo"
|
||||
assert called_kwargs["file_id"] == "file-provider-123"
|
||||
assert "llm_router" not in called_kwargs
|
||||
assert "litellm_parent_otel_span" not in called_kwargs
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_get_file_content_streaming_uses_litellm_for_managed_model_branch(
|
||||
mocker: MockerFixture, monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
encoded_file_id = encode_file_id_with_model("file-abc123", "gpt-3.5-turbo")
|
||||
hook = _build_managed_files_hook()
|
||||
_setup_streaming_proxy(monkeypatch, llm_router, hook)
|
||||
_patch_common_processing(monkeypatch, model="gpt-4o")
|
||||
|
||||
stream_mock = mocker.patch(
|
||||
"litellm.afile_content_streaming",
|
||||
new=mocker.AsyncMock(return_value=_stream_bytes(b"model-", b"stream")),
|
||||
)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
f"/v2/files/{encoded_file_id}/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.content == b"model-stream"
|
||||
stream_mock.assert_awaited_once()
|
||||
called_kwargs = stream_mock.await_args.kwargs
|
||||
assert called_kwargs["model"] == "gpt-4o"
|
||||
assert called_kwargs["file_id"] == "file-abc123"
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_get_file_content_streaming_routes_model_request(
|
||||
mocker: MockerFixture, monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
hook = _build_managed_files_hook()
|
||||
_setup_streaming_proxy(monkeypatch, llm_router, hook)
|
||||
_patch_common_processing(monkeypatch)
|
||||
|
||||
handle_routing_mock = mocker.patch(
|
||||
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
|
||||
return_value=(
|
||||
True,
|
||||
"routed-model",
|
||||
"file-original-123",
|
||||
{
|
||||
"custom_llm_provider": "azure",
|
||||
"api_key": "routed-api-key",
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
def fake_prepare_data_with_credentials(data, credentials, file_id):
|
||||
data.update(credentials)
|
||||
data["file_id"] = file_id
|
||||
|
||||
mocker.patch(
|
||||
"litellm.proxy.openai_files_endpoints.files_endpoints.prepare_data_with_credentials",
|
||||
side_effect=fake_prepare_data_with_credentials,
|
||||
)
|
||||
stream_mock = mocker.patch(
|
||||
"litellm.afile_content_streaming",
|
||||
new=mocker.AsyncMock(return_value=_stream_bytes(b"routed-", b"stream")),
|
||||
)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
"/v2/files/file-plain-123/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.content == b"routed-stream"
|
||||
handle_routing_mock.assert_called_once()
|
||||
stream_mock.assert_awaited_once()
|
||||
called_kwargs = stream_mock.await_args.kwargs
|
||||
assert called_kwargs["custom_llm_provider"] == "azure"
|
||||
assert called_kwargs["file_id"] == "file-original-123"
|
||||
assert called_kwargs["api_key"] == "routed-api-key"
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_get_file_content_streaming_falls_back_to_provider(
|
||||
mocker: MockerFixture, monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
hook = _build_managed_files_hook()
|
||||
_setup_streaming_proxy(monkeypatch, llm_router, hook)
|
||||
_patch_managed_file_detection(monkeypatch, False)
|
||||
_patch_common_processing(monkeypatch)
|
||||
|
||||
mocker.patch(
|
||||
"litellm.proxy.openai_files_endpoints.files_endpoints.handle_model_based_routing",
|
||||
return_value=(False, None, None, {}),
|
||||
)
|
||||
stream_mock = mocker.patch(
|
||||
"litellm.afile_content_streaming",
|
||||
new=mocker.AsyncMock(return_value=_stream_bytes(b"fallback-", b"stream")),
|
||||
)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
"/v2/files/file-plain-456/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.content == b"fallback-stream"
|
||||
stream_mock.assert_awaited_once()
|
||||
called_kwargs = stream_mock.await_args.kwargs
|
||||
assert called_kwargs["custom_llm_provider"] == "openai"
|
||||
assert called_kwargs["file_id"] == "file-plain-456"
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_get_file_content_streaming_errors_when_managed_files_hook_missing(
|
||||
monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
encoded_file_id = encode_file_id_with_model("file-abc123", "gpt-3.5-turbo")
|
||||
_setup_streaming_proxy(monkeypatch, llm_router, None)
|
||||
_patch_managed_file_detection(monkeypatch, True)
|
||||
_patch_common_processing(monkeypatch)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
f"/v2/files/{encoded_file_id}/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 500
|
||||
assert "Managed files hook not found" in response.text
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_get_file_content_streaming_errors_when_managed_files_hook_is_wrong_type(
|
||||
monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
encoded_file_id = encode_file_id_with_model("file-abc123", "gpt-3.5-turbo")
|
||||
_setup_streaming_proxy(monkeypatch, llm_router, object())
|
||||
_patch_managed_file_detection(monkeypatch, True)
|
||||
_patch_common_processing(monkeypatch)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
f"/v2/files/{encoded_file_id}/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 500
|
||||
assert "Managed files hook is not a BaseFileEndpoints" in response.text
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_get_file_content_streaming_errors_when_router_is_missing(
|
||||
monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
encoded_file_id = encode_file_id_with_model("file-abc123", "gpt-3.5-turbo")
|
||||
hook = _build_managed_files_hook()
|
||||
_setup_streaming_proxy(monkeypatch, llm_router, hook)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
_patch_managed_file_detection(monkeypatch, True)
|
||||
_patch_common_processing(monkeypatch)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
f"/v2/files/{encoded_file_id}/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 500
|
||||
assert "LLM Router not found" in response.text
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_get_file_content_streaming_errors_when_storage_backend_is_invalid(
|
||||
mocker: MockerFixture, monkeypatch, llm_router: Router
|
||||
):
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
encoded_file_id = encode_file_id_with_model("file-abc123", "gpt-3.5-turbo")
|
||||
hook = _build_managed_files_hook()
|
||||
_setup_streaming_proxy(monkeypatch, llm_router, hook)
|
||||
_patch_managed_file_detection(monkeypatch, True)
|
||||
_patch_common_processing(monkeypatch)
|
||||
|
||||
hook.prisma_client = SimpleNamespace(
|
||||
db=SimpleNamespace(
|
||||
litellm_managedfiletable=SimpleNamespace(
|
||||
find_first=mocker.AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
storage_backend="bad-backend",
|
||||
storage_url="https://storage.example/file.bin",
|
||||
)
|
||||
)
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
mocker.patch(
|
||||
"litellm.llms.base_llm.files.storage_backend_factory.get_storage_backend",
|
||||
side_effect=ValueError("backend missing"),
|
||||
)
|
||||
|
||||
try:
|
||||
response = client.get(
|
||||
f"/v2/files/{encoded_file_id}/content",
|
||||
headers={"Authorization": "Bearer test-key"},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert "Storage backend error" in response.text
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue