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.
This commit is contained in:
teknium1
2026-09-15 14:20:11 -07:00
committed by Teknium
parent e1114bdcf9
commit 55e2986dfd
2 changed files with 19 additions and 3 deletions

View File

@@ -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
# ---------------------------------------------------------------------------

View File

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