mirror of
https://github.com/usestrix/strix.git
synced 2026-10-02 02:13:43 +00:00
fix: finalize cancelled fix agents
This commit is contained in:
parent
60d4ce15e5
commit
3763a677f9
4 changed files with 166 additions and 23 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue