fix(code-exec): preserve interrupt cancellation source
This commit is contained in:
+20
-8
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user