From c1c0efa375cb34da85936676c63fbcc3da50e320 Mon Sep 17 00:00:00 2001 From: fangliquanflq Date: Mon, 24 Aug 2026 01:38:06 +0800 Subject: [PATCH] fix(code-exec): preserve interrupt cancellation source --- agent/tool_executor.py | 28 ++++++++---- agent/turn_context.py | 7 ++- run_agent.py | 21 ++++++++- tests/run_agent/test_interrupt_propagation.py | 19 +++++++- .../run_agent/test_sequential_tool_timeout.py | 43 +++++++++++++++++++ .../test_approved_command_clean_slate.py | 3 ++ tests/tools/test_code_execution.py | 28 ++++++++++++ tools/code_execution_tool.py | 19 ++++++-- tools/interrupt.py | 25 ++++++++++- 9 files changed, 175 insertions(+), 18 deletions(-) diff --git a/agent/tool_executor.py b/agent/tool_executor.py index 1209749e47..69bcca6075 100644 --- a/agent/tool_executor.py +++ b/agent/tool_executor.py @@ -905,7 +905,11 @@ def _run_sequential_tool_execution_middleware( # tids, but the worker may have registered after the fan-out ran. for tid in worker_tid: try: - _ra()._set_interrupt(True, tid) + _ra()._set_interrupt( + True, + tid, + reason=getattr(agent, "_tool_interrupt_reason", None), + ) except Exception: pass # Give a cooperative tool a moment to notice its per-thread @@ -916,13 +920,17 @@ def _run_sequential_tool_execution_middleware( return future.result() timed_out = True # reuse the abandon-shutdown path in finally future.cancel() + interrupt_reason = ( + getattr(agent, "_tool_interrupt_reason", None) + or "interrupt requested" + ) message = ( - f"[Tool execution cancelled — {function_name} was abandoned " - "after user interrupt]" + f"[Tool execution cancelled — {function_name} was abandoned: " + f"{interrupt_reason}]" ) logger.info( - "sequential tool %s abandoned after user interrupt (%.1fs elapsed)", - function_name, time.monotonic() - started, + "sequential tool %s abandoned due to %s (%.1fs elapsed)", + function_name, interrupt_reason, time.monotonic() - started, ) trace = middleware_trace if middleware_trace is not None else [] _emit_terminal_post_tool_call( @@ -934,8 +942,8 @@ def _run_sequential_tool_execution_middleware( tool_call_id=tool_call_id, duration_ms=int((time.monotonic() - started) * 1000), status="cancelled", - error_type="keyboard_interrupt", - error_message="Tool execution cancelled by user interrupt", + error_type="tool_interrupted", + error_message=f"Tool execution cancelled: {interrupt_reason}", middleware_trace=list(trace), ) return _ManagedToolResult( @@ -1313,7 +1321,11 @@ def execute_tool_calls_concurrent(agent, assistant_message, messages: list, effe # the tool returns True on the next poll. if agent._interrupt_requested: try: - _ra()._set_interrupt(True, _worker_tid) + _ra()._set_interrupt( + True, + _worker_tid, + reason=getattr(agent, "_tool_interrupt_reason", None), + ) except Exception: pass # Set the activity callback on THIS worker thread so diff --git a/agent/turn_context.py b/agent/turn_context.py index f9bc38f126..9df563d4ec 100644 --- a/agent/turn_context.py +++ b/agent/turn_context.py @@ -1331,10 +1331,15 @@ def build_turn_context( # Clear stale per-thread interrupt state, preserving a pending interrupt. ra()._set_interrupt(False, agent._execution_thread_id) if agent._interrupt_requested: - ra()._set_interrupt(True, agent._execution_thread_id) + ra()._set_interrupt( + True, + agent._execution_thread_id, + reason=getattr(agent, "_tool_interrupt_reason", None), + ) agent._interrupt_thread_signal_pending = False else: agent._interrupt_message = None + agent._tool_interrupt_reason = None agent._interrupt_thread_signal_pending = False # Notify memory providers of the new turn (BEFORE prefetch_all). diff --git a/run_agent.py b/run_agent.py index d5c2447860..8cc82617a5 100644 --- a/run_agent.py +++ b/run_agent.py @@ -3314,17 +3314,28 @@ class AIAgent: ) event.set() + # Keep tool cancellation attribution separate from _interrupt_message: + # ordinary interrupts may carry the user's full next message, which + # must not be copied into tool output. + tool_interrupt_reason = ( + (message or "explicit stop requested") + if hard_cancel + else ("user sent a new message" if message else "user interrupt") + ) + _redirect_lock = getattr(self, "_pending_redirect_lock", None) if _redirect_lock is not None: with _redirect_lock: self._interrupt_requested = True self._interrupt_message = message + self._tool_interrupt_reason = tool_interrupt_reason if hard_cancel: _admit_hard_cancel() self._pending_redirect = None else: self._interrupt_requested = True self._interrupt_message = message + self._tool_interrupt_reason = tool_interrupt_reason if hard_cancel: _admit_hard_cancel() self._pending_redirect = None @@ -3357,7 +3368,11 @@ class AIAgent: # Scope the interrupt to this agent's execution thread so other # agents running in the same process (gateway) are not affected. if self._execution_thread_id is not None: - _set_interrupt(True, self._execution_thread_id) + _set_interrupt( + True, + self._execution_thread_id, + reason=tool_interrupt_reason, + ) self._interrupt_thread_signal_pending = False else: # The interrupt arrived before run_conversation() finished @@ -3381,7 +3396,7 @@ class AIAgent: _worker_tids = list(_tracker) for _wtid in _worker_tids: try: - _set_interrupt(True, _wtid) + _set_interrupt(True, _wtid, reason=tool_interrupt_reason) except Exception: pass # Propagate interrupt to any running child agents (subagent delegation) @@ -3423,6 +3438,7 @@ class AIAgent: return False self._interrupt_requested = False self._interrupt_message = None + self._tool_interrupt_reason = None getattr(self, "_hard_interrupt_requested", threading.Event()).clear() if not preserve_redirect: self._pending_redirect = None @@ -3431,6 +3447,7 @@ class AIAgent: return False self._interrupt_requested = False self._interrupt_message = None + self._tool_interrupt_reason = None getattr(self, "_hard_interrupt_requested", threading.Event()).clear() if not preserve_redirect: self._pending_redirect = None diff --git a/tests/run_agent/test_interrupt_propagation.py b/tests/run_agent/test_interrupt_propagation.py index f53afe7186..fccf28124d 100644 --- a/tests/run_agent/test_interrupt_propagation.py +++ b/tests/run_agent/test_interrupt_propagation.py @@ -9,7 +9,7 @@ import time import unittest from unittest.mock import MagicMock -from tools.interrupt import set_interrupt, is_interrupted +from tools.interrupt import get_interrupt_reason, set_interrupt, is_interrupted class TestInterruptPropagationToChild(unittest.TestCase): @@ -75,6 +75,23 @@ class TestInterruptPropagationToChild(unittest.TestCase): assert agent._interrupt_requested is True assert not agent._hard_interrupt_requested.is_set() + def test_message_interrupt_records_source_without_user_text(self): + agent = self._make_bare_agent() + agent._execution_thread_id = threading.current_thread().ident + + agent.interrupt("private follow-up text") + + assert get_interrupt_reason() == "user sent a new message" + assert "private follow-up text" not in get_interrupt_reason() + + def test_hard_interrupt_records_control_reason(self): + agent = self._make_bare_agent() + agent._execution_thread_id = threading.current_thread().ident + + agent.hard_interrupt("superseded by a new live turn") + + assert get_interrupt_reason() == "superseded by a new live turn" + def test_active_turn_redirect_does_not_set_hard_cancel(self): agent = self._make_bare_agent() agent._model_request_active = threading.Event() diff --git a/tests/run_agent/test_sequential_tool_timeout.py b/tests/run_agent/test_sequential_tool_timeout.py index b22895c55a..d517af8f11 100644 --- a/tests/run_agent/test_sequential_tool_timeout.py +++ b/tests/run_agent/test_sequential_tool_timeout.py @@ -215,6 +215,49 @@ def test_sequential_tool_timeout_suppresses_late_terminal_event(tmp_path, monkey ] +def test_sequential_tool_interrupt_reports_hard_cancel_reason(tmp_path, monkeypatch): + agent = _make_agent(tmp_path) + first_started = threading.Event() + release_first = threading.Event() + terminal_events: list[dict] = [] + + def _dispatch(*_args, **_kwargs): + first_started.set() + release_first.wait() + return "late result" + + def _interrupt(): + assert first_started.wait(timeout=5) + agent.hard_interrupt("superseded by a new live turn") + + interrupter = threading.Thread(target=_interrupt, daemon=True) + interrupter.start() + messages: list[dict] = [] + monkeypatch.setenv("HERMES_CONCURRENT_TOOL_TIMEOUT_S", "30") + + try: + with ( + patch("run_agent.handle_function_call", side_effect=_dispatch), + patch( + "agent.tool_executor._emit_terminal_post_tool_call", + side_effect=lambda *_args, **kwargs: terminal_events.append(kwargs), + ), + ): + execute_tool_calls_sequential( + agent, + SimpleNamespace(tool_calls=[_tool_call("hung")]), + messages, + "task", + ) + finally: + release_first.set() + interrupter.join(timeout=2) + + assert "superseded by a new live turn" in messages[0]["content"] + assert "user interrupt" not in messages[0]["content"] + assert terminal_events[0]["error_type"] == "tool_interrupted" + + @pytest.mark.parametrize( "clarify_timeout", [resolve_clarify_timeout({}), 0], diff --git a/tests/tools/test_approved_command_clean_slate.py b/tests/tools/test_approved_command_clean_slate.py index 4c06830236..04a924fe57 100644 --- a/tests/tools/test_approved_command_clean_slate.py +++ b/tests/tools/test_approved_command_clean_slate.py @@ -232,5 +232,8 @@ def test_execute_code_non_approved_still_interrupts_on_stale_bit(monkeypatch): # Killed on the first poll before the script can print. assert "CODE_DONE" not in result["output"], result + assert result["status"] == "interrupted", result + assert result["output"] == "[execution interrupted]" + assert "user sent a new message" not in result["output"] diff --git a/tests/tools/test_code_execution.py b/tests/tools/test_code_execution.py index 8e60206e38..3c03c320a6 100644 --- a/tests/tools/test_code_execution.py +++ b/tests/tools/test_code_execution.py @@ -46,6 +46,7 @@ from tools.code_execution_tool import ( EXECUTE_CODE_SCHEMA, _TOOL_DOC_LINES, _execute_remote, + _format_interrupted_output, ) from tools.registry import registry @@ -81,6 +82,33 @@ class TestSandboxRequirements(unittest.TestCase): self.assertIn("code", EXECUTE_CODE_SCHEMA["parameters"]["required"]) +class TestInterruptedOutput(unittest.TestCase): + def tearDown(self): + from tools.interrupt import set_interrupt + + set_interrupt(False) + + def test_uses_recorded_interrupt_source(self): + from tools.interrupt import set_interrupt + + set_interrupt(True, reason="superseded by a new live turn") + + self.assertEqual( + _format_interrupted_output("partial output"), + "partial output\n[execution interrupted — superseded by a new live turn]", + ) + + def test_unknown_interrupt_source_is_neutral(self): + from tools.interrupt import set_interrupt + + set_interrupt(True) + + self.assertEqual( + _format_interrupted_output(""), + "[execution interrupted]", + ) + + class TestHermesToolsGeneration(unittest.TestCase): def test_generates_all_allowed_tools(self): src = generate_hermes_tools_module(list(SANDBOX_ALLOWED_TOOLS)) diff --git a/tools/code_execution_tool.py b/tools/code_execution_tool.py index 81cdbcf745..f15b06223a 100644 --- a/tools/code_execution_tool.py +++ b/tools/code_execution_tool.py @@ -1060,6 +1060,19 @@ def _rpc_poll_loop( stop_event.wait(poll_interval) +def _format_interrupted_output(stdout_text: str) -> str: + """Append an interruption marker without guessing who caused it.""" + from tools.interrupt import get_interrupt_reason + + reason = get_interrupt_reason() + marker = ( + f"[execution interrupted — {reason}]" + if reason + else "[execution interrupted]" + ) + return f"{stdout_text}\n{marker}" if stdout_text else marker + + def _execute_remote( code: str, task_id: Optional[str], @@ -1238,9 +1251,7 @@ def _execute_remote( duration, timeout, tool_call_counter[0], ) elif status == "interrupted": - result["output"] = ( - stdout_text + "\n[execution interrupted — user sent a new message]" - ) + result["output"] = _format_interrupted_output(stdout_text) elif exit_code != 0: result["status"] = "error" result["error"] = f"Script exited with code {exit_code}" @@ -1710,7 +1721,7 @@ def execute_code( duration, timeout, tool_call_counter[0], ) elif status == "interrupted": - result["output"] = stdout_text + "\n[execution interrupted — user sent a new message]" + result["output"] = _format_interrupted_output(stdout_text) elif exit_code != 0: result["status"] = "error" result["error"] = stderr_text or f"Script exited with code {exit_code}" diff --git a/tools/interrupt.py b/tools/interrupt.py index da31dcfeb7..7ec1db39d8 100644 --- a/tools/interrupt.py +++ b/tools/interrupt.py @@ -31,25 +31,39 @@ if _DEBUG_INTERRUPT: # Force our own logger back to INFO so the trace is visible in agent.log. logger.setLevel(logging.INFO) -# Set of thread idents that have been interrupted. +# Set of thread idents that have been interrupted, plus an optional +# user-safe cause for each signal. The cause deliberately does not contain an +# incoming user's message text. _interrupted_threads: set[int] = set() +_interrupt_reasons: dict[int, str] = {} _lock = threading.Lock() -def set_interrupt(active: bool, thread_id: int | None = None) -> None: +def set_interrupt( + active: bool, + thread_id: int | None = None, + *, + reason: str | None = None, +) -> None: """Set or clear interrupt for a specific thread. Args: active: True to signal interrupt, False to clear it. thread_id: Target thread ident. When None, targets the current thread (backward compat for CLI/tests). + reason: Optional user-safe cause for the interrupt. """ tid = thread_id if thread_id is not None else threading.current_thread().ident with _lock: if active: _interrupted_threads.add(tid) + if reason: + _interrupt_reasons[tid] = reason + else: + _interrupt_reasons.pop(tid, None) else: _interrupted_threads.discard(tid) + _interrupt_reasons.pop(tid, None) _snapshot = set(_interrupted_threads) if _DEBUG_INTERRUPT else None if _DEBUG_INTERRUPT: logger.info( @@ -70,6 +84,13 @@ def is_interrupted() -> bool: return tid in _interrupted_threads +def get_interrupt_reason() -> str | None: + """Return the user-safe interrupt cause for the current thread, if known.""" + tid = threading.current_thread().ident + with _lock: + return _interrupt_reasons.get(tid) + + def clear_current_thread_interrupt() -> None: """Clear any interrupt bit on the CURRENT thread.