refactor(mcp): one cycle-safe exception walker for every connect-error scan

The salvaged fix gave `_find_missing` and `_flatten_messages` each their own
visited-set loop, next to the one `_is_session_expired_error` already had —
three copies of the same idiom in one module. Collapse them into
`_iter_exception_nodes` (pre-order, left-to-right, each node once, bounded by
`_EXC_TRAVERSAL_MAX_NODES`) and read all three scans off that list. Acyclic
output is byte-identical: the missing-executable search keeps its depth-first
order and a message-less leaf still renders as its class name.

Tests move from the issue-numbered file into `tests/tools/test_mcp_tool_errors.py`
(mirror of the source module): a two-node cycle renders the real messages, and a
missing stdio binary wrapped deeper than the recursion limit with the chain
looping back to the top is still reported as the missing executable. Both are
red on origin/main (RecursionError).

Co-authored-by: Stephan Mongstad <stephan@users.noreply.github.com>
This commit is contained in:
teknium1
2026-09-15 11:50:09 -07:00
committed by Teknium
parent 030d4caa0e
commit e1114bdcf9
3 changed files with 67 additions and 66 deletions

View File

@@ -0,0 +1,33 @@
"""Invariants for ``tools/mcp_tool_errors._format_connect_error`` on malformed exception chains.
``__cause__``/``__context__`` can form a cycle (the same OAuth error re-raised on the SSE fallback,
a raised-and-caught pair) and stdio failures can nest deeper than the recursion limit; either used
to turn ``hermes mcp test`` into a RecursionError that hid the real connect error (#111952, #111997).
"""
import sys
from tools.mcp_tool_errors import _format_connect_error
def test_format_connect_error_reports_real_messages_on_cyclic_chain():
"""A two-node ``__cause__``/``__context__`` cycle renders every distinct message, once, in chain order."""
first = RuntimeError("first failure")
second = RuntimeError("second failure")
first.__cause__ = second
second.__context__ = first
assert _format_connect_error(first) == "first failure; second failure"
def test_format_connect_error_finds_missing_executable_through_deep_cyclic_chain():
"""A missing stdio binary wrapped deeper than the recursion limit, with the chain looping back to the top,
is still reported as the missing executable rather than as a RecursionError."""
missing = FileNotFoundError(2, "No such file or directory", "/opt/homebrew/bin/removed-mcp-server")
current = missing
for _ in range(sys.getrecursionlimit() + 10):
wrapper = RuntimeError("stdio startup failed")
wrapper.__cause__ = current
current = wrapper
missing.__context__ = current
assert _format_connect_error(current) == "missing executable '/opt/homebrew/bin/removed-mcp-server'"

View File

@@ -20,28 +20,6 @@ if not _MCP_AVAILABLE:
_mcp_mod.ClientSession = MagicMock
def test_format_connect_error_finds_missing_executable_in_deep_exception_chain():
"""A deeply wrapped missing binary must not recurse while formatting it."""
missing = FileNotFoundError(2, "No such file or directory", "/opt/bin/removed-mcp")
current = missing
for _ in range(sys.getrecursionlimit() + 10):
wrapper = RuntimeError("stdio startup failed")
wrapper.__cause__ = current
current = wrapper
assert _format_connect_error(current) == "missing executable '/opt/bin/removed-mcp'"
def test_format_connect_error_handles_cyclic_exception_chain():
"""Malformed exception chains must fall back to a finite, sanitized message."""
first = RuntimeError("first failure")
second = RuntimeError("second failure")
first.__cause__ = second
second.__cause__ = first
assert _format_connect_error(first) == "first failure; second failure"
def test_resolve_stdio_command_falls_back_to_hermes_node_bin(tmp_path):
node_bin = tmp_path / "node" / "bin"
node_bin.mkdir(parents=True)

View File

@@ -299,55 +299,59 @@ def _make_mcp_body_cap_transport(httpx_mod, inner_transport, limit: int = _MCP_H
return _BodyCapTransport(inner_transport)
# Node budget for ``_iter_exception_nodes`` (the visited set breaks cycles; this bounds acyclic blow-ups).
# Well above ``sys.getrecursionlimit()`` so deep task-group nesting is fully scanned.
_EXC_TRAVERSAL_MAX_NODES = 10_000
def _exc_children(exc: BaseException) -> List[BaseException]:
"""Sub-exceptions of a group, else ``__cause__``/``__context__`` when they are exceptions."""
nested = getattr(exc, "exceptions", None)
return list(nested) if nested else [c for c in (exc.__cause__, exc.__context__) if isinstance(c, BaseException)]
def _iter_exception_nodes(exc: BaseException) -> List[BaseException]:
"""Pre-order, left-to-right walk of an exception tree/chain, each node once. ``__cause__``/``__context__``
can point back at an ancestor (a raised-and-caught pair does this routinely, e.g. the same OAuth error
raised on the Streamable-HTTP attempt and again on the SSE fallback), so a naive recursive walk dies with
RecursionError and hides the real connect error; the visited set breaks cycles, the budget bounds acyclic
blow-ups."""
stack = [exc]
seen: set[int] = set()
ordered: List[BaseException] = []
while stack and len(ordered) < _EXC_TRAVERSAL_MAX_NODES:
current = stack.pop()
if id(current) in seen:
continue
seen.add(id(current))
ordered.append(current)
stack.extend(reversed(_exc_children(current)))
return ordered
def _format_connect_error(exc: BaseException) -> str:
"""Render nested MCP connection errors into an actionable short message."""
nodes = _iter_exception_nodes(exc)
def _find_missing() -> Optional[str]:
"""Find a missing executable without recursing through malformed chains."""
stack = [exc]
seen: set[int] = set()
budget = _EXC_TRAVERSAL_MAX_NODES
while stack and budget > 0:
current = stack.pop()
if id(current) in seen:
continue
seen.add(id(current))
budget -= 1
for current in nodes:
if isinstance(current, FileNotFoundError):
if getattr(current, "filename", None):
return str(current.filename)
match = re.search(r"No such file or directory: '([^']+)'", str(current))
if match:
return match.group(1)
# Reverse preserves the former left-to-right depth-first order.
stack.extend(reversed(_exc_children(current)))
return None
def _flatten_messages() -> List[str]:
"""Collect a short, cycle-safe rendering of an exception chain."""
stack = [exc]
seen: set[int] = set()
messages: List[str] = []
budget = _EXC_TRAVERSAL_MAX_NODES
while stack and budget > 0:
current = stack.pop()
if id(current) in seen:
continue
seen.add(id(current))
budget -= 1
children = _exc_children(current)
# A group's own str() is opaque — only its children speak.
for current in nodes:
# A group's own str() is opaque — only its children speak; a message-less leaf still names its type.
text = "" if getattr(current, "exceptions", None) else str(current).strip()
if text:
messages.append(text)
elif not children:
elif not _exc_children(current):
messages.append(current.__class__.__name__)
stack.extend(reversed(children))
return messages or [exc.__class__.__name__]
missing = _find_missing()
@@ -407,34 +411,20 @@ _SESSION_EXPIRED_MARKERS: tuple = (
"unknown session", "session terminated", "closedresourceerror", "closed resource",
"transport is closed", "connection closed", "broken pipe", "end of file")
# Node budget for ``_is_session_expired_error`` (the visited set breaks cycles; this bounds acyclic blow-ups).
# Well above ``sys.getrecursionlimit()`` so deep task-group nesting is fully scanned.
_EXC_TRAVERSAL_MAX_NODES = 10_000
def _is_session_expired_error(exc: BaseException) -> bool:
"""True if ``exc`` looks like a transport session expiry (Streamable-HTTP servers GC session state on idle TTL /
restart / pod rotation while the OAuth token stays valid) — the fix is a transport reconnect, not an OAuth
refresh. Iterative walk over ``exceptions`` / ``__cause__`` / ``__context__`` with a visited set AND a node
budget; every reachable node is inspected so an InterruptedError anywhere overrides transport markers, and the
chain walk matters because SDK wrappers raise a generic RuntimeError *from* a message-less ClosedResourceError."""
refresh. Every node ``_iter_exception_nodes`` reaches is inspected so an InterruptedError anywhere overrides
transport markers; the chain walk matters because SDK wrappers raise a generic RuntimeError *from* a
message-less ClosedResourceError."""
# AnyIO stream exceptions are often message-less, so type checks complement marker matching.
transport_error_types = tuple(_optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream"))
stack: "list[BaseException | None]" = [exc]
seen: set[int] = set()
found = False
budget = _EXC_TRAVERSAL_MAX_NODES
while stack and budget > 0:
current = stack.pop()
if current is None or id(current) in seen:
continue
seen.add(id(current))
budget -= 1
for current in _iter_exception_nodes(exc):
if isinstance(current, InterruptedError):
return False
# Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids false positives.
msg = str(current).lower()
found = found or isinstance(current, transport_error_types) or any(m in msg for m in _SESSION_EXPIRED_MARKERS)
stack.extend((*getattr(current, "exceptions", ()), getattr(current, "__cause__", None),
getattr(current, "__context__", None)))
return found