diff --git a/strix/core/execution.py b/strix/core/execution.py index c0cd1f865..15097b6cf 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -21,6 +21,7 @@ from openai import ( ) from strix.config import codex +from strix.core.agents import TERMINAL_STATUSES, Status from strix.core.hooks import ( BudgetExceededError, BudgetPausedError, @@ -47,7 +48,7 @@ if TYPE_CHECKING: from agents.memory import Session, SQLiteSession from agents.result import RunResultBase - from strix.core.agents import AgentCoordinator, Status + from strix.core.agents import AgentCoordinator logger = logging.getLogger(__name__) @@ -1104,6 +1105,8 @@ async def _start_child_runner( # spurious "Task exception was never retrieved" warning. The root agent # hits the same limit on its next call and tears the scan down. result = None + terminal_status: Status = "completed" + terminal_error = None try: result = await run_agent_loop( agent=child_agent, @@ -1119,18 +1122,49 @@ async def _start_child_runner( event_sink=event_sink, hooks=hooks, ) + except asyncio.CancelledError: + terminal_status = "stopped" + raise except BudgetExceededError: + terminal_status = "stopped" logger.info("child %s stopped after reaching the scan budget limit", child_id) except SubagentBudgetReservedError: + terminal_status = "stopped" logger.info("child %s stopped at the sub-agent budget reserve", child_id) + except Exception as error: + terminal_status = "crashed" + terminal_error = request_log.failure_text(error) + raise finally: - if on_complete is not None: - try: - await on_complete(result, session) - except Exception: - logger.exception("child %s completion delivery failed", child_id) - if not coordinator.is_shutting_down: - await _notify_parent_on_exit(coordinator, child_id) + try: + if on_complete is not None: + try: + await on_complete(result, session) + except asyncio.CancelledError: + terminal_status = "stopped" + raise + except Exception: + logger.exception("child %s completion delivery failed", child_id) + finally: + await _settle_child_exit( + coordinator, + child_id, + terminal_status, + terminal_error, + ) task_handle = asyncio.create_task(_child_loop(), name=f"agent-{name}-{child_id}") await coordinator.attach_runtime(child_id, task=task_handle) + + +async def _settle_child_exit( + coordinator: AgentCoordinator, + child_id: str, + terminal_status: Status, + terminal_error: str | None, +) -> None: + status = await _agent_status(coordinator, child_id) + if status not in TERMINAL_STATUSES: + await coordinator.set_status(child_id, terminal_status, error=terminal_error) + if not coordinator.is_shutting_down: + await _notify_parent_on_exit(coordinator, child_id) diff --git a/strix/fix/scan.py b/strix/fix/scan.py index 91b15ea6c..9b756bcf1 100644 --- a/strix/fix/scan.py +++ b/strix/fix/scan.py @@ -143,10 +143,7 @@ class ScanFixes: parent_history=[], ) except Exception as error: # noqa: BLE001 - report launch failure to the scan - active = self.tasks.get(finding_id) - if active and not active.done(): - active.cancel() - await asyncio.gather(active, return_exceptions=True) + await self._cancel_active(finding_id) logger.warning("fix.dispatch finding=%s rejected=%s", finding_id, error) await self.coordinator.send( self._parent_ctx["agent_id"], @@ -235,8 +232,7 @@ class ScanFixes: previous.get("reason") or "The previous Fix attempt failed; it was not restarted." ) if running and not running.done(): - running.cancel() - await asyncio.gather(running, return_exceptions=True) + await self._cancel_active(finding_id) used = int(previous.get("turns", 0)) if used >= 300: raise ValueError("This finding has exhausted its 300-turn Fix allowance.") @@ -421,10 +417,18 @@ class ScanFixes: for task in self.dispatches: task.cancel() await asyncio.gather(*self.dispatches, return_exceptions=True) - for task in self.tasks.values(): - if not task.done(): - task.cancel() - await asyncio.gather(*self.tasks.values(), return_exceptions=True) + for finding_id in list(self.tasks): + await self._cancel_active(finding_id) + + async def _cancel_active(self, finding_id: str) -> None: + task = self.tasks.get(finding_id) + if task is None or task.done(): + return + agent_id = self.records.get(finding_id, {}).get("agent_id") + if agent_id: + await self.coordinator.set_status(agent_id, "stopped") + task.cancel() + await asyncio.gather(task, return_exceptions=True) async def _stage_base(self, source: Path) -> str: key = hashlib.sha256(str(source).encode()).hexdigest()[:16] diff --git a/tests/test_execution.py b/tests/test_execution.py index 8fbf18ff5..1020879cf 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -5,10 +5,11 @@ from __future__ import annotations import asyncio import contextlib import json -from typing import Any, cast +from typing import TYPE_CHECKING, Any, cast from unittest.mock import MagicMock import pytest +from agents import RunConfig from agents.exceptions import MaxTurnsExceeded from agents.items import MessageOutputItem from agents.memory import SQLiteSession @@ -19,6 +20,7 @@ from strix.core import execution from strix.core.agents import AgentCoordinator from strix.core.execution import ( _notify_root_on_budget_reserve, + _start_child_runner, notify_parent_on_terminal, ) from strix.core.sessions import seed_initial_input @@ -26,6 +28,10 @@ from strix.tools.agents_graph.tools import agent_finish, stop_agent from strix.tools.finish.tool import finish_scan +if TYPE_CHECKING: + from pathlib import Path + + _NO_STREAM_EVENTS: list[Any] = [] @@ -143,6 +149,85 @@ async def test_concurrent_reserve_claims_yield_single_root() -> None: assert all(r is None for r in results if r != "root") +@pytest.mark.asyncio +async def test_cancelled_native_child_is_always_terminal( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "fix", parent_id="root") + started = asyncio.Event() + + async def run_forever(**_kwargs: Any) -> None: + started.set() + await asyncio.Event().wait() + + monkeypatch.setattr(execution, "run_agent_loop", run_forever) + sessions: list[SQLiteSession] = [] + await _start_child_runner( + parent_ctx={"agent_id": "root", "parent_id": None}, + coordinator=coordinator, + agents_db_path=tmp_path / "agents.sqlite", + sessions_to_close=sessions, + run_config=RunConfig(tracing_disabled=True), + max_turns=10, + interactive=False, + child_agent=MagicMock(), + child_id="child", + name="fix", + parent_id="root", + task="repair", + initial_input=[], + ) + await started.wait() + task = coordinator.runtimes["child"].task + assert task is not None + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + assert coordinator.statuses["child"] == "stopped" + for session in sessions: + session.close() + + +@pytest.mark.asyncio +async def test_crashed_native_child_is_always_terminal( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + await coordinator.register("child", "fix", parent_id="root") + + async def crash(**_kwargs: Any) -> None: + raise RuntimeError("repair crashed") + + monkeypatch.setattr(execution, "run_agent_loop", crash) + sessions: list[SQLiteSession] = [] + await _start_child_runner( + parent_ctx={"agent_id": "root", "parent_id": None}, + coordinator=coordinator, + agents_db_path=tmp_path / "agents.sqlite", + sessions_to_close=sessions, + run_config=RunConfig(tracing_disabled=True), + max_turns=10, + interactive=False, + child_agent=MagicMock(), + child_id="child", + name="fix", + parent_id="root", + task="repair", + initial_input=[], + ) + task = coordinator.runtimes["child"].task + assert task is not None + await asyncio.gather(task, return_exceptions=True) + + assert coordinator.statuses["child"] == "crashed" + assert coordinator.errors["child"] == "repair crashed" + for session in sessions: + session.close() + + @pytest.mark.asyncio async def test_claim_reserve_sets_flag_and_wakes_parked_agents() -> None: coordinator = AgentCoordinator() diff --git a/tests/test_scan_fixes.py b/tests/test_scan_fixes.py index d2bdd1c76..9ed9d53d7 100644 --- a/tests/test_scan_fixes.py +++ b/tests/test_scan_fixes.py @@ -7,6 +7,7 @@ from __future__ import annotations import asyncio from pathlib import Path from types import SimpleNamespace +from typing import Any from unittest.mock import AsyncMock, Mock import pytest @@ -17,7 +18,11 @@ from agents.tool_context import ToolContext from strix.core.agents import AgentCoordinator from strix.core.execution import spawn_child_agent from strix.core.hooks import ReportUsageHooks -from strix.fix import FindingContext +from strix.fix import ( + FindingContext, + FixPreparationResultV1, + PreparationState, +) from strix.fix import scan as scan_module from strix.fix.scan import ScanFixes from strix.fix.session import WorktreeSession @@ -188,12 +193,27 @@ async def test_finding_revision_invalidates_active_completion(tmp_path): assert not fixes._current("finding", digest) +@pytest.mark.asyncio +async def test_superseded_fix_agent_is_marked_stopped_before_cancellation(tmp_path): + fixes, _, _, _, _, _, _ = setup(tmp_path) + await fixes.coordinator.register("old-fix", "Fix", "reporter", skills=["fix_task"]) + await fixes.coordinator.mark_running("old-fix") + task = asyncio.create_task(asyncio.Event().wait()) + fixes.records["finding"] = {"agent_id": "old-fix", "status": "running"} + fixes.tasks["finding"] = task + + await fixes._cancel_active("finding") + + assert fixes.coordinator.statuses["old-fix"] == "stopped" + assert task.cancelled() + + @pytest.mark.parametrize("change", ["revised", "withdrawn", "unconfirmed"]) async def test_finding_changed_before_delivery_discards_reviewed_patch( tmp_path, monkeypatch, change ): fixes, report, _, _, reports, context, _ = setup(tmp_path) - callback = None + callback: Any = None async def spawn(**kwargs): nonlocal callback @@ -211,8 +231,8 @@ async def test_finding_changed_before_delivery_discards_reviewed_patch( reports.clear() else: report["validation_status"] = "unconfirmed" - return scan_module.FixPreparationResultV1( - state="ready", + return FixPreparationResultV1( + state=PreparationState.READY, stop_reason="Approved.", source_identity=request.candidate.source_identity, candidate=request.candidate,