From 20d59bd726e105ee1d9cbd001a3ceee2fe359321 Mon Sep 17 00:00:00 2001 From: byte1024 Date: Mon, 7 Sep 2026 09:13:15 +0800 Subject: [PATCH] fix: preserve Claude terminal outcomes --- nerve/agent/backends/claude.py | 26 +++++++++++++++++++- tests/test_engine.py | 44 +++++++++++++++++++++++++++++----- 2 files changed, 63 insertions(+), 7 deletions(-) diff --git a/nerve/agent/backends/claude.py b/nerve/agent/backends/claude.py index 99ff3caa..546eaf88 100644 --- a/nerve/agent/backends/claude.py +++ b/nerve/agent/backends/claude.py @@ -208,6 +208,7 @@ def translate_message(message: Any) -> list[ev.AgentEvent]: ev.NormalizedUsage.from_anthropic(message.usage) if message.usage else None ) + status, error = _result_outcome(message) out.append(ev.TurnCompleted( native_session_id=message.session_id, model=None, # claude reports the model per AssistantMessage @@ -218,12 +219,35 @@ def translate_message(message: Any) -> list[ev.AgentEvent]: duration_ms=getattr(message, "duration_ms", None), duration_api_ms=getattr(message, "duration_api_ms", None), num_turns=getattr(message, "num_turns", None), - status="completed", + status=status, + error=error, )) return out +def _result_outcome(message: ResultMessage) -> tuple[ev.TurnStatus, str | None]: + reason = getattr(message, "terminal_reason", None) + if reason in {"aborted_streaming", "aborted_tools"}: + return "interrupted", reason.replace("_", " ") + + subtype = getattr(message, "subtype", "") + if not getattr(message, "is_error", False) and subtype == "success": + return "completed", None + + errors = getattr(message, "errors", None) + if errors: + detail = "; ".join(errors) + elif reason == "max_turns": + turns = getattr(message, "num_turns", None) + detail = f"max turns ({turns}) exhausted" if turns is not None else "max turns exhausted" + elif status := getattr(message, "api_error_status", None): + detail = f"API error (HTTP {status})" + else: + detail = (reason or subtype or "Claude turn failed").replace("_", " ") + return "failed", detail + + def _translate_tool_result( block: ToolResultBlock, parent_id: str | None, ) -> ev.ToolResult: diff --git a/tests/test_engine.py b/tests/test_engine.py index 63b8f53e..a9e05cf1 100644 --- a/tests/test_engine.py +++ b/tests/test_engine.py @@ -202,13 +202,15 @@ def _translated(messages: list) -> list: return [event for m in messages for event in translate_message(m)] -def _result_msg(session_id: str = "sdk-1") -> ResultMessage: +def _result_msg(session_id: str = "sdk-1", **overrides) -> ResultMessage: """A terminal ResultMessage: translates to one TurnCompleted.""" - return ResultMessage( - subtype="success", duration_ms=1, duration_api_ms=1, - is_error=False, num_turns=1, session_id=session_id, - total_cost_usd=0.5, usage={"input_tokens": 1}, - ) + values = { + "subtype": "success", "duration_ms": 1, "duration_api_ms": 1, + "is_error": False, "num_turns": 1, "session_id": session_id, + "total_cost_usd": 0.5, "usage": {"input_tokens": 1}, + } + values.update(overrides) + return ResultMessage(**values) @pytest.mark.asyncio @@ -412,6 +414,36 @@ async def test_receive_turn_completes_on_result_without_raising(): assert sdk.aclose_calls == 1 +@pytest.mark.parametrize( + ("message", "status", "error"), + [ + ( + _result_msg( + subtype="error_max_turns", is_error=True, + terminal_reason="max_turns", num_turns=50, + ), + "failed", + "max turns (50) exhausted", + ), + ( + _result_msg(is_error=True, api_error_status=529), + "failed", + "API error (HTTP 529)", + ), + ( + _result_msg(is_error=True, terminal_reason="aborted_streaming"), + "interrupted", + "aborted streaming", + ), + ], +) +def test_result_message_preserves_abnormal_terminal_state(message, status, error): + event = translate_message(message)[0] + + assert isinstance(event, ev.TurnCompleted) + assert (event.status, event.error) == (status, error) + + @pytest.mark.asyncio async def test_receive_turn_idle_timeout_is_not_a_transport_death(): """A hung (but live) CLI must stay distinguishable from a dead one.