fix: make every tool interruptible — sequential executor abandons on user interrupt
The sequential tool path only noticed a user interrupt after the running tool returned: with the deadline disabled it ran the tool inline (fully blocking), and with a deadline it waited in 5s slices without ever checking agent._interrupt_requested. Any tool without cooperative is_interrupted() polling (image_generate, tts, transcription, skills sync, ...) held the whole turn hostage — the reported symptom was a redirect queued ~40s behind a FAL image generation + upscale pass. Executor backstop (class fix, covers ALL tools): - _run_sequential_tool_execution_middleware always dispatches on the daemon worker (timeout None no longer means inline blocking) and polls the interrupt flag every 1s. - On interrupt: 3s cooperative grace (mirrors the concurrent path), then synthesize a cancelled tool result (_ToolCancelledResult), emit the terminal post_tool_call with status=cancelled, and abandon the worker. - _ToolCancelledResult suppresses downstream post-hook double emission exactly like _ToolTimeoutResult, so an abandoned worker finishing late cannot report success for a cancelled call. - clarify (interactive, _NEVER_PARALLEL_TOOLS) keeps the inline path — it owns its own human wait. Cooperative layer in the reported offender: - image_generation_tool: blind handler.get() (generation + Clarity upscale) replaced with _wait_fal_result(), which polls is_interrupted() in 0.5s slices and raises ImageGenerationInterrupted immediately. - _upscale_image propagates the interrupt instead of swallowing it into the "upscale failed, use original" fallback. Message alternation is preserved: the cancelled result is a normal tool result for the call_id. Sabotage-verified: with the old wait loop restored, the new tests fail (tool blocks full runtime); with the fix they pass in ~4s.
This commit is contained in:
+90
-11
@@ -392,6 +392,16 @@ class _ToolTimeoutResult(str):
|
||||
"""Marker for a synthesized sequential-tool timeout result."""
|
||||
|
||||
|
||||
class _ToolCancelledResult(str):
|
||||
"""Marker for a synthesized sequential-tool user-interrupt result.
|
||||
|
||||
Like ``_ToolTimeoutResult``, the executor already emitted the terminal
|
||||
post_tool_call event for this call (status="cancelled"), so downstream
|
||||
emission must be suppressed — an abandoned worker finishing late must not
|
||||
report success for a call the user already cancelled.
|
||||
"""
|
||||
|
||||
|
||||
class _ConcurrentToolAuthorizationGate:
|
||||
"""Serialize policy prompts and exclude human approval waits from batch deadlines.
|
||||
|
||||
@@ -735,6 +745,12 @@ def _run_agent_tool_execution_middleware(
|
||||
)
|
||||
|
||||
|
||||
# How often the sequential-tool wait loop wakes to check for a user
|
||||
# interrupt while the worker runs. Short enough that /stop or a redirect
|
||||
# lands within ~1s even when the tool itself never polls is_interrupted().
|
||||
_SEQUENTIAL_INTERRUPT_POLL_SECONDS = 1.0
|
||||
|
||||
|
||||
def _resolve_sequential_tool_timeout() -> float | None:
|
||||
"""Deadline for one sequential tool call (#85125 Phase 2a).
|
||||
|
||||
@@ -788,7 +804,7 @@ def _run_sequential_tool_execution_middleware(
|
||||
"display_index": display_index,
|
||||
"middleware_trace": middleware_trace,
|
||||
}
|
||||
if timeout_s is None or function_name in _NEVER_PARALLEL_TOOLS:
|
||||
if function_name in _NEVER_PARALLEL_TOOLS:
|
||||
return _run_agent_tool_execution_middleware(agent, **kwargs)
|
||||
|
||||
from tools.daemon_pool import DaemonThreadPoolExecutor
|
||||
@@ -815,26 +831,87 @@ def _run_sequential_tool_execution_middleware(
|
||||
|
||||
executor = DaemonThreadPoolExecutor(max_workers=1)
|
||||
future = executor.submit(propagate_context_to_thread(_run))
|
||||
deadline = time.monotonic() + timeout_s
|
||||
# ``timeout_s`` disabled (None) still runs on the worker: the wait loop
|
||||
# below is what makes a non-cooperative tool interruptible at all, so
|
||||
# "no deadline" must not mean "no interrupt checks" (#86xxx class fix —
|
||||
# sequential path previously blocked until the tool returned).
|
||||
deadline = time.monotonic() + timeout_s if timeout_s is not None else None
|
||||
started = time.monotonic()
|
||||
timed_out = False
|
||||
interrupted = False
|
||||
_last_heartbeat = 0
|
||||
try:
|
||||
while True:
|
||||
remaining = (
|
||||
deadline + authorization_gate.excluded_seconds() - time.monotonic()
|
||||
)
|
||||
if remaining <= 0:
|
||||
timed_out = True
|
||||
break
|
||||
wait_slice = _SEQUENTIAL_INTERRUPT_POLL_SECONDS
|
||||
if deadline is not None:
|
||||
remaining = (
|
||||
deadline + authorization_gate.excluded_seconds() - time.monotonic()
|
||||
)
|
||||
if remaining <= 0:
|
||||
timed_out = True
|
||||
break
|
||||
wait_slice = min(wait_slice, remaining)
|
||||
try:
|
||||
return future.result(timeout=min(5.0, remaining))
|
||||
return future.result(timeout=wait_slice)
|
||||
except concurrent.futures.TimeoutError:
|
||||
if agent._interrupt_requested:
|
||||
interrupted = True
|
||||
break
|
||||
elapsed = int(time.monotonic() - started)
|
||||
if elapsed > 0 and elapsed % 30 < 5:
|
||||
if elapsed - _last_heartbeat >= 30:
|
||||
_last_heartbeat = elapsed
|
||||
agent._touch_activity(
|
||||
f"sequential tool running ({elapsed}s): {function_name}"
|
||||
)
|
||||
|
||||
if interrupted:
|
||||
# Belt-and-braces: interrupt() already fans out to tracked worker
|
||||
# tids, but the worker may have registered after the fan-out ran.
|
||||
for tid in worker_tid:
|
||||
try:
|
||||
_ra()._set_interrupt(True, tid)
|
||||
except Exception:
|
||||
pass
|
||||
# Give a cooperative tool a moment to notice its per-thread
|
||||
# interrupt bit and return a real result (mirrors the concurrent
|
||||
# path's 3s grace).
|
||||
concurrent.futures.wait([future], timeout=3.0)
|
||||
if future.done() and not future.cancelled():
|
||||
return future.result()
|
||||
timed_out = True # reuse the abandon-shutdown path in finally
|
||||
future.cancel()
|
||||
message = (
|
||||
f"[Tool execution cancelled — {function_name} was abandoned "
|
||||
"after user interrupt]"
|
||||
)
|
||||
logger.info(
|
||||
"sequential tool %s abandoned after user interrupt (%.1fs elapsed)",
|
||||
function_name, time.monotonic() - started,
|
||||
)
|
||||
trace = middleware_trace if middleware_trace is not None else []
|
||||
_emit_terminal_post_tool_call(
|
||||
agent,
|
||||
function_name=function_name,
|
||||
function_args=function_args,
|
||||
result=message,
|
||||
effective_task_id=effective_task_id,
|
||||
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",
|
||||
middleware_trace=list(trace),
|
||||
)
|
||||
return _ManagedToolResult(
|
||||
result=_ToolCancelledResult(message),
|
||||
args=function_args,
|
||||
middleware_trace=trace,
|
||||
blocked=False,
|
||||
dispatched=True,
|
||||
)
|
||||
|
||||
# Only reachable when a deadline exists (interrupted returns above).
|
||||
assert timeout_s is not None
|
||||
message = (
|
||||
f"Error executing tool '{function_name}': "
|
||||
f"timed out after {timeout_s:.1f}s"
|
||||
@@ -2420,7 +2497,9 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe
|
||||
logger.error("handle_function_call raised for %s: %s", function_name, tool_error, exc_info=True)
|
||||
tool_duration = time.time() - tool_start_time
|
||||
|
||||
_execution_timed_out = isinstance(function_result, _ToolTimeoutResult)
|
||||
_execution_timed_out = isinstance(
|
||||
function_result, (_ToolTimeoutResult, _ToolCancelledResult)
|
||||
)
|
||||
if isinstance(function_result, str):
|
||||
result_preview = function_result if agent.verbose_logging else (
|
||||
function_result[:200] if len(function_result) > 200 else function_result
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
"""Sequential tool execution must abandon the wait when the user interrupts.
|
||||
|
||||
Regression tests for the "interrupt doesn't end a running tool" class:
|
||||
the sequential executor path previously ran the tool inline (when the
|
||||
deadline was disabled) or waited in 5s slices without checking
|
||||
``agent._interrupt_requested`` — a non-cooperative tool (e.g. a blocking
|
||||
FAL ``handler.get()``) held the whole turn hostage until it returned.
|
||||
|
||||
Now the wait loop polls the interrupt flag every
|
||||
``_SEQUENTIAL_INTERRUPT_POLL_SECONDS`` and, after a 3s cooperative grace,
|
||||
synthesizes a cancelled tool result and abandons the worker.
|
||||
"""
|
||||
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
import agent.tool_executor as tool_executor
|
||||
from agent.tool_executor import (
|
||||
_ManagedToolResult,
|
||||
_ToolCancelledResult,
|
||||
_run_sequential_tool_execution_middleware,
|
||||
)
|
||||
|
||||
|
||||
class _FakeAgent:
|
||||
def __init__(self):
|
||||
self._tool_worker_threads = set()
|
||||
self._tool_worker_threads_lock = threading.Lock()
|
||||
self._interrupt_requested = False
|
||||
self.activity = []
|
||||
|
||||
def _touch_activity(self, msg):
|
||||
self.activity.append(msg)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def fake_agent():
|
||||
return _FakeAgent()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _fast_polls(monkeypatch):
|
||||
# Keep the test fast: short poll slice, no config lookups.
|
||||
monkeypatch.setattr(tool_executor, "_SEQUENTIAL_INTERRUPT_POLL_SECONDS", 0.05)
|
||||
emitted = []
|
||||
monkeypatch.setattr(
|
||||
tool_executor,
|
||||
"_emit_terminal_post_tool_call",
|
||||
lambda agent, **kw: emitted.append(kw),
|
||||
)
|
||||
yield emitted
|
||||
|
||||
|
||||
def test_interrupt_abandons_noncooperative_tool(monkeypatch, fake_agent, _fast_polls):
|
||||
"""A blocking tool is abandoned within ~poll+grace once interrupted."""
|
||||
|
||||
started = threading.Event()
|
||||
|
||||
def _fake_middleware(agent_arg, **kwargs):
|
||||
started.set()
|
||||
time.sleep(30) # non-cooperative: never checks is_interrupted()
|
||||
return _ManagedToolResult(
|
||||
result="late result", args={}, middleware_trace=[],
|
||||
blocked=False, dispatched=True,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
tool_executor, "_run_agent_tool_execution_middleware", _fake_middleware
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
tool_executor, "_resolve_sequential_tool_timeout", lambda: None
|
||||
)
|
||||
|
||||
def _interrupt_soon():
|
||||
started.wait(5)
|
||||
time.sleep(0.1)
|
||||
fake_agent._interrupt_requested = True
|
||||
|
||||
threading.Thread(target=_interrupt_soon, daemon=True).start()
|
||||
|
||||
t0 = time.monotonic()
|
||||
managed = _run_sequential_tool_execution_middleware(
|
||||
fake_agent,
|
||||
function_name="image_generate",
|
||||
function_args={"prompt": "x"},
|
||||
effective_task_id="t",
|
||||
tool_call_id="call_1",
|
||||
execute=lambda a: "unused",
|
||||
)
|
||||
elapsed = time.monotonic() - t0
|
||||
|
||||
assert isinstance(managed.result, _ToolCancelledResult)
|
||||
assert "cancelled" in str(managed.result)
|
||||
# poll (0.05s) + interrupt delay (0.1s) + grace (3s) + slack — nowhere
|
||||
# near the 30s tool runtime.
|
||||
assert elapsed < 10.0
|
||||
# The executor emitted the terminal post_tool_call itself.
|
||||
assert any(kw.get("status") == "cancelled" for kw in _fast_polls)
|
||||
|
||||
|
||||
def test_interrupt_prefers_real_result_from_cooperative_tool(
|
||||
monkeypatch, fake_agent, _fast_polls
|
||||
):
|
||||
"""A tool that finishes within the grace window returns its real result."""
|
||||
|
||||
def _fake_middleware(agent_arg, **kwargs):
|
||||
# Cooperative-ish: returns quickly once running (well inside grace).
|
||||
time.sleep(0.3)
|
||||
return _ManagedToolResult(
|
||||
result="real result", args={}, middleware_trace=[],
|
||||
blocked=False, dispatched=True,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
tool_executor, "_run_agent_tool_execution_middleware", _fake_middleware
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
tool_executor, "_resolve_sequential_tool_timeout", lambda: None
|
||||
)
|
||||
fake_agent._interrupt_requested = True # interrupted before first poll
|
||||
|
||||
managed = _run_sequential_tool_execution_middleware(
|
||||
fake_agent,
|
||||
function_name="web_search",
|
||||
function_args={},
|
||||
effective_task_id="t",
|
||||
tool_call_id="call_2",
|
||||
execute=lambda a: "unused",
|
||||
)
|
||||
|
||||
assert managed.result == "real result"
|
||||
assert not isinstance(managed.result, _ToolCancelledResult)
|
||||
|
||||
|
||||
def test_no_deadline_still_runs_on_worker(monkeypatch, fake_agent):
|
||||
"""timeout disabled (None) must not fall back to inline blocking."""
|
||||
|
||||
seen_thread = []
|
||||
|
||||
def _fake_middleware(agent_arg, **kwargs):
|
||||
seen_thread.append(threading.current_thread().ident)
|
||||
return _ManagedToolResult(
|
||||
result="ok", args={}, middleware_trace=[],
|
||||
blocked=False, dispatched=True,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
tool_executor, "_run_agent_tool_execution_middleware", _fake_middleware
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
tool_executor, "_resolve_sequential_tool_timeout", lambda: None
|
||||
)
|
||||
|
||||
managed = _run_sequential_tool_execution_middleware(
|
||||
fake_agent,
|
||||
function_name="read_file",
|
||||
function_args={},
|
||||
effective_task_id="t",
|
||||
tool_call_id="call_3",
|
||||
execute=lambda a: "unused",
|
||||
)
|
||||
|
||||
assert managed.result == "ok"
|
||||
assert seen_thread and seen_thread[0] != threading.current_thread().ident
|
||||
|
||||
|
||||
def test_never_parallel_tools_stay_inline(monkeypatch, fake_agent):
|
||||
"""clarify (interactive) keeps the inline path — it owns its own wait."""
|
||||
|
||||
seen_thread = []
|
||||
|
||||
def _fake_middleware(agent_arg, **kwargs):
|
||||
seen_thread.append(threading.current_thread().ident)
|
||||
return _ManagedToolResult(
|
||||
result="ok", args={}, middleware_trace=[],
|
||||
blocked=False, dispatched=True,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
tool_executor, "_run_agent_tool_execution_middleware", _fake_middleware
|
||||
)
|
||||
|
||||
managed = _run_sequential_tool_execution_middleware(
|
||||
fake_agent,
|
||||
function_name="clarify",
|
||||
function_args={},
|
||||
effective_task_id="t",
|
||||
tool_call_id="call_4",
|
||||
execute=lambda a: "unused",
|
||||
)
|
||||
|
||||
assert managed.result == "ok"
|
||||
assert seen_thread and seen_thread[0] == threading.current_thread().ident
|
||||
@@ -0,0 +1,82 @@
|
||||
"""_wait_fal_result must notice a user interrupt while the FAL job runs."""
|
||||
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
import tools.image_generation_tool as image_tool
|
||||
from tools.interrupt import set_interrupt
|
||||
|
||||
|
||||
class _SlowHandler:
|
||||
"""Fake FAL handler whose get() blocks like the real SDK."""
|
||||
|
||||
def __init__(self, delay=30.0, result=None):
|
||||
self.delay = delay
|
||||
self._result = result if result is not None else {"images": []}
|
||||
|
||||
def get(self):
|
||||
time.sleep(self.delay)
|
||||
return self._result
|
||||
|
||||
|
||||
class _FastHandler:
|
||||
def __init__(self, result):
|
||||
self._result = result
|
||||
|
||||
def get(self):
|
||||
return self._result
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_interrupt():
|
||||
set_interrupt(False)
|
||||
yield
|
||||
set_interrupt(False)
|
||||
|
||||
|
||||
def test_wait_fal_result_returns_result():
|
||||
result = image_tool._wait_fal_result(_FastHandler({"images": [{"url": "u"}]}))
|
||||
assert result == {"images": [{"url": "u"}]}
|
||||
|
||||
|
||||
def test_wait_fal_result_raises_on_interrupt():
|
||||
def _interrupt_soon(tid):
|
||||
time.sleep(0.2)
|
||||
set_interrupt(True, tid)
|
||||
|
||||
tid = threading.current_thread().ident
|
||||
threading.Thread(target=_interrupt_soon, args=(tid,), daemon=True).start()
|
||||
|
||||
t0 = time.monotonic()
|
||||
with pytest.raises(image_tool.ImageGenerationInterrupted):
|
||||
image_tool._wait_fal_result(_SlowHandler(delay=30.0), poll_seconds=0.05)
|
||||
assert time.monotonic() - t0 < 5.0
|
||||
|
||||
|
||||
def test_wait_fal_result_propagates_handler_error():
|
||||
class _ErrHandler:
|
||||
def get(self):
|
||||
raise ValueError("boom")
|
||||
|
||||
with pytest.raises(ValueError, match="boom"):
|
||||
image_tool._wait_fal_result(_ErrHandler())
|
||||
|
||||
|
||||
def test_upscale_interrupt_propagates(monkeypatch):
|
||||
"""_upscale_image must NOT swallow the interrupt into a None fallback."""
|
||||
|
||||
monkeypatch.setattr(
|
||||
image_tool, "_submit_fal_request", lambda *a, **k: _SlowHandler(30.0)
|
||||
)
|
||||
|
||||
def _interrupt_soon(tid):
|
||||
time.sleep(0.2)
|
||||
set_interrupt(True, tid)
|
||||
|
||||
tid = threading.current_thread().ident
|
||||
threading.Thread(target=_interrupt_soon, args=(tid,), daemon=True).start()
|
||||
|
||||
with pytest.raises(image_tool.ImageGenerationInterrupted):
|
||||
image_tool._upscale_image("https://example.com/x.png", "prompt")
|
||||
@@ -726,6 +726,46 @@ def _get_managed_fal_client(managed_gateway):
|
||||
return _managed_fal_client
|
||||
|
||||
|
||||
class ImageGenerationInterrupted(Exception):
|
||||
"""Raised when the user interrupts while a FAL job is in flight."""
|
||||
|
||||
|
||||
def _wait_fal_result(handler, *, poll_seconds: float = 0.5):
|
||||
"""Interrupt-aware replacement for a blind ``handler.get()``.
|
||||
|
||||
``handler.get()`` blocks inside the FAL SDK until the remote job
|
||||
finishes — a 30-60s window where a user interrupt was previously
|
||||
invisible (the reported symptom: redirects queued behind a running
|
||||
generation). Run the blocking get on a daemon worker and poll the
|
||||
per-thread interrupt bit between join slices; on interrupt, abandon
|
||||
the worker (daemon thread, remote job keeps running server-side but
|
||||
we stop waiting) and raise ``ImageGenerationInterrupted``.
|
||||
"""
|
||||
from tools.interrupt import is_interrupted
|
||||
|
||||
result_box: list = []
|
||||
error_box: list = []
|
||||
|
||||
def _get():
|
||||
try:
|
||||
result_box.append(handler.get())
|
||||
except BaseException as exc: # noqa: BLE001 — re-raised on the caller thread
|
||||
error_box.append(exc)
|
||||
|
||||
worker = threading.Thread(target=_get, daemon=True, name="fal-result-wait")
|
||||
worker.start()
|
||||
while worker.is_alive():
|
||||
if is_interrupted():
|
||||
raise ImageGenerationInterrupted(
|
||||
"Image generation interrupted by user — abandoned the "
|
||||
"in-flight FAL job."
|
||||
)
|
||||
worker.join(timeout=poll_seconds)
|
||||
if error_box:
|
||||
raise error_box[0]
|
||||
return result_box[0] if result_box else None
|
||||
|
||||
|
||||
def _submit_fal_request(model: str, arguments: Dict[str, Any]):
|
||||
"""Submit a FAL request using direct credentials or the managed queue gateway."""
|
||||
# Trigger the lazy import on first call. Idempotent.
|
||||
@@ -938,7 +978,7 @@ def _upscale_image(image_url: str, original_prompt: str) -> Optional[Dict[str, A
|
||||
}
|
||||
|
||||
handler = _submit_fal_request(UPSCALER_MODEL, arguments=upscaler_arguments)
|
||||
result = handler.get()
|
||||
result = _wait_fal_result(handler)
|
||||
|
||||
if result and "image" in result:
|
||||
upscaled_image = result["image"]
|
||||
@@ -957,6 +997,10 @@ def _upscale_image(image_url: str, original_prompt: str) -> Optional[Dict[str, A
|
||||
logger.error("Upscaler returned invalid response")
|
||||
return None
|
||||
|
||||
except ImageGenerationInterrupted:
|
||||
# Propagate: the user interrupt must not degrade into a silent
|
||||
# "upscale failed, use original" fallback that keeps the turn alive.
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("Error upscaling image: %s", e, exc_info=True)
|
||||
return None
|
||||
@@ -1204,7 +1248,7 @@ def image_generate_tool(
|
||||
)
|
||||
|
||||
handler = _submit_fal_request(endpoint, arguments=arguments)
|
||||
result = handler.get()
|
||||
result = _wait_fal_result(handler)
|
||||
|
||||
generation_time = (datetime.datetime.now() - start_time).total_seconds()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user