diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 2330120adad..0fd249b6bbf 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3935,6 +3935,7 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, add_prefix: bool = True, raw_headers: dict[str, str] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> list[Prompt]: """ Helper method to get prompts from a single MCP server with prefixed names. @@ -3967,6 +3968,7 @@ class MCPServerManager: extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, + user_api_key_auth=user_api_key_auth, ) prompts: Final = await client.list_prompts() @@ -3986,6 +3988,7 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, add_prefix: bool = True, raw_headers: dict[str, str] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> list[Resource]: """Fetch available resources from a single MCP server.""" @@ -4009,6 +4012,7 @@ class MCPServerManager: extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, + user_api_key_auth=user_api_key_auth, ) resources: Final = await client.list_resources() @@ -4028,6 +4032,7 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, add_prefix: bool = True, raw_headers: dict[str, str] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> list[ResourceTemplate]: """Fetch available resource templates from a single MCP server.""" @@ -4051,6 +4056,7 @@ class MCPServerManager: extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, + user_api_key_auth=user_api_key_auth, ) resource_templates: Final = await client.list_resource_templates() @@ -4072,6 +4078,7 @@ class MCPServerManager: mcp_auth_header: str | dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> ReadResourceResult: """Read resource contents from a specific MCP server.""" @@ -4092,6 +4099,7 @@ class MCPServerManager: extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, + user_api_key_auth=user_api_key_auth, ) return await client.read_resource(url) @@ -4104,6 +4112,7 @@ class MCPServerManager: mcp_auth_header: str | dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, ) -> GetPromptResult: """Fetch a specific prompt definition from a single MCP server.""" @@ -4124,6 +4133,7 @@ class MCPServerManager: extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, + user_api_key_auth=user_api_key_auth, ) get_prompt_request_params: Final = GetPromptRequestParams( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index c6b2ac489bb..0b6e5c9df1b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2184,6 +2184,7 @@ if MCP_AVAILABLE: extra_headers=extra_headers, add_prefix=True, # Always add server prefix raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) all_prompts.extend(prompts) @@ -2237,6 +2238,7 @@ if MCP_AVAILABLE: extra_headers=extra_headers, add_prefix=True, # Always add server prefix raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) all_resources.extend(resources) @@ -2288,6 +2290,7 @@ if MCP_AVAILABLE: extra_headers=extra_headers, add_prefix=True, # Always add server prefix raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) all_resource_templates.extend(resource_templates) verbose_logger.debug( @@ -3207,6 +3210,7 @@ if MCP_AVAILABLE: mcp_auth_header=server_auth_header, extra_headers=extra_headers, raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) async def mcp_read_resource( @@ -3256,6 +3260,7 @@ if MCP_AVAILABLE: mcp_auth_header=server_auth_header, extra_headers=extra_headers, raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) def _get_standard_logging_mcp_tool_call( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 82f74cda835..bef3935ba18 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -919,6 +919,7 @@ async def test_mcp_get_prompt_success(): mcp_auth_header={"Authorization": "token"}, extra_headers={"X-Test": "1"}, raw_headers=None, + user_api_key_auth=user_api_key_auth, ) assert result is prompt_result @@ -980,10 +981,60 @@ async def test_mcp_read_resource_success(): mcp_auth_header={"Authorization": "token"}, extra_headers={"X-Test": "1"}, raw_headers=None, + user_api_key_auth=user_api_key_auth, ) assert result is read_result +@pytest.mark.asyncio +async def test_prompt_and_resource_manager_paths_pass_user_api_key_auth(): + """Regression for GitHub #37497: prompt/resource clients get the same auth context as tools.""" + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + from pydantic import AnyUrl + + manager = MCPServerManager() + server = MCPServer( + server_id="protected-mcp", + name="protected-mcp", + server_name="protected-mcp", + url="https://mcp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + sentinel = UserAPIKeyAuth(api_key="sk-test", user_id="user-1") + + mock_client = MagicMock() + mock_client.list_prompts = AsyncMock(return_value=[]) + mock_client.list_resources = AsyncMock(return_value=[]) + mock_client.list_resource_templates = AsyncMock(return_value=[]) + mock_client.get_prompt = AsyncMock(return_value=MagicMock()) + mock_client.read_resource = AsyncMock(return_value=MagicMock()) + + with patch.object( + manager, "_create_mcp_client", new=AsyncMock(return_value=mock_client) + ) as create_mock: + await manager.get_prompts_from_server(server=server, user_api_key_auth=sentinel) + await manager.get_resources_from_server(server=server, user_api_key_auth=sentinel) + await manager.get_resource_templates_from_server(server=server, user_api_key_auth=sentinel) + await manager.get_prompt_from_server(server=server, prompt_name="hello", user_api_key_auth=sentinel) + await manager.read_resource_from_server( + server=server, + url=AnyUrl("https://example.com/r"), + user_api_key_auth=sentinel, + ) + + assert create_mock.await_count == 5 + for call in create_mock.await_args_list: + assert call.kwargs["user_api_key_auth"] is sentinel + + def test_normalize_resource_contents_passes_metadata(): """Test that _normalize_resource_contents preserves meta from ResourceContents (MCP 1.26.0+).""" try: