diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 0bc81ece5f0..3cb94e4bc11 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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]: """ diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index f215b7fa28b..78cecfed0bd 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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: diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index dee689708c3..c9e500b4a5b 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -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__]) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 6433e0f6360..caff9ea2d28 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -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 ``). + """ + + 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() diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index b205d9e6582..018fd292902 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -44,7 +44,7 @@ const reduceStaticHeaders = (list: unknown): Record => { if (!Array.isArray(list)) return {}; return list.reduce((acc: Record, entry: Record) => { const header = entry?.header?.trim(); - if (header) acc[header] = entry?.value ?? ""; + if (header) acc[header] = (entry?.value ?? "").trim(); return acc; }, {}); }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index 9445b4226d6..f6ded2287e3 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -118,7 +118,7 @@ const MCPServerEdit: React.FC = ({ if (!header) { return acc; } - acc[header] = entry?.value ?? ""; + acc[header] = (entry?.value ?? "").trim(); return acc; }, {}) : ({} as Record); @@ -426,7 +426,7 @@ const MCPServerEdit: React.FC = ({ if (!header) { return acc; } - acc[header] = entry?.value ?? ""; + acc[header] = (entry?.value ?? "").trim(); return acc; }, {}) : ({} as Record);