Files
hermes-agent/tests/tools/test_interrupt.py
teknium1 5f6b1d251f test: purge low-value tests, lane py17 (495 removed)
Change-detectors, tautologies, source-reading tests, redundant duplicates,
mock-echo tests and dead/unrunnable tests. Per-test rationale in the lane
ledger (category + reason for every removal).
2026-09-23 03:15:26 -07:00

348 lines
13 KiB
Python

"""Tests for the interrupt system.
Run with: python -m pytest tests/test_interrupt.py -v
"""
import threading
import time
import pytest
# ---------------------------------------------------------------------------
# Unit tests: shared interrupt module
# ---------------------------------------------------------------------------
class TestInterruptModule:
"""Tests for tools/interrupt.py"""
def test_set_and_check(self):
from tools.interrupt import set_interrupt, is_interrupted
set_interrupt(False)
assert not is_interrupted()
set_interrupt(True)
assert is_interrupted()
set_interrupt(False)
assert not is_interrupted()
def test_is_thread_interrupted_checks_target_tid_not_caller(self):
from tools.interrupt import (
set_interrupt, is_interrupted, is_thread_interrupted, _interrupted_threads, _lock,
)
with _lock:
_interrupted_threads.clear()
other_tid = threading.get_ident() + 1
set_interrupt(True, thread_id=other_tid)
assert not is_interrupted()
assert is_thread_interrupted(other_tid)
assert is_thread_interrupted(None) is False
set_interrupt(False, thread_id=other_tid)
assert not is_thread_interrupted(other_tid)
def test_clear_current_thread_interrupt_leaves_other_threads(self):
"""clear_current_thread_interrupt only touches the calling thread."""
from tools.interrupt import (
set_interrupt, is_interrupted, clear_current_thread_interrupt,
_interrupted_threads, _lock,
)
with _lock:
_interrupted_threads.clear()
other_tid = threading.get_ident() + 1 # an ident that isn't us
set_interrupt(True, thread_id=other_tid)
set_interrupt(True) # current thread
assert is_interrupted()
clear_current_thread_interrupt()
assert not is_interrupted() # ours cleared
with _lock:
assert other_tid in _interrupted_threads # other thread untouched
_interrupted_threads.discard(other_tid)
def test_run_if_not_interrupted_skips_callback_when_already_interrupted(self):
from tools.interrupt import run_if_not_interrupted, set_interrupt
callbacks = []
set_interrupt(True)
try:
assert run_if_not_interrupted(lambda: callbacks.append("claimed")) is False
finally:
set_interrupt(False)
assert callbacks == []
@pytest.mark.parametrize("callback_should_fail", [False, True])
def test_run_if_not_interrupted_orders_callback_before_concurrent_interrupt(
self, callback_should_fail, monkeypatch
):
import tools.interrupt as interrupt
class CallbackFailure(Exception):
pass
original_lock = interrupt._lock
attempting_interrupt_lock = threading.Event()
interrupt_published = threading.Event()
publisher_lock_contention = []
callback_observations = []
setters = []
setter_tids = []
class ObservedLock:
def __enter__(self):
if (
threading.current_thread() in setters
and not attempting_interrupt_lock.is_set()
):
acquired = original_lock.acquire(blocking=False)
publisher_lock_contention.append(not acquired)
attempting_interrupt_lock.set()
if acquired:
return self
original_lock.acquire()
return self
def __exit__(self, exc_type, exc_value, traceback):
original_lock.release()
interrupt.set_interrupt(False)
with original_lock:
baseline = (
set(interrupt._interrupted_threads),
dict(interrupt._interrupt_reasons),
)
monkeypatch.setattr(interrupt, "_lock", ObservedLock())
def publish_interrupt():
setter_tids.append(threading.get_ident())
try:
interrupt.set_interrupt(True)
interrupt_published.set()
finally:
interrupt.set_interrupt(False)
def callback():
setter = threading.Thread(target=publish_interrupt)
setters.append(setter)
setter.start()
assert attempting_interrupt_lock.wait(5)
assert publisher_lock_contention == [True]
callback_observations.append(interrupt_published.is_set())
if callback_should_fail:
raise CallbackFailure
try:
if callback_should_fail:
with pytest.raises(CallbackFailure):
interrupt.run_if_not_interrupted(callback)
else:
assert interrupt.run_if_not_interrupted(callback) is True
assert interrupt_published.wait(5)
finally:
for setter in setters:
if setter.ident is not None:
setter.join(timeout=5)
interrupt.set_interrupt(False)
assert setters
assert all(not setter.is_alive() for setter in setters)
assert setter_tids
assert callback_observations == [False]
assert interrupt_published.is_set()
with original_lock:
final_state = (
set(interrupt._interrupted_threads),
dict(interrupt._interrupt_reasons),
)
assert final_state == baseline
assert all(setter_tid not in final_state[0] for setter_tid in setter_tids)
assert all(setter_tid not in final_state[1] for setter_tid in setter_tids)
# ---------------------------------------------------------------------------
# Unit tests: pre-tool interrupt check
# ---------------------------------------------------------------------------
class TestPreToolCheck:
"""Verify that _execute_tool_calls skips all tools when interrupted."""
def test_all_tools_skipped_when_interrupted(self):
"""Mock an interrupted agent and verify no tools execute."""
from unittest.mock import MagicMock
# Build a fake assistant_message with 3 tool calls
tc1 = MagicMock()
tc1.id = "tc_1"
tc1.function.name = "terminal"
tc1.function.arguments = '{"command": "rm -rf /"}'
tc2 = MagicMock()
tc2.id = "tc_2"
tc2.function.name = "terminal"
tc2.function.arguments = '{"command": "echo hello"}'
tc3 = MagicMock()
tc3.id = "tc_3"
tc3.function.name = "web_search"
tc3.function.arguments = '{"query": "test"}'
assistant_msg = MagicMock()
assistant_msg.tool_calls = [tc1, tc2, tc3]
messages = []
# Create a minimal mock agent with _interrupt_requested = True
agent = MagicMock()
agent._interrupt_requested = True
agent.log_prefix = ""
agent._persist_session = MagicMock()
# PR #72425: execute_tool_calls_* read _incremental_persistence_failed
# via getattr at loop top. A bare MagicMock auto-creates a truthy value
# for any attribute access, which would short-circuit the interrupt
# skip path before any cancelled-tool messages are appended.
agent._incremental_persistence_failed = False
# Import and call the method
import types
from run_agent import AIAgent
# Bind the real methods to our mock so dispatch works correctly
agent._execute_tool_calls_sequential = types.MethodType(AIAgent._execute_tool_calls_sequential, agent)
agent._execute_tool_calls_concurrent = types.MethodType(AIAgent._execute_tool_calls_concurrent, agent)
AIAgent._execute_tool_calls(agent, assistant_msg, messages, "default")
# All 3 should be skipped
assert len(messages) == 3
for msg in messages:
assert msg["role"] == "tool"
assert "cancelled" in msg["content"].lower() or "interrupted" in msg["content"].lower()
# No actual tool handlers should have been called
# (handle_function_call should NOT have been invoked)
# ---------------------------------------------------------------------------
# Unit tests: message combining
# ---------------------------------------------------------------------------
# ---------------------------------------------------------------------------
# Integration tests (require local terminal)
# ---------------------------------------------------------------------------
class TestSIGKILLEscalation:
"""Test that SIGTERM-resistant processes get SIGKILL'd."""
@pytest.mark.skipif(
not __import__("shutil").which("bash"),
reason="Requires bash"
)
def test_sigterm_trap_killed_within_2s(self):
"""A process that traps SIGTERM should be SIGKILL'd after 1s grace."""
from tools.interrupt import set_interrupt
from tools.environments.local import LocalEnvironment
set_interrupt(False)
env = LocalEnvironment(cwd="/tmp", timeout=30)
# Start execution in a thread, interrupt after 0.5s
result_holder = {"value": None}
def _run():
result_holder["value"] = env.execute(
"trap '' TERM; sleep 60",
timeout=30,
)
t = threading.Thread(target=_run)
t.start()
time.sleep(0.5)
set_interrupt(True, thread_id=t.ident)
t.join(timeout=5)
set_interrupt(False, thread_id=t.ident)
assert result_holder["value"] is not None
assert result_holder["value"]["returncode"] == 130
assert "interrupted" in result_holder["value"]["output"].lower()
# ---------------------------------------------------------------------------
# Regression: _run_tool cleanup on BaseException (issue #35309)
# ---------------------------------------------------------------------------
class TestRunToolCleanupOnBaseException:
"""Verify that _run_tool cleans up _interrupted_threads even when
_invoke_tool raises a BaseException (e.g. CancelledError).
Regression test for #35309: without the finally block, a BaseException
bypasses ``except Exception``, leaking the worker tid into
_interrupted_threads. ThreadPoolExecutor recycles tids, so the next
tool scheduled on the same thread is instantly "interrupted".
"""
def test_cleanup_on_base_exception(self):
from unittest.mock import MagicMock
import types
from tools.interrupt import set_interrupt, _interrupted_threads, _lock
# Clear global state
with _lock:
_interrupted_threads.clear()
# Build a minimal mock agent with the attributes _run_tool needs
agent = MagicMock()
agent._interrupt_requested = False
agent._tool_worker_threads = set()
agent._tool_worker_threads_lock = threading.Lock()
# _set_interrupt delegates to the real module
def _mock_set_interrupt(active, tid=None):
set_interrupt(active, tid)
agent._set_interrupt = _mock_set_interrupt
# _invoke_tool raises BaseException (simulating CancelledError)
agent._invoke_tool = MagicMock(side_effect=BaseException("simulated CancelledError"))
# Bind the real concurrent method so we get _run_tool
from run_agent import AIAgent
agent._execute_tool_calls_concurrent = types.MethodType(
AIAgent._execute_tool_calls_concurrent, agent
)
# Build a single tool call
tc = MagicMock()
tc.id = "tc_base_exc"
tc.function.name = "dummy_tool"
tc.function.arguments = "{}"
assistant_msg = MagicMock()
assistant_msg.tool_calls = [tc]
# _execute_tool_calls_concurrent will submit _run_tool to a
# ThreadPoolExecutor. The BaseException propagates out of the
# worker, but the finally block should still clean up.
try:
agent._execute_tool_calls_concurrent(assistant_msg, [], "default")
except Exception:
pass # ThreadPoolExecutor may re-raise
# After the worker finishes (even with BaseException), the worker
# tid should have been removed from _interrupted_threads and
# _tool_worker_threads.
assert len(agent._tool_worker_threads) == 0, (
f"_tool_worker_threads not cleaned up: {agent._tool_worker_threads}"
)
# Verify no stale tid is left in the global interrupt set. The
# worker thread is recycled by ThreadPoolExecutor, so a leaked tid
# would poison the next task on that thread. We cleared the set at
# the start and never set any interrupt ourselves, so a leak from
# _run_tool is the only way an entry could land here.
with _lock:
leaked = set(_interrupted_threads)
assert leaked == set(), f"leaked tids in _interrupted_threads: {leaked}"