From 55e2986dfd8826e6dc6ddee70068bdd3aace0d24 Mon Sep 17 00:00:00 2001 From: teknium1 <127238744+teknium1@users.noreply.github.com> Date: Tue, 15 Sep 2026 14:20:11 -0700 Subject: [PATCH] fix: walk a group's __cause__/__context__ in the exception node walker Review finding: _exc_children returned only .exceptions for a group, so _is_session_expired_error missed a session-expiry marker (or the InterruptedError override) hanging off a group's __cause__/__context__ that main used to inspect. Groups now yield nested + chain like every other node; _flatten_messages' "group str() is opaque" rule is unchanged. --- tests/tools/test_mcp_tool_session_expired.py | 14 ++++++++++++++ tools/mcp_tool_errors.py | 8 +++++--- 2 files changed, 19 insertions(+), 3 deletions(-) diff --git a/tests/tools/test_mcp_tool_session_expired.py b/tests/tools/test_mcp_tool_session_expired.py index 4f078a3a51..34e5158dd5 100644 --- a/tests/tools/test_mcp_tool_session_expired.py +++ b/tests/tools/test_mcp_tool_session_expired.py @@ -66,6 +66,20 @@ def test_is_session_expired_traversal_is_budget_bounded(): assert _is_session_expired_error(exc) is False +def test_is_session_expired_walks_group_chain(): + """A group's own ``__cause__``/``__context__`` are inspected like any node's: a marker there classifies + as expired, and an InterruptedError there still overrides a marker inside the group.""" + from tools.mcp_tool_errors import _is_session_expired_error + + group = ExceptionGroup("task group", [ValueError("unrelated")]) + group.__context__ = RuntimeError("session terminated") + assert _is_session_expired_error(group) is True + + group = ExceptionGroup("task group", [RuntimeError("session terminated")]) + group.__context__ = InterruptedError() + assert _is_session_expired_error(group) is False + + # --------------------------------------------------------------------------- # Handler integration — verify the recovery plumbing wires end-to-end # --------------------------------------------------------------------------- diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index d705f03598..a70dac9aa8 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -305,9 +305,11 @@ _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)] + """A group's sub-exceptions (if any) followed by ``__cause__``/``__context__`` when they are exceptions — a + group raised inside an ``except`` block carries the caught error as ``__context__``, so the chain is never + skipped.""" + nested = getattr(exc, "exceptions", None) or () + return [*nested, *(c for c in (exc.__cause__, exc.__context__) if isinstance(c, BaseException))] def _iter_exception_nodes(exc: BaseException) -> List[BaseException]: