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:
33
tests/tools/test_mcp_tool_errors.py
Normal file
33
tests/tools/test_mcp_tool_errors.py
Normal 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'"
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user