fix(mcp): recognize per-server auth header at the connect-time preemptive 401

The preemptive 401 for true_passthrough and oauth_delegate only inspected the
request-wide Authorization, so a caller who bound the upstream token via the
per-server x-mcp-{alias}-authorization header (the required shape in a
multi-server aggregate, where the request-wide Authorization is withheld) was
spuriously 401'd at initialize even though egress already honors that header.
The gate now recognizes the per-server header for both modes via a shared
helper, mode-correctly: true_passthrough treats any Authorization or the
per-server header as the upstream token, oauth_delegate keeps requiring a
distinct x-litellm-api-key so a lone Authorization consumed for admission is
never mistaken for an upstream token. The preemptive raise is also gated to
single-server scopes so a multi-server aggregate degrades gracefully instead
of one missing token 401-ing the whole connect.
This commit is contained in:
Tin 2026-07-08 15:44:36 -07:00
parent 4a25cce114
commit edf00bbe23
2 changed files with 139 additions and 20 deletions

View file

@ -1441,6 +1441,35 @@ if MCP_AVAILABLE:
return allowed_mcp_servers
def _client_has_per_server_auth_header(
server: MCPServer,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
) -> bool:
"""True if the request carries a per-server ``x-mcp-{alias}-authorization``
header for this server. This is the multi-server binding: it names one
upstream, so it is unambiguously the caller's upstream token regardless of
auth mode (never the LiteLLM admission credential).
"""
if not mcp_server_auth_headers:
return False
for key in (server.alias, server.server_name, server.name):
if not key:
continue
server_headers = None
for k, v in mcp_server_auth_headers.items():
if k.lower() == key.lower():
server_headers = v
break
if server_headers is None:
continue
if isinstance(server_headers, str) and server_headers.strip():
return True
if isinstance(server_headers, dict):
for hk in server_headers.keys():
if hk.lower() == "authorization":
return True
return False
def _client_has_passthrough_authorization(
server: MCPServer,
oauth2_headers: Optional[Dict[str, str]],
@ -1458,24 +1487,7 @@ if MCP_AVAILABLE:
for k in oauth2_headers.keys():
if k.lower() == "authorization":
return True
if mcp_server_auth_headers:
for key in (server.alias, server.server_name, server.name):
if not key:
continue
server_headers = None
for k, v in mcp_server_auth_headers.items():
if k.lower() == key.lower():
server_headers = v
break
if server_headers is None:
continue
if isinstance(server_headers, str) and server_headers.strip():
return True
if isinstance(server_headers, dict):
for hk in server_headers.keys():
if hk.lower() == "authorization":
return True
return False
return _client_has_per_server_auth_header(server, mcp_server_auth_headers)
async def _get_user_oauth_extra_headers_from_db(
server: MCPServer,
@ -3540,7 +3552,13 @@ if MCP_AVAILABLE:
headers={"www-authenticate": www_authenticate},
)
if server and server.is_oauth_delegate and _get_forwarded_auth_from_scope(scope) is None:
if (
server
and server.is_oauth_delegate
and len(mcp_servers or []) == 1
and _get_forwarded_auth_from_scope(scope) is None
and not _client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
www_authenticate = _get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
@ -3551,7 +3569,13 @@ if MCP_AVAILABLE:
headers={"www-authenticate": www_authenticate},
)
if server and server.is_true_passthrough and not _scope_has_authorization_header(scope):
if (
server
and server.is_true_passthrough
and len(mcp_servers or []) == 1
and not _scope_has_authorization_header(scope)
and not _client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "")
if upstream_status == 401 and upstream_www_authenticate:
raise HTTPException(

View file

@ -1216,6 +1216,101 @@ async def test_handle_streamable_http_mcp_oauth_delegate_with_forwarded_token_sk
assert mock_handle_request.await_count == 1
async def _run_passthrough_connect(
*,
auth_type,
server_names,
mcp_server_auth_headers,
scope_extra_headers=None,
):
"""Drive handle_streamable_http_mcp through the preemptive-401 gate and report whether it
challenged (raised) or forwarded to the session manager. Returns (challenged, www_authenticate)."""
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
session_manager_stateless,
)
scope = _passthrough_mode_scope(server_names[0], extra_headers=scope_extra_headers)
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 = "u1"
server = _build_passthrough_mode_server(server_names[0], auth_type)
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(user_auth, None, server_names, mcp_server_auth_headers, 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._check_passthrough_upstream_auth",
new_callable=AsyncMock,
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name",
return_value=server,
),
patch.object(session_manager_stateless, "handle_request", new_callable=AsyncMock) as mock_handle_request,
patch.object(session_manager_stateless, "_server_instances", {}),
):
try:
await handle_streamable_http_mcp(scope, receive, send)
except HTTPException as exc:
return True, (exc.headers or {}).get("www-authenticate")
return mock_handle_request.await_count == 0, None
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth_delegate, MCPAuth.true_passthrough])
async def test_handle_streamable_http_mcp_per_server_header_skips_preemptive_challenge(auth_type):
"""A per-server x-mcp-{alias}-authorization header binds the upstream token to one server; the
connect gate must recognize it and forward instead of spuriously 401-ing, since egress already
honors it. Without this, the mandatory multi-server binding is unusable at connect."""
try:
from litellm.proxy._experimental.mcp_server.server import handle_streamable_http_mcp # noqa: F401
except ImportError:
pytest.skip("MCP server not available")
challenged, _ = await _run_passthrough_connect(
auth_type=auth_type,
server_names=["pt_server"],
mcp_server_auth_headers={"pt_server": {"Authorization": "Bearer upstream-token"}},
)
assert challenged is False
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth_delegate, MCPAuth.true_passthrough])
async def test_handle_streamable_http_mcp_aggregate_does_not_preemptively_challenge(auth_type):
"""A multi-server aggregate must degrade gracefully: the preemptive 401 is single-server only, so
one server missing a token cannot 401 the whole connect (the listing absorbs per-server failures)."""
try:
from litellm.proxy._experimental.mcp_server.server import handle_streamable_http_mcp # noqa: F401
except ImportError:
pytest.skip("MCP server not available")
challenged, _ = await _run_passthrough_connect(
auth_type=auth_type,
server_names=["pt_server", "pt_server_2"],
mcp_server_auth_headers=None,
)
assert challenged is False
@pytest.mark.asyncio
async def test_handle_streamable_http_mcp_true_passthrough_without_token_surfaces_verbatim_upstream_challenge():
"""true_passthrough is a transparent proxy: with no client Authorization the