fix(mcp): pass user auth into prompt and resource clients

Tools already forwarded per-user OAuth context into MCP client
construction. Prompt and resource paths omitted it, so protected
prompts/list and resources/list returned 401 after tools/list succeeded.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Huanyi Xie 2026-08-20 14:01:49 +03:00
parent 6d47468dae
commit e0473fa6d5
3 changed files with 66 additions and 0 deletions

View file

@ -3943,6 +3943,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.
@ -3975,6 +3976,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()
@ -3994,6 +3996,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."""
@ -4017,6 +4020,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()
@ -4036,6 +4040,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."""
@ -4059,6 +4064,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()
@ -4080,6 +4086,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."""
@ -4100,6 +4107,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)
@ -4112,6 +4120,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."""
@ -4132,6 +4141,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(

View file

@ -2181,6 +2181,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)
@ -2234,6 +2235,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)
@ -2285,6 +2287,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(
@ -3204,6 +3207,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(
@ -3253,6 +3257,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(

View file

@ -918,6 +918,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
@ -979,10 +980,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: