mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
abf04d03c3
commit
b9c0a35636
5 changed files with 224 additions and 29 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)),
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue