fix(mcp): widen the SSE fallback trigger and harden its guards (salvage #53764)

Relocated onto the decomposed module layout and hardened:

- Trigger covers the rejection CLASS, not just literal 400: SSE-only
  servers' load balancers answer the chunked Streamable HTTP initialize
  POST with 400/405/406/411, and the mcp>=2.0 SDK surfaces many such
  rejections as an opaque -32603 'Server returned an error response'
  (error class per #104363 by @RohithPariki). Timeouts and 5xx never
  trigger the fallback: they are not transport mismatches.
- Reconnect exclusion via _ever_connected instead of _ready: run()
  clears _ready before re-entering the transport, so the original guard
  also fired on reconnects after a proven session.
- Successful fallback latches _sse_fallback so reconnects go straight
  to SSE, and logs a warning suggesting the user pin transport: sse.
- Both transports failing raises a ConnectionError naming both errors
  and suggesting transport: sse / checking the URL.
- No fallback with strict_redirect_headers (SSE cannot enforce that
  boundary) or when transport is explicitly configured.
- Tests trimmed to 3 invariant contracts (proven red on base): fallback
  connects + latches; no fallback on reconnect/timeout/5xx; both-fail
  error is actionable.

The extracted SSE path reuses _sse_transport/_serve_transport from main,
preserving the bounded handshake timeout and reconnect-retry semantics.

Fixes #53676
This commit is contained in:
Teknium
2026-09-09 02:30:49 -07:00
parent b2465f1608
commit a565e2d493
4 changed files with 115 additions and 367 deletions

View File

@@ -1,380 +1,86 @@
"""Tests for automatic SSE fallback when Streamable HTTP returns 400.
"""Invariant tests: automatic Streamable HTTP -> SSE transport fallback (#53676, #104343).
When an MCP server (e.g. WigAI) only implements the SSE transport and
rejects Streamable HTTP initialize requests with 400 Bad Request, the
client should fall back to SSE transport automatically on the initial
connect — without requiring the user to set ``transport: sse``.
An SSE-only MCP server rejects the Streamable HTTP ``initialize`` POST (400-family status,
or the SDK's opaque -32603 "Server returned an error response"); the client must retry over
SSE on the initial connect only — never on reconnect after a proven session, never on
timeout, and a both-transports failure must say so actionably.
"""
from __future__ import annotations
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_400_error(url="https://example.com/mcp"):
"""Build an httpx.HTTPStatusError for 400 Bad Request."""
request = httpx.Request("POST", url)
response = httpx.Response(400, request=request)
return httpx.HTTPStatusError("Bad Request", request=request, response=response)
from tools.mcp_tool import MCPServerTask
def _make_500_error(url="https://example.com/mcp"):
"""Build an httpx.HTTPStatusError for 500 Internal Server Error."""
request = httpx.Request("POST", url)
response = httpx.Response(500, request=request)
return httpx.HTTPStatusError("Server Error", request=request, response=response)
def _http_400(status=400):
request = httpx.Request("POST", "http://127.0.0.1:1/mcp")
return httpx.HTTPStatusError("Bad Request", request=request,
response=httpx.Response(status, request=request))
def _build_server(name="sse-fallback-test"):
"""Create an MCPServerTask with mocks for transport testing."""
from tools.mcp_tool import MCPServerTask
server = MCPServerTask(name)
server._auth_type = ""
server._sampling = None
server._elicitation = None
return server
class _FakeStream:
"""Mock async context manager yielding (read, write) streams."""
class _SdkInternalError(Exception):
"""Shape of mcp.shared.exceptions.MCPError for an opaque initialize rejection."""
def __init__(self):
self._read = AsyncMock()
self._write = AsyncMock()
async def __aenter__(self):
return (self._read, self._write)
async def __aexit__(self, *a):
return False
super().__init__("Server returned an error response")
self.error = type("E", (), {"code": -32603})()
class _FakeSession:
"""Mock MCP ClientSession."""
def _task(monkeypatch, http_exc, sse_result="shutdown", sse_exc=None):
"""MCPServerTask whose transports are recorded fakes: HTTP raises, SSE serves or raises."""
task = MCPServerTask("t")
task._config = {}
calls = []
def __init__(self, *args, **kwargs):
pass
async def fake_serve(self, cm, label, timeout):
calls.append(label)
if label != "SSE":
raise http_exc
if sse_exc is not None:
raise sse_exc
self._ever_connected = True
return sse_result
async def __aenter__(self):
mock_session = MagicMock()
mock_session.initialize = AsyncMock()
return mock_session
async def __aexit__(self, *a):
return False
monkeypatch.setattr(MCPServerTask, "_serve_transport", fake_serve)
monkeypatch.setattr(MCPServerTask, "_streamable_http_transport", lambda self, *a, **k: object())
monkeypatch.setattr(MCPServerTask, "_sse_transport", lambda self, *a, **k: object())
monkeypatch.setattr(MCPServerTask, "_build_oauth_auth", lambda self, *a: None)
return task, calls
class _FakeHTTPTransport:
"""Mock streamable_http_client that records calls and can raise."""
def __init__(self, side_effect=None):
self._side_effect = side_effect
self.called = False
def __call__(self, url, http_client=None):
self.called = True
if self._side_effect is not None:
raise self._side_effect
return _FakeStream()
_CONFIG = {"url": "http://127.0.0.1:1/mcp", "connect_timeout": 1}
class _FakeSSETransport:
"""Mock sse_client that records calls."""
def __init__(self):
self.called = False
self.kwargs = {}
def __call__(self, **kwargs):
self.called = True
self.kwargs.update(kwargs)
return _FakeStream()
@pytest.mark.parametrize("exc", [_http_400(), _http_400(405),
ExceptionGroup("g", [_SdkInternalError()])])
def test_sse_only_server_connects_via_fallback(monkeypatch, exc):
"""Initial connect: a Streamable HTTP rejection falls back to SSE and serves; the latch
routes subsequent reconnects straight to SSE without re-trying Streamable HTTP."""
task, calls = _task(monkeypatch, exc)
assert asyncio.run(task._run_http(dict(_CONFIG))) == "shutdown"
assert calls[-1] == "SSE" and len(calls) == 2
assert asyncio.run(task._run_http(dict(_CONFIG))) == "shutdown" # reconnect after latch
assert calls[2:] == ["SSE"]
class _FakeAsyncClient:
"""Minimal httpx.AsyncClient mock for the Streamable HTTP path."""
def __init__(self, **kwargs):
pass
async def __aenter__(self):
return self
async def __aexit__(self, *a):
return False
@pytest.mark.parametrize("exc,ever_connected", [
(_http_400(), True), # reconnect after a proven session: never mask the 400
(asyncio.TimeoutError(), False), # timeout is not a transport mismatch
(_http_400(500), False), # 5xx is a broken server, not SSE-only
])
def test_no_fallback_on_reconnect_timeout_or_server_error(monkeypatch, exc, ever_connected):
task, calls = _task(monkeypatch, exc)
task._ever_connected = ever_connected
with pytest.raises(type(exc)):
asyncio.run(task._run_http(dict(_CONFIG)))
assert "SSE" not in calls
class _NO_SSE:
"""Sentinel: simulate sse_client not being installed (set to None)."""
pass
_NO_SSE_SENTINEL = _NO_SSE()
def _http_patches(server, *, http_side_effect=None, sse_transport=None,
extra_patches=None):
"""Return a combined context manager with all needed patches.
Uses ``create=True`` for attributes that only exist when the MCP SDK
is installed (``_MCP_NEW_HTTP``, ``streamable_http_client``).
``sse_transport`` controls what ``tools.mcp_tool.sse_client`` is set to:
- A callable/mock: used as the SSE client (default: _FakeSSETransport)
- ``_NO_SSE_SENTINEL``: set sse_client to None (simulates missing SDK)
"""
from contextlib import ExitStack
stack = ExitStack()
stack.enter_context(patch("tools.mcp_tool._MCP_HTTP_AVAILABLE", True))
stack.enter_context(patch("tools.mcp_tool._MCP_NEW_HTTP", True, create=True))
if http_side_effect is not None:
stack.enter_context(patch(
"tools.mcp_tool.streamable_http_client",
_FakeHTTPTransport(side_effect=http_side_effect),
create=True,
))
else:
stack.enter_context(patch(
"tools.mcp_tool.streamable_http_client",
_FakeHTTPTransport(),
create=True,
))
if sse_transport is _NO_SSE_SENTINEL:
stack.enter_context(patch(
"tools.mcp_tool.sse_client", new=None, create=True,
))
elif sse_transport is not None:
stack.enter_context(patch(
"tools.mcp_tool.sse_client", new=sse_transport, create=True,
))
else:
stack.enter_context(patch(
"tools.mcp_tool.sse_client", new=_FakeSSETransport(), create=True,
))
stack.enter_context(patch("tools.mcp_tool.ClientSession", new=_FakeSession, create=True))
stack.enter_context(patch("httpx.AsyncClient", new=_FakeAsyncClient))
stack.enter_context(patch.object(type(server), "_discover_tools",
new=AsyncMock()))
stack.enter_context(patch.object(type(server), "_wait_for_lifecycle_event",
new=AsyncMock(return_value="shutdown")))
if extra_patches:
for p in extra_patches:
stack.enter_context(p)
return stack
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestStreamableHTTP400Fallback:
"""When Streamable HTTP returns 400 on initial connect, fall back to SSE."""
def test_streamable_http_400_falls_back_to_sse(self):
"""Streamable HTTP 400 -> SSE fallback -> success."""
server = _build_server()
fake_sse = _FakeSSETransport()
async def drive():
with _http_patches(server, http_side_effect=_make_400_error(),
sse_transport=fake_sse):
await asyncio.wait_for(
server._run_http({
"url": "https://example.com/mcp",
"timeout": 60,
}),
timeout=5.0,
)
asyncio.run(drive())
assert fake_sse.called, "sse_client was NOT called — SSE fallback did not trigger"
def test_streamable_http_non_400_does_not_fallback(self):
"""Streamable HTTP 500 -> error propagates, NO SSE fallback."""
server = _build_server()
fake_sse = _FakeSSETransport()
async def drive():
with _http_patches(server, http_side_effect=_make_500_error(),
sse_transport=fake_sse):
with pytest.raises(httpx.HTTPStatusError) as exc_info:
await asyncio.wait_for(
server._run_http({
"url": "https://example.com/mcp",
"timeout": 60,
}),
timeout=5.0,
)
assert exc_info.value.response.status_code == 500
asyncio.run(drive())
assert not fake_sse.called, "sse_client was called on a non-400 error"
def test_streamable_http_400_logs_warning(self):
"""400 fallback should log a warning mentioning the server name."""
server = _build_server("wigai")
fake_sse = _FakeSSETransport()
async def drive():
with _http_patches(server, http_side_effect=_make_400_error(),
sse_transport=fake_sse):
await asyncio.wait_for(
server._run_http({
"url": "https://example.com/mcp",
"timeout": 60,
}),
timeout=5.0,
)
asyncio.run(drive())
# The warning is logged at WARNING level — we verify the fallback
# happened (sse_client called) as a proxy for the log being emitted.
assert fake_sse.called
def test_streamable_http_400_fallback_forwards_headers_to_sse(self):
"""SSE fallback receives the same headers dict built for Streamable HTTP."""
server = _build_server()
fake_sse = _FakeSSETransport()
custom_headers = {"X-Custom": "value"}
async def drive():
with _http_patches(server, http_side_effect=_make_400_error(),
sse_transport=fake_sse):
await asyncio.wait_for(
server._run_http({
"url": "https://example.com/mcp",
"headers": custom_headers,
"timeout": 60,
}),
timeout=5.0,
)
asyncio.run(drive())
assert fake_sse.called
# headers should include both the user's custom header and the
# auto-injected mcp-protocol-version
sse_headers = fake_sse.kwargs.get("headers") or {}
assert sse_headers.get("X-Custom") == "value"
assert "mcp-protocol-version" in sse_headers
def test_streamable_http_400_fallback_forwards_oauth_to_sse(self):
"""SSE fallback receives the OAuth auth provider when configured."""
server = _build_server()
server._auth_type = "oauth"
fake_sse = _FakeSSETransport()
fake_oauth = MagicMock(name="fake_oauth_provider")
fake_manager = MagicMock()
fake_manager.get_or_build_provider.return_value = fake_oauth
async def drive():
with _http_patches(
server, http_side_effect=_make_400_error(),
sse_transport=fake_sse,
extra_patches=[
patch("tools.mcp_oauth_manager.get_manager",
return_value=fake_manager),
],
):
await asyncio.wait_for(
server._run_http({
"url": "https://example.com/mcp",
"timeout": 60,
}),
timeout=5.0,
)
asyncio.run(drive())
assert fake_sse.called
assert "auth" in fake_sse.kwargs, "OAuth auth not forwarded to SSE fallback"
assert fake_sse.kwargs["auth"] is fake_oauth
def test_sse_fallback_still_fails_error_propagates(self):
"""Both Streamable HTTP (400) and SSE fail -> SSE error propagates."""
server = _build_server()
class _FailingSSEReturn:
"""sse_client replacement that raises on enter."""
def __init__(self, **kwargs):
pass
def __call__(self, **kwargs):
return self
async def __aenter__(self):
raise ConnectionRefusedError("SSE also refused")
async def __aexit__(self, *a):
return False
async def drive():
with _http_patches(server, http_side_effect=_make_400_error(),
sse_transport=_FailingSSEReturn()):
with pytest.raises(ConnectionRefusedError, match="SSE also refused"):
await asyncio.wait_for(
server._run_http({
"url": "https://example.com/mcp",
"timeout": 60,
}),
timeout=5.0,
)
asyncio.run(drive())
def test_explicit_sse_transport_not_affected_by_fallback(self):
"""When transport: sse is explicit, Streamable HTTP is never attempted."""
server = _build_server()
fake_sse = _FakeSSETransport()
fake_http = _FakeHTTPTransport()
async def drive():
with _http_patches(server, sse_transport=fake_sse):
# Override the streamable_http_client patch to use our
# trackable fake instead of the default one.
import tools.mcp_tool as m
original = m.streamable_http_client
m.streamable_http_client = fake_http
try:
await asyncio.wait_for(
server._run_http({
"url": "https://example.com/mcp",
"transport": "sse",
"timeout": 60,
}),
timeout=5.0,
)
finally:
m.streamable_http_client = original
asyncio.run(drive())
assert fake_sse.called, "SSE path should have been called"
assert not fake_http.called, "Streamable HTTP should NOT be called when transport=sse"
def test_sse_unavailable_during_fallback_raises_import_error(self):
"""When Streamable HTTP 400s and sse_client is None, ImportError raised."""
server = _build_server()
async def drive():
with _http_patches(server, http_side_effect=_make_400_error(),
sse_transport=_NO_SSE_SENTINEL):
with pytest.raises(ImportError, match="SSE transport"):
await asyncio.wait_for(
server._run_http({
"url": "https://example.com/mcp",
"timeout": 60,
}),
timeout=5.0,
)
asyncio.run(drive())
def test_both_transports_failing_names_both_and_suggests_config(monkeypatch):
task, calls = _task(monkeypatch, _http_400(), sse_exc=ConnectionRefusedError("no sse"))
with pytest.raises(ConnectionError, match="both Streamable HTTP and SSE.*transport: sse"):
asyncio.run(task._run_http(dict(_CONFIG)))
assert calls == ["HTTP", "SSE"] or calls == ["legacy HTTP", "SSE"]
assert task._sse_fallback is False # failed fallback must not latch

View File

@@ -316,7 +316,7 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM
"_recycled_reason", "initialize_result", "_ping_unsupported", "_list_cache_meta",
"_reconnect_retries", "_session_proven", "_was_parked", "_inflight_tasks", "_reconnecting",
"_suspect_reason", "_teardown_race", "_permanent_grace_used", "_stdio_child_pids",
"_ever_connected")
"_ever_connected", "_sse_fallback")
def __init__(self, name: str):
self.name = name
@@ -345,6 +345,8 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM
self._session_proven: bool = False
# Never cleared (unlike _ready): separates first-connect from reconnect failures.
self._ever_connected: bool = False
# Latched when the Streamable HTTP -> SSE fallback connects: reconnects reuse SSE directly.
self._sse_fallback: bool = False
# True from park until proven healthy again; logs the revival once.
self._was_parked: bool = False
# In-flight RPC tasks so a deliberate teardown fails them fast; _reconnecting is True

View File

@@ -62,6 +62,25 @@ class NonMcpEndpointError(ConnectionError):
so broad catches still see a connection problem."""
# Streamable-HTTP rejection statuses an SSE-only server (or its load balancer) produces for the
# chunked ``initialize`` POST: Bad Request, Method Not Allowed, Not Acceptable, Length Required.
_STREAMABLE_REJECT_STATUSES = (400, 405, 406, 411)
def _is_streamable_http_rejection(exc: BaseException) -> bool:
"""True when a Streamable-HTTP connect failure looks like a transport mismatch rather than a
broken server: a 400-family rejection of the initialize POST, or the SDK's opaque INTERNAL_ERROR
(-32603 ``Server returned an error response``) it maps such rejections to on mcp >= 2.0 (error
class per PR #104363, @RohithPariki). Timeouts and auth errors never qualify — neither carries
these markers — so a slow or 401ing server is not retried on the wrong transport.
"""
root = _unwrap_exception_group(exc)
if getattr(getattr(root, "response", None), "status_code", None) in _STREAMABLE_REJECT_STATUSES:
return True
code = getattr(getattr(root, "error", None), "code", None)
return code == -32603 and "server returned an error response" in str(root).lower()
def _unwrap_exception_group(exc: BaseException) -> BaseException:
"""Root-cause leaf of anyio ``(Base)ExceptionGroup`` wrappers (group ``str()`` is opaque). A
``KeyboardInterrupt``/``SystemExit`` leaf anywhere is re-raised, never flattened into a loggable

View File

@@ -7,7 +7,7 @@ import asyncio
import os
from contextlib import asynccontextmanager
from typing import Dict, Optional, Set
from tools.mcp_tool_errors import NonMcpEndpointError, _apply_identity_header, _handshake_rejected_as_modern, _make_redirect_header_stripper, _resolve_client_cert, _unwrap_exception_group
from tools.mcp_tool_errors import NonMcpEndpointError, _apply_identity_header, _handshake_rejected_as_modern, _is_streamable_http_rejection, _make_redirect_header_stripper, _resolve_client_cert, _unwrap_exception_group
from tools.mcp_tool_lifecycle import _filter_mcp_children, _orphan_stdio_pid_servers, _orphan_stdio_pids, _stdio_pgids, _stdio_pids
from tools.mcp_tool_common import _core
from tools import mcp_tool_config as _config
@@ -411,23 +411,44 @@ class MCPServerTransportMixin:
self._build_oauth_auth(url, config), bool(config.get("strict_redirect_headers")))
if config.get("transport") == "sse":
return await self._serve_transport(self._sse_transport(*common), "SSE", float(connect_timeout))
if self._sse_fallback:
# A prior connect already proved this server SSE-only: skip the doomed Streamable
# HTTP attempt on reconnects instead of flapping into the retry budget.
logger.info("MCP server '%s': using latched SSE fallback transport", self.name)
return await self._serve_transport(self._sse_transport(*common), "SSE", float(connect_timeout))
transport = self._streamable_http_transport(*common, configured_header_names)
label = "HTTP" if _core._MCP_NEW_HTTP else "legacy HTTP"
try:
return await self._serve_transport(transport, label, float(connect_timeout))
except Exception as exc:
# SSE-only servers (e.g. WigAI for Bitwig Studio) reject the Streamable HTTP
# initialize request with 400 Bad Request, previously a permanent failure with
# 0 active tools unless the user set ``transport: sse`` (#53676). Fall back to
# SSE automatically on the initial connect; reconnects are excluded so a genuine
# 400 on an established transport is not silently masked.
root = _unwrap_exception_group(exc) if isinstance(exc, BaseExceptionGroup) else exc
if (self._ready.is_set()
or getattr(getattr(root, "response", None), "status_code", None) != 400):
# SSE-only servers (or their load balancers) reject the Streamable HTTP chunked
# ``initialize`` POST — with a 400-family status or an opaque SDK INTERNAL_ERROR —
# previously a permanent failure with 0 active tools unless the user set
# ``transport: sse`` (#53676, #104343). Retry over SSE on the initial connect, as
# the MCP spec's transport-fallback behavior describes. Never on reconnect after a
# proven session (``_ever_connected``: a genuine rejection on an established
# transport must not silently switch transports), never on a timeout (not a
# transport mismatch — ``_is_streamable_http_rejection`` matches neither), and never
# with ``strict_redirect_headers`` (SSE cannot enforce that boundary).
if (self._ever_connected or common[-1] or not _is_streamable_http_rejection(exc)):
raise
logger.warning("MCP server '%s': Streamable HTTP returned 400, "
"falling back to SSE transport", self.name)
return await self._serve_transport(self._sse_transport(*common), "SSE", float(connect_timeout))
logger.warning(
"MCP server '%s': Streamable HTTP rejected the initial connect (%s) — retrying "
"over SSE. If this connects, set `transport: sse` for this server in config.yaml "
"to skip the failed attempt on future startups.",
self.name, _unwrap_exception_group(exc))
try:
self._sse_fallback = True
return await self._serve_transport(self._sse_transport(*common), "SSE", float(connect_timeout))
except Exception as sse_exc:
if self._ever_connected: # SSE session was live and dropped: transient, keep the latch
raise
self._sse_fallback = False
raise ConnectionError(
f"MCP server '{self.name}': both Streamable HTTP and SSE transports failed "
f"(Streamable HTTP: {_unwrap_exception_group(exc)}; SSE: "
f"{_unwrap_exception_group(sse_exc)}). Check the URL points at an MCP "
"endpoint, or pin `transport: sse` if the server is SSE-only.") from sse_exc
# -------------------------------------------------------------- discovery