fix(mcp): trim static header whitespace and surface real connection errors

A leading or trailing space in a static header value made httpx/h11 reject it
as an illegal header value, which aborted the MCP transport task group,
cancelled session.initialize(), and surfaced only as an opaque "Failed to
connect to MCP server". Strip surrounding whitespace from header names and
values in _get_auth_headers so every path (test connection, runtime, config,
DB) is protected, and mirror the trim in the UI static-header reducers so
stored values stay clean.

When the transport task group does fail, unwrap the ExceptionGroup raised on
transport context exit and re-raise the real cause (httpx ConnectError,
LocalProtocolError, ...) instead of the cancellation, so the proxy logs and
the admin test endpoint explain why the connection failed. The test endpoint
maps known causes to an actionable, secret-safe message and stops swallowing
genuine cancellation
This commit is contained in:
mateo-berri 2026-06-05 14:04:20 -07:00
parent 02344057a2
commit 587829c3a0
6 changed files with 266 additions and 9 deletions

View file

@ -60,6 +60,27 @@ def to_basic_auth(auth_value: str) -> str:
return base64.b64encode(auth_value.encode("utf-8")).decode()
def _strip_header_whitespace(headers: Dict[str, str]) -> Dict[str, str]:
return {
(key.strip() if isinstance(key, str) else key): (
value.strip() if isinstance(value, str) else value
)
for key, value in headers.items()
}
def _first_non_cancelled_cause(exc: BaseException) -> Optional[BaseException]:
queue: List[BaseException] = [exc]
while queue:
current = queue.pop(0)
nested = getattr(current, "exceptions", None)
if nested:
queue.extend(nested)
elif not isinstance(current, asyncio.CancelledError):
return current
return None
TSessionResult = TypeVar("TSessionResult")
@ -335,6 +356,7 @@ class MCPClient:
user input (elicitation), or send log messages.
"""
transport = await transport_ctx.__aenter__()
in_flight_error: Optional[BaseException] = None
try:
read_stream, write_stream = transport[0], transport[1]
# Build session kwargs with optional callbacks
@ -360,11 +382,21 @@ class MCPClient:
await session_ctx.__aexit__(None, None, None)
except BaseException as e:
verbose_logger.debug(f"Error during session context exit: {e}")
except BaseException as e:
in_flight_error = e
raise
finally:
try:
await transport_ctx.__aexit__(None, None, None)
except BaseException as e:
verbose_logger.debug(f"Error during transport context exit: {e}")
except BaseException as exit_error:
verbose_logger.debug(
f"Error during transport context exit: {exit_error}"
)
root_cause = _first_non_cancelled_cause(exit_error)
if root_cause is not None and isinstance(
in_flight_error, asyncio.CancelledError
):
raise root_cause from in_flight_error
async def run_with_session(
self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]
@ -426,7 +458,7 @@ class MCPClient:
# update the headers with the extra headers
if self.extra_headers:
headers.update(self.extra_headers)
return headers
return _strip_header_whitespace(headers)
def _create_httpx_client_factory(self) -> Callable[..., httpx.AsyncClient]:
"""

View file

@ -1,3 +1,4 @@
import asyncio
import importlib
from datetime import datetime
from typing import (
@ -13,6 +14,7 @@ from typing import (
Union,
)
import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from litellm._logging import verbose_logger
@ -44,6 +46,28 @@ router = APIRouter(
tags=["mcp"],
)
def _connection_error_message(exc: BaseException) -> str:
if isinstance(exc, httpx.LocalProtocolError):
return (
"Failed to connect to MCP server: a request header is malformed. "
"Check static headers for leading/trailing spaces or illegal characters."
)
if isinstance(exc, (httpx.ConnectError, httpx.ConnectTimeout)):
return (
"Failed to connect to MCP server: the server is unreachable. "
"Check the URL and that the server is running."
)
if isinstance(exc, httpx.TimeoutException):
return "Failed to connect to MCP server: the connection timed out."
if isinstance(exc, httpx.HTTPStatusError):
return (
f"Failed to connect to MCP server: it returned HTTP "
f"{exc.response.status_code}."
)
return "Failed to connect to MCP server. Check proxy logs for details."
if MCP_AVAILABLE:
from mcp.types import Tool as MCPTool
@ -981,14 +1005,14 @@ if MCP_AVAILABLE:
return await operation(client)
except (KeyboardInterrupt, SystemExit):
except (KeyboardInterrupt, SystemExit, asyncio.CancelledError):
raise
except BaseException as e:
verbose_logger.error("Error in MCP operation: %s", e, exc_info=True)
return {
"status": "error",
"error": True,
"message": "Failed to connect to MCP server. Check proxy logs for details.",
"message": _connection_error_message(e),
}
async def _preview_openapi_tools(spec_path: str) -> dict:

View file

@ -1,3 +1,4 @@
import asyncio
import os
import ssl
import sys
@ -10,10 +11,26 @@ import pytest
sys.path.insert(0, "../../../")
import litellm.experimental_mcp_client.client as mcp_client_module
from litellm.experimental_mcp_client.client import MCPClient
from litellm.experimental_mcp_client.client import (
MCPClient,
_first_non_cancelled_cause,
)
from litellm.types.mcp import MCPAuth, MCPStdioConfig, MCPTransport
class _FakeExceptionGroup(Exception):
"""Duck-typed stand-in for an anyio/builtin ExceptionGroup.
The production unwrapper reads ``.exceptions`` rather than depending on the
builtin ``ExceptionGroup`` type, so this exercises the same code path on
every Python version.
"""
def __init__(self, message, exceptions):
super().__init__(message)
self.exceptions = tuple(exceptions)
class TestMCPClient:
"""Test MCP Client stdio functionality"""
@ -307,6 +324,26 @@ class TestMCPClient:
assert headers["Authorization"] == "token my-token"
assert headers["X-Custom-Header"] == "custom-value"
def test_get_auth_headers_strips_static_header_whitespace(self):
"""
Static header names/values must be stripped of surrounding whitespace.
h11 rejects header values with leading/trailing whitespace as an
"Illegal header value", which silently aborts the MCP connection. A
stray space in a configured static header value would otherwise make
every request to that server fail with an opaque error.
"""
client = MCPClient(
server_url="http://example.com/mcp",
transport_type="http",
extra_headers={"X-Db-Url": " mew://host ", " X-Pad ": "v"},
)
headers = client._get_auth_headers()
assert headers["X-Db-Url"] == "mew://host"
assert headers["X-Pad"] == "v"
def test_token_auth_enum_value(self):
"""Test that MCPAuth.token enum exists and has correct value"""
assert hasattr(MCPAuth, "token")
@ -388,5 +425,123 @@ class TestMCPClientInstructionsCapture:
assert client._last_initialize_instructions is None
# ---------------------------------------------------------------------------
# Transport error surfacing
# ---------------------------------------------------------------------------
class TestFirstNonCancelledCause:
"""Unwrapping the real cause out of a (possibly nested) exception group."""
def test_returns_plain_non_cancelled(self):
err = ValueError("boom")
assert _first_non_cancelled_cause(err) is err
def test_returns_none_for_plain_cancelled(self):
assert _first_non_cancelled_cause(asyncio.CancelledError()) is None
def test_unwraps_group_to_non_cancelled_leaf(self):
target = httpx.ConnectError("refused")
group = _FakeExceptionGroup("g", [asyncio.CancelledError(), target])
assert _first_non_cancelled_cause(group) is target
def test_unwraps_nested_group(self):
target = httpx.LocalProtocolError("Illegal header value")
inner = _FakeExceptionGroup("inner", [asyncio.CancelledError(), target])
outer = _FakeExceptionGroup("outer", [asyncio.CancelledError(), inner])
assert _first_non_cancelled_cause(outer) is target
def test_all_cancelled_returns_none(self):
group = _FakeExceptionGroup(
"g", [asyncio.CancelledError(), asyncio.CancelledError()]
)
assert _first_non_cancelled_cause(group) is None
@pytest.mark.skipif(
sys.version_info < (3, 11), reason="builtin ExceptionGroup requires 3.11+"
)
def test_unwraps_builtin_exception_group(self):
target = httpx.ConnectError("refused")
group = ExceptionGroup("transport failed", [target]) # noqa: F821
assert _first_non_cancelled_cause(group) is target
class TestExecuteSessionOperationSurfacesTransportError:
"""_execute_session_operation should surface the real transport failure.
When the upstream transport's task group fails (illegal header, connection
refused, ...), the in-flight ``session.initialize()`` is cancelled and the
real error only appears when the transport context exits. The opaque
``CancelledError`` must be replaced with that real cause.
"""
def _make_session(self, mock_session_cls, initialize):
mock_session = AsyncMock()
mock_session.initialize = initialize
session_ctx = MagicMock()
session_ctx.__aenter__ = AsyncMock(return_value=mock_session)
session_ctx.__aexit__ = AsyncMock(return_value=False)
mock_session_cls.return_value = session_ctx
def _make_transport(self, aexit_side_effect):
transport_ctx = MagicMock()
transport_ctx.__aenter__ = AsyncMock(return_value=(MagicMock(), MagicMock()))
transport_ctx.__aexit__ = AsyncMock(side_effect=aexit_side_effect)
return transport_ctx
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_surfaces_connect_error_over_cancelled(self, mock_session_cls):
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
self._make_session(
mock_session_cls,
AsyncMock(side_effect=asyncio.CancelledError("cancelled by group")),
)
connect_error = httpx.ConnectError("All connection attempts failed")
transport_ctx = self._make_transport(
_FakeExceptionGroup("transport", [connect_error])
)
async def _op(session):
return "done"
with pytest.raises(httpx.ConnectError):
await client._execute_session_operation(transport_ctx, _op)
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_genuine_cancellation_is_not_replaced(self, mock_session_cls):
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
self._make_session(
mock_session_cls, AsyncMock(side_effect=asyncio.CancelledError())
)
transport_ctx = self._make_transport(
_FakeExceptionGroup("teardown", [asyncio.CancelledError()])
)
async def _op(session):
return "done"
with pytest.raises(asyncio.CancelledError):
await client._execute_session_operation(transport_ctx, _op)
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_cleanup_error_after_success_is_swallowed(self, mock_session_cls):
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
init_result = MagicMock()
init_result.instructions = None
self._make_session(mock_session_cls, AsyncMock(return_value=init_result))
transport_ctx = self._make_transport(
_FakeExceptionGroup("late", [httpx.ConnectError("late cleanup error")])
)
async def _op(session):
return "done"
result = await client._execute_session_operation(transport_ctx, _op)
assert result == "done"
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -2,6 +2,7 @@ import json
from typing import Any, Dict, Optional
from unittest.mock import MagicMock
import httpx
import pytest
from fastapi import HTTPException
from starlette.requests import Request
@ -1629,3 +1630,48 @@ class TestPreviewOpenAPITools:
"order is out of sync, so collision suffixes (_2, _3, ...) "
"land on different operations"
)
class TestConnectionErrorMessage:
"""The test-connection endpoints turn raw transport errors into messages.
The message is returned to an admin in an API response, so it must explain
the failure without echoing the raw header value, which can carry a secret
(e.g. ``Authorization: Bearer <token>``).
"""
def test_local_protocol_error_is_actionable_and_redacted(self):
secret = "Bearer sk-super-secret-token"
exc = httpx.LocalProtocolError(f"Illegal header value b' {secret}'")
message = rest_endpoints._connection_error_message(exc)
assert "header" in message.lower()
assert secret not in message
def test_connect_error_points_at_reachability(self):
message = rest_endpoints._connection_error_message(
httpx.ConnectError("All connection attempts failed")
)
assert "unreachable" in message.lower()
def test_timeout_error_message(self):
message = rest_endpoints._connection_error_message(
httpx.ConnectTimeout("timed out")
)
assert "unreachable" in message.lower()
def test_http_status_error_includes_status_code(self):
response = httpx.Response(status_code=503)
exc = httpx.HTTPStatusError(
"server error",
request=httpx.Request("POST", "http://x/"),
response=response,
)
message = rest_endpoints._connection_error_message(exc)
assert "503" in message
def test_unknown_error_falls_back_to_generic(self):
message = rest_endpoints._connection_error_message(RuntimeError("weird"))
assert "weird" not in message
assert "proxy logs" in message.lower()

View file

@ -44,7 +44,7 @@ const reduceStaticHeaders = (list: unknown): Record<string, string> => {
if (!Array.isArray(list)) return {};
return list.reduce((acc: Record<string, string>, entry: Record<string, string>) => {
const header = entry?.header?.trim();
if (header) acc[header] = entry?.value ?? "";
if (header) acc[header] = (entry?.value ?? "").trim();
return acc;
}, {});
};

View file

@ -118,7 +118,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
if (!header) {
return acc;
}
acc[header] = entry?.value ?? "";
acc[header] = (entry?.value ?? "").trim();
return acc;
}, {})
: ({} as Record<string, string>);
@ -426,7 +426,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
if (!header) {
return acc;
}
acc[header] = entry?.value ?? "";
acc[header] = (entry?.value ?? "").trim();
return acc;
}, {})
: ({} as Record<string, string>);