fix(delegation): restore parent cancellation for rejected async units
This commit is contained in:
304
tests/tools/test_delegate_capacity_interrupt.py
Normal file
304
tests/tools/test_delegate_capacity_interrupt.py
Normal file
@@ -0,0 +1,304 @@
|
||||
"""A rejected background batch retains synchronous cancellation ownership."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import Future
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.interrupt_control import InterruptControlMixin
|
||||
from agent.turn_context import _bind_interrupt_scope
|
||||
from tools import async_delegation
|
||||
from tools.delegate_tool_dispatch import _Batch, _dispatch_background
|
||||
from tools.interrupt import is_interrupted, set_interrupt
|
||||
from tools.process_registry import process_registry
|
||||
|
||||
|
||||
class _Parent(InterruptControlMixin):
|
||||
def __init__(self):
|
||||
self.session_id = "capacity-interrupt-parent"
|
||||
self._active_children = []
|
||||
self._active_children_lock = threading.Lock()
|
||||
self._execution_thread_id = None
|
||||
self._interrupt_requested = False
|
||||
self._hard_interrupt_requested = threading.Event()
|
||||
self.quiet_mode = True
|
||||
|
||||
|
||||
class _ControlledChild(_Parent):
|
||||
"""Replace model work while retaining real interrupt and worker-start semantics."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.session_id = "capacity-interrupt-child"
|
||||
self._delegate_role = "leaf"
|
||||
self._delegate_depth = 1
|
||||
self._delegate_saved_tool_names = []
|
||||
self._credential_pool = None
|
||||
self._subagent_id = None
|
||||
self.tool_progress_callback = None
|
||||
self.model = "test-model"
|
||||
self.started = threading.Event()
|
||||
self.stop_received = threading.Event()
|
||||
self.unwinding = threading.Event()
|
||||
self.allow_finish = threading.Event()
|
||||
self.finished = threading.Event()
|
||||
self.closed = threading.Event()
|
||||
self.close_count = 0
|
||||
self.closed_while_running = False
|
||||
self.observed_interrupt = None
|
||||
|
||||
def interrupt(self, message=None, **kwargs):
|
||||
accepted = super().interrupt(message, **kwargs)
|
||||
self.stop_received.set()
|
||||
return accepted
|
||||
|
||||
def hard_interrupt(self, message=None, **kwargs):
|
||||
super().hard_interrupt(message, **kwargs)
|
||||
self.stop_received.set()
|
||||
|
||||
def run_conversation(self, **_kwargs):
|
||||
# A stop can arrive before this thread exists. Use the real turn-start
|
||||
# binding so the pending agent interrupt must reach the tool thread too.
|
||||
_bind_interrupt_scope(self, lambda: SimpleNamespace(_set_interrupt=set_interrupt))
|
||||
self.started.set()
|
||||
try:
|
||||
assert self.stop_received.wait(30), "child never received cancellation"
|
||||
assert self._interrupt_requested
|
||||
assert is_interrupted(), "stop did not reach the child execution thread"
|
||||
self.observed_interrupt = (
|
||||
self._interrupt_message, self._hard_interrupt_requested.is_set(),
|
||||
)
|
||||
self.unwinding.set()
|
||||
assert self.allow_finish.wait(30), "test did not release child cleanup"
|
||||
return {
|
||||
"final_response": "", "completed": False, "interrupted": True,
|
||||
"api_calls": 0, "messages": [],
|
||||
}
|
||||
finally:
|
||||
self.clear_interrupt()
|
||||
self.finished.set()
|
||||
|
||||
def get_activity_summary(self):
|
||||
return {"api_call_count": 0}
|
||||
|
||||
def close(self):
|
||||
self.closed_while_running |= not self.finished.is_set()
|
||||
self.close_count += 1
|
||||
self.closed.set()
|
||||
|
||||
|
||||
def _batch(parent, *children):
|
||||
tasks = [{"goal": f"wait until cancelled {i}"} for i in range(len(children))]
|
||||
parent._active_children.extend(children)
|
||||
return _Batch(
|
||||
task_list=tasks, children=[(i, tasks[i], child) for i, child in enumerate(children)], parent_agent=parent,
|
||||
creds={"model": children[0].model}, context=None, top_role="leaf", max_children=len(children),
|
||||
live_deleg_id=None, live_writers=[], live_paths=[], origin_wake_sid="",
|
||||
origin_ui_session_id="", origin_owner_transport=None,
|
||||
origin_owner_session_record=None, overall_start=time.monotonic(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def registry_state(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
monkeypatch.delenv("HERMES_IGNORE_USER_CONFIG", raising=False)
|
||||
(tmp_path / "config.yaml").write_text(
|
||||
"delegation:\n max_concurrent_children: 1\n worktree_isolation: false\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
async_delegation._reset_for_tests()
|
||||
completion_queue = queue.Queue()
|
||||
monkeypatch.setattr(process_registry, "completion_queue", completion_queue)
|
||||
yield completion_queue
|
||||
# Test bodies release their gates and join their workers before registry teardown.
|
||||
if async_delegation._executor is not None:
|
||||
async_delegation._executor.shutdown(wait=True)
|
||||
async_delegation._reset_for_tests()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rejection", ["capacity", "schedule_failure", "partial_schedule_failure"])
|
||||
@pytest.mark.parametrize("stop_timing", ["running", "during_admission"])
|
||||
@pytest.mark.parametrize("stop_kind", ["soft", "hard"])
|
||||
def test_rejected_background_child_stops_with_parent(
|
||||
registry_state, monkeypatch, tmp_path, rejection, stop_timing, stop_kind,
|
||||
):
|
||||
parent, child = _Parent(), _ControlledChild()
|
||||
background_child = pending_child = None
|
||||
if rejection == "partial_schedule_failure":
|
||||
# Three independent units: one accepted, one rejected, one not yet submitted.
|
||||
# The model-facing batch width is legal under the configured limit.
|
||||
(tmp_path / "config.yaml").write_text(
|
||||
"delegation:\n max_concurrent_children: 3\n worktree_isolation: false\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
background_child, pending_child = _ControlledChild(), _ControlledChild()
|
||||
background_child.session_id += "-background"
|
||||
pending_child.session_id += "-pending"
|
||||
batch = _batch(parent, background_child, child, pending_child)
|
||||
else:
|
||||
batch = _batch(parent, child)
|
||||
occupied = threading.Event()
|
||||
release_occupier = threading.Event()
|
||||
admission_started = threading.Event()
|
||||
continue_admission = threading.Event()
|
||||
outcome = Future()
|
||||
|
||||
def occupy_slot():
|
||||
occupied.set()
|
||||
assert release_occupier.wait(30)
|
||||
return {"status": "completed", "summary": "slot released"}
|
||||
|
||||
if rejection == "capacity":
|
||||
accepted = async_delegation.dispatch_async_delegation(
|
||||
goal="occupy the only slot", context=None, toolsets=None, role="leaf",
|
||||
model=child.model, session_key="other-session", runner=occupy_slot,
|
||||
max_async_children=1,
|
||||
)
|
||||
assert accepted["status"] == "dispatched"
|
||||
assert occupied.wait(5)
|
||||
elif rejection == "schedule_failure":
|
||||
class RejectingExecutor:
|
||||
def submit(self, *_args, **_kwargs):
|
||||
raise RuntimeError("executor shut down")
|
||||
|
||||
monkeypatch.setattr(async_delegation, "_get_executor", lambda _n: RejectingExecutor())
|
||||
else:
|
||||
executor = async_delegation._get_executor(3)
|
||||
|
||||
class PartiallyRejectingExecutor:
|
||||
submitted = 0
|
||||
|
||||
def submit(self, *args, **kwargs):
|
||||
self.submitted += 1
|
||||
if self.submitted == 2:
|
||||
raise RuntimeError("unit submission failed")
|
||||
return executor.submit(*args, **kwargs)
|
||||
|
||||
partial_executor = PartiallyRejectingExecutor()
|
||||
monkeypatch.setattr(async_delegation, "_get_executor", lambda _n: partial_executor)
|
||||
|
||||
dispatch = async_delegation.dispatch_async_delegation_batch
|
||||
admissions = 0
|
||||
accepted_ids = []
|
||||
|
||||
def pause_admission(**kwargs):
|
||||
nonlocal admissions
|
||||
admissions += 1
|
||||
rejected_admission = 2 if background_child is not None else 1
|
||||
if admissions == rejected_admission:
|
||||
admission_started.set()
|
||||
assert continue_admission.wait(5)
|
||||
result = dispatch(**kwargs)
|
||||
if result.get("status") == "dispatched":
|
||||
accepted_ids.append(result["delegation_id"])
|
||||
return result
|
||||
|
||||
monkeypatch.setattr(async_delegation, "dispatch_async_delegation_batch", pause_admission)
|
||||
|
||||
def run_dispatch():
|
||||
try:
|
||||
outcome.set_result(json.loads(_dispatch_background(batch)))
|
||||
except BaseException as exc:
|
||||
outcome.set_exception(exc)
|
||||
|
||||
worker = threading.Thread(target=run_dispatch, daemon=True)
|
||||
worker.start()
|
||||
try:
|
||||
assert admission_started.wait(5)
|
||||
if background_child is not None:
|
||||
assert background_child.started.wait(5)
|
||||
request_stop = parent.hard_interrupt if stop_kind == "hard" else parent.interrupt
|
||||
stop_message = "user correction or stop request"
|
||||
if stop_timing == "during_admission":
|
||||
request_stop(stop_message)
|
||||
continue_admission.set()
|
||||
assert child.started.wait(5)
|
||||
if stop_timing == "running":
|
||||
request_stop(stop_message)
|
||||
|
||||
assert child.stop_received.wait(5), "fallback lost parent cancellation ownership"
|
||||
assert child.unwinding.wait(5)
|
||||
assert child.observed_interrupt == (stop_message, stop_kind == "hard")
|
||||
if background_child is not None:
|
||||
assert not background_child.stop_received.is_set()
|
||||
assert pending_child.stop_received.is_set(), "unsubmitted unit lost parent cancellation ownership"
|
||||
assert not pending_child.started.is_set()
|
||||
assert not outcome.done(), "dispatch returned while its child still owned resources"
|
||||
assert child.close_count == 0
|
||||
child.allow_finish.set()
|
||||
result = outcome.result(timeout=5)
|
||||
if background_child is None:
|
||||
assert "SYNCHRONOUSLY" in result["note"]
|
||||
assert result["results"][0]["status"] == "interrupted"
|
||||
else:
|
||||
assert result["status"] == "dispatched"
|
||||
assert result["inline_results"][0]["status"] == "interrupted"
|
||||
assert not background_child.stop_received.is_set()
|
||||
assert pending_child.unwinding.wait(5)
|
||||
assert pending_child.observed_interrupt == (stop_message, stop_kind == "hard")
|
||||
assert child.finished.is_set()
|
||||
assert child.close_count == 1
|
||||
assert not child.closed_while_running
|
||||
assert parent._active_children == []
|
||||
finally:
|
||||
continue_admission.set()
|
||||
if not child.finished.is_set():
|
||||
child.hard_interrupt("test teardown")
|
||||
child.allow_finish.set()
|
||||
worker.join(timeout=5)
|
||||
release_occupier.set()
|
||||
if background_child is not None:
|
||||
async_delegation.interrupt_for_session(parent_session_id=parent.session_id)
|
||||
for extra in (background_child, pending_child):
|
||||
extra.allow_finish.set()
|
||||
assert extra.closed.wait(5)
|
||||
assert extra.finished.is_set()
|
||||
assert extra.close_count == 1
|
||||
assert not extra.closed_while_running
|
||||
completed_ids = {registry_state.get(timeout=5)["delegation_id"] for _ in accepted_ids}
|
||||
assert completed_ids == set(accepted_ids)
|
||||
if rejection == "capacity":
|
||||
completion = registry_state.get(timeout=5)
|
||||
assert completion["delegation_id"] == accepted["delegation_id"]
|
||||
assert not worker.is_alive()
|
||||
|
||||
|
||||
def test_accepted_background_child_keeps_registry_cancellation_ownership(registry_state):
|
||||
parent, child = _Parent(), _ControlledChild()
|
||||
try:
|
||||
result = json.loads(_dispatch_background(_batch(parent, child)))
|
||||
assert result["status"] == "dispatched"
|
||||
assert child.started.wait(5)
|
||||
parent.interrupt()
|
||||
# Parent interrupt fan-out is synchronous; observing it return establishes
|
||||
# that a detached child did not receive it without a timing-based wait.
|
||||
assert parent._interrupt_requested
|
||||
assert not child.stop_received.is_set()
|
||||
assert not child.finished.is_set()
|
||||
parent.hard_interrupt("stop the current parent turn")
|
||||
assert parent._hard_interrupt_requested.is_set()
|
||||
assert not child.stop_received.is_set()
|
||||
assert async_delegation.interrupt_for_session(parent_session_id=parent.session_id) == 1
|
||||
assert child.unwinding.wait(5)
|
||||
assert child.observed_interrupt[1] is True
|
||||
assert child.close_count == 0
|
||||
child.allow_finish.set()
|
||||
completion = registry_state.get(timeout=5)
|
||||
assert completion["delegation_id"] == result["delegation_id"]
|
||||
assert completion["results"][0]["status"] == "interrupted"
|
||||
assert child.finished.is_set()
|
||||
assert child.close_count == 1
|
||||
assert not child.closed_while_running
|
||||
assert parent._active_children == []
|
||||
finally:
|
||||
if not child.finished.is_set():
|
||||
child.hard_interrupt("test teardown")
|
||||
child.allow_finish.set()
|
||||
assert child.closed.wait(5)
|
||||
@@ -14,7 +14,7 @@ from dataclasses import dataclass, replace
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from tools.async_delegation import _new_delegation_id, record_unit_child
|
||||
from tools.delegate_tool_child_run import _detach_child, _fabricated_entry, _signal_child_stop
|
||||
from tools.delegate_tool_child_run import _attach_child, _detach_child, _fabricated_entry, _signal_child_stop
|
||||
from tools.delegate_tool_progress import (
|
||||
SUBAGENT_FAILURE_STATUSES, _clean_error_text, _print_completion_line, _quiet, format_batch_tag,
|
||||
)
|
||||
@@ -374,6 +374,21 @@ def _dispatch_unit(unit: _Batch, unit_id: Optional[str], slot_key: Optional[str]
|
||||
progress_fn=lambda: _batch_progress_token(child_agents), **routing,
|
||||
)
|
||||
|
||||
def _restore_parent_cancellation(unit: _Batch) -> None:
|
||||
# Rejected children remain owned by the parent. Attach before replaying a
|
||||
# cancellation that may have arrived while async admission had them detached.
|
||||
parent = unit.parent_agent
|
||||
for _, _, child in unit.children:
|
||||
_attach_child(parent, child)
|
||||
if getattr(parent, "_interrupt_requested", False) is True:
|
||||
hard_stop = getattr(parent, "_hard_interrupt_requested", None)
|
||||
for _, _, child in unit.children:
|
||||
if hard_stop is not None and hard_stop.is_set():
|
||||
_signal_child_stop(child, getattr(parent, "_interrupt_message", None))
|
||||
else:
|
||||
with _quiet("Failed to propagate interrupt to fallback child: %s"):
|
||||
child.interrupt(getattr(parent, "_interrupt_message", None))
|
||||
|
||||
def _dispatch_background(batch: _Batch) -> str:
|
||||
"""Dispatch the call as independent async units (see ``_units_of``) and return the tool result JSON. Every unit
|
||||
of one call shares ONE pool slot (``slot_key``), so grouping never changes capacity accounting. Falls back to
|
||||
@@ -387,10 +402,6 @@ def _dispatch_background(batch: _Batch) -> str:
|
||||
|
||||
parent_agent = batch.parent_agent
|
||||
session_key, origin_ui_session_id = _resolve_async_session_key(parent_agent, batch.origin_ui_session_id)
|
||||
# The children's lifecycle is owned by the async registry now: drop them from the parent's
|
||||
# interrupt-propagation list (_build_child_agent attached them, which is correct for sync runs).
|
||||
for (_, _, c) in batch.children:
|
||||
_detach_child(parent_agent, c)
|
||||
routing = dict(
|
||||
session_key=session_key, origin_ui_session_id=origin_ui_session_id, origin_session_id=wake_sid,
|
||||
parent_session_id=getattr(parent_agent, "session_id", None), max_async_children=_get_max_async_children(),
|
||||
@@ -405,11 +416,16 @@ def _dispatch_background(batch: _Batch) -> str:
|
||||
# cache/delegation/live/<id>/; several units suffix it (-1, -2, ...) and the call keeps the bare id.
|
||||
unit_id = batch.live_deleg_id if len(units) == 1 else (f"{batch.live_deleg_id}-{k + 1}" if batch.live_deleg_id else None)
|
||||
unit.unit_id = unit_id = unit_id or _new_delegation_id() # fixed before the runner can start
|
||||
# The worker can start before admission returns. Detach only this unit:
|
||||
# unsubmitted units must still receive parent stops while a fallback runs.
|
||||
for _, _, child in unit.children:
|
||||
_detach_child(parent_agent, child)
|
||||
dispatch = _dispatch_unit(unit, unit_id, slot_key, routing)
|
||||
if dispatch.get("status") == "dispatched":
|
||||
slot_key = slot_key or dispatch["delegation_id"]
|
||||
dispatched.append((unit, dispatch["delegation_id"]))
|
||||
continue
|
||||
_restore_parent_cancellation(unit)
|
||||
if not dispatched:
|
||||
logger.info(
|
||||
"delegate_task: async pool at capacity (%s); running the whole batch synchronously instead.",
|
||||
@@ -419,7 +435,7 @@ def _dispatch_background(batch: _Batch) -> str:
|
||||
# Later units of an admitted call share its slot and cannot be capacity-rejected; a scheduler failure runs
|
||||
# the unit inline so no task is silently dropped.
|
||||
logger.warning("delegate_task: unit %d/%d not accepted (%s); running it inline.", k + 1, len(units), dispatch.get("error"))
|
||||
inline_results.extend(_execute_and_aggregate(unit, honor_parent_interrupt=False)["results"])
|
||||
inline_results.extend(_execute_and_aggregate(unit)["results"])
|
||||
payload = _dispatched_payload(batch, dispatched)
|
||||
if inline_results:
|
||||
payload["inline_results"] = inline_results
|
||||
|
||||
Reference in New Issue
Block a user