fix(mcp): surface upstream challenges for delegated OAuth (#30124)

* fix(mcp): surface upstream challenges for delegated OAuth

* docs(mcp): clarify delegated upstream auth comments
This commit is contained in:
King Star 2026-06-11 18:54:34 +08:00 • committed by GitHub
parent abf04d03c3
commit b9c0a35636
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 224 additions and 29 deletions

View file

@ -10,8 +10,9 @@ class MCPUpstreamAuthError(Exception):
(typically HTTP 401) and the gateway should surface it transparently to
the client instead of swallowing it.
Only relevant for pass-through MCP servers (see
``MCPServer.is_oauth_passthrough``). The gateway converts this exception
Relevant for MCP servers that delegate OAuth to the upstream server,
including pass-through servers and OAuth2 servers with
``delegate_auth_to_upstream`` enabled. The gateway converts this exception
into an HTTP 401 response on single-server routes, preserving any
``WWW-Authenticate`` challenge emitted by the upstream so standards-
compliant MCP clients can trigger the upstream OAuth flow.

View file

@ -2777,28 +2777,40 @@ class MCPServerManager:
Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts
with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details.
For pass-through MCP servers (``MCPServer.is_oauth_passthrough``) an
For OAuth pass-through and upstream-delegated OAuth2 MCP servers, an
upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError`
instead of being swallowed to an empty tool list. That lets the
single-server HTTP routes surface a proper 401 + ``WWW-Authenticate``
challenge so standards-compliant MCP clients trigger the upstream
OAuth flow. Non-pass-through servers keep today's swallow-and-log
behaviour so the multi-server ``/mcp`` aggregator doesn't get
tainted by a single bad server.
OAuth flow. Other servers keep today's swallow-and-log behaviour so
the multi-server ``/mcp`` aggregator doesn't get tainted by a single
bad server.
Args:
client: MCP client instance
server_name: Name of the server for logging
server: Optional MCPServer; when pass-through, auth errors are
re-raised as :class:`MCPUpstreamAuthError`.
server: Optional MCPServer; when upstream auth is delegated, auth
errors are re-raised as :class:`MCPUpstreamAuthError`.
Returns:
List of tools from the server
"""
is_passthrough = bool(server is not None and server.is_oauth_passthrough)
should_surface_upstream_auth = bool(
server is not None
and (
server.is_oauth_passthrough
or (
server.auth_type == MCPAuth.oauth2
and getattr(server, "delegate_auth_to_upstream", False) is True
and not server.has_client_credentials
)
)
)
try:
with anyio.fail_after(MCP_TOOL_LISTING_TIMEOUT):
tools = await client.list_tools(raise_on_error=is_passthrough)
tools = await client.list_tools(
raise_on_error=should_surface_upstream_auth
)
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
except TimeoutError:
@ -2815,12 +2827,12 @@ class MCPServerManager:
)
return []
except Exception as e:
if is_passthrough:
if should_surface_upstream_auth:
auth_info = _extract_upstream_auth_failure(e)
if auth_info is not None:
status_code, www_authenticate = auth_info
verbose_logger.info(
f"Upstream auth failure from pass-through MCP server "
f"Upstream auth failure from MCP server "
f"{server_name}: HTTP {status_code}"
)
raise MCPUpstreamAuthError(

View file

@ -3427,6 +3427,8 @@ if MCP_AVAILABLE:
)
if stored_oauth_headers:
continue
if getattr(server, "delegate_auth_to_upstream", False) is True:
continue
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
@ -3961,7 +3963,7 @@ if MCP_AVAILABLE:
):
_stateful_session_locks.pop(active_request_session_id, None)
except MCPUpstreamAuthError as e:
# Pass-through server returned 401 — surface it to the client so
# Upstream delegated auth returned 401; surface it to the client so
# standards-compliant MCP clients trigger the upstream OAuth flow.
raise e.to_http_exception(
base_url=get_request_base_url(StarletteRequest(scope)),
@ -4077,7 +4079,7 @@ if MCP_AVAILABLE:
):
await sse_session_manager.handle_request(scope, receive, send)
except MCPUpstreamAuthError as e:
# Pass-through server returned 401 — surface it to the client so
# Upstream delegated auth returned 401; surface it to the client so
# standards-compliant MCP clients trigger the upstream OAuth flow.
raise e.to_http_exception(
base_url=get_request_base_url(StarletteRequest(scope)),

View file

@ -88,6 +88,76 @@ async def test_fetch_tools_from_passthrough_raises_on_upstream_401():
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
@pytest.mark.asyncio
async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401():
manager = MCPServerManager()
delegated_server = MCPServer(
server_id="oauth1",
name="delegated_docs",
url="https://upstream/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
)
response = httpx.Response(
status_code=401,
headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'},
request=httpx.Request("GET", "https://upstream/mcp"),
)
upstream_error = httpx.HTTPStatusError(
"401", request=response.request, response=response
)
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(
mock_client, delegated_server.name, server=delegated_server
)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == (
'Bearer resource_metadata="https://upstream"'
)
assert exc_info.value.server_name == "delegated_docs"
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
@pytest.mark.asyncio
async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior():
manager = MCPServerManager()
m2m_server = MCPServer(
server_id="oauth-m2m",
name="m2m_docs",
url="https://upstream/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=True,
oauth2_flow="client_credentials",
)
response = httpx.Response(
status_code=401,
headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'},
request=httpx.Request("GET", "https://upstream/mcp"),
)
upstream_error = httpx.HTTPStatusError(
"401", request=response.request, response=response
)
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
tools = await manager._fetch_tools_with_timeout(
mock_client, m2m_server.name, server=m2m_server
)
assert tools == []
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
@pytest.mark.asyncio
async def test_fetch_tools_from_passthrough_returns_tools_on_success():
manager = MCPServerManager()

View file

@ -673,6 +673,110 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
assert "Bearer authorization_uri=" in exc_info.value.headers["www-authenticate"]
@pytest.mark.asyncio
async def test_handle_streamable_http_mcp_delegated_server_surfaces_upstream_challenge():
"""
OAuth2 server with ``delegate_auth_to_upstream=True`` should let the
upstream MCP server's RFC 9728 challenge reach the client instead of
pre-emptively returning LiteLLM's gateway authorization_uri challenge.
"""
from fastapi import HTTPException
try:
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
session_manager_stateful,
)
except ImportError:
pytest.skip("MCP server not available")
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/delegated_oauth_server",
"scheme": "https",
"query_string": b"",
"root_path": "",
"server": ("litellm.example.com", 443),
"headers": [
(b"content-type", b"application/json"),
(b"host", b"litellm.example.com"),
],
}
receive = AsyncMock(
return_value={
"type": "http.request",
"body": b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}',
"more_body": False,
}
)
send = AsyncMock()
user_auth = MagicMock()
user_auth.user_id = None
delegated_server = MagicMock()
delegated_server.auth_type = MCPAuth.oauth2
delegated_server.delegate_auth_to_upstream = True
delegated_server.needs_user_oauth_token = True
delegated_server.server_id = "delegated-oauth-server"
upstream_challenge = (
'Bearer resource_metadata="https://upstream.example.com/.well-known/oauth-protected-resource"'
)
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(
user_auth,
None,
["delegated_oauth_server"],
None,
None,
None,
),
),
patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch(
"litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session",
new_callable=AsyncMock,
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
new_callable=AsyncMock,
return_value=None,
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
return_value=delegated_server,
),
patch.object(
session_manager_stateful,
"handle_request",
new_callable=AsyncMock,
side_effect=MCPUpstreamAuthError(
status_code=401,
www_authenticate=upstream_challenge,
server_name="delegated_oauth_server",
),
) as mock_handle_request,
):
with pytest.raises(HTTPException) as exc_info:
await handle_streamable_http_mcp(scope, receive, send)
assert mock_handle_request.await_count == 1
assert exc_info.value.status_code == 401
assert exc_info.value.headers == {"www-authenticate": upstream_challenge}
@pytest.mark.asyncio
async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
"""
@ -759,19 +863,16 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
@pytest.mark.asyncio
async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without_token():
async def test_handle_streamable_http_mcp_delegated_server_without_token_reaches_session_manager():
"""
OAuth2 server with ``delegate_auth_to_upstream=True`` and no Authorization
header must still emit a pre-emptive 401 with WWW-Authenticate so the
client kicks off PKCE. The 401 points at LiteLLM's discovery shim, which
in turn delegates to the upstream OAuth issuer.
OAuth2 server with ``delegate_auth_to_upstream=True`` and no stored token
should not receive LiteLLM's gateway authorization_uri challenge. The
request continues so the upstream MCP server can emit its RFC 9728 challenge.
"""
from fastapi import HTTPException
try:
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
session_manager,
session_manager_stateless,
)
except ImportError:
pytest.skip("MCP server not available")
@ -785,7 +886,13 @@ async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without
(b"host", b"litellm.example.com"),
],
}
receive = AsyncMock()
receive = AsyncMock(
return_value={
"type": "http.request",
"body": b'{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}',
"more_body": False,
}
)
send = AsyncMock()
user_auth = MagicMock()
user_auth.user_id = None
@ -819,19 +926,22 @@ async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without
new_callable=AsyncMock,
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db",
new_callable=AsyncMock,
return_value=None,
) as mock_get_stored_token,
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
return_value=delegated_server,
),
patch.object(
session_manager,
session_manager_stateless,
"handle_request",
new_callable=AsyncMock,
) as mock_handle_request,
):
with pytest.raises(HTTPException) as exc_info:
await handle_streamable_http_mcp(scope, receive, send)
await handle_streamable_http_mcp(scope, receive, send)
assert exc_info.value.status_code == 401
assert "www-authenticate" in exc_info.value.headers
assert mock_handle_request.await_count == 0
assert mock_get_stored_token.await_count == 1
assert mock_handle_request.await_count == 1