mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
* ci: fix the litellm-tests unit job with sysmon coverage, an env allowlist and coverage upload on failure * test: replace key-dependent proxy, enterprise and mcp unit tests with synthetic values and integration and e2e coverage * test: drop key reads at the legacy proxy, enterprise and mcp paths and wire the gemini pass-through split * ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests under their legacy flags * ci: move caching, proxy-extras, gateway and enterprise tests into tests/unit and run them from litellm-tests under their legacy flags * ci: fail the unit shard when circleci tests split errors * test: drop restating comments from the gemini pass-through split * ci: exit the unit shard cleanly when circleci tests split assigns it no files --------- Co-authored-by: yuneng <yuneng@berri.ai>
3707 lines
126 KiB
Python
3707 lines
126 KiB
Python
import base64
|
|
import json
|
|
from typing import cast
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
from fastapi import HTTPException
|
|
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
|
|
|
|
from litellm.caching import DualCache
|
|
from litellm.proxy._types import CallTypes
|
|
from litellm.proxy.openai_files_endpoints.common_utils import (
|
|
BATCH_CREATE_HIDDEN_PARAM,
|
|
_is_base64_encoded_unified_file_id,
|
|
encode_file_id_with_model,
|
|
)
|
|
|
|
|
|
def test_get_file_ids_from_messages():
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{"type": "text", "text": "What is in this recording?"},
|
|
{
|
|
"type": "file",
|
|
"file": {
|
|
"file_id": "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCxmYzdmMmVhNS0wZjUwLTQ5ZjYtODljMS03ZTZhNTRiMTIxMzg",
|
|
},
|
|
},
|
|
],
|
|
},
|
|
]
|
|
file_ids = proxy_managed_files.get_file_ids_from_messages(messages)
|
|
assert file_ids == [
|
|
"bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCxmYzdmMmVhNS0wZjUwLTQ5ZjYtODljMS03ZTZhNTRiMTIxMzg"
|
|
]
|
|
|
|
|
|
def test_get_file_ids_from_messages_skips_bedrock_content_blocks_without_type():
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{"text": "What is Apptio?"},
|
|
{
|
|
"toolResult": {
|
|
"toolUseId": "tooluse_123",
|
|
"status": "success",
|
|
"content": [
|
|
{
|
|
"searchResult": {
|
|
"source": "source",
|
|
"title": "title",
|
|
"content": [{"text": "snippet"}],
|
|
"citations": {"enabled": True},
|
|
}
|
|
}
|
|
],
|
|
}
|
|
},
|
|
{"type": "file", "file": {"file_id": "file-keep"}},
|
|
],
|
|
}
|
|
]
|
|
file_ids = proxy_managed_files.get_file_ids_from_messages(messages)
|
|
assert file_ids == ["file-keep"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_pre_call_hook_batch_retrieve():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
return_value = MagicMock()
|
|
return_value.created_by = "123"
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
data = {
|
|
"user_api_key_dict": UserAPIKeyAuth(
|
|
user_id="123", parent_otel_span=MagicMock()
|
|
),
|
|
"data": {
|
|
"batch_id": "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1nZW5lcmFsLWF6dXJlLWRlcGxveW1lbnQ7bGxtX2JhdGNoX2lkOmJhdGNoX2EzMjJiNmJhLWFjN2UtNDg4OC05MjljLTFhZDM0NDJmMDZlZA",
|
|
},
|
|
"call_type": "aretrieve_batch",
|
|
"cache": MagicMock(),
|
|
}
|
|
response = await proxy_managed_files.async_pre_call_hook(**data)
|
|
assert response["batch_id"] == "batch_a322b6ba-ac7e-4888-929c-1ad3442f06ed"
|
|
assert response["model"] == "my-general-azure-deployment"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_user_batches_limit_zero_returns_empty_page_without_db_query():
|
|
"""OpenAI parity for GET /v1/batches?limit=0: an empty page, never the
|
|
default page of 20 (issue #37149). `min(limit or 20, 100)` treated 0 as
|
|
unset before this regression guard existed."""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = MagicMock()
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache(), prisma_client=prisma_client)
|
|
|
|
page = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="123"),
|
|
limit=0,
|
|
)
|
|
|
|
assert page == {"object": "list", "data": [], "first_id": None, "last_id": None, "has_more": False}
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_pre_call_deployment_hook_resolves_model_id_from_litellm_metadata():
|
|
"""
|
|
For batch operations the router stores model_info under
|
|
kwargs["litellm_metadata"]["model_info"] (not top-level kwargs["model_info"]).
|
|
async_pre_call_deployment_hook must check both locations so the managed
|
|
file ID is resolved to the provider-specific file ID.
|
|
"""
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
|
|
managed_file_id = "managed-file-abc"
|
|
model_id = "deployment-xyz"
|
|
provider_file_id = "gs://bucket/path/to/file.jsonl"
|
|
|
|
# model_info is nested under litellm_metadata (batch path)
|
|
kwargs = {
|
|
"input_file_id": managed_file_id,
|
|
"model_file_id_mapping": {
|
|
managed_file_id: {model_id: provider_file_id},
|
|
},
|
|
"litellm_metadata": {
|
|
"model_info": {"id": model_id},
|
|
},
|
|
}
|
|
|
|
result = await proxy_managed_files.async_pre_call_deployment_hook(
|
|
kwargs=kwargs, call_type=CallTypes.acreate_batch
|
|
)
|
|
|
|
assert (
|
|
result["input_file_id"] == provider_file_id
|
|
), f"Expected provider file ID '{provider_file_id}', got '{result['input_file_id']}'"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_pre_call_deployment_hook_prefers_top_level_model_info():
|
|
"""
|
|
When model_info exists at top-level kwargs, async_pre_call_deployment_hook
|
|
should use it without falling back to litellm_metadata.
|
|
"""
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
|
|
managed_file_id = "managed-file-abc"
|
|
top_level_model_id = "deployment-top"
|
|
nested_model_id = "deployment-nested"
|
|
top_level_provider_file = "file-top-123"
|
|
nested_provider_file = "file-nested-456"
|
|
|
|
kwargs = {
|
|
"input_file_id": managed_file_id,
|
|
"model_file_id_mapping": {
|
|
managed_file_id: {
|
|
top_level_model_id: top_level_provider_file,
|
|
nested_model_id: nested_provider_file,
|
|
},
|
|
},
|
|
"model_info": {"id": top_level_model_id},
|
|
"litellm_metadata": {
|
|
"model_info": {"id": nested_model_id},
|
|
},
|
|
}
|
|
|
|
result = await proxy_managed_files.async_pre_call_deployment_hook(
|
|
kwargs=kwargs, call_type=CallTypes.acreate_batch
|
|
)
|
|
|
|
assert (
|
|
result["input_file_id"] == top_level_provider_file
|
|
), "Should prefer top-level model_info over litellm_metadata"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_unchanged():
|
|
"""
|
|
When model_info is absent from both top-level and litellm_metadata,
|
|
the managed file ID should remain unchanged.
|
|
"""
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
|
|
managed_file_id = "managed-file-abc"
|
|
|
|
kwargs = {
|
|
"input_file_id": managed_file_id,
|
|
"model_file_id_mapping": {
|
|
managed_file_id: {"some-model": "provider-file-xyz"},
|
|
},
|
|
}
|
|
|
|
result = await proxy_managed_files.async_pre_call_deployment_hook(
|
|
kwargs=kwargs, call_type=CallTypes.acreate_batch
|
|
)
|
|
|
|
assert (
|
|
result["input_file_id"] == managed_file_id
|
|
), "File ID should remain unchanged when model_info is not available"
|
|
|
|
|
|
# def test_list_managed_files():
|
|
# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache())
|
|
|
|
# # Create some test files
|
|
# file1 = proxy_managed_files.create_file(
|
|
# file=("test1.txt", b"test content 1", "text/plain"),
|
|
# purpose="assistants"
|
|
# )
|
|
# file2 = proxy_managed_files.create_file(
|
|
# file=("test2.pdf", b"test content 2", "application/pdf"),
|
|
# purpose="assistants"
|
|
# )
|
|
|
|
# # List all files
|
|
# files = proxy_managed_files.list_files()
|
|
|
|
# # Verify response
|
|
# assert len(files) == 2
|
|
# assert all(f.id.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value) for f in files)
|
|
# assert any(f.filename == "test1.txt" for f in files)
|
|
# assert any(f.filename == "test2.pdf" for f in files)
|
|
# assert all(f.purpose == "assistants" for f in files)
|
|
|
|
# def test_retrieve_managed_file():
|
|
# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache())
|
|
|
|
# # Create a test file
|
|
# test_content = b"test content for retrieve"
|
|
# created_file = proxy_managed_files.create_file(
|
|
# file=("test.txt", test_content, "text/plain"),
|
|
# purpose="assistants"
|
|
# )
|
|
|
|
# # Retrieve the file
|
|
# retrieved_file = proxy_managed_files.retrieve_file(created_file.id)
|
|
|
|
# # Verify response
|
|
# assert retrieved_file.id == created_file.id
|
|
# assert retrieved_file.filename == "test.txt"
|
|
# assert retrieved_file.purpose == "assistants"
|
|
# assert retrieved_file.bytes == len(test_content)
|
|
# assert retrieved_file.status == "uploaded"
|
|
|
|
# def test_delete_managed_file():
|
|
# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache())
|
|
|
|
# # Create a test file
|
|
# created_file = proxy_managed_files.create_file(
|
|
# file=("test.txt", b"test content", "text/plain"),
|
|
# purpose="assistants"
|
|
# )
|
|
|
|
# # Delete the file
|
|
# deleted_file = proxy_managed_files.delete_file(created_file.id)
|
|
|
|
# # Verify deletion
|
|
# assert deleted_file.id == created_file.id
|
|
# assert deleted_file.deleted == True
|
|
|
|
# # Verify file is no longer retrievable
|
|
# with pytest.raises(Exception):
|
|
# proxy_managed_files.retrieve_file(created_file.id)
|
|
|
|
# # Verify file is not in list
|
|
# files = proxy_managed_files.list_files()
|
|
# assert created_file.id not in [f.id for f in files]
|
|
|
|
# def test_retrieve_nonexistent_file():
|
|
# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache())
|
|
|
|
# # Try to retrieve a non-existent file
|
|
# with pytest.raises(Exception):
|
|
# proxy_managed_files.retrieve_file("nonexistent-file-id")
|
|
|
|
# def test_delete_nonexistent_file():
|
|
# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache())
|
|
|
|
# # Try to delete a non-existent file
|
|
# with pytest.raises(Exception):
|
|
# proxy_managed_files.delete_file("nonexistent-file-id")
|
|
|
|
# def test_list_files_with_purpose_filter():
|
|
# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache())
|
|
|
|
# # Create files with different purposes
|
|
# file1 = proxy_managed_files.create_file(
|
|
# file=("test1.txt", b"test content 1", "text/plain"),
|
|
# purpose="assistants"
|
|
# )
|
|
# file2 = proxy_managed_files.create_file(
|
|
# file=("test2.pdf", b"test content 2", "application/pdf"),
|
|
# purpose="batch"
|
|
# )
|
|
|
|
# # List files with purpose filter
|
|
# assistant_files = proxy_managed_files.list_files(purpose="assistants")
|
|
# batch_files = proxy_managed_files.list_files(purpose="batch")
|
|
|
|
# # Verify filtering
|
|
# assert len(assistant_files) == 1
|
|
# assert len(batch_files) == 1
|
|
# assert assistant_files[0].id == file1.id
|
|
# assert batch_files[0].id == file2.id
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_post_call_success_hook_for_unified_finetuning_job():
|
|
from litellm.types.utils import LiteLLMFineTuningJob
|
|
|
|
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxiZTQ0ZDVlYi1mNDU3LTRiNzktOWM4My01N2QxMTMxYWM0YzY7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00LjEtb3BlbmFpO2xsbV9vdXRwdXRfZmlsZV9pZCxmaWxlLURKMnQ0OWZlQ2NTQk5vNG9oekZ6NGc7bGxtX291dHB1dF9maWxlX21vZGVsX2lkLGRiNjY5ODcwNzdkZTdmYzZjNzAzY2Y1MDczMGU2MmNkOWQ3YTU1N2NlNjVmMDUzNTFkYTM4YTA3ZjBlZDEyNzQ"
|
|
provider_ft_job = LiteLLMFineTuningJob(
|
|
object="fine_tuning.job",
|
|
id="ftjob-0kEBV5b4sPrFcMnuzmYSzU1G",
|
|
model="gpt-3.5-turbo-0613",
|
|
created_at=1692779769,
|
|
finished_at=None,
|
|
fine_tuned_model=None,
|
|
organization_id="org-dUVLhaAQ37YCGwVC2QVY8sdB",
|
|
result_files=[],
|
|
status="validating_files",
|
|
validation_file=None,
|
|
training_file="file-azQuKMLAmiFdEjxpCcbI11zF",
|
|
hyperparameters={"n_epochs": 8},
|
|
trained_tokens=None,
|
|
seed=0,
|
|
)
|
|
provider_ft_job._hidden_params = {
|
|
"unified_file_id": unified_file_id,
|
|
"model_id": "gpt-3.5-turbo-0613",
|
|
}
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=AsyncMock()
|
|
)
|
|
data = {
|
|
"user_api_key_dict": {"parent_otel_span": MagicMock()},
|
|
}
|
|
|
|
response = await proxy_managed_files.async_post_call_success_hook(
|
|
data=data,
|
|
user_api_key_dict=MagicMock(),
|
|
response=provider_ft_job,
|
|
)
|
|
|
|
assert isinstance(response, LiteLLMFineTuningJob)
|
|
assert _is_base64_encoded_unified_file_id(response.id)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_pre_call_hook_for_unified_finetuning_job():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
return_value = MagicMock()
|
|
return_value.created_by = "123"
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
data = {
|
|
"user_api_key_dict": UserAPIKeyAuth(
|
|
user_id="123", parent_otel_span=MagicMock()
|
|
),
|
|
"data": {
|
|
"fine_tuning_job_id": "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDo0OTIxODU4MWY3OGViZTllZjE4NDE0ZmE0ZjdmYjlmYTc0YzA5NWVkMTEyY2E4NDBkZDU2ZGZmZTliZDMwZGQxO2dlbmVyaWNfcmVzcG9uc2VfaWQ6ZnRqb2ItalRCeXM3YlZzYnlaRE93TDlHbHBZcVhS",
|
|
},
|
|
"call_type": "acancel_fine_tuning_job",
|
|
"cache": MagicMock(),
|
|
}
|
|
|
|
response = await proxy_managed_files.async_pre_call_hook(**data)
|
|
assert response["fine_tuning_job_id"] == "ftjob-jTBys7bVsbyZDOwL9GlpYqXR"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"call_type", ["afile_content", "afile_delete", "afile_retrieve"]
|
|
)
|
|
async def test_can_user_call_unified_file_id(call_type):
|
|
"""
|
|
Test that on file retrieve, delete, and content we check if the user has access to the file
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
return_value = MagicMock()
|
|
return_value.created_by = "123"
|
|
prisma_client.db.litellm_managedfiletable.find_first.return_value = return_value
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(), prisma_client=prisma_client
|
|
)
|
|
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxmMTNlNDAzZS01YWM3LTRhZjktOGQzNS0wNDgwZDMxOTgyYTg7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00by1taW5pLW9wZW5haTtsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1Ib3UxZDFXc3c1SDNKcjFMYllpZDJiO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxmODBiNWU2NzQ1NzdkNjkyMjM4YmVhNTIxZDdiMGI5ZGYyY2FmMTEwMTU2YmU5YzBjM2NjMmNkNTBjOTM1ZDI0"
|
|
|
|
with pytest.raises(HTTPException):
|
|
await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="456", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"file_id": unified_file_id},
|
|
call_type=call_type,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_router_acreate_batch_only_selects_from_file_id_mapping(monkeypatch):
|
|
"""
|
|
Test that router.acreate_batch only selects model_id from the file_id_mapping
|
|
"""
|
|
import litellm
|
|
|
|
prisma_client = AsyncMock()
|
|
return_value = MagicMock()
|
|
return_value.created_by = "123"
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
monkeypatch.setattr(
|
|
litellm,
|
|
"callbacks",
|
|
[proxy_managed_files],
|
|
)
|
|
|
|
router = litellm.Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-5-mini",
|
|
"litellm_params": {"model": "gpt-5-mini"},
|
|
"model_info": {"id": "1234"},
|
|
},
|
|
{
|
|
"model_name": "gpt-5-mini",
|
|
"litellm_params": {"model": "gpt-5-mini"},
|
|
"model_info": {"id": "5678"},
|
|
},
|
|
],
|
|
)
|
|
|
|
file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCw2YmQ4ZjhhYS02NmEzLTRmY2MtOTIxZS1lMTYwYzIzZWZjNzU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1MTENVRkI1MnVUTWE5aE5ZanRldzlWO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxmMzJlNWQ0OC05YWZmLTQ5YjMtOWE1Ny0zYzJhN2JjN2NjMmE"
|
|
|
|
model_file_id_mapping = {file_id: {"5678": "file-LLCUFB52uTMa9hNYjtew9V"}}
|
|
|
|
with patch.object(
|
|
litellm, "acreate_batch", return_value=AsyncMock()
|
|
) as mock_acreate_batch:
|
|
for _ in range(1000):
|
|
await router.acreate_batch(
|
|
model="gpt-5-mini",
|
|
input_file_id=file_id,
|
|
model_file_id_mapping=model_file_id_mapping,
|
|
)
|
|
|
|
mock_acreate_batch.assert_called()
|
|
assert "5678" in json.dumps(mock_acreate_batch.call_args.kwargs)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_output_file_id_for_batch_retrieve():
|
|
"""
|
|
Test that the output file id is the same as the input file id
|
|
"""
|
|
from typing import cast
|
|
|
|
from openai.types.batch import BatchRequestCounts
|
|
|
|
from litellm.types.utils import LiteLLMBatch
|
|
|
|
batch = LiteLLMBatch(
|
|
id="bGl0ZWxsbV9wcm94eTttb2RlbF9pZDoxMjM0NTY3OTtsbG1fYmF0Y2hfaWQ6YmF0Y2hfNjg1YzVlNWQ2Mzk4ODE5MGI4NWJkYjIxNDdiYTEzMWQ",
|
|
completion_window="24h",
|
|
created_at=1750883933,
|
|
endpoint="/v1/chat/completions",
|
|
input_file_id="file-8ci8gux8s7oES7GydYvnMG",
|
|
object="batch",
|
|
status="completed",
|
|
cancelled_at=None,
|
|
cancelling_at=None,
|
|
completed_at=1750883939,
|
|
error_file_id=None,
|
|
errors=None,
|
|
expired_at=None,
|
|
expires_at=1750970333,
|
|
failed_at=None,
|
|
finalizing_at=1750883938,
|
|
in_progress_at=1750883934,
|
|
metadata={"description": "nightly eval job"},
|
|
output_file_id="file-3BZYhmdJQ3V2oZPAtQsEax",
|
|
request_counts=BatchRequestCounts(completed=1, failed=0, total=1),
|
|
usage=None,
|
|
)
|
|
|
|
batch._hidden_params = {
|
|
"litellm_call_id": "dcd789e0-c0ad-4244-9564-4e611448d650",
|
|
"api_base": "https://api.openai.com",
|
|
"model_id": "12345679",
|
|
"response_cost": 0.0,
|
|
"additional_headers": {},
|
|
"litellm_model_name": "gpt-5.5",
|
|
"unified_batch_id": "litellm_proxy;model_id:12345679;llm_batch_id:batch_685c5e5d63988190b85bdb2147ba131d",
|
|
}
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=AsyncMock()
|
|
)
|
|
|
|
response = await proxy_managed_files.async_post_call_success_hook(
|
|
data={},
|
|
user_api_key_dict=MagicMock(),
|
|
response=batch,
|
|
)
|
|
|
|
assert not cast(LiteLLMBatch, response).output_file_id.startswith("file-")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_output_file_id_preserves_target_model_names_when_model_name_missing():
|
|
"""
|
|
Regression test: when provider response does not include _hidden_params.model_name
|
|
(e.g. Vertex batch retrieve), unified output_file_id should still include
|
|
target_model_names from the managed input file ID.
|
|
"""
|
|
from openai.types.batch import BatchRequestCounts
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
from litellm.types.llms.openai import OpenAIFileObject
|
|
from litellm.types.utils import LiteLLMBatch
|
|
|
|
batch = LiteLLMBatch(
|
|
id="batch_123",
|
|
completion_window="24h",
|
|
created_at=1750883933,
|
|
endpoint="/v1/chat/completions",
|
|
input_file_id="file-input-provider-id",
|
|
object="batch",
|
|
status="completed",
|
|
output_file_id="file-provider-output-id",
|
|
request_counts=BatchRequestCounts(completed=1, failed=0, total=1),
|
|
usage=None,
|
|
)
|
|
|
|
# Build a valid managed input id string and base64 encode it.
|
|
managed_input_file_payload = (
|
|
"litellm_proxy:application/octet-stream;"
|
|
"unified_id,test-uuid;"
|
|
"target_model_names,gemini-2.5-pro;"
|
|
"llm_output_file_id,file-input-1;"
|
|
"llm_output_file_model_id,model-id-1"
|
|
)
|
|
managed_input_file_id = (
|
|
base64.urlsafe_b64encode(managed_input_file_payload.encode())
|
|
.decode()
|
|
.rstrip("=")
|
|
)
|
|
|
|
batch._hidden_params = {
|
|
"model_id": "model-id-1",
|
|
"unified_batch_id": "litellm_proxy;model_id:model-id-1;llm_batch_id:batch_123",
|
|
"unified_file_id": managed_input_file_id,
|
|
# Intentionally omit model_name to mimic Vertex issue.
|
|
}
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=AsyncMock()
|
|
)
|
|
|
|
provider_output_file = OpenAIFileObject(
|
|
id="file-provider-output-id",
|
|
object="file",
|
|
bytes=1,
|
|
created_at=1,
|
|
filename="predictions.jsonl",
|
|
purpose="batch_output",
|
|
)
|
|
|
|
with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_retrieve:
|
|
mock_retrieve.return_value = provider_output_file
|
|
response = await proxy_managed_files.async_post_call_success_hook(
|
|
data={},
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
|
response=batch,
|
|
)
|
|
|
|
decoded_output_file_id = _is_base64_encoded_unified_file_id(
|
|
cast(LiteLLMBatch, response).output_file_id
|
|
)
|
|
assert decoded_output_file_id
|
|
assert "target_model_names,gemini-2.5-pro" in cast(str, decoded_output_file_id)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_error_file_id_for_failed_batch():
|
|
"""
|
|
Test that the error_file_id is properly managed when a batch fails
|
|
"""
|
|
from typing import cast
|
|
|
|
from openai.types.batch import BatchRequestCounts
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
from litellm.types.llms.openai import OpenAIFileObject
|
|
from litellm.types.utils import LiteLLMBatch
|
|
|
|
batch = LiteLLMBatch(
|
|
id="bGl0ZWxsbV9wcm94eTttb2RlbF9pZDoxMjM0NTY3OTtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz",
|
|
completion_window="24h",
|
|
created_at=1714508499,
|
|
endpoint="/v1/chat/completions",
|
|
input_file_id="file-abc123",
|
|
object="batch",
|
|
status="failed",
|
|
cancelled_at=None,
|
|
cancelling_at=None,
|
|
completed_at=None,
|
|
error_file_id="error-abc123",
|
|
errors=None,
|
|
expired_at=None,
|
|
expires_at=1714536634,
|
|
failed_at=None,
|
|
finalizing_at=None,
|
|
in_progress_at=None,
|
|
metadata=None,
|
|
output_file_id=None,
|
|
request_counts=BatchRequestCounts(completed=0, failed=0, total=0),
|
|
usage=None,
|
|
)
|
|
|
|
batch._hidden_params = {
|
|
"litellm_call_id": "test-call-id",
|
|
"api_base": "https://api.openai.com",
|
|
"model_id": "test-model-id",
|
|
"model_name": "gpt-5.5",
|
|
"response_cost": 0.0,
|
|
"additional_headers": {},
|
|
"litellm_model_name": "gpt-5.5",
|
|
"unified_batch_id": "litellm_proxy;model_id:test-model-id;llm_batch_id:batch_abc123",
|
|
}
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=AsyncMock()
|
|
)
|
|
|
|
# Create a proper OpenAIFileObject for the error file
|
|
error_file_object = OpenAIFileObject(
|
|
id="error-abc123",
|
|
object="file",
|
|
bytes=1234,
|
|
created_at=1714508500,
|
|
filename="error.jsonl",
|
|
purpose="batch_output",
|
|
status="processed",
|
|
)
|
|
|
|
# Mock the afile_retrieve to simulate retrieving error file metadata
|
|
with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_retrieve:
|
|
mock_retrieve.return_value = error_file_object
|
|
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
user_id="test-user-123", parent_otel_span=MagicMock()
|
|
)
|
|
|
|
response = await proxy_managed_files.async_post_call_success_hook(
|
|
data={},
|
|
user_api_key_dict=user_api_key_dict,
|
|
response=batch,
|
|
)
|
|
|
|
# Verify that error_file_id was transformed to a managed file ID
|
|
assert cast(LiteLLMBatch, response).error_file_id is not None
|
|
assert not cast(LiteLLMBatch, response).error_file_id.startswith("error-")
|
|
# Verify it's a base64 encoded managed file ID
|
|
assert _is_base64_encoded_unified_file_id(
|
|
cast(LiteLLMBatch, response).error_file_id
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_post_call_success_hook_twice_assert_no_unique_violation():
|
|
import asyncio
|
|
|
|
from openai.types.batch import BatchRequestCounts
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
from litellm.types.utils import LiteLLMBatch
|
|
|
|
# Use AsyncMock instead of real database connection
|
|
prisma_client = AsyncMock()
|
|
|
|
batch = LiteLLMBatch(
|
|
id="bGl0ZWxsbV9wcm94eTttb2RlbF9pZDoxMjM0NTY3OTtsbG1fYmF0Y2hfaWQ6YmF0Y2hfNjg1YzVlNWQ2Mzk4ODE5MGI4NWJkYjIxNDdiYTEzMWQ",
|
|
completion_window="24h",
|
|
created_at=1750883933,
|
|
endpoint="/v1/chat/completions",
|
|
input_file_id="file-8ci8gux8s7oES7GydYvnMG",
|
|
object="batch",
|
|
status="completed",
|
|
metadata={"description": "nightly eval job"},
|
|
request_counts=BatchRequestCounts(completed=1, failed=0, total=1),
|
|
usage=None,
|
|
)
|
|
|
|
batch._hidden_params = {
|
|
"model_id": "12345679",
|
|
"response_cost": 0.0,
|
|
"litellm_model_name": "gpt-5.5",
|
|
"unified_batch_id": "litellm_proxy;model_id:12345679;llm_batch_id:batch_685c5e5d63988190b85bdb2147ba131d",
|
|
}
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
# first retrieve batch
|
|
tasks = []
|
|
first_create_task = asyncio.create_task
|
|
with patch("asyncio.create_task") as mock_create_task:
|
|
mock_create_task.side_effect = (
|
|
lambda coro: tasks.append(first_create_task(coro)) or tasks[-1]
|
|
)
|
|
|
|
response = await proxy_managed_files.async_post_call_success_hook(
|
|
data={},
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="default_id"),
|
|
response=batch.copy(),
|
|
)
|
|
|
|
if tasks:
|
|
# make sure asyncio(db create) is finished
|
|
await asyncio.sleep(0.02)
|
|
await asyncio.gather(*tasks, return_exceptions=True)
|
|
for task in tasks:
|
|
assert task.exception() is None, f"Error: {task.exception()}"
|
|
|
|
assert isinstance(response, LiteLLMBatch)
|
|
assert _is_base64_encoded_unified_file_id(response.id)
|
|
|
|
# second retrieve batch
|
|
tasks = []
|
|
second_create_task = asyncio.create_task
|
|
with patch("asyncio.create_task") as mock_create_task:
|
|
mock_create_task.side_effect = (
|
|
lambda coro: tasks.append(second_create_task(coro)) or tasks[-1]
|
|
)
|
|
|
|
await proxy_managed_files.async_post_call_success_hook(
|
|
data={},
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="default_id"),
|
|
response=batch.copy(),
|
|
)
|
|
|
|
if tasks:
|
|
await asyncio.sleep(0.01)
|
|
await asyncio.gather(*tasks, return_exceptions=True)
|
|
for task in tasks:
|
|
assert task.exception() is None, f"Error: {task.exception()}"
|
|
|
|
|
|
def test_update_responses_input_with_unified_file_id():
|
|
"""
|
|
Test that update_responses_input_with_model_file_ids correctly decodes
|
|
unified file IDs and extracts llm_output_file_id from responses API input.
|
|
"""
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
update_responses_input_with_model_file_ids,
|
|
)
|
|
|
|
# Create a base64-encoded unified file ID
|
|
# This decodes to: litellm_proxy:application/pdf;unified_id,6c0b5890-8914-48e0-b8f4-0ae5ed3c14a5;target_model_names,gpt-4o;llm_output_file_id,file-ECBPW7ML9g7XHdwGgUPZaM;llm_output_file_model_id,e26453f9e76e7993680d0068d98c1f4cc205bbad0967a33c664893568ca743c2
|
|
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
|
|
|
|
# Test input with unified file ID in content array
|
|
input_data = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "input_file",
|
|
"file_id": unified_file_id,
|
|
},
|
|
{
|
|
"type": "input_text",
|
|
"text": "What is the first dragon in the book?",
|
|
},
|
|
],
|
|
}
|
|
]
|
|
|
|
# Update the input
|
|
updated_input = update_responses_input_with_model_file_ids(input=input_data)
|
|
|
|
# Verify the file_id was updated to the provider-specific file ID
|
|
assert updated_input[0]["content"][0]["type"] == "input_file"
|
|
assert updated_input[0]["content"][0]["file_id"] == "file-ECBPW7ML9g7XHdwGgUPZaM"
|
|
assert updated_input[0]["content"][1]["type"] == "input_text"
|
|
assert (
|
|
updated_input[0]["content"][1]["text"]
|
|
== "What is the first dragon in the book?"
|
|
)
|
|
|
|
|
|
def test_update_responses_input_with_regular_file_id():
|
|
"""
|
|
Test that update_responses_input_with_model_file_ids keeps regular
|
|
OpenAI file IDs unchanged.
|
|
"""
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
update_responses_input_with_model_file_ids,
|
|
)
|
|
|
|
# Regular OpenAI file ID (not a unified file ID)
|
|
regular_file_id = "file-abc123xyz"
|
|
|
|
input_data = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "input_file",
|
|
"file_id": regular_file_id,
|
|
},
|
|
{
|
|
"type": "input_text",
|
|
"text": "What is this file?",
|
|
},
|
|
],
|
|
}
|
|
]
|
|
|
|
# Update the input
|
|
updated_input = update_responses_input_with_model_file_ids(input=input_data)
|
|
|
|
# Verify the file_id was kept unchanged (regular OpenAI file ID)
|
|
assert updated_input[0]["content"][0]["type"] == "input_file"
|
|
assert updated_input[0]["content"][0]["file_id"] == regular_file_id
|
|
assert updated_input[0]["content"][1]["type"] == "input_text"
|
|
|
|
|
|
def test_update_responses_input_with_string_input():
|
|
"""
|
|
Test that update_responses_input_with_model_file_ids returns string input unchanged.
|
|
"""
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
update_responses_input_with_model_file_ids,
|
|
)
|
|
|
|
input_data = "What is AI?"
|
|
|
|
updated_input = update_responses_input_with_model_file_ids(input=input_data)
|
|
|
|
assert updated_input == input_data
|
|
assert isinstance(updated_input, str)
|
|
|
|
|
|
def test_update_responses_input_with_multiple_file_ids():
|
|
"""
|
|
Test that update_responses_input_with_model_file_ids handles multiple file IDs
|
|
(both unified and regular) in the same input.
|
|
"""
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
update_responses_input_with_model_file_ids,
|
|
)
|
|
|
|
# Unified file ID
|
|
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
|
|
# Regular OpenAI file ID
|
|
regular_file_id = "file-regular123"
|
|
|
|
input_data = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "input_file",
|
|
"file_id": unified_file_id,
|
|
},
|
|
{
|
|
"type": "input_text",
|
|
"text": "Compare these files",
|
|
},
|
|
{
|
|
"type": "input_file",
|
|
"file_id": regular_file_id,
|
|
},
|
|
],
|
|
}
|
|
]
|
|
|
|
updated_input = update_responses_input_with_model_file_ids(input=input_data)
|
|
|
|
# Verify unified file ID was updated
|
|
assert updated_input[0]["content"][0]["file_id"] == "file-ECBPW7ML9g7XHdwGgUPZaM"
|
|
# Verify regular file ID was kept unchanged
|
|
assert updated_input[0]["content"][2]["file_id"] == regular_file_id
|
|
# Verify text content was preserved
|
|
assert updated_input[0]["content"][1]["text"] == "Compare these files"
|
|
|
|
|
|
def test_update_responses_input_with_model_file_id_mapping():
|
|
"""
|
|
Test that update_responses_input_with_model_file_ids correctly uses
|
|
model_file_id_mapping to map managed file IDs to provider-specific file IDs.
|
|
"""
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
update_responses_input_with_model_file_ids,
|
|
)
|
|
|
|
# Managed file ID (unified)
|
|
managed_file_id = "litellm_proxy_file_123"
|
|
|
|
# Model file ID mapping
|
|
model_file_id_mapping = {
|
|
managed_file_id: {
|
|
"model_id_1": "openai_file_abc",
|
|
"model_id_2": "azure_file_xyz",
|
|
}
|
|
}
|
|
|
|
input_data = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "input_file",
|
|
"file_id": managed_file_id,
|
|
},
|
|
{
|
|
"type": "input_text",
|
|
"text": "Analyze this file",
|
|
},
|
|
],
|
|
}
|
|
]
|
|
|
|
# Update input with model_id_1 mapping
|
|
updated_input = update_responses_input_with_model_file_ids(
|
|
input=input_data,
|
|
model_id="model_id_1",
|
|
model_file_id_mapping=model_file_id_mapping,
|
|
)
|
|
|
|
# Verify the file_id was mapped to the correct provider-specific file ID
|
|
assert updated_input[0]["content"][0]["file_id"] == "openai_file_abc"
|
|
|
|
# Test with different model_id
|
|
updated_input_2 = update_responses_input_with_model_file_ids(
|
|
input=input_data,
|
|
model_id="model_id_2",
|
|
model_file_id_mapping=model_file_id_mapping,
|
|
)
|
|
|
|
assert updated_input_2[0]["content"][0]["file_id"] == "azure_file_xyz"
|
|
|
|
|
|
def test_update_responses_tools_with_model_file_id_mapping():
|
|
"""
|
|
Test that update_responses_tools_with_model_file_ids correctly maps
|
|
file IDs in code_interpreter tools with container.file_ids.
|
|
|
|
This is a regression test for the issue where managed file IDs in
|
|
tools.container.file_ids were not being replaced with provider-specific
|
|
file IDs, causing "string too long" errors from OpenAI.
|
|
"""
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
update_responses_tools_with_model_file_ids,
|
|
)
|
|
|
|
# Managed file IDs
|
|
managed_file_id_1 = "litellm_proxy_file_123"
|
|
managed_file_id_2 = "litellm_proxy_file_456"
|
|
|
|
# Model file ID mapping
|
|
model_file_id_mapping = {
|
|
managed_file_id_1: {
|
|
"model_id_1": "openai_file_abc",
|
|
},
|
|
managed_file_id_2: {
|
|
"model_id_1": "openai_file_def",
|
|
},
|
|
}
|
|
|
|
tools = [
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {
|
|
"type": "auto",
|
|
"file_ids": [managed_file_id_1, managed_file_id_2],
|
|
},
|
|
}
|
|
]
|
|
|
|
# Update tools with model mapping
|
|
updated_tools = update_responses_tools_with_model_file_ids(
|
|
tools=tools,
|
|
model_id="model_id_1",
|
|
model_file_id_mapping=model_file_id_mapping,
|
|
)
|
|
|
|
# Verify the file IDs were mapped to provider-specific file IDs
|
|
assert updated_tools[0]["type"] == "code_interpreter"
|
|
assert updated_tools[0]["container"]["file_ids"] == [
|
|
"openai_file_abc",
|
|
"openai_file_def",
|
|
]
|
|
|
|
|
|
def test_update_responses_tools_without_mapping():
|
|
"""
|
|
Test that update_responses_tools_with_model_file_ids keeps file IDs
|
|
unchanged when no mapping is provided.
|
|
"""
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
update_responses_tools_with_model_file_ids,
|
|
)
|
|
|
|
regular_file_id = "file-abc123"
|
|
|
|
tools = [
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {
|
|
"type": "auto",
|
|
"file_ids": [regular_file_id],
|
|
},
|
|
}
|
|
]
|
|
|
|
# Update tools without mapping
|
|
updated_tools = update_responses_tools_with_model_file_ids(
|
|
tools=tools,
|
|
model_id=None,
|
|
model_file_id_mapping=None,
|
|
)
|
|
|
|
# Verify the file ID was kept unchanged
|
|
assert updated_tools[0]["container"]["file_ids"] == [regular_file_id]
|
|
|
|
|
|
def test_update_responses_tools_with_mixed_file_ids():
|
|
"""
|
|
Test that update_responses_tools_with_model_file_ids correctly handles
|
|
a mix of managed and regular file IDs.
|
|
"""
|
|
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|
update_responses_tools_with_model_file_ids,
|
|
)
|
|
|
|
managed_file_id = "litellm_proxy_file_123"
|
|
regular_file_id = "file-abc123"
|
|
|
|
model_file_id_mapping = {
|
|
managed_file_id: {
|
|
"model_id_1": "openai_file_abc",
|
|
},
|
|
}
|
|
|
|
tools = [
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {
|
|
"type": "auto",
|
|
"file_ids": [managed_file_id, regular_file_id],
|
|
},
|
|
}
|
|
]
|
|
|
|
# Update tools
|
|
updated_tools = update_responses_tools_with_model_file_ids(
|
|
tools=tools,
|
|
model_id="model_id_1",
|
|
model_file_id_mapping=model_file_id_mapping,
|
|
)
|
|
|
|
# Verify managed file ID was mapped and regular file ID was kept
|
|
assert updated_tools[0]["container"]["file_ids"] == [
|
|
"openai_file_abc",
|
|
regular_file_id,
|
|
]
|
|
|
|
|
|
def test_get_file_ids_from_responses_tools():
|
|
"""
|
|
Test that get_file_ids_from_responses_tools correctly extracts
|
|
file IDs from the tools parameter.
|
|
"""
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
|
|
tools = [
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {
|
|
"type": "auto",
|
|
"file_ids": ["file-123", "file-456"],
|
|
},
|
|
}
|
|
]
|
|
|
|
file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools)
|
|
|
|
assert file_ids == ["file-123", "file-456"]
|
|
|
|
|
|
def test_get_file_ids_from_responses_tools_multiple_tools():
|
|
"""
|
|
Test that get_file_ids_from_responses_tools handles multiple tools.
|
|
"""
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
|
|
tools = [
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {
|
|
"type": "auto",
|
|
"file_ids": ["file-123"],
|
|
},
|
|
},
|
|
{
|
|
"type": "file_search",
|
|
},
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {
|
|
"type": "auto",
|
|
"file_ids": ["file-456", "file-789"],
|
|
},
|
|
},
|
|
]
|
|
|
|
file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools)
|
|
|
|
# Should extract file IDs only from code_interpreter tools
|
|
assert file_ids == ["file-123", "file-456", "file-789"]
|
|
|
|
|
|
def test_get_file_ids_from_responses_tools_empty():
|
|
"""
|
|
Test that get_file_ids_from_responses_tools handles empty or None tools.
|
|
"""
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
|
|
# Test with None
|
|
file_ids = proxy_managed_files.get_file_ids_from_responses_tools(None)
|
|
assert file_ids == []
|
|
|
|
# Test with empty list
|
|
file_ids = proxy_managed_files.get_file_ids_from_responses_tools([])
|
|
assert file_ids == []
|
|
|
|
# Test with tools without file_ids
|
|
tools = [{"type": "file_search"}]
|
|
file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools)
|
|
assert file_ids == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_check_file_ids_access_with_unified_file_ids():
|
|
"""
|
|
Test that check_file_ids_access validates user access to managed file IDs.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
# Create a unified file ID
|
|
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
|
|
regular_file_id = "file-abc123"
|
|
|
|
# Mock the access check to return True
|
|
prisma_client = AsyncMock()
|
|
internal_usage_cache = MagicMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock can_user_call_unified_file_id to return True
|
|
proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True)
|
|
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
user_id="test_user_123",
|
|
parent_otel_span=MagicMock(),
|
|
)
|
|
|
|
# Should not raise an exception for accessible files
|
|
await proxy_managed_files.check_file_ids_access(
|
|
[unified_file_id, regular_file_id],
|
|
user_api_key_dict,
|
|
)
|
|
|
|
# Verify can_user_call_unified_file_id was called for the unified file ID
|
|
proxy_managed_files.can_user_call_unified_file_id.assert_called_once_with(
|
|
unified_file_id, user_api_key_dict
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_check_file_ids_access_denied():
|
|
"""
|
|
Test that check_file_ids_access raises HTTPException when user doesn't have access.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
|
|
|
|
prisma_client = AsyncMock()
|
|
internal_usage_cache = MagicMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock can_user_call_unified_file_id to return False (access denied)
|
|
proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=False)
|
|
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
user_id="test_user_123",
|
|
parent_otel_span=MagicMock(),
|
|
)
|
|
|
|
# Should raise HTTPException with 403 status code
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.check_file_ids_access(
|
|
[unified_file_id],
|
|
user_api_key_dict,
|
|
)
|
|
|
|
assert exc_info.value.status_code == 403
|
|
assert "does not have access to the file" in exc_info.value.detail
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_check_file_ids_access_with_regular_files_only():
|
|
"""
|
|
Test that check_file_ids_access doesn't check access for regular (non-unified) file IDs.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
regular_file_id_1 = "file-abc123"
|
|
regular_file_id_2 = "file-xyz789"
|
|
|
|
prisma_client = AsyncMock()
|
|
internal_usage_cache = MagicMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock can_user_call_unified_file_id (should not be called for regular files)
|
|
proxy_managed_files.can_user_call_unified_file_id = AsyncMock()
|
|
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
user_id="test_user_123",
|
|
parent_otel_span=MagicMock(),
|
|
)
|
|
|
|
# Should not raise exception and should not call can_user_call_unified_file_id
|
|
await proxy_managed_files.check_file_ids_access(
|
|
[regular_file_id_1, regular_file_id_2],
|
|
user_api_key_dict,
|
|
)
|
|
|
|
# Verify can_user_call_unified_file_id was NOT called
|
|
proxy_managed_files.can_user_call_unified_file_id.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_completion_with_file_access_check():
|
|
"""
|
|
Test that completion call type checks file access before processing.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
|
|
|
|
internal_usage_cache = MagicMock()
|
|
internal_usage_cache.async_get_cache = AsyncMock(return_value=None)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock the get_model_file_id_mapping to return empty dict
|
|
proxy_managed_files.get_model_file_id_mapping = AsyncMock(return_value={})
|
|
|
|
# Mock access check to allow access
|
|
proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True)
|
|
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
user_id="test_user_123",
|
|
parent_otel_span=MagicMock(),
|
|
)
|
|
|
|
data = {
|
|
"messages": [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{"type": "text", "text": "What's in this file?"},
|
|
{
|
|
"type": "file",
|
|
"file": {"file_id": unified_file_id},
|
|
},
|
|
],
|
|
}
|
|
],
|
|
"model": "gpt-5.5",
|
|
}
|
|
|
|
# Should not raise exception
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=user_api_key_dict,
|
|
cache=DualCache(),
|
|
data=data,
|
|
call_type="acompletion",
|
|
)
|
|
|
|
# Verify access check was called
|
|
proxy_managed_files.can_user_call_unified_file_id.assert_called_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_responses_with_file_access_check():
|
|
"""
|
|
Test that responses API checks file access for files in both input and tools.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
unified_file_id_1 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
|
|
unified_file_id_2 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsNzc3Nzc3Nzc7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1YWVo7bGxtX291dHB1dF9maWxlX21vZGVsX2lkLG1vZGVsXzEyMw"
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
|
|
|
|
internal_usage_cache = MagicMock()
|
|
internal_usage_cache.async_get_cache = AsyncMock(return_value=None)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock the get_model_file_id_mapping to return empty dict
|
|
proxy_managed_files.get_model_file_id_mapping = AsyncMock(return_value={})
|
|
|
|
# Mock access check to allow access
|
|
proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True)
|
|
|
|
user_api_key_dict = UserAPIKeyAuth(
|
|
user_id="test_user_123",
|
|
parent_otel_span=MagicMock(),
|
|
)
|
|
|
|
data = {
|
|
"input": [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{"type": "input_text", "text": "Analyze this"},
|
|
{"type": "input_file", "file_id": unified_file_id_1},
|
|
],
|
|
}
|
|
],
|
|
"tools": [
|
|
{
|
|
"type": "code_interpreter",
|
|
"container": {
|
|
"type": "auto",
|
|
"file_ids": [unified_file_id_2],
|
|
},
|
|
}
|
|
],
|
|
"model": "gpt-5.5",
|
|
}
|
|
|
|
# Should not raise exception
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=user_api_key_dict,
|
|
cache=DualCache(),
|
|
data=data,
|
|
call_type="aresponses",
|
|
)
|
|
|
|
# Verify access check was called for both file IDs
|
|
assert proxy_managed_files.can_user_call_unified_file_id.call_count == 2
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_store_unified_file_id_with_none_file_object():
|
|
"""
|
|
Test that store_unified_file_id works when file_object is None
|
|
(e.g., for batch output files that are stored before file metadata is available).
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.upsert = AsyncMock(
|
|
return_value=MagicMock()
|
|
)
|
|
internal_usage_cache = MagicMock()
|
|
internal_usage_cache.async_set_cache = AsyncMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Store with file_object=None (simulating batch output file storage)
|
|
await proxy_managed_files.store_unified_file_id(
|
|
file_id="test-unified-file-id",
|
|
file_object=None,
|
|
litellm_parent_otel_span=None,
|
|
model_mappings={"model-123": "file-provider-xyz"},
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
|
)
|
|
|
|
# Verify DB upsert was called idempotently with expected create data (without file_object)
|
|
prisma_client.db.litellm_managedfiletable.upsert.assert_called_once()
|
|
call_args = prisma_client.db.litellm_managedfiletable.upsert.call_args
|
|
assert call_args.kwargs["where"] == {"unified_file_id": "test-unified-file-id"}
|
|
create_data = call_args.kwargs["data"]["create"]
|
|
assert create_data["unified_file_id"] == "test-unified-file-id"
|
|
assert "file_object" not in create_data
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_store_unified_file_id_updates_file_metadata_on_existing_row():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
from litellm.types.llms.openai import OpenAIFileObject
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.upsert = AsyncMock(
|
|
return_value=MagicMock()
|
|
)
|
|
internal_usage_cache = MagicMock()
|
|
internal_usage_cache.async_set_cache = AsyncMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
user_api_key_dict = UserAPIKeyAuth(user_id="test-user")
|
|
|
|
await proxy_managed_files.store_unified_file_id(
|
|
file_id="test-unified-file-id",
|
|
file_object=None,
|
|
litellm_parent_otel_span=None,
|
|
model_mappings={"model-123": "file-provider-xyz"},
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
|
|
file_object = OpenAIFileObject(
|
|
id="file-provider-xyz",
|
|
object="file",
|
|
bytes=1234,
|
|
created_at=1234567890,
|
|
filename="output.jsonl",
|
|
purpose="batch_output",
|
|
status="processed",
|
|
)
|
|
file_object._hidden_params = {
|
|
"storage_backend": "s3",
|
|
"storage_url": "s3://bucket/output.jsonl",
|
|
}
|
|
|
|
await proxy_managed_files.store_unified_file_id(
|
|
file_id="test-unified-file-id",
|
|
file_object=file_object,
|
|
litellm_parent_otel_span=None,
|
|
model_mappings={"model-123": "file-provider-xyz"},
|
|
user_api_key_dict=user_api_key_dict,
|
|
)
|
|
|
|
first_update = prisma_client.db.litellm_managedfiletable.upsert.await_args_list[
|
|
0
|
|
].kwargs["data"]["update"]
|
|
second_update = prisma_client.db.litellm_managedfiletable.upsert.await_args_list[
|
|
1
|
|
].kwargs["data"]["update"]
|
|
assert "file_object" not in first_update
|
|
assert second_update["file_object"] == file_object.model_dump_json()
|
|
assert second_update["storage_backend"] == "s3"
|
|
assert second_update["storage_url"] == "s3://bucket/output.jsonl"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_afile_delete_returns_provider_response_when_stored_file_object_none():
|
|
"""
|
|
Test that afile_delete returns the provider's delete response when the
|
|
stored file_object is None (e.g., for batch output files).
|
|
"""
|
|
from litellm.types.llms.openai import OpenAIFileObject
|
|
|
|
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsdGVzdC1pZDt0YXJnZXRfbW9kZWxfbmFtZXMsZ3B0LTRvO2xsbV9vdXRwdXRfZmlsZV9pZCxmaWxlLXByb3ZpZGVyLXh5ejtsbG1fb3V0cHV0X2ZpbGVfbW9kZWxfaWQsbW9kZWwtMTIz"
|
|
|
|
prisma_client = AsyncMock()
|
|
db_record = MagicMock()
|
|
db_record.model_mappings = '{"model-123": "file-provider-xyz"}'
|
|
prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(
|
|
return_value=db_record
|
|
)
|
|
prisma_client.db.litellm_managedfiletable.delete = AsyncMock()
|
|
|
|
internal_usage_cache = MagicMock()
|
|
internal_usage_cache.async_get_cache = AsyncMock(
|
|
return_value={
|
|
"unified_file_id": unified_file_id,
|
|
"model_mappings": {"model-123": "file-provider-xyz"},
|
|
"flat_model_file_ids": ["file-provider-xyz"],
|
|
"file_object": None,
|
|
"created_by": "test-user",
|
|
"updated_by": "test-user",
|
|
}
|
|
)
|
|
internal_usage_cache.async_set_cache = AsyncMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock the delete_unified_file_id to return None (simulating file_object=None)
|
|
proxy_managed_files.delete_unified_file_id = AsyncMock(return_value=None)
|
|
|
|
# Mock router response
|
|
provider_delete_response = OpenAIFileObject(
|
|
id="file-provider-xyz",
|
|
object="file",
|
|
bytes=1234,
|
|
created_at=1234567890,
|
|
filename="test.jsonl",
|
|
purpose="batch",
|
|
)
|
|
|
|
mock_router = MagicMock()
|
|
mock_router.afile_delete = AsyncMock(return_value=provider_delete_response)
|
|
|
|
result = await proxy_managed_files.afile_delete(
|
|
file_id=unified_file_id,
|
|
litellm_parent_otel_span=None,
|
|
llm_router=mock_router,
|
|
)
|
|
|
|
# Should return the provider response with the unified file ID
|
|
assert result is not None
|
|
assert result.id == unified_file_id
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_afile_retrieve_fetches_from_provider_when_file_object_none():
|
|
"""
|
|
Test that afile_retrieve fetches from the provider when the stored
|
|
file_object is None (e.g., for batch output files).
|
|
"""
|
|
from litellm.types.llms.openai import OpenAIFileObject
|
|
|
|
prisma_client = AsyncMock()
|
|
internal_usage_cache = MagicMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock get_unified_file_id to return a stored object with file_object=None
|
|
stored_file = MagicMock()
|
|
stored_file.file_object = None
|
|
stored_file.model_mappings = {"model-123": "file-provider-xyz"}
|
|
proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file)
|
|
|
|
# Mock the router and provider response
|
|
provider_file_response = OpenAIFileObject(
|
|
id="file-provider-xyz",
|
|
object="file",
|
|
bytes=5678,
|
|
created_at=1234567890,
|
|
filename="output.jsonl",
|
|
purpose="batch_output",
|
|
)
|
|
|
|
mock_router = MagicMock()
|
|
mock_router.get_deployment_credentials_with_provider = MagicMock(
|
|
return_value={
|
|
"api_key": "test-key",
|
|
"api_base": "https://api.openai.com",
|
|
}
|
|
)
|
|
|
|
with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_afile_retrieve:
|
|
mock_afile_retrieve.return_value = provider_file_response
|
|
|
|
unified_file_id = "test-unified-file-id"
|
|
result = await proxy_managed_files.afile_retrieve(
|
|
file_id=unified_file_id,
|
|
litellm_parent_otel_span=None,
|
|
llm_router=mock_router,
|
|
)
|
|
|
|
# Should return the provider response with the unified file ID
|
|
assert result is not None
|
|
assert result.id == unified_file_id
|
|
mock_afile_retrieve.assert_called_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none():
|
|
"""
|
|
Test that afile_retrieve raises an appropriate error when file_object is None
|
|
and no llm_router is provided to fetch from the provider.
|
|
"""
|
|
prisma_client = AsyncMock()
|
|
internal_usage_cache = MagicMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock get_unified_file_id to return a stored object with file_object=None
|
|
stored_file = MagicMock()
|
|
stored_file.file_object = None
|
|
stored_file.model_mappings = {"model-123": "file-provider-xyz"}
|
|
proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file)
|
|
|
|
unified_file_id = "test-unified-file-id"
|
|
|
|
with pytest.raises(Exception, match='LiteLLM Managed File object with id=test-unified-file-id') as exc_info:
|
|
await proxy_managed_files.afile_retrieve(
|
|
file_id=unified_file_id,
|
|
litellm_parent_otel_span=None,
|
|
llm_router=None,
|
|
)
|
|
|
|
assert "llm_router is required" in str(exc_info.value)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_afile_retrieve_returns_stored_file_object_when_exists():
|
|
"""
|
|
Test that afile_retrieve returns the stored file_object directly when it exists
|
|
(the normal case for user-uploaded files).
|
|
"""
|
|
from litellm.types.llms.openai import OpenAIFileObject
|
|
|
|
prisma_client = AsyncMock()
|
|
internal_usage_cache = MagicMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock get_unified_file_id to return a stored object WITH file_object
|
|
stored_file_object = OpenAIFileObject(
|
|
id="test-unified-file-id",
|
|
object="file",
|
|
bytes=1234,
|
|
created_at=1234567890,
|
|
filename="input.jsonl",
|
|
purpose="batch",
|
|
)
|
|
stored_file = MagicMock()
|
|
stored_file.file_object = stored_file_object
|
|
proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file)
|
|
|
|
result = await proxy_managed_files.afile_retrieve(
|
|
file_id="test-unified-file-id",
|
|
litellm_parent_otel_span=None,
|
|
llm_router=None,
|
|
)
|
|
|
|
# Should return the stored file object directly
|
|
assert result == stored_file_object
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_afile_retrieve_raises_error_for_non_managed_file():
|
|
"""
|
|
Test that afile_retrieve raises an error when the file_id is not found
|
|
in the managed files table.
|
|
"""
|
|
prisma_client = AsyncMock()
|
|
internal_usage_cache = MagicMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
internal_usage_cache=internal_usage_cache,
|
|
prisma_client=prisma_client,
|
|
)
|
|
|
|
# Mock get_unified_file_id to return None (file not found)
|
|
proxy_managed_files.get_unified_file_id = AsyncMock(return_value=None)
|
|
|
|
with pytest.raises(Exception, match='LiteLLM Managed File object with id=non-existent-file-id') as exc_info:
|
|
await proxy_managed_files.afile_retrieve(
|
|
file_id="non-existent-file-id",
|
|
litellm_parent_otel_span=None,
|
|
)
|
|
|
|
assert "not found" in str(exc_info.value)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_from_managed_objects_table():
|
|
from openai.types.batch import BatchRequestCounts
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
batch_record_1 = MagicMock()
|
|
batch_record_1.unified_object_id = "unified-batch-id-1"
|
|
batch_record_1.file_object = json.dumps(
|
|
{
|
|
"id": "batch_abc123",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "completed",
|
|
"created_at": 1234567890,
|
|
"input_file_id": "file-input-1",
|
|
"request_counts": {"total": 1, "completed": 1, "failed": 0},
|
|
}
|
|
)
|
|
|
|
batch_record_2 = MagicMock()
|
|
batch_record_2.unified_object_id = "unified-batch-id-2"
|
|
batch_record_2.file_object = json.dumps(
|
|
{
|
|
"id": "batch_xyz789",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "in_progress",
|
|
"created_at": 1234567891,
|
|
"input_file_id": "file-input-2",
|
|
"request_counts": {"total": 5, "completed": 2, "failed": 0},
|
|
}
|
|
)
|
|
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [
|
|
batch_record_1,
|
|
batch_record_2,
|
|
]
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
|
limit=10,
|
|
)
|
|
|
|
assert result["object"] == "list"
|
|
assert len(result["data"]) == 2
|
|
assert result["data"][0].id == "unified-batch-id-1"
|
|
assert result["data"][1].id == "unified-batch-id-2"
|
|
assert result["first_id"] == "unified-batch-id-1"
|
|
assert result["last_id"] == "unified-batch-id-2"
|
|
|
|
# Should filter by user_id (created_by)
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
|
|
where={"file_purpose": "batch", "created_by": "test-user"},
|
|
take=11,
|
|
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_from_managed_objects_table_empty_list():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = []
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
|
)
|
|
|
|
assert result["object"] == "list"
|
|
assert len(result["data"]) == 0
|
|
assert result["first_id"] is None
|
|
assert result["last_id"] is None
|
|
assert result["has_more"] is False
|
|
|
|
# Verify where clause includes created_by filter
|
|
# Default take is 20 when no limit is provided
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
|
|
where={"file_purpose": "batch", "created_by": "test-user"},
|
|
take=21,
|
|
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
|
)
|
|
|
|
|
|
def _create_unified_batch_id(model_id: str, batch_id: str) -> str:
|
|
import base64
|
|
|
|
unified_str = f"litellm_proxy;model_id:{model_id};llm_batch_id:{batch_id}"
|
|
return base64.urlsafe_b64encode(unified_str.encode()).decode().rstrip("=")
|
|
|
|
|
|
def _decode_unified_id(b64_id: str) -> str:
|
|
return base64.urlsafe_b64decode(b64_id + "=" * (-len(b64_id) % 4)).decode()
|
|
|
|
|
|
def _terminal_batch_record(
|
|
unified_batch_uid: str,
|
|
raw_input_file_id: str,
|
|
raw_output_file_id: str,
|
|
raw_error_file_id: str,
|
|
):
|
|
record = MagicMock()
|
|
record.unified_object_id = unified_batch_uid
|
|
record.created_by = "owner-user"
|
|
record.team_id = "owner-team"
|
|
record.status = "cancelled"
|
|
record.file_object = json.dumps(
|
|
{
|
|
"id": "batch-raw-456",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "cancelled",
|
|
"created_at": 1234567890,
|
|
"input_file_id": raw_input_file_id,
|
|
"output_file_id": raw_output_file_id,
|
|
"error_file_id": raw_error_file_id,
|
|
}
|
|
)
|
|
return record
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_registers_and_returns_unified_output_file_ids():
|
|
"""A stored batch blob with raw provider file IDs (e.g. persisted by the cost
|
|
poller for a cancelled batch) must be listed with unified managed IDs, and the
|
|
output/error files must be registered in the managed file table so GET
|
|
/files/{id}/content can route them."""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
unified_batch_uid = _create_unified_batch_id("model-123", "batch-456")
|
|
raw_input_file_id = "file-list-in-1"
|
|
raw_output_file_id = "file-list-out-1"
|
|
raw_error_file_id = "file-list-err-1"
|
|
unified_input_file_id = base64.urlsafe_b64encode(
|
|
b"litellm_proxy:application/octet-stream;unified_id,in-1;target_model_names,gpt-5-batch"
|
|
).decode()
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [
|
|
_terminal_batch_record(
|
|
unified_batch_uid, raw_input_file_id, raw_output_file_id, raw_error_file_id
|
|
)
|
|
]
|
|
|
|
input_file_row = MagicMock()
|
|
input_file_row.unified_file_id = unified_input_file_id
|
|
input_file_row.flat_model_file_ids = [raw_input_file_id]
|
|
|
|
prisma_client.db.litellm_managedfiletable.find_many = AsyncMock(
|
|
return_value=[input_file_row]
|
|
)
|
|
prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="owner-user"),
|
|
limit=10,
|
|
)
|
|
|
|
listed = result["data"][0]
|
|
assert listed.id == unified_batch_uid
|
|
assert listed.input_file_id == unified_input_file_id
|
|
|
|
bulk_lookup = prisma_client.db.litellm_managedfiletable.find_many.await_args
|
|
assert set(bulk_lookup.kwargs["where"]["flat_model_file_ids"]["hasSome"]) == {
|
|
raw_input_file_id,
|
|
raw_output_file_id,
|
|
raw_error_file_id,
|
|
}
|
|
|
|
decoded_output = _decode_unified_id(listed.output_file_id)
|
|
assert decoded_output.startswith("litellm_proxy")
|
|
assert f"llm_output_file_id,{raw_output_file_id}" in decoded_output
|
|
assert "llm_output_file_model_id,model-123" in decoded_output
|
|
assert "target_model_names,gpt-5-batch" in decoded_output
|
|
|
|
decoded_error = _decode_unified_id(listed.error_file_id)
|
|
assert f"llm_output_file_id,{raw_error_file_id}" in decoded_error
|
|
|
|
upsert_calls = prisma_client.db.litellm_managedfiletable.upsert.await_args_list
|
|
stored_raw_ids = {
|
|
c.kwargs["data"]["create"]["flat_model_file_ids"][0] for c in upsert_calls
|
|
}
|
|
assert stored_raw_ids == {raw_output_file_id, raw_error_file_id}
|
|
for c in upsert_calls:
|
|
assert c.kwargs["data"]["create"]["created_by"] == "owner-user"
|
|
assert c.kwargs["data"]["create"]["team_id"] == "owner-team"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_resolves_existing_managed_rows_without_minting():
|
|
"""When the raw provider file IDs already have managed file rows, listing must
|
|
swap in the existing unified IDs via one bulk lookup for the whole page, with
|
|
no per-row queries and no duplicate upserts."""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
unified_input_file_id = base64.urlsafe_b64encode(
|
|
b"litellm_proxy:application/octet-stream;unified_id,in-9;target_model_names,gpt-5-batch"
|
|
).decode()
|
|
raw_output_file_ids = ["file-list-out-existing-1", "file-list-out-existing-2"]
|
|
existing_unified_output_ids = [
|
|
base64.urlsafe_b64encode(
|
|
f"litellm_proxy:application/json;unified_id,u-{i};llm_output_file_id,{raw_id}".encode()
|
|
).decode()
|
|
for i, raw_id in enumerate(raw_output_file_ids)
|
|
]
|
|
|
|
records = [
|
|
_terminal_batch_record(
|
|
_create_unified_batch_id("model-123", f"batch-{i}"),
|
|
unified_input_file_id,
|
|
raw_id,
|
|
"",
|
|
)
|
|
for i, raw_id in enumerate(raw_output_file_ids)
|
|
]
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = records
|
|
|
|
existing_rows = [
|
|
MagicMock(unified_file_id=unified_id, flat_model_file_ids=[raw_id])
|
|
for raw_id, unified_id in zip(raw_output_file_ids, existing_unified_output_ids)
|
|
]
|
|
|
|
prisma_client.db.litellm_managedfiletable.find_many = AsyncMock(
|
|
return_value=existing_rows
|
|
)
|
|
prisma_client.db.litellm_managedfiletable.find_first = AsyncMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="owner-user"),
|
|
limit=10,
|
|
)
|
|
|
|
assert [b.output_file_id for b in result["data"]] == existing_unified_output_ids
|
|
prisma_client.db.litellm_managedfiletable.find_many.assert_awaited_once()
|
|
prisma_client.db.litellm_managedfiletable.find_first.assert_not_awaited()
|
|
prisma_client.db.litellm_managedfiletable.upsert.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_caps_page_size_at_100():
|
|
"""The list page size must be capped at 100 rows (matching OpenAI's limit)
|
|
even when the caller asks for more, so one request cannot fan out into an
|
|
unbounded scan."""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = []
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="owner-user"),
|
|
limit=100000,
|
|
)
|
|
|
|
assert (
|
|
prisma_client.db.litellm_managedobjecttable.find_many.await_args.kwargs["take"]
|
|
== 101
|
|
)
|
|
assert result["data"] == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_from_managed_objects_table_provider_filter_raises_exception():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
# Filtering by provider should raise Exception
|
|
with pytest.raises(Exception, match="Filtering by 'provider' is not supported when using managed") as exc_info:
|
|
await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
|
limit=10,
|
|
provider="openai",
|
|
)
|
|
|
|
assert str(exc_info.value) == (
|
|
"Filtering by 'provider' is not supported when using managed batches."
|
|
)
|
|
|
|
# Verify find_many was NOT called since exception is raised before database query
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_from_managed_objects_table_target_model_name_filter_raises_exception():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
# Filtering by provider should raise Exception
|
|
with pytest.raises(Exception, match="Filtering by 'target_model_names' is not supported when") as exc_info:
|
|
await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
|
limit=10,
|
|
target_model_names="gpt-5.5,gpt-3.5",
|
|
)
|
|
|
|
assert str(exc_info.value) == (
|
|
"Filtering by 'target_model_names' is not supported when using managed batches."
|
|
)
|
|
|
|
# Verify find_many was NOT called since exception is raised before database query
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_from_managed_objects_table_filters_by_created_by():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Create batch for user1
|
|
batch_user1 = MagicMock()
|
|
batch_user1.unified_object_id = "unified-batch-user1"
|
|
batch_user1.file_object = json.dumps(
|
|
{
|
|
"id": "batch_user1_abc",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "completed",
|
|
"created_at": 1234567890,
|
|
"input_file_id": "file-input-user1",
|
|
"request_counts": {"total": 1, "completed": 1, "failed": 0},
|
|
}
|
|
)
|
|
|
|
# Create batch for user2
|
|
batch_user2 = MagicMock()
|
|
batch_user2.unified_object_id = "unified-batch-user2"
|
|
batch_user2.file_object = json.dumps(
|
|
{
|
|
"id": "batch_user2_xyz",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "completed",
|
|
"created_at": 1234567891,
|
|
"input_file_id": "file-input-user2",
|
|
"request_counts": {"total": 2, "completed": 2, "failed": 0},
|
|
}
|
|
)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
# Query with user1's API key - should only return user1's batch
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user1]
|
|
result_user1 = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="user1"),
|
|
limit=10,
|
|
)
|
|
|
|
assert len(result_user1["data"]) == 1
|
|
assert result_user1["data"][0].id == "unified-batch-user1"
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with(
|
|
where={"file_purpose": "batch", "created_by": "user1"},
|
|
take=11,
|
|
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
|
)
|
|
|
|
# Query with user2's API key - should only return user2's batch
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user2]
|
|
result_user2 = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="user2"),
|
|
limit=10,
|
|
)
|
|
|
|
assert len(result_user2["data"]) == 1
|
|
assert result_user2["data"][0].id == "unified-batch-user2"
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_with(
|
|
where={"file_purpose": "batch", "created_by": "user2"},
|
|
take=11,
|
|
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_pagination_uses_unified_object_id_cursor():
|
|
"""Regression for LIT-4678.
|
|
|
|
The ``after`` cursor a client sends back is a batch's ``unified_object_id``
|
|
(that is what is returned as ``.id`` / ``last_id``). Paginating must use a
|
|
Prisma cursor on the unique ``unified_object_id`` column, not a
|
|
``where id > after`` filter against the random-uuid primary key.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = MagicMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = []
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
|
limit=5,
|
|
after="unified-batch-id-7",
|
|
)
|
|
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
|
|
where={"file_purpose": "batch", "created_by": "test-user"},
|
|
take=6,
|
|
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
|
cursor={"unified_object_id": "unified-batch-id-7"},
|
|
skip=1,
|
|
)
|
|
|
|
_, call_kwargs = prisma_client.db.litellm_managedobjecttable.find_many.call_args
|
|
assert "id" not in call_kwargs["where"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_pagination_walks_all_pages_without_loops_or_gaps():
|
|
"""Regression for LIT-4678.
|
|
|
|
Simulates the managed-objects table (random-uuid ``id`` primary key,
|
|
base64 ``unified_object_id``, reverse-chronological ``created_at``) and
|
|
walks every page the way a client would, feeding ``last_id`` back as
|
|
``after``. With the old ``where id > after`` cursor this loops and drops
|
|
batches; the fixed cursor returns each batch exactly once, newest first.
|
|
"""
|
|
import uuid as _uuid
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
def _unified_id(i: int) -> str:
|
|
raw = f"litellm_proxy;model_id:gpt-4o-batch;llm_batch_id:batch_{i:03d}"
|
|
return base64.urlsafe_b64encode(raw.encode()).decode().rstrip("=")
|
|
|
|
total = 10
|
|
rows = []
|
|
for i in range(total):
|
|
row = MagicMock()
|
|
row.id = str(_uuid.uuid4())
|
|
row.unified_object_id = _unified_id(i)
|
|
row.created_at = 1_000_000 + i
|
|
row.file_object = json.dumps(
|
|
{
|
|
"id": f"batch_provider_{i:03d}",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "completed",
|
|
"created_at": 1_000_000 + i,
|
|
"input_file_id": f"file-input-{i:03d}",
|
|
"request_counts": {"total": 1, "completed": 1, "failed": 0},
|
|
}
|
|
)
|
|
rows.append(row)
|
|
|
|
async def fake_find_many(where, take, order, cursor=None, skip=0):
|
|
result = list(rows)
|
|
id_filter = where.get("id")
|
|
if isinstance(id_filter, dict) and "gt" in id_filter:
|
|
result = [r for r in result if r.id > id_filter["gt"]]
|
|
order_keys = order if isinstance(order, list) else [order]
|
|
for clause in reversed(order_keys):
|
|
(order_field, direction), = clause.items()
|
|
result.sort(
|
|
key=lambda r: getattr(r, order_field), reverse=(direction == "desc")
|
|
)
|
|
if cursor is not None:
|
|
(cur_field, cur_val), = cursor.items()
|
|
idx = next(
|
|
(i for i, r in enumerate(result) if getattr(r, cur_field) == cur_val),
|
|
None,
|
|
)
|
|
if idx is None:
|
|
return []
|
|
result = result[idx + skip:]
|
|
return result[:take]
|
|
|
|
async def fake_find_first(where):
|
|
return next(
|
|
(r for r in rows if r.unified_object_id == where.get("unified_object_id")),
|
|
None,
|
|
)
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
|
side_effect=fake_find_many
|
|
)
|
|
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
|
|
side_effect=fake_find_first
|
|
)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
user = UserAPIKeyAuth(user_id="test-user")
|
|
|
|
seen: list = []
|
|
after = None
|
|
for _ in range(total + 5):
|
|
resp = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=user, limit=3, after=after
|
|
)
|
|
page_ids = [b.id for b in resp["data"]]
|
|
seen.extend(page_ids)
|
|
if not resp["has_more"]:
|
|
break
|
|
assert page_ids, "has_more was true but the page was empty"
|
|
assert resp["last_id"] != after, "cursor did not advance (pagination loop)"
|
|
after = resp["last_id"]
|
|
|
|
expected = [_unified_id(i) for i in reversed(range(total))]
|
|
assert seen == expected
|
|
assert len(seen) == len(set(seen))
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_pagination_stable_when_created_at_ties():
|
|
"""Regression for LIT-4678.
|
|
|
|
Cursor pagination is only well-defined when the ``order`` fully determines
|
|
row order. If listing ordered by non-unique ``created_at`` alone, batches
|
|
sharing a timestamp come back in an arbitrary order that can shift between
|
|
page requests, so a cursor row's neighbours change and batches get skipped
|
|
or duplicated. Listing must add the unique ``unified_object_id`` as a
|
|
tie-breaker so the order is total and pagination is stable.
|
|
"""
|
|
import itertools
|
|
import uuid as _uuid
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
def _unified_id(i: int) -> str:
|
|
raw = f"litellm_proxy;model_id:gpt-4o-batch;llm_batch_id:batch_{i:03d}"
|
|
return base64.urlsafe_b64encode(raw.encode()).decode().rstrip("=")
|
|
|
|
total = 6
|
|
shared_created_at = 1_000_000
|
|
rows = []
|
|
for i in range(total):
|
|
row = MagicMock()
|
|
row.id = str(_uuid.uuid4())
|
|
row.unified_object_id = _unified_id(i)
|
|
row.created_at = shared_created_at
|
|
row.file_object = json.dumps(
|
|
{
|
|
"id": f"batch_provider_{i:03d}",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "completed",
|
|
"created_at": shared_created_at,
|
|
"input_file_id": f"file-input-{i:03d}",
|
|
"request_counts": {"total": 1, "completed": 1, "failed": 0},
|
|
}
|
|
)
|
|
rows.append(row)
|
|
|
|
call_counter = itertools.count()
|
|
|
|
async def fake_find_many(where, take, order, cursor=None, skip=0):
|
|
call = next(call_counter)
|
|
order_keys = order if isinstance(order, list) else [order]
|
|
fields = [next(iter(clause)) for clause in order_keys]
|
|
result = list(rows)
|
|
for clause in reversed(order_keys):
|
|
field, direction = next(iter(clause.items()))
|
|
result.sort(
|
|
key=lambda r: getattr(r, field), reverse=(direction == "desc")
|
|
)
|
|
|
|
def order_key(r):
|
|
return tuple(getattr(r, f) for f in fields)
|
|
|
|
stabilized = []
|
|
i = 0
|
|
while i < len(result):
|
|
j = i
|
|
while j < len(result) and order_key(result[j]) == order_key(result[i]):
|
|
j += 1
|
|
group = result[i:j]
|
|
if len(group) > 1:
|
|
rot = call % len(group)
|
|
group = group[rot:] + group[:rot]
|
|
stabilized.extend(group)
|
|
i = j
|
|
result = stabilized
|
|
|
|
if cursor is not None:
|
|
cur_field, cur_val = next(iter(cursor.items()))
|
|
idx = next(
|
|
(k for k, r in enumerate(result) if getattr(r, cur_field) == cur_val),
|
|
None,
|
|
)
|
|
if idx is None:
|
|
return []
|
|
result = result[idx + skip:]
|
|
return result[:take]
|
|
|
|
async def fake_find_first(where):
|
|
return next(
|
|
(r for r in rows if r.unified_object_id == where.get("unified_object_id")),
|
|
None,
|
|
)
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
|
side_effect=fake_find_many
|
|
)
|
|
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
|
|
side_effect=fake_find_first
|
|
)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
user = UserAPIKeyAuth(user_id="test-user")
|
|
|
|
seen: list = []
|
|
after = None
|
|
for _ in range(total + 5):
|
|
resp = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=user, limit=2, after=after
|
|
)
|
|
page_ids = [b.id for b in resp["data"]]
|
|
seen.extend(page_ids)
|
|
if not resp["has_more"]:
|
|
break
|
|
assert page_ids, "has_more was true but the page was empty"
|
|
assert resp["last_id"] != after, "cursor did not advance (pagination loop)"
|
|
after = resp["last_id"]
|
|
|
|
assert sorted(seen) == sorted(_unified_id(i) for i in range(total))
|
|
assert len(seen) == len(set(seen)), "a tied batch was returned more than once"
|
|
|
|
|
|
def _managed_batch_row(index, file_object=None):
|
|
row = MagicMock()
|
|
row.id = f"pk-{index:03d}"
|
|
raw = f"litellm_proxy;model_id:gpt-4o-batch;llm_batch_id:batch_{index:03d}"
|
|
row.unified_object_id = base64.urlsafe_b64encode(raw.encode()).decode().rstrip("=")
|
|
row.created_at = 1_000_000 + index
|
|
row.file_object = (
|
|
file_object
|
|
if file_object is not None
|
|
else json.dumps(
|
|
{
|
|
"id": f"batch_provider_{index:03d}",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "completed",
|
|
"created_at": 1_000_000 + index,
|
|
"input_file_id": f"file-input-{index:03d}",
|
|
"request_counts": {"total": 1, "completed": 1, "failed": 0},
|
|
}
|
|
)
|
|
)
|
|
return row
|
|
|
|
|
|
def _fake_managed_object_table(rows):
|
|
async def find_many(where, take, order, cursor=None, skip=0):
|
|
result = sorted(
|
|
rows, key=lambda r: (r.created_at, r.unified_object_id), reverse=True
|
|
)
|
|
if cursor is not None:
|
|
(cur_field, cur_val), = cursor.items()
|
|
idx = next(
|
|
(i for i, r in enumerate(result) if getattr(r, cur_field) == cur_val),
|
|
None,
|
|
)
|
|
if idx is None:
|
|
return []
|
|
result = result[idx + skip:]
|
|
return result[:take]
|
|
|
|
async def find_first(where):
|
|
return next(
|
|
(r for r in rows if r.unified_object_id == where.get("unified_object_id")),
|
|
None,
|
|
)
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
|
side_effect=find_many
|
|
)
|
|
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
|
|
side_effect=find_first
|
|
)
|
|
return prisma_client
|
|
|
|
|
|
async def _walk_batch_pages(proxy_managed_files, user, limit, max_pages=20):
|
|
pages = []
|
|
after = None
|
|
for _ in range(max_pages):
|
|
resp = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=user, limit=limit, after=after
|
|
)
|
|
pages.append(resp)
|
|
if not resp["has_more"]:
|
|
break
|
|
assert resp["last_id"] is not None, "has_more was true but there is no cursor"
|
|
after = resp["last_id"]
|
|
return pages
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_rejects_unknown_after_cursor():
|
|
"""An ``after`` that does not resolve to a batch the caller can see is a
|
|
client error, not an empty page.
|
|
|
|
Returning ``[]`` for an unresolvable cursor is indistinguishable from
|
|
"you have reached the end of the list", so a client walking pages with a
|
|
stale or malformed cursor silently sees a truncated batch list instead of
|
|
an error it can act on.
|
|
"""
|
|
from fastapi import HTTPException
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
|
|
return_value=None
|
|
)
|
|
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[])
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
|
|
limit=3,
|
|
after="does-not-exist-xyz",
|
|
)
|
|
|
|
assert exc_info.value.status_code == 400
|
|
assert "does-not-exist-xyz" in str(exc_info.value.detail)
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
|
|
prisma_client.db.litellm_managedobjecttable.find_first.assert_called_once_with(
|
|
where={
|
|
"file_purpose": "batch",
|
|
"created_by": "test-user",
|
|
"unified_object_id": "does-not-exist-xyz",
|
|
}
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_treats_empty_after_as_no_cursor():
|
|
"""``?after=`` means "start from the beginning", as it always has.
|
|
|
|
Only a cursor the client actually sent is validated, so an SDK that always
|
|
emits the query parameter does not get a 400 on its first page.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
rows = [_managed_batch_row(i) for i in range(2)]
|
|
prisma_client = _fake_managed_object_table(rows)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
page = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), limit=2, after=""
|
|
)
|
|
|
|
assert [batch.id for batch in page["data"]] == [
|
|
rows[1].unified_object_id,
|
|
rows[0].unified_object_id,
|
|
]
|
|
prisma_client.db.litellm_managedobjecttable.find_first.assert_not_called()
|
|
_, call_kwargs = prisma_client.db.litellm_managedobjecttable.find_many.call_args
|
|
assert "cursor" not in call_kwargs
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_rejects_after_cursor_owned_by_another_user():
|
|
"""The cursor lookup must be scoped to the rows the caller can list.
|
|
|
|
A Prisma cursor resolves by unique column regardless of the ``where``
|
|
filter, so an unscoped cursor would let one user anchor their page window
|
|
to another user's batch and learn when it was created.
|
|
"""
|
|
from fastapi import HTTPException
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
other_users_batch = _managed_batch_row(0)
|
|
|
|
async def find_first(where):
|
|
if where.get("created_by") != "user-b":
|
|
return None
|
|
return other_users_batch
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
|
|
side_effect=find_first
|
|
)
|
|
prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[])
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="user-a"),
|
|
limit=3,
|
|
after=other_users_batch.unified_object_id,
|
|
)
|
|
|
|
assert exc_info.value.status_code == 400
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_has_more_false_on_exactly_full_final_page():
|
|
"""``has_more`` must mean "another row exists", not "this page is full".
|
|
|
|
With a batch count that is an exact multiple of ``limit``, reporting
|
|
``has_more`` off page fullness makes every client fetch one extra empty
|
|
page before it can stop.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
rows = [_managed_batch_row(i) for i in range(4)]
|
|
prisma_client = _fake_managed_object_table(rows)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
pages = await _walk_batch_pages(
|
|
proxy_managed_files, UserAPIKeyAuth(user_id="test-user"), limit=2
|
|
)
|
|
|
|
seen = [batch.id for page in pages for batch in page["data"]]
|
|
assert seen == [r.unified_object_id for r in reversed(rows)]
|
|
assert [page["has_more"] for page in pages] == [True, False]
|
|
assert prisma_client.db.litellm_managedobjecttable.find_many.call_count == 2
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_unparseable_row_does_not_truncate_pagination():
|
|
"""A row that fails to parse must not end pagination early.
|
|
|
|
Skipping a corrupt row shortens the page, so deriving ``has_more`` from
|
|
the number of returned batches reports "no more results" while older
|
|
batches are still unread, silently hiding them from the caller.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
rows = [_managed_batch_row(i) for i in range(4)]
|
|
rows[2].file_object = "{ not valid json"
|
|
prisma_client = _fake_managed_object_table(rows)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
pages = await _walk_batch_pages(
|
|
proxy_managed_files, UserAPIKeyAuth(user_id="test-user"), limit=2
|
|
)
|
|
|
|
seen = [batch.id for page in pages for batch in page["data"]]
|
|
assert seen == [rows[3].unified_object_id, rows[1].unified_object_id, rows[0].unified_object_id]
|
|
assert len(seen) == len(set(seen))
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_fills_a_page_past_a_full_page_of_unparseable_rows():
|
|
"""A page whose rows all fail to parse must still let the caller advance.
|
|
|
|
``has_more`` came from the raw fetch while ``last_id`` came from the parsed
|
|
survivors, so a full page of corrupt rows answered ``data: []``,
|
|
``last_id: None``, ``has_more: True``, and a client following ``last_id``
|
|
could not move past them.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
rows = [_managed_batch_row(i) for i in range(5)]
|
|
for corrupt_row in rows[2:4]:
|
|
corrupt_row.file_object = "{ not valid json"
|
|
prisma_client = _fake_managed_object_table(rows)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
pages = await _walk_batch_pages(
|
|
proxy_managed_files, UserAPIKeyAuth(user_id="test-user"), limit=1
|
|
)
|
|
|
|
assert [[batch.id for batch in page["data"]] for page in pages] == [
|
|
[rows[4].unified_object_id],
|
|
[rows[1].unified_object_id],
|
|
[rows[0].unified_object_id],
|
|
]
|
|
assert [page["has_more"] for page in pages] == [True, True, False]
|
|
|
|
|
|
_DEEP_BATCH_SCAN_ROW_COUNT = 2000
|
|
_DEEP_BATCH_SCAN_QUERY_BUDGET = 10
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_bounds_the_queries_a_deep_unparseable_run_costs():
|
|
"""A tiny limit behind thousands of corrupt rows must not turn one request into thousands of queries."""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
rows = [_managed_batch_row(0)] + [
|
|
_managed_batch_row(index, file_object="{ not valid json")
|
|
for index in range(1, _DEEP_BATCH_SCAN_ROW_COUNT + 1)
|
|
]
|
|
prisma_client = _fake_managed_object_table(rows)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
page = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), limit=1
|
|
)
|
|
|
|
assert [batch.id for batch in page["data"]] == [rows[0].unified_object_id]
|
|
assert page["has_more"] is False
|
|
assert (
|
|
prisma_client.db.litellm_managedobjecttable.find_many.call_count
|
|
<= _DEEP_BATCH_SCAN_QUERY_BUDGET
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_reads_one_chunk_when_the_first_one_fills_the_page():
|
|
"""The widened chunk must stay off the common path, where the newest rows already fill the page."""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
rows = [_managed_batch_row(index) for index in range(_DEEP_BATCH_SCAN_ROW_COUNT)]
|
|
prisma_client = _fake_managed_object_table(rows)
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
page = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), limit=2
|
|
)
|
|
|
|
assert [batch.id for batch in page["data"]] == [
|
|
rows[-1].unified_object_id,
|
|
rows[-2].unified_object_id,
|
|
]
|
|
assert page["has_more"] is True
|
|
assert prisma_client.db.litellm_managedobjecttable.find_many.call_count == 1
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_return_unified_file_id_includes_expires_at():
|
|
from litellm.types.llms.openai import OpenAIFileObject
|
|
|
|
# Create a mock file object with expires_at set
|
|
file_object = OpenAIFileObject(
|
|
id="file-abc123",
|
|
object="file",
|
|
bytes=1234,
|
|
created_at=1234567890,
|
|
filename="test.jsonl",
|
|
purpose="batch",
|
|
status="uploaded",
|
|
expires_at=1234657890,
|
|
)
|
|
file_object._hidden_params = {"model_id": "test-model-id"}
|
|
|
|
create_file_request = {
|
|
"file": ("test.jsonl", b"test content", "application/jsonl"),
|
|
"purpose": "batch",
|
|
}
|
|
|
|
internal_usage_cache = MagicMock()
|
|
|
|
result = await _PROXY_LiteLLMManagedFiles.return_unified_file_id(
|
|
file_objects=[file_object],
|
|
create_file_request=create_file_request,
|
|
internal_usage_cache=internal_usage_cache,
|
|
litellm_parent_otel_span=None,
|
|
target_model_names_list=["gpt-5.5"],
|
|
)
|
|
|
|
# Verify expires_at is passed through
|
|
assert result.expires_at == 1234657890
|
|
assert result.purpose == "batch"
|
|
assert result.filename == "test.jsonl"
|
|
assert result.bytes == 1234
|
|
assert result.created_at == 1234567890
|
|
assert _is_base64_encoded_unified_file_id(result.id)
|
|
|
|
|
|
# ============================================================================
|
|
# Permission Tests - Cross-User Batch Access
|
|
# ============================================================================
|
|
# These tests verify that batches and files created by one user
|
|
# cannot be accessed, modified, or cancelled by a different user.
|
|
# Reference: https://github.com/BerriAI/litellm/pull/17401/files
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_user_b_cannot_retrieve_user_a_batch():
|
|
"""
|
|
Test that User B cannot retrieve a batch created by User A.
|
|
|
|
This verifies batch isolation between users at the database/hook level.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Mock database to return User A as the creator
|
|
batch_record = MagicMock()
|
|
batch_record.created_by = "user_a_id"
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
# User B tries to retrieve User A's batch
|
|
unified_batch_id = (
|
|
"bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz"
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b_id", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"batch_id": unified_batch_id},
|
|
call_type="aretrieve_batch",
|
|
)
|
|
|
|
# Should raise 403 Permission Denied
|
|
assert exc_info.value.status_code == 403
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_user_b_cannot_cancel_user_a_batch():
|
|
"""
|
|
Test that User B cannot cancel a batch created by User A.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Mock database to return User A as the creator
|
|
batch_record = MagicMock()
|
|
batch_record.created_by = "user_a_id"
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
# User B tries to cancel User A's batch
|
|
unified_batch_id = (
|
|
"bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz"
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b_id", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"batch_id": unified_batch_id},
|
|
call_type="acancel_batch",
|
|
)
|
|
|
|
# Should raise 403 Permission Denied
|
|
assert exc_info.value.status_code == 403
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_user_a_can_retrieve_own_batch():
|
|
"""
|
|
Test that User A can successfully retrieve their own batch.
|
|
|
|
This is a positive test case to ensure permission checks don't block
|
|
legitimate access.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Mock database to return User A as the creator
|
|
batch_record = MagicMock()
|
|
batch_record.created_by = "user_a_id"
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
# User A retrieves their own batch
|
|
unified_batch_id = (
|
|
"bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz"
|
|
)
|
|
|
|
# Should not raise an exception
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_a_id", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"batch_id": unified_batch_id},
|
|
call_type="aretrieve_batch",
|
|
)
|
|
|
|
# Should successfully return the decoded batch_id
|
|
assert "batch_id" in result
|
|
assert result["model"] == "my-model"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_user_b_cannot_retrieve_user_a_file():
|
|
"""
|
|
Test that User B cannot retrieve a file created by User A.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Mock database to return User A as the creator
|
|
file_record = MagicMock()
|
|
file_record.created_by = "user_a_id"
|
|
prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(), prisma_client=prisma_client
|
|
)
|
|
|
|
# User B tries to retrieve User A's file
|
|
unified_file_id = (
|
|
"bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM"
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b_id", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"file_id": unified_file_id},
|
|
call_type="afile_retrieve",
|
|
)
|
|
|
|
# Should raise 403 Permission Denied
|
|
assert exc_info.value.status_code == 403
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_user_b_cannot_download_user_a_file_content():
|
|
"""
|
|
Test that User B cannot download file content for User A's file.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Mock database to return User A as the creator
|
|
file_record = MagicMock()
|
|
file_record.created_by = "user_a_id"
|
|
prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(), prisma_client=prisma_client
|
|
)
|
|
|
|
# User B tries to download User A's file content
|
|
unified_file_id = (
|
|
"bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM"
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b_id", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"file_id": unified_file_id},
|
|
call_type="afile_content",
|
|
)
|
|
|
|
# Should raise 403 Permission Denied
|
|
assert exc_info.value.status_code == 403
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_user_b_cannot_delete_user_a_file():
|
|
"""
|
|
Test that User B cannot delete a file created by User A.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Mock database to return User A as the creator
|
|
file_record = MagicMock()
|
|
file_record.created_by = "user_a_id"
|
|
prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(), prisma_client=prisma_client
|
|
)
|
|
|
|
# User B tries to delete User A's file
|
|
unified_file_id = (
|
|
"bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM"
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b_id", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"file_id": unified_file_id},
|
|
call_type="afile_delete",
|
|
)
|
|
|
|
# Should raise 403 Permission Denied
|
|
assert exc_info.value.status_code == 403
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_user_a_can_retrieve_own_file():
|
|
"""
|
|
Test that User A can successfully retrieve their own file.
|
|
|
|
Positive test case to ensure permission checks work correctly for the owner.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Mock database to return User A as the creator
|
|
file_record = MagicMock()
|
|
file_record.created_by = "user_a_id"
|
|
file_record.model_mappings = '{"model-123": "file-abc123"}'
|
|
file_record.file_object = json.dumps(
|
|
{
|
|
"id": "file-abc123",
|
|
"object": "file",
|
|
"bytes": 1234,
|
|
"created_at": 1234567890,
|
|
"filename": "test.jsonl",
|
|
"purpose": "batch",
|
|
}
|
|
)
|
|
prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(), prisma_client=prisma_client
|
|
)
|
|
|
|
# User A retrieves their own file
|
|
unified_file_id = (
|
|
"bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsZmlsZS1hYmMxMjM"
|
|
)
|
|
|
|
# Should not raise an exception
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_a_id", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"file_id": unified_file_id},
|
|
call_type="afile_retrieve",
|
|
)
|
|
|
|
# Should successfully return the decoded file_id
|
|
assert "file_id" in result
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_batches_only_returns_user_own_batches():
|
|
"""
|
|
Test that list_user_batches only returns batches created by the requesting user.
|
|
|
|
This ensures users cannot see other users' batches in list operations.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Create batches for User A
|
|
batch_user_a = MagicMock()
|
|
batch_user_a.unified_object_id = "batch-user-a"
|
|
batch_user_a.file_object = json.dumps(
|
|
{
|
|
"id": "batch_a",
|
|
"object": "batch",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
"status": "completed",
|
|
"created_at": 1234567890,
|
|
"input_file_id": "file-a",
|
|
"request_counts": {"total": 1, "completed": 1, "failed": 0},
|
|
}
|
|
)
|
|
|
|
# Mock database to only return User A's batches
|
|
prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user_a]
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
# User A requests their batches
|
|
result = await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="user_a_id"),
|
|
limit=10,
|
|
)
|
|
|
|
# Should only return User A's batches
|
|
assert len(result["data"]) == 1
|
|
assert result["data"][0].id == "batch-user-a"
|
|
|
|
# Verify the database query filtered by user_id
|
|
prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with(
|
|
where={"file_purpose": "batch", "created_by": "user_a_id"},
|
|
take=11,
|
|
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_same_user_different_keys_can_access_batch():
|
|
"""
|
|
Test that different API keys for the same user can access the same batch.
|
|
|
|
This verifies that permission checks are based on user_id, not API key,
|
|
allowing users to have multiple keys that can all access their resources.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
|
|
# Mock database to return the user_id as creator
|
|
batch_record = MagicMock()
|
|
batch_record.created_by = "user_a_id"
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
unified_batch_id = (
|
|
"bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1tb2RlbDtsbG1fYmF0Y2hfaWQ6YmF0Y2hfYWJjMTIz"
|
|
)
|
|
|
|
# First API key for User A retrieves the batch
|
|
result1 = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_a_id", api_key="key-1", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"batch_id": unified_batch_id},
|
|
call_type="aretrieve_batch",
|
|
)
|
|
|
|
assert "batch_id" in result1
|
|
|
|
# Second API key for the same User A retrieves the batch
|
|
result2 = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_a_id", api_key="key-2", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"batch_id": unified_batch_id},
|
|
call_type="aretrieve_batch",
|
|
)
|
|
|
|
assert "batch_id" in result2
|
|
# Both keys should get the same result
|
|
assert result1["batch_id"] == result2["batch_id"]
|
|
|
|
|
|
MODEL_ENCODED_BATCH_ID = encode_file_id_with_model(
|
|
"batch_provider123", "gpt-4o-team-alias", id_type="batch"
|
|
)
|
|
MODEL_ENCODED_OUTPUT_FILE_ID = encode_file_id_with_model(
|
|
"file-output456", "gpt-4o-team-alias", id_type="file"
|
|
)
|
|
RAW_PROVIDER_BATCH_ID = "batch_provider123"
|
|
RAW_PROVIDER_FILE_ID = "file-output456"
|
|
|
|
|
|
def _owned_record(created_by, team_id):
|
|
record = MagicMock()
|
|
record.created_by = created_by
|
|
record.team_id = team_id
|
|
return record
|
|
|
|
|
|
def _batch_response(batch_id, output_file_id=None, is_create=False):
|
|
from litellm.types.utils import LiteLLMBatch
|
|
|
|
batch = LiteLLMBatch(
|
|
id=batch_id,
|
|
completion_window="24h",
|
|
created_at=1700000000,
|
|
endpoint="/v1/chat/completions",
|
|
input_file_id="file-input789",
|
|
object="batch",
|
|
status="completed",
|
|
output_file_id=output_file_id,
|
|
)
|
|
if is_create:
|
|
batch._hidden_params[BATCH_CREATE_HIDDEN_PARAM] = True
|
|
return batch
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("call_type", ["aretrieve_batch", "acancel_batch"])
|
|
@pytest.mark.parametrize(
|
|
"batch_id", [MODEL_ENCODED_BATCH_ID, RAW_PROVIDER_BATCH_ID]
|
|
)
|
|
async def test_team_b_cannot_access_team_a_provider_format_batch(
|
|
call_type, batch_id
|
|
):
|
|
"""
|
|
Cross-team retrieve/cancel of a model-encoded or raw provider batch id
|
|
must 403 when an ownership row exists for another team.
|
|
|
|
Regression test: before this check only unified (litellm_proxy-prefixed)
|
|
batch ids were enforced, so any key could read any model-encoded or raw
|
|
provider batch.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
|
_owned_record(created_by="user_a", team_id="team_a")
|
|
)
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"batch_id": batch_id},
|
|
call_type=call_type,
|
|
)
|
|
|
|
assert exc_info.value.status_code == 403
|
|
prisma_client.db.litellm_managedobjecttable.find_first.assert_awaited_once_with(
|
|
where={"OR": [{"unified_object_id": batch_id}, {"model_object_id": batch_id}]}
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"caller_kwargs",
|
|
[
|
|
{"user_id": "user_a", "team_id": "team_a"},
|
|
{"user_id": "teammate_of_a", "team_id": "team_a"},
|
|
{"user_id": "admin_user", "user_role": "proxy_admin"},
|
|
],
|
|
)
|
|
async def test_authorized_callers_can_access_provider_format_batch(caller_kwargs):
|
|
"""
|
|
The creator, a same-team member, and a proxy admin can all retrieve a
|
|
model-encoded batch owned by team_a. Data must pass through unmodified so
|
|
the endpoint's own routing still applies.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
|
_owned_record(created_by="user_a", team_id="team_a")
|
|
)
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
parent_otel_span=MagicMock(), **caller_kwargs
|
|
),
|
|
cache=MagicMock(),
|
|
data={"batch_id": MODEL_ENCODED_BATCH_ID},
|
|
call_type="aretrieve_batch",
|
|
)
|
|
|
|
assert result["batch_id"] == MODEL_ENCODED_BATCH_ID
|
|
assert "model" not in result
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_provider_format_batch_without_ownership_row_stays_accessible():
|
|
"""
|
|
A provider-format batch id with no ownership row (created before ownership
|
|
tracking, or directly on the provider account) must stay retrievable.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = None
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"batch_id": RAW_PROVIDER_BATCH_ID},
|
|
call_type="aretrieve_batch",
|
|
)
|
|
|
|
assert result["batch_id"] == RAW_PROVIDER_BATCH_ID
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_fine_tuning_provider_format_id_not_enforced():
|
|
"""
|
|
Provider-format fine-tuning job ids are deliberately out of scope for
|
|
ownership enforcement; only unified fine-tuning ids are checked.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
|
_owned_record(created_by="user_a", team_id="team_a")
|
|
)
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"fine_tuning_job_id": "ftjob-abc123"},
|
|
call_type="aretrieve_fine_tuning_job",
|
|
)
|
|
|
|
assert result["fine_tuning_job_id"] == "ftjob-abc123"
|
|
prisma_client.db.litellm_managedobjecttable.find_first.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"call_type", ["afile_content", "afile_retrieve", "afile_delete"]
|
|
)
|
|
@pytest.mark.parametrize(
|
|
"file_id", [MODEL_ENCODED_OUTPUT_FILE_ID, RAW_PROVIDER_FILE_ID]
|
|
)
|
|
async def test_team_b_cannot_access_team_a_provider_format_file(
|
|
call_type, file_id
|
|
):
|
|
"""
|
|
Cross-team content/retrieve/delete of a model-encoded or raw provider
|
|
file id must 403 when an ownership row exists for another team.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.find_first.return_value = (
|
|
_owned_record(created_by="user_a", team_id="team_a")
|
|
)
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(), prisma_client=prisma_client
|
|
)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"file_id": file_id},
|
|
call_type=call_type,
|
|
)
|
|
|
|
assert exc_info.value.status_code == 403
|
|
prisma_client.db.litellm_managedfiletable.find_first.assert_awaited_once_with(
|
|
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_same_team_can_access_provider_format_file():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.find_first.return_value = (
|
|
_owned_record(created_by="user_a", team_id="team_a")
|
|
)
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="teammate_of_a", team_id="team_a", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"file_id": MODEL_ENCODED_OUTPUT_FILE_ID},
|
|
call_type="afile_content",
|
|
)
|
|
|
|
assert result["file_id"] == MODEL_ENCODED_OUTPUT_FILE_ID
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_provider_format_file_without_ownership_row_stays_accessible():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.find_first.return_value = None
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(), prisma_client=prisma_client
|
|
)
|
|
|
|
result = await proxy_managed_files.async_pre_call_hook(
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
|
),
|
|
cache=MagicMock(),
|
|
data={"file_id": RAW_PROVIDER_FILE_ID},
|
|
call_type="afile_content",
|
|
)
|
|
|
|
assert result["file_id"] == RAW_PROVIDER_FILE_ID
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("batch_id", [MODEL_ENCODED_BATCH_ID, RAW_PROVIDER_BATCH_ID])
|
|
async def test_post_call_batch_create_stores_ownership_row(batch_id):
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
|
)
|
|
|
|
await proxy_managed_files.async_post_call_success_hook(
|
|
data={
|
|
"input_file_id": "file-input789",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
},
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
|
|
),
|
|
response=_batch_response(batch_id, is_create=True),
|
|
)
|
|
|
|
upsert_call = prisma_client.db.litellm_managedobjecttable.upsert.await_args
|
|
assert upsert_call.kwargs["where"] == {"unified_object_id": batch_id}
|
|
create_data = upsert_call.kwargs["data"]["create"]
|
|
assert create_data["created_by"] == "user_a"
|
|
assert create_data["team_id"] == "team_a"
|
|
prisma_client.db.litellm_managedobjecttable.update_many.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_post_call_batch_sync_does_not_claim_ownership():
|
|
"""
|
|
Retrieve/cancel of a batch with no ownership row must NOT create one:
|
|
otherwise the first foreign key to touch a legacy batch would become its
|
|
owner and lock out the real creator once enforcement is on.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 0
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
|
)
|
|
|
|
await proxy_managed_files.async_post_call_success_hook(
|
|
data={"batch_id": MODEL_ENCODED_BATCH_ID},
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
|
),
|
|
response=_batch_response(MODEL_ENCODED_BATCH_ID),
|
|
)
|
|
|
|
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
|
|
prisma_client.db.litellm_managedobjecttable.update_many.assert_awaited_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_post_call_batch_sync_updates_existing_row():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
|
_owned_record(created_by="user_a", team_id="team_a")
|
|
)
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
|
)
|
|
|
|
await proxy_managed_files.async_post_call_success_hook(
|
|
data={"batch_id": MODEL_ENCODED_BATCH_ID},
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
|
|
),
|
|
response=_batch_response(MODEL_ENCODED_BATCH_ID),
|
|
)
|
|
|
|
update_call = prisma_client.db.litellm_managedobjecttable.update_many.await_args
|
|
assert update_call.kwargs["where"] == {
|
|
"unified_object_id": MODEL_ENCODED_BATCH_ID
|
|
}
|
|
assert update_call.kwargs["data"]["status"] == "completed"
|
|
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_post_call_batch_sync_stores_output_file_ownership_from_batch_row():
|
|
"""
|
|
When a synced batch reports a provider-format output file id, an
|
|
ownership row for that file must be written with the BATCH row's
|
|
created_by/team_id, not the caller's identity.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1
|
|
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
|
_owned_record(created_by="user_a", team_id="team_a")
|
|
)
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
|
)
|
|
|
|
await proxy_managed_files.async_post_call_success_hook(
|
|
data={"batch_id": MODEL_ENCODED_BATCH_ID},
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="admin_user",
|
|
user_role="proxy_admin",
|
|
parent_otel_span=MagicMock(),
|
|
),
|
|
response=_batch_response(
|
|
MODEL_ENCODED_BATCH_ID, output_file_id=MODEL_ENCODED_OUTPUT_FILE_ID
|
|
),
|
|
)
|
|
|
|
file_upsert = prisma_client.db.litellm_managedfiletable.upsert.await_args
|
|
assert file_upsert.kwargs["where"] == {
|
|
"unified_file_id": MODEL_ENCODED_OUTPUT_FILE_ID
|
|
}
|
|
create_data = file_upsert.kwargs["data"]["create"]
|
|
assert create_data["created_by"] == "user_a"
|
|
assert create_data["team_id"] == "team_a"
|
|
assert create_data["flat_model_file_ids"] == ["file-output456"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_post_call_batch_create_does_not_store_output_file_ownership():
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
prisma_client = AsyncMock()
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
|
)
|
|
|
|
await proxy_managed_files.async_post_call_success_hook(
|
|
data={
|
|
"input_file_id": "file-input789",
|
|
"endpoint": "/v1/chat/completions",
|
|
"completion_window": "24h",
|
|
},
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
|
|
),
|
|
response=_batch_response(
|
|
MODEL_ENCODED_BATCH_ID,
|
|
output_file_id=MODEL_ENCODED_OUTPUT_FILE_ID,
|
|
is_create=True,
|
|
),
|
|
)
|
|
|
|
prisma_client.db.litellm_managedfiletable.upsert.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_file_list_cursors_are_scoped_to_the_caller():
|
|
"""A non-owner must not learn other callers' file ids through the page cursors."""
|
|
from openai.pagination import AsyncCursorPage
|
|
from openai.types import FileObject
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
owner_file = FileObject(
|
|
id="file-owner-1",
|
|
bytes=100,
|
|
created_at=1,
|
|
filename="owner.jsonl",
|
|
object="file",
|
|
purpose="batch",
|
|
status="processed",
|
|
)
|
|
upstream_page = AsyncCursorPage[FileObject].construct(
|
|
data=[owner_file],
|
|
has_more=True,
|
|
first_id=owner_file.id,
|
|
last_id=owner_file.id,
|
|
object="list",
|
|
)
|
|
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.find_many.return_value = []
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
response = await proxy_managed_files.async_post_call_success_hook(
|
|
data={},
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="other-user", team_id="other-team", parent_otel_span=MagicMock()
|
|
),
|
|
response=upstream_page,
|
|
)
|
|
|
|
assert response.data == []
|
|
assert response.first_id is None
|
|
assert response.last_id is None
|
|
assert response.has_more is False
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_file_list_cursors_follow_the_owner_scoped_page():
|
|
from openai.pagination import AsyncCursorPage
|
|
from openai.types import FileObject
|
|
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
def _raw_file(file_id: str) -> FileObject:
|
|
return FileObject(
|
|
id=file_id,
|
|
bytes=100,
|
|
created_at=1,
|
|
filename=f"{file_id}.jsonl",
|
|
object="file",
|
|
purpose="batch",
|
|
status="processed",
|
|
)
|
|
|
|
upstream_page = AsyncCursorPage[FileObject].construct(
|
|
data=[_raw_file("file-someone-else"), _raw_file("file-mine")],
|
|
has_more=True,
|
|
first_id="file-someone-else",
|
|
last_id="file-mine",
|
|
object="list",
|
|
)
|
|
|
|
managed_row = MagicMock()
|
|
managed_row.unified_file_id = "litellm_proxy:mine"
|
|
managed_row.file_object = {
|
|
"id": "file-mine",
|
|
"bytes": 100,
|
|
"created_at": 1,
|
|
"filename": "mine.jsonl",
|
|
"object": "file",
|
|
"purpose": "batch",
|
|
"status": "processed",
|
|
}
|
|
prisma_client = AsyncMock()
|
|
prisma_client.db.litellm_managedfiletable.find_many.return_value = [managed_row]
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=prisma_client
|
|
)
|
|
|
|
response = await proxy_managed_files.async_post_call_success_hook(
|
|
data={},
|
|
user_api_key_dict=UserAPIKeyAuth(
|
|
user_id="mine-user", parent_otel_span=MagicMock()
|
|
),
|
|
response=upstream_page,
|
|
)
|
|
|
|
assert [file_object.id for file_object in response.data] == ["litellm_proxy:mine"]
|
|
assert response.first_id == "litellm_proxy:mine"
|
|
assert response.last_id == "litellm_proxy:mine"
|
|
assert response.has_more is False
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_user_batches_provider_filter_rejected_with_400():
|
|
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
|
|
with pytest.raises(ProxyException) as exc:
|
|
await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="123"),
|
|
provider="openai",
|
|
)
|
|
|
|
assert exc.value.code == "400"
|
|
assert exc.value.type == "invalid_request_error"
|
|
assert exc.value.param == "provider"
|
|
assert exc.value.message == "Filtering by 'provider' is not supported when using managed batches."
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_list_user_batches_target_model_names_filter_rejected_with_400():
|
|
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
|
|
|
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
|
DualCache(), prisma_client=MagicMock()
|
|
)
|
|
|
|
with pytest.raises(ProxyException) as exc:
|
|
await proxy_managed_files.list_user_batches(
|
|
user_api_key_dict=UserAPIKeyAuth(user_id="123"),
|
|
target_model_names="gpt-4o",
|
|
)
|
|
|
|
assert exc.value.code == "400"
|
|
assert exc.value.type == "invalid_request_error"
|
|
assert exc.value.param == "target_model_names"
|
|
assert exc.value.message == "Filtering by 'target_model_names' is not supported when using managed batches."
|