fix(code-exec): preserve interrupt cancellation source

This commit is contained in:
fangliquanflq
2026-08-24 01:38:06 +08:00
committed by Teknium
parent ee8a66233f
commit c1c0efa375
9 changed files with 175 additions and 18 deletions
+20 -8
View File
@@ -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
+6 -1
View File
@@ -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).
+19 -2
View File
@@ -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
+18 -1
View File
@@ -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()
@@ -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],
@@ -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"]
+28
View File
@@ -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))
+15 -4
View File
@@ -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}"
+23 -2
View File
@@ -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.