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."