fix: finalize cancelled fix agents

This commit is contained in:
yoni 2026-10-01 13:42:30 +00:00
parent 60d4ce15e5
commit 3763a677f9
4 changed files with 166 additions and 23 deletions

View file

@ -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)

View file

@ -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]

View file

@ -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()

View file

@ -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,