From e1114bdcf95680e2499b8979bbb52e9bd922cbee Mon Sep 17 00:00:00 2001 From: teknium1 <127238744+teknium1@users.noreply.github.com> Date: Tue, 15 Sep 2026 11:50:09 -0700 Subject: [PATCH] refactor(mcp): one cycle-safe exception walker for every connect-error scan MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- tests/tools/test_mcp_tool_errors.py | 33 +++++++++++ tests/tools/test_mcp_tool_issue_948.py | 22 -------- tools/mcp_tool_errors.py | 78 +++++++++++--------------- 3 files changed, 67 insertions(+), 66 deletions(-) create mode 100644 tests/tools/test_mcp_tool_errors.py diff --git a/tests/tools/test_mcp_tool_errors.py b/tests/tools/test_mcp_tool_errors.py new file mode 100644 index 0000000000..ddf1e01260 --- /dev/null +++ b/tests/tools/test_mcp_tool_errors.py @@ -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'" diff --git a/tests/tools/test_mcp_tool_issue_948.py b/tests/tools/test_mcp_tool_issue_948.py index d31c02be7c..230668b4a1 100644 --- a/tests/tools/test_mcp_tool_issue_948.py +++ b/tests/tools/test_mcp_tool_issue_948.py @@ -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) diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index c0cb5b299f..d705f03598 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -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