mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
4a25cce114
commit
edf00bbe23
2 changed files with 139 additions and 20 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue