From 8b3213ce5c07a502be356480eaadfcd2cd6c3365 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 4 Feb 2026 13:12:45 +0530 Subject: [PATCH] Add mapping for responses tools in file ids --- .../proxy/hooks/managed_files.py | 89 +++- .../prompt_templates/common_utils.py | 115 +++- litellm/responses/main.py | 21 +- .../proxy/hooks/test_managed_files.py | 493 ++++++++++++++++++ 4 files changed, 695 insertions(+), 23 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index fe476c0ba2..569ea17f6d 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -358,6 +358,31 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) return False + async def check_file_ids_access( + self, file_ids: List[str], user_api_key_dict: UserAPIKeyAuth + ) -> None: + """ + Check if the user has access to a list of file IDs. + Only checks managed (unified) file IDs. + + Args: + file_ids: List of file IDs to check access for + user_api_key_dict: User API key authentication details + + Raises: + HTTPException: If user doesn't have access to any of the files + """ + for file_id in file_ids: + is_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + if is_unified_file_id: + if not await self.can_user_call_unified_file_id( + file_id, user_api_key_dict + ): + raise HTTPException( + status_code=403, + detail=f"User {user_api_key_dict.user_id} does not have access to the file {file_id}", + ) + async def async_pre_call_hook( # noqa: PLR0915 self, user_api_key_dict: UserAPIKeyAuth, @@ -391,6 +416,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if messages: file_ids = self.get_file_ids_from_messages(messages) if file_ids: + # Check user has access to all managed files + await self.check_file_ids_access(file_ids, user_api_key_dict) + # Check if any files are stored in storage backends and need base64 conversion # This is needed for Vertex AI/Gemini which requires base64 content is_vertex_ai = model and ("vertex_ai" in model or "gemini" in model.lower()) @@ -406,15 +434,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) data["model_file_id_mapping"] = model_file_id_mapping elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value: - # Handle managed files in responses API input + # Handle managed files in responses API input and tools + file_ids = [] + + # Extract file IDs from input parameter input_data = data.get("input") if input_data: - file_ids = self.get_file_ids_from_responses_input(input_data) - if file_ids: - model_file_id_mapping = await self.get_model_file_id_mapping( - file_ids, user_api_key_dict.parent_otel_span - ) - data["model_file_id_mapping"] = model_file_id_mapping + file_ids.extend(self.get_file_ids_from_responses_input(input_data)) + + # Extract file IDs from tools parameter (e.g., code_interpreter container) + tools = data.get("tools") + if tools: + file_ids.extend(self.get_file_ids_from_responses_tools(tools)) + + if file_ids: + # Check user has access to all managed files + await self.check_file_ids_access(file_ids, user_api_key_dict) + + model_file_id_mapping = await self.get_model_file_id_mapping( + file_ids, user_api_key_dict.parent_otel_span + ) + data["model_file_id_mapping"] = model_file_id_mapping elif call_type == CallTypes.afile_content.value: retrieve_file_id = cast(Optional[str], data.get("file_id")) potential_file_id = ( @@ -616,6 +656,41 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): return file_ids + def get_file_ids_from_responses_tools( + self, tools: List[Dict[str, Any]] + ) -> List[str]: + """ + Gets file ids from responses API tools parameter. + + The tools can contain code_interpreter with container.file_ids: + [ + { + "type": "code_interpreter", + "container": {"type": "auto", "file_ids": ["file-123", "file-456"]} + } + ] + """ + file_ids: List[str] = [] + + if not isinstance(tools, list): + return file_ids + + for tool in tools: + if not isinstance(tool, dict): + continue + + # Check for code_interpreter with container file_ids + if tool.get("type") == "code_interpreter": + container = tool.get("container") + if isinstance(container, dict): + container_file_ids = container.get("file_ids") + if isinstance(container_file_ids, list): + for file_id in container_file_ids: + if isinstance(file_id, str): + file_ids.append(file_id) + + return file_ids + async def get_model_file_id_mapping( self, file_ids: List[str], litellm_parent_otel_span: Span ) -> dict: diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 7790fb8336..e588ca1383 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -443,13 +443,21 @@ def update_messages_with_model_file_ids( def update_responses_input_with_model_file_ids( input: Any, + model_id: Optional[str] = None, + model_file_id_mapping: Optional[Dict[str, Dict[str, str]]] = None, ) -> Union[str, List[Dict[str, Any]]]: """ Updates responses API input with provider-specific file IDs. File IDs are always inside the content array, not as direct input_file items. - For managed files (unified file IDs), decodes the base64-encoded unified file ID - and extracts the llm_output_file_id directly. + For managed files (unified file IDs), uses model_file_id_mapping if provided, + otherwise decodes the base64-encoded unified file ID and extracts the llm_output_file_id directly. + + Args: + input: The responses API input parameter + model_id: The model ID to use for looking up provider-specific file IDs + model_file_id_mapping: Dictionary mapping litellm file IDs to provider file IDs + Format: {"litellm_file_id": {"model_id": "provider_file_id"}} """ from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, @@ -479,22 +487,35 @@ def update_responses_input_with_model_file_ids( ): file_id = content_item.get("file_id") if file_id: - # Check if this is a managed file ID (base64-encoded unified file ID) - is_unified_file_id = _is_base64_encoded_unified_file_id(file_id) - if is_unified_file_id: - unified_file_id = convert_b64_uid_to_unified_uid(file_id) - if "llm_output_file_id," in unified_file_id: - provider_file_id = unified_file_id.split( - "llm_output_file_id," - )[1].split(";")[0] - else: - # Fallback: keep original if we can't extract - provider_file_id = file_id + provider_file_id = file_id # Default to original + + # Check if we have a mapping for this file ID + if model_file_id_mapping and model_id and file_id in model_file_id_mapping: + # Use the model-specific file ID from mapping + provider_file_id = ( + model_file_id_mapping.get(file_id, {}).get(model_id) + or file_id + ) updated_content_item = content_item.copy() updated_content_item["file_id"] = provider_file_id updated_content.append(updated_content_item) else: - updated_content.append(content_item) + # Check if this is a base64-encoded unified file ID without mapping + is_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + if is_unified_file_id: + # Fallback: decode unified file ID + unified_file_id = convert_b64_uid_to_unified_uid(file_id) + if "llm_output_file_id," in unified_file_id: + provider_file_id = unified_file_id.split( + "llm_output_file_id," + )[1].split(";")[0] + + updated_content_item = content_item.copy() + updated_content_item["file_id"] = provider_file_id + updated_content.append(updated_content_item) + else: + # Not a managed file, keep as-is + updated_content.append(content_item) else: updated_content.append(content_item) else: @@ -506,6 +527,72 @@ def update_responses_input_with_model_file_ids( return updated_input +def update_responses_tools_with_model_file_ids( + tools: Optional[List[Dict[str, Any]]], + model_id: Optional[str] = None, + model_file_id_mapping: Optional[Dict[str, Dict[str, str]]] = None, +) -> Optional[List[Dict[str, Any]]]: + """ + Updates responses API tools with provider-specific file IDs. + + Handles code_interpreter tools with container.file_ids. + + Args: + tools: The responses API tools parameter + model_id: The model ID to use for looking up provider-specific file IDs + model_file_id_mapping: Dictionary mapping litellm file IDs to provider file IDs + Format: {"litellm_file_id": {"model_id": "provider_file_id"}} + """ + from litellm.proxy.openai_files_endpoints.common_utils import ( + _is_base64_encoded_unified_file_id, + ) + + if not tools or not isinstance(tools, list): + return tools + + if not model_file_id_mapping or not model_id: + return tools + + updated_tools = [] + for tool in tools: + if not isinstance(tool, dict): + updated_tools.append(tool) + continue + + updated_tool = tool.copy() + + # Handle code_interpreter with container file_ids + if tool.get("type") == "code_interpreter": + container = tool.get("container") + if isinstance(container, dict): + container_file_ids = container.get("file_ids") + if isinstance(container_file_ids, list): + updated_file_ids = [] + for file_id in container_file_ids: + if isinstance(file_id, str): + # Check if we have a mapping for this file ID + if file_id in model_file_id_mapping: + # Map to provider-specific file ID + provider_file_id = ( + model_file_id_mapping.get(file_id, {}).get(model_id) + or file_id + ) + updated_file_ids.append(provider_file_id) + else: + updated_file_ids.append(file_id) + else: + updated_file_ids.append(file_id) + + # Update the tool with new file IDs + updated_container = container.copy() + updated_container["file_ids"] = updated_file_ids + updated_tool["container"] = updated_container + + updated_tools.append(updated_tool) + + return updated_tools + + def extract_file_data(file_data: FileTypes) -> ExtractedFileData: """ Extracts and processes file data from various input formats. diff --git a/litellm/responses/main.py b/litellm/responses/main.py index b2c2493c81..efe1607331 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -24,6 +24,7 @@ from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.prompt_templates.common_utils import ( update_responses_input_with_model_file_ids, + update_responses_tools_with_model_file_ids, ) from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler @@ -595,13 +596,29 @@ def responses( litellm_params.api_base = dynamic_api_base ######################################################### - # Update input with provider-specific file IDs if managed files are used + # Update input and tools with provider-specific file IDs if managed files are used ######################################################### + model_file_id_mapping = kwargs.get("model_file_id_mapping") + model_info_id = kwargs.get("model_info", {}).get("id") if isinstance(kwargs.get("model_info"), dict) else None + input = cast( Union[str, ResponseInputParam], - update_responses_input_with_model_file_ids(input=input), + update_responses_input_with_model_file_ids( + input=input, + model_id=model_info_id, + model_file_id_mapping=model_file_id_mapping, + ), ) local_vars["input"] = input + + # Update tools with provider-specific file IDs if needed + if tools: + tools = update_responses_tools_with_model_file_ids( + tools=tools, + model_id=model_info_id, + model_file_id_mapping=model_file_id_mapping, + ) + local_vars["tools"] = tools ######################################################### # Native MCP Responses API 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 3fd19cfa18..946c5ad172 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -587,6 +587,499 @@ def test_update_responses_input_with_multiple_file_ids(): 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-4", + } + + # 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-4", + } + + # 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(): """