mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
parent
02344057a2
commit
587829c3a0
6 changed files with 266 additions and 9 deletions
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}, {});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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>);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue