diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index a83d7e224b5..445d2b242b4 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cas from fastapi import HTTPException +import litellm from litellm import Router, verbose_logger from litellm._uuid import uuid from litellm.caching.caching import DualCache @@ -836,15 +837,36 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): return response async def afile_retrieve( - self, file_id: str, litellm_parent_otel_span: Optional[Span] + self, file_id: str, litellm_parent_otel_span: Optional[Span], llm_router=None ) -> OpenAIFileObject: stored_file_object = await self.get_unified_file_id( file_id, litellm_parent_otel_span ) - if stored_file_object: - return stored_file_object.file_object - else: + + # Case 1 : This is not a managed file + if not stored_file_object: raise Exception(f"LiteLLM Managed File object with id={file_id} not found") + + # Case 2: Managed file and the file object exists in the database + if stored_file_object and stored_file_object.file_object: + return stored_file_object.file_object + + # Case 3: Managed file exists in the database but not the file object (for. e.g the batch task might not have run) + # So we fetch the file object from the provider. We deliberately do not store the result to avoid interfering with batch cost tracking code. + if not llm_router: + raise Exception( + f"LiteLLM Managed File object with id={file_id} has no file_object " + f"and llm_router is required to fetch from provider" + ) + + try: + model_id, model_file_id = next(iter(stored_file_object.model_mappings.items())) + credentials = llm_router.get_deployment_credentials_with_provider(model_id) or {} + response = await litellm.afile_retrieve(file_id=model_file_id, **credentials) + response.id = file_id # Replace with unified ID + return response + except Exception as e: + raise Exception(f"Failed to retrieve file {file_id} from provider: {str(e)}") from e async def afile_list( self, @@ -868,10 +890,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): [file_id], litellm_parent_otel_span ) + delete_response = None specific_model_file_id_mapping = model_file_id_mapping.get(file_id) if specific_model_file_id_mapping: for model_id, model_file_id in specific_model_file_id_mapping.items(): - await llm_router.afile_delete(model=model_id, file_id=model_file_id, **data) # type: ignore + delete_response = await llm_router.afile_delete(model=model_id, file_id=model_file_id, **data) # type: ignore stored_file_object = await self.delete_unified_file_id( file_id, litellm_parent_otel_span @@ -879,6 +902,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if stored_file_object: return stored_file_object + elif delete_response: + delete_response.id = file_id + return delete_response else: raise Exception(f"LiteLLM Managed File object with id={file_id} not found") diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py index 7b0a1868f19..58df15f0c46 100644 --- a/litellm/llms/base_llm/files/transformation.py +++ b/litellm/llms/base_llm/files/transformation.py @@ -227,6 +227,7 @@ class BaseFileEndpoints(ABC): self, file_id: str, litellm_parent_otel_span: Optional[Span], + llm_router: Optional[Router] = None, ) -> OpenAIFileObject: pass diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8c0a2e6ae2a..559d70ab797 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3805,7 +3805,7 @@ class SpendUpdateQueueItem(TypedDict, total=False): class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase): unified_file_id: str - file_object: OpenAIFileObject + file_object: Optional[OpenAIFileObject] = None model_mappings: Dict[str, str] flat_model_file_ids: List[str] created_by: Optional[str] diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 810f5c62720..7e3f5820814 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -885,6 +885,7 @@ async def get_file( response = await managed_files_obj.afile_retrieve( file_id=file_id, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + llm_router=llm_router, ) else: response = await litellm.afile_retrieve( diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 5f66b03aad4..9a6e153a22b 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -1,18 +1,14 @@ import json -import os -import sys from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException -from fastapi.testclient import TestClient from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles from litellm.caching import DualCache from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, ) -from litellm.types.utils import SpecialEnums def test_get_file_ids_from_messages(): @@ -255,7 +251,7 @@ async def test_can_user_call_unified_file_id(call_type): ) unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxmMTNlNDAzZS01YWM3LTRhZjktOGQzNS0wNDgwZDMxOTgyYTg7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00by1taW5pLW9wZW5haTtsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1Ib3UxZDFXc3c1SDNKcjFMYllpZDJiO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxmODBiNWU2NzQ1NzdkNjkyMjM4YmVhNTIxZDdiMGI5ZGYyY2FmMTEwMTU2YmU5YzBjM2NjMmNkNTBjOTM1ZDI0" - with pytest.raises(HTTPException) as e: + with pytest.raises(HTTPException): await proxy_managed_files.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth( user_id="456", parent_otel_span=MagicMock() @@ -310,7 +306,7 @@ async def test_router_acreate_batch_only_selects_from_file_id_mapping(monkeypatc litellm, "acreate_batch", return_value=AsyncMock() ) as mock_acreate_batch: for _ in range(1000): - response = await router.acreate_batch( + await router.acreate_batch( model="gpt-3.5-turbo", input_file_id=file_id, model_file_id_mapping=model_file_id_mapping, @@ -329,7 +325,6 @@ async def test_output_file_id_for_batch_retrieve(): from openai.types.batch import BatchRequestCounts - from litellm.proxy._types import UserAPIKeyAuth from litellm.types.utils import LiteLLMBatch batch = LiteLLMBatch( @@ -381,8 +376,6 @@ async def test_output_file_id_for_batch_retrieve(): @pytest.mark.asyncio async def test_async_post_call_success_hook_twice_assert_no_unique_violation(): import asyncio - from litellm.proxy.proxy_server import proxy_logging_obj - from litellm.proxy.utils import PrismaClient from litellm.types.utils import LiteLLMBatch from litellm.proxy._types import UserAPIKeyAuth from openai.types.batch import BatchRequestCounts @@ -590,3 +583,247 @@ def test_update_responses_input_with_multiple_file_ids(): 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" + + +@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.create = 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 create was called with expected data (without file_object) + prisma_client.db.litellm_managedfiletable.create.assert_called_once() + call_args = prisma_client.db.litellm_managedfiletable.create.call_args + assert call_args.kwargs["data"]["unified_file_id"] == "test-unified-file-id" + assert "file_object" not in call_args.kwargs["data"] + + +@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) 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) 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) diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 36f0ab5097d..4651bf59b40 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -158,16 +158,16 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router: status="uploaded", ) - async def afile_retrieve(self, file_id, litellm_parent_otel_span): + async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") async def afile_list(self, purpose, litellm_parent_otel_span): raise NotImplementedError("Not implemented for test") - async def afile_delete(self, file_id, litellm_parent_otel_span): + async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data): raise NotImplementedError("Not implemented for test") - async def afile_content(self, file_id, litellm_parent_otel_span): + async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data): raise NotImplementedError("Not implemented for test") # Manually add the hook to the proxy_hook_mapping @@ -607,16 +607,16 @@ def test_create_file_with_expires_after(mocker: MockerFixture, monkeypatch, llm_ status="uploaded", ) - async def afile_retrieve(self, file_id, litellm_parent_otel_span): + async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") async def afile_list(self, purpose, litellm_parent_otel_span): raise NotImplementedError("Not implemented for test") - async def afile_delete(self, file_id, litellm_parent_otel_span): + async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data): raise NotImplementedError("Not implemented for test") - async def afile_content(self, file_id, litellm_parent_otel_span): + async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data): raise NotImplementedError("Not implemented for test") proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles() @@ -747,16 +747,16 @@ def test_create_file_with_expires_after_valid_values(mocker: MockerFixture, monk status="uploaded", ) - async def afile_retrieve(self, file_id, litellm_parent_otel_span): + async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") async def afile_list(self, purpose, litellm_parent_otel_span): raise NotImplementedError("Not implemented for test") - async def afile_delete(self, file_id, litellm_parent_otel_span): + async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data): raise NotImplementedError("Not implemented for test") - async def afile_content(self, file_id, litellm_parent_otel_span): + async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data): raise NotImplementedError("Not implemented for test") proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles() @@ -820,16 +820,16 @@ def test_create_file_without_expires_after(mocker: MockerFixture, monkeypatch, l status="uploaded", ) - async def afile_retrieve(self, file_id, litellm_parent_otel_span): + async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") async def afile_list(self, purpose, litellm_parent_otel_span): raise NotImplementedError("Not implemented for test") - async def afile_delete(self, file_id, litellm_parent_otel_span): + async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data): raise NotImplementedError("Not implemented for test") - async def afile_content(self, file_id, litellm_parent_otel_span): + async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data): raise NotImplementedError("Not implemented for test") proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()