From 098fa3659c06b068f6a063c704a050bf5061e7f9 Mon Sep 17 00:00:00 2001 From: Jonathan Singer Date: Thu, 1 Oct 2026 14:45:24 -0400 Subject: [PATCH] fix: guard resumed assessments and ignore withdrawn fix records --- strix/core/execution.py | 10 ++++++++++ strix/fix/scan.py | 4 +++- tests/test_execution.py | 32 ++++++++++++++++++++++++++++++++ tests/test_runner_teardown.py | 30 ++++++++++++++++++++++++++---- tests/test_scan_fixes.py | 35 ++++++++++++++++++++++++++++++++--- 5 files changed, 103 insertions(+), 8 deletions(-) diff --git a/strix/core/execution.py b/strix/core/execution.py index e308ffa5f..a9f3a72cd 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -251,6 +251,16 @@ async def _run_agent_loop( # noqa: PLR0912 - interactive completion and cancell ) result: RunResultBase | None = None + if ( + interactive + and return_on_completion + and await _agent_status(coordinator, agent_id) == "completed" + ): + # A restored assessment can already be final while fixes are pending. + # Finalize it without another model cycle or reopening the report. + await coordinator.attach_runtime(agent_id, resumable=False) + return result + first_cycle_input = await _seed_and_prepare_first_input( session, initial_input, start_parked=start_parked ) diff --git a/strix/fix/scan.py b/strix/fix/scan.py index 9c44d002b..5e4859eb1 100644 --- a/strix/fix/scan.py +++ b/strix/fix/scan.py @@ -441,7 +441,9 @@ class ScanFixes: for report in self.report_state.get_existing_vulnerabilities() } for finding_id, record in sorted(self.records.items()): - title = titles.get(finding_id, finding_id) + if finding_id not in titles: + continue + title = titles[finding_id] if record.get("status") != "done" or not record.get("artifact"): errors.append( { diff --git a/tests/test_execution.py b/tests/test_execution.py index 381e36e0b..911dd34a0 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -1056,6 +1056,38 @@ async def test_run_agent_loop_seeds_identity_before_first_cycle( session.close() +@pytest.mark.asyncio +@pytest.mark.parametrize("start_parked", [False, True]) +async def test_completed_interactive_assessment_resumes_without_another_cycle( + monkeypatch: pytest.MonkeyPatch, start_parked: bool +) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "Strix", parent_id=None) + await coordinator.set_status("root", "completed") + restored = AgentCoordinator() + await restored.restore(await coordinator.snapshot()) + + async def unexpected_cycle(*_args: Any, **_kwargs: Any) -> None: + raise AssertionError("A completed assessment must not call the model on resume") + + monkeypatch.setattr(execution, "_run_until_lifecycle", unexpected_cycle) + result = await execution.run_agent_loop( + agent=MagicMock(), + initial_input=[], + run_config=MagicMock(), + context={"agent_id": "root", "parent_id": None}, + max_turns=5, + coordinator=restored, + agent_id="root", + interactive=True, + start_parked=start_parked, + return_on_completion=True, + ) + assert result is None + assert restored.statuses["root"] == "completed" + assert not await restored.send("root", {"from": "user", "content": "too late"}) + + @pytest.mark.asyncio async def test_interactive_completion_returns_for_finalization_but_waiting_can_resume( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_runner_teardown.py b/tests/test_runner_teardown.py index 63df28b33..ccede4109 100644 --- a/tests/test_runner_teardown.py +++ b/tests/test_runner_teardown.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import json import types from typing import TYPE_CHECKING, Any @@ -9,7 +10,7 @@ from agents import ModelSettings import strix.tools.notes.tools as notes_tools import strix.tools.todo.tools as todo_tools -from strix.core import runner +from strix.core import execution, runner from strix.core.agents import AgentCoordinator from strix.runtime import session_manager from tests.test_fix_reliability import LocalSandbox @@ -104,10 +105,15 @@ async def test_a_live_child_is_settled_before_sessions_close( @pytest.mark.asyncio @pytest.mark.parametrize( - "interactive,local_branches", [(False, False), (False, True), (True, True)] + "interactive,local_branches,is_resume", + [(False, False, False), (False, True, False), (True, True, False), (True, True, True)], ) async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown( - monkeypatch: pytest.MonkeyPatch, tmp_path: Path, interactive: bool, local_branches: bool + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + interactive: bool, + local_branches: bool, + is_resume: bool, ) -> None: _wire_runner(monkeypatch, tmp_path) events: list[str] = [] @@ -121,6 +127,9 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown( def get_existing_vulnerabilities(self) -> list[Any]: return [] + def get_total_llm_cost(self) -> float: + return 0.0 + def save_run_data(self, **_: Any) -> None: events.append("save") @@ -153,9 +162,22 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown( async def cleanup(*_: Any) -> None: events.append("sandbox deleted") + if is_resume: + coordinator = AgentCoordinator() + await coordinator.register("root", "Root Agent", parent_id=None) + await coordinator.set_status("root", "completed") + (tmp_path / "agents.json").write_text(json.dumps(await coordinator.snapshot())) + (tmp_path / "agents.db").touch() + + async def unexpected_cycle(*_args: Any, **_kwargs: Any) -> None: + raise AssertionError("Resume must finish pending fixes without restarting root") + + monkeypatch.setattr(execution, "_run_until_lifecycle", unexpected_cycle) + monkeypatch.setattr(runner, "get_global_report_state", State) monkeypatch.setattr(runner, "ScanFixes", Fixes) - monkeypatch.setattr(runner, "run_agent_loop", root) + if not is_resume: + monkeypatch.setattr(runner, "run_agent_loop", root) monkeypatch.setattr(session_manager, "cleanup", cleanup) await runner.run_strix_scan( scan_config={ diff --git a/tests/test_scan_fixes.py b/tests/test_scan_fixes.py index 615e24d5d..4ba10c9d4 100644 --- a/tests/test_scan_fixes.py +++ b/tests/test_scan_fixes.py @@ -228,10 +228,11 @@ async def test_blocked_native_child_has_no_patch(tmp_path, monkeypatch): @pytest.mark.asyncio +@pytest.mark.parametrize("withdrawn", [False, True]) async def test_cancelled_verification_is_reported_without_publishing_a_branch( - tmp_path, monkeypatch + tmp_path, monkeypatch, withdrawn ): - fixes, _, source, _, _, context, sessions = setup(tmp_path, interactive=True) + fixes, _, source, _, reports, context, sessions = setup(tmp_path, interactive=True) fixes.publish_local_branches = True reviewing = asyncio.Event() @@ -251,10 +252,15 @@ async def test_cancelled_verification_is_reported_without_publishing_a_branch( ) assert (await delegate(context))["success"] await asyncio.wait_for(reviewing.wait(), timeout=5) + if withdrawn: + reports.clear() await fixes.close() branches, errors = await fixes.wait() assert branches == [] - assert "interrupted" in errors[0]["error"] + if withdrawn: + assert errors == [] + else: + assert "interrupted" in errors[0]["error"] assert fixes.records["finding"]["status"] == "stopped" assert _git(source, "branch", "--list", "strix/fix-*") == "" assert not list((tmp_path / "state/fixes").glob("*/prepared-fix.zip")) @@ -262,6 +268,29 @@ async def test_cancelled_verification_is_reported_without_publishing_a_branch( session.close() +@pytest.mark.asyncio +@pytest.mark.parametrize("status", ["running", "stopped", "failed", "done"]) +async def test_wait_omits_withdrawn_records_but_keeps_current_fix_errors(tmp_path, status): + fixes, _, source, _, _, _, _ = setup(tmp_path) + fixes.publish_local_branches = True + digest = fixes._finding("finding")[1].digest() + fixes.records = { + "finding": {"digest": digest, "status": "stopped", "reason": "Tests failed"}, + "withdrawn": { + "digest": digest, + "status": status, + "artifact": str(tmp_path / "withdrawn.zip"), + "reason": "The finding was withdrawn", + }, + } + + branches, errors = await fixes.wait() + + assert branches == [] + assert errors == [{"finding_id": "finding", "title": "finding", "error": "Tests failed"}] + assert _git(source, "branch", "--list", "strix/fix-*") == "" + + @pytest.mark.asyncio async def test_dispatch_failure_is_reported_before_an_agent_exists(tmp_path, monkeypatch): fixes, _, _, _, _, context, _ = setup(tmp_path)