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]: