diff --git a/tests/tools/test_delegate_capacity_interrupt.py b/tests/tools/test_delegate_capacity_interrupt.py new file mode 100644 index 0000000000..071cc81343 --- /dev/null +++ b/tests/tools/test_delegate_capacity_interrupt.py @@ -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) diff --git a/tools/delegate_tool_dispatch.py b/tools/delegate_tool_dispatch.py index da75fd9d8a..7f768416b8 100644 --- a/tools/delegate_tool_dispatch.py +++ b/tools/delegate_tool_dispatch.py @@ -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//; 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