diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 276e0cb9a04..69c7f1504d0 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -177,6 +177,18 @@ async def aresponses_api_with_mcp( "litellm_metadata", {} ).get("user_api_key_auth") + # Extract MCP auth headers from the request for tool list retrieval + secret_fields: Optional[Dict[str, Any]] = kwargs.get("secret_fields") + ( + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers_from_request, + ) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request( + secret_fields=secret_fields, + tools=tools, + ) + # Get original MCP tools (for events) and OpenAI tools (for LLM) by reusing existing methods ( original_mcp_tools, @@ -185,6 +197,10 @@ async def aresponses_api_with_mcp( user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, litellm_trace_id=kwargs.get("litellm_trace_id"), + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers_from_request, ) openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( original_mcp_tools @@ -288,18 +304,6 @@ async def aresponses_api_with_mcp( "user_api_key_auth" ) - # Extract MCP auth headers from the request to pass to MCP server - secret_fields: Optional[Dict[str, Any]] = kwargs.get("secret_fields") - ( - mcp_auth_header, - mcp_server_auth_headers, - oauth2_headers, - raw_headers_from_request, - ) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request( - secret_fields=secret_fields, - tools=tools, - ) - tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=tool_server_map, tool_calls=tool_calls, diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 377ce396457..ae1c3905a3e 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -120,7 +120,18 @@ async def acompletion_with_mcp( # noqa: PLR0915 (kwargs.get("metadata", {}) or {}).get("user_api_key_auth") ) - # Process MCP tools + # Extract MCP auth headers early so they can be used for both tool fetching and execution + ( + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request( + secret_fields=kwargs.get("secret_fields"), + tools=tools, + ) + + # Process MCP tools (with auth headers for servers that require authentication) ( deduplicated_mcp_tools, tool_server_map, @@ -128,6 +139,10 @@ async def acompletion_with_mcp( # noqa: PLR0915 user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, litellm_trace_id=kwargs.get("litellm_trace_id"), + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, ) openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( @@ -143,17 +158,6 @@ async def acompletion_with_mcp( # noqa: PLR0915 mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy ) - # Extract MCP auth headers - ( - mcp_auth_header, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - ) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request( - secret_fields=kwargs.get("secret_fields"), - tools=tools, - ) - # Prepare call parameters # Remove keys that shouldn't be passed to acompletion clean_kwargs = {k: v for k, v in kwargs.items() if k not in ["acompletion"]} diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 805a1958552..6d677677047 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -99,6 +99,10 @@ class LiteLLM_Proxy_MCP_Handler: user_api_key_auth: Any, mcp_tools_with_litellm_proxy: Optional[Iterable[ToolParam]], litellm_trace_id: Optional[str] = None, + mcp_auth_header: Optional[str] = None, + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, + oauth2_headers: Optional[Dict[str, str]] = None, + raw_headers: Optional[Dict[str, str]] = None, ) -> tuple[List[MCPTool], List[str]]: """ Get available tools from the MCP server manager. @@ -106,6 +110,10 @@ class LiteLLM_Proxy_MCP_Handler: Args: user_api_key_auth: User authentication info for access control mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy" + mcp_auth_header: Legacy auth header for MCP server (from x-mcp-auth) + mcp_server_auth_headers: Server-specific auth headers (from x-mcp--* headers) + oauth2_headers: OAuth2 headers for MCP server + raw_headers: Raw HTTP headers from the original request Returns: List of MCP tools @@ -133,9 +141,11 @@ class LiteLLM_Proxy_MCP_Handler: tools = await _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, - mcp_auth_header=None, + mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, - mcp_server_auth_headers=None, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, log_list_tools_to_spendlogs=True, list_tools_log_source="responses", litellm_trace_id=litellm_trace_id, @@ -245,6 +255,10 @@ class LiteLLM_Proxy_MCP_Handler: user_api_key_auth: Any, mcp_tools_with_litellm_proxy: List[ToolParam], litellm_trace_id: Optional[str] = None, + mcp_auth_header: Optional[str] = None, + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, + oauth2_headers: Optional[Dict[str, str]] = None, + raw_headers: Optional[Dict[str, str]] = None, ) -> tuple[List[Any], dict[str, str]]: """ Centralized method to process MCP tools through the complete pipeline. @@ -253,6 +267,10 @@ class LiteLLM_Proxy_MCP_Handler: user_api_key_auth: User authentication info for access control mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy" litellm_trace_id: Optional trace ID for linking list_mcp_tools spend logs to parent request + mcp_auth_header: Legacy auth header for MCP server (from x-mcp-auth) + mcp_server_auth_headers: Server-specific auth headers (from x-mcp--* headers) + oauth2_headers: OAuth2 headers for MCP server + raw_headers: Raw HTTP headers from the original request Returns: List of tools in OpenAI format ready to be sent to the LLM @@ -265,6 +283,10 @@ class LiteLLM_Proxy_MCP_Handler: user_api_key_auth, mcp_tools_with_litellm_proxy, litellm_trace_id=litellm_trace_id, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, ) openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( @@ -278,6 +300,10 @@ class LiteLLM_Proxy_MCP_Handler: user_api_key_auth: Any, mcp_tools_with_litellm_proxy: List[ToolParam], litellm_trace_id: Optional[str] = None, + mcp_auth_header: Optional[str] = None, + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, + oauth2_headers: Optional[Dict[str, str]] = None, + raw_headers: Optional[Dict[str, str]] = None, ) -> tuple[List[Any], dict[str, str]]: """ Process MCP tools through filtering and deduplication pipeline without OpenAI transformation. @@ -286,6 +312,10 @@ class LiteLLM_Proxy_MCP_Handler: Args: user_api_key_auth: User authentication info for access control mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy" + mcp_auth_header: Legacy auth header for MCP server (from x-mcp-auth) + mcp_server_auth_headers: Server-specific auth headers (from x-mcp--* headers) + oauth2_headers: OAuth2 headers for MCP server + raw_headers: Raw HTTP headers from the original request Returns: List of filtered and deduplicated MCP tools in their original format @@ -301,6 +331,10 @@ class LiteLLM_Proxy_MCP_Handler: user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, litellm_trace_id=litellm_trace_id, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, ) # Step 2: Filter tools based on allowed_tools parameter diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index 15fdc7bd0c4..ea6e1561d5e 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -402,3 +402,87 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch assert mock_get_tools.await_args is not None assert mock_get_tools.await_args.kwargs["log_list_tools_to_spendlogs"] is True assert mock_get_tools.await_args.kwargs["list_tools_log_source"] == "responses" + + +@pytest.mark.asyncio +async def test_get_mcp_tools_from_manager_forwards_auth_headers(monkeypatch): + """ + Verify that _get_mcp_tools_from_manager forwards x-mcp-auth and + x-mcp--* auth headers to _get_tools_from_mcp_servers. + + Previously these were hardcoded to None, meaning MCP servers that + require authentication would fail during tool list retrieval even + though the caller supplied the correct headers. + """ + mock_get_tools = AsyncMock(return_value=[]) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._get_tools_from_mcp_servers", + mock_get_tools, + ) + + fake_manager = types.SimpleNamespace( + get_allowed_mcp_servers=AsyncMock(return_value=[]), + get_mcp_servers_from_ids=MagicMock(return_value=[]), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + fake_manager, + ) + + user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") + + test_mcp_auth_header = "Bearer my-mcp-token" + test_server_auth_headers = {"deepwiki": {"Authorization": "Bearer deepwiki-token"}} + test_oauth2_headers = {"Authorization": "Bearer oauth2-token"} + test_raw_headers = {"x-mcp-auth": "Bearer my-mcp-token"} + + await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + user_api_key_auth=user_auth, + mcp_tools_with_litellm_proxy=[ + {"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"} + ], + mcp_auth_header=test_mcp_auth_header, + mcp_server_auth_headers=test_server_auth_headers, + oauth2_headers=test_oauth2_headers, + raw_headers=test_raw_headers, + ) + + assert mock_get_tools.await_count == 1 + call_kwargs = mock_get_tools.await_args.kwargs + assert call_kwargs["mcp_auth_header"] == test_mcp_auth_header + assert call_kwargs["mcp_server_auth_headers"] == test_server_auth_headers + assert call_kwargs["oauth2_headers"] == test_oauth2_headers + assert call_kwargs["raw_headers"] == test_raw_headers + + +@pytest.mark.asyncio +async def test_process_mcp_tools_without_openai_transform_forwards_auth_headers( + monkeypatch, +): + """ + Verify that _process_mcp_tools_without_openai_transform passes auth + headers through to _get_mcp_tools_from_manager. + """ + mock_get_manager = AsyncMock(return_value=([], [])) + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_get_mcp_tools_from_manager", + mock_get_manager, + ) + + test_mcp_auth = "Bearer token-abc" + test_server_headers = {"myserver": {"x-api-key": "key123"}} + + await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + user_api_key_auth=None, + mcp_tools_with_litellm_proxy=[ + {"type": "mcp", "server_url": "litellm_proxy/mcp/myserver"} + ], + mcp_auth_header=test_mcp_auth, + mcp_server_auth_headers=test_server_headers, + ) + + assert mock_get_manager.await_count == 1 + call_kwargs = mock_get_manager.await_args.kwargs + assert call_kwargs["mcp_auth_header"] == test_mcp_auth + assert call_kwargs["mcp_server_auth_headers"] == test_server_headers