fix: complete interactive autofix scans before cleanup

This commit is contained in:
Jonathan Singer 2026-10-01 14:23:56 -04:00
parent 1fa7d211b4
commit f14468565a
17 changed files with 467 additions and 30 deletions

View file

@ -199,6 +199,7 @@ async def run_agent_loop(
interactive: bool,
session: Session | None = None,
start_parked: bool = False,
return_on_completion: bool = False,
event_sink: StreamEventSink | None = None,
hooks: RunHooks[dict[str, Any]] | None = None,
) -> RunResultBase | None:
@ -218,6 +219,7 @@ async def run_agent_loop(
interactive=interactive,
session=session,
start_parked=start_parked,
return_on_completion=return_on_completion,
event_sink=event_sink,
hooks=hooks,
)
@ -225,7 +227,7 @@ async def run_agent_loop(
request_log.reset_call_context(token)
async def _run_agent_loop(
async def _run_agent_loop( # noqa: PLR0912 - interactive completion and cancellation differ
*,
agent: Any,
initial_input: Any,
@ -237,6 +239,7 @@ async def _run_agent_loop(
interactive: bool,
session: Session | None,
start_parked: bool,
return_on_completion: bool,
event_sink: StreamEventSink | None,
hooks: RunHooks[dict[str, Any]] | None,
) -> RunResultBase | None:
@ -284,10 +287,17 @@ async def _run_agent_loop(
return result
while True:
if return_on_completion and await _agent_status(coordinator, agent_id) == "completed":
# The assessment is final. Let the caller await fixes and publish
# branches while the interactive UI remains open to display them.
await coordinator.attach_runtime(agent_id, resumable=False)
return result
timeout = await _plain_waiting_timeout(coordinator, agent_id)
try:
woke = await coordinator.wait_for_message(agent_id, timeout=timeout)
except asyncio.CancelledError:
if return_on_completion:
raise
return result
if coordinator.budget_stopped:

View file

@ -625,7 +625,13 @@ async def run_strix_scan(
agent_id=root_id,
interactive=interactive,
session=root_session,
start_parked=bool(interactive and is_resume and root_status != "running"),
return_on_completion=fixes is not None,
start_parked=bool(
interactive
and is_resume
and root_status != "running"
and not (fixes is not None and root_status == "completed")
),
event_sink=event_sink,
hooks=hooks,
)

View file

@ -61,7 +61,9 @@ class ScanFixes:
self.directory.mkdir(parents=True, exist_ok=True, mode=0o700)
self.directory.chmod(0o700)
self.path = self.directory / "tasks.json"
self.records = json.loads(self.path.read_text()) if self.path.exists() else {}
self.records: dict[str, dict[str, Any]] = (
json.loads(self.path.read_text()) if self.path.exists() else {}
)
self.sources = [
Path(s["source_path"]).resolve()
for s in local_sources
@ -135,8 +137,9 @@ class ScanFixes:
await asyncio.gather(*pending, return_exceptions=True)
async def _dispatch(self, finding_id: str) -> None:
candidate = None
try:
report, _ = self._finding(finding_id)
report, candidate = self._finding(finding_id)
parent_id = report.get("agent_id") or self._parent_ctx["agent_id"]
await self.spawn(
finding_id,
@ -149,6 +152,10 @@ class ScanFixes:
)
except Exception as error: # noqa: BLE001 - report launch failure to the scan
await self._cancel_active(finding_id)
if candidate is not None and self.publish_local_branches:
record = self.records.setdefault(finding_id, {})
record.update(digest=candidate.digest(), status="stopped", reason=str(error))
self._save()
logger.warning("fix.dispatch finding=%s rejected=%s", finding_id, error)
await self.coordinator.send(
self._parent_ctx["agent_id"],
@ -346,6 +353,10 @@ class ScanFixes:
record["reason"] = prepared.stop_reason
if prepared.state == "ready":
record["artifact"] = str(artifact)
except asyncio.CancelledError:
record["status"] = "stopped"
record["reason"] = "Fix preparation was interrupted before completion."
raise
except Exception as error:
record["status"] = "stopped"
record["reason"] = str(error)
@ -363,6 +374,7 @@ class ScanFixes:
**kwargs["parent_ctx"],
"sandbox_session": borrowed,
"before_agent_finish": hooks.before_finish,
"interactive": False,
}
assignment = _untrusted_prompt_data(
{
@ -379,6 +391,9 @@ class ScanFixes:
"parent_ctx": parent_ctx,
"task": kwargs["task"] + "\n\n" + assignment,
"skills": ["fix_task"],
# Fix completion starts the verifier; it must not park for
# terminal messages after agent_finish in an interactive scan.
"interactive": False,
"factory": lambda **kw: build_fix_agent(name=kw["name"], workspace_root=root),
"run_config": _run_config(environment),
"hooks": hooks,
@ -421,10 +436,23 @@ class ScanFixes:
branches: list[dict[str, str]] = []
errors: list[dict[str, str]] = []
titles: dict[str, str] = {
str(report["id"]): str(report.get("title") or report["id"])
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 record.get("status") != "done" or not record.get("artifact"):
errors.append(
{
"finding_id": finding_id,
"title": title,
"error": str(
record.get("reason") or "Fix did not produce a verified artifact."
),
}
)
continue
title = finding_id
try:
report, candidate = self._finding(finding_id)
assert candidate.source_identity is not None

View file

@ -8,6 +8,7 @@ import asyncio
import contextlib
import sys
from pathlib import Path
from typing import cast
from rich.console import Console
from rich.panel import Panel
@ -264,6 +265,28 @@ def display_completion_message(args: argparse.Namespace, results_path: Path) ->
if stats_text.plain:
panel_parts.extend(["\n", stats_text])
results = (report_state.scan_results or {}) if report_state is not None else {}
for item in cast("list[dict[str, str]]", results.get("fix_branches") or []):
panel_parts.extend(
[
"\n\n",
Text(f"Fix {item['title']}", style="bold #22c55e"),
"\n",
Text(f" {item['branch']}"),
"\n",
Text(f" {item['source_path']}", style="dim"),
]
)
for item in cast("list[dict[str, str]]", results.get("fix_branch_errors") or []):
panel_parts.extend(
[
"\n\n",
Text(f"Fix unavailable {item['title']}", style="bold #eab308"),
"\n",
Text(f" {item['error']}"),
]
)
results_text = Text()
results_text.append("\n")
results_text.append("Output", style="dim")

View file

@ -67,6 +67,13 @@ class TuiLiveView(BaseLiveView):
current["updated_at"] = now
return changed
def record_runtime_message(self, agent_id: str, content: str) -> None:
self._append_event(
agent_id,
"chat",
{"role": "assistant", "content": content, "metadata": {"source": "runtime"}},
)
def _append_event(
self,
agent_id: str,

View file

@ -1104,6 +1104,7 @@ func TestTerminalSnapshotWithoutAgentsDoesNotKeepLoading(t *testing.T) {
{state: "stopped", want: "Scan stopped"},
{state: "completed", want: "Scan completed"},
{state: "preparing", want: "Preparing scan..."},
{state: "preparing_fixes", want: "Fixes in progress"},
}
for _, tt := range tests {
@ -1126,6 +1127,16 @@ func TestTerminalSnapshotWithoutAgentsDoesNotKeepLoading(t *testing.T) {
}
}
func TestCompletedRootShowsPendingFixStatus(t *testing.T) {
model := New(nil)
model.snapshot.ScanState = "preparing_fixes"
model.snapshot.Agents = []protocol.Agent{{ID: "root", Status: "completed"}}
out := model.statusView(100)
if !strings.Contains(out, "Fixes in progress") || strings.Contains(out, "Agent completed") {
t.Fatalf("premature completion status: %s", out)
}
}
func TestCrashedAndBudgetPausedAgentStatusParity(t *testing.T) {
model := New(nil)
model.width = 100

View file

@ -80,6 +80,8 @@ func (m *Model) chatContent() string {
return centeredPlaceholder("Scan completed", m.viewport.Width, m.viewport.Height)
case "preparing":
return centeredPlaceholder("Preparing scan...", m.viewport.Width, m.viewport.Height)
case "preparing_fixes":
return centeredPlaceholder("Assessment complete · Fixes in progress", m.viewport.Width, m.viewport.Height)
default:
return centeredPlaceholder("Loading...", m.viewport.Width, m.viewport.Height)
}
@ -828,6 +830,11 @@ func (m Model) statusView(width int) string {
left = statusMessage(msg, red, " · Send message to resume", width)
}
}
if m.snapshot.ScanState == "preparing_fixes" {
left = lipgloss.NewStyle().Foreground(amber).Render("Assessment complete · Fixes in progress")
} else if m.snapshot.ScanState == "completed" {
left = lipgloss.NewStyle().Foreground(mid).Render("Scan completed")
}
if m.errorText != "" {
left = statusMessage(m.errorText, red, "", width-lipgloss.Width(right))
}

View file

@ -97,7 +97,7 @@ func Tool(data map[string]any) string {
case "respond_to_user":
return renderRespondToUser(args)
case "finish_scan":
return renderFinishScan(args)
return renderFinishScan(args, result, status)
case "think":
return renderThink(args)
case "web_search":

View file

@ -33,6 +33,16 @@ func TestChatUserMessage(t *testing.T) {
requireContains(t, out, "You:", "hello", "world")
}
func TestFinishScanDoesNotClaimPendingFixesAreComplete(t *testing.T) {
for _, result := range []any{map[string]any{"fixes_pending": true}, `{"fixes_pending":true}`} {
out := Tool(tool("finish_scan", map[string]any{"executive_summary": "Assessment"}, result, "completed"))
requireContains(t, out, "Assessment complete", "Fixes in progress")
if strings.Contains(out, "Penetration test completed") {
t.Fatalf("pending fixes rendered as complete: %s", out)
}
}
}
func TestChatAssistantMarkdown(t *testing.T) {
out := Chat(map[string]any{"role": "assistant", "content": "# Heading\n\nSome **bold** text"})
requireContains(t, out, "Heading", "bold")

View file

@ -1,6 +1,7 @@
package render
import (
"encoding/json"
"strings"
)
@ -8,9 +9,21 @@ import (
// Finish scan (finish_renderer.py)
// ---------------------------------------------------------------------------
func renderFinishScan(args map[string]any) string {
func renderFinishScan(args map[string]any, result any, status string) string {
var b strings.Builder
b.WriteString(Col(Green).Render("◆ ") + Bold(Green).Render("Penetration test completed"))
label := "Penetration test completed"
resultMap, _ := resultMapOf(result)
if raw, ok := result.(string); ok {
_ = json.Unmarshal([]byte(raw), &resultMap)
}
if status != "completed" {
label = "Finalizing assessment"
} else if truthy(resultMap["fixes_pending"]) {
label = "Assessment complete · Fixes in progress"
} else if success, ok := resultMap["success"].(bool); ok && !success {
label = "Assessment not completed"
}
b.WriteString(Col(Green).Render("◆ ") + Bold(Green).Render(label))
section := func(label, value string) {
if value != "" {
b.WriteString("\n\n" + Bold(Field).Render(label) + "\n" + value)

View file

@ -258,11 +258,12 @@ class GoTuiRuntime:
if self.report_state is not None:
results = self.report_state.scan_results or {}
for item in results.get("fix_branches") or []:
self.controller.add_message(
f"Prepared fix branch for {item['title']}: {item['branch']}"
self._show_fix_result(
f"Prepared fix branch for {item['title']}: {item['branch']} "
f"in {item['source_path']}"
)
for item in results.get("fix_branch_errors") or []:
self.controller.add_message(
self._show_fix_result(
f"Could not create fix branch for {item['title']}: {item['error']}",
"error",
)
@ -288,6 +289,21 @@ class GoTuiRuntime:
await self._sync_agent_state()
self.controller.notify_changed()
def _show_fix_result(self, text: str, level: str = "info") -> None:
self.controller.add_message(text, level)
root_id = next(
(
agent_id
for agent_id, agent in self.live_view.agents.items()
if agent.get("parent_id") is None
),
None,
)
if root_id is not None:
# Controller messages are setup-only in the Go UI. Publish final
# fixes into the assessment transcript so they are visible live.
self.live_view.record_runtime_message(root_id, text)
def capture_event(self, agent_id: str, event: Any) -> None:
self.live_view.ingest_sdk_event(agent_id, event)
if getattr(getattr(event, "item", None), "type", "") == "tool_call_output_item":
@ -359,9 +375,16 @@ class GoTuiRuntime:
scan_state = "completed"
elif root_status == "stopped":
scan_state = "stopped"
elif root_status == "completed":
scan_state = "failed"
self.controller.error = "Scan ended without a completed report"
elif root_status == "completed" and scan_state != "stopped":
fixes_pending = bool(
self.report_state is not None
and self.report_state.defer_completion
and self.report_state.scan_results
)
scan_state = "preparing_fixes" if fixes_pending else "failed"
self.controller.error = (
None if fixes_pending else "Scan ended without a completed report"
)
if scan_state != self.controller.scan_state:
self.controller.scan_state = scan_state
changed = True

View file

@ -79,6 +79,9 @@ def _do_finish(
"message": "Scan completed successfully",
"vulnerabilities_found": vuln_count,
}
if report_state.defer_completion:
result["fixes_pending"] = True
result["message"] = "Assessment complete; fixes are still being prepared."
result.update(coverage_summary)
return result

View file

@ -1056,6 +1056,125 @@ async def test_run_agent_loop_seeds_identity_before_first_cycle(
session.close()
@pytest.mark.asyncio
async def test_interactive_completion_returns_for_finalization_but_waiting_can_resume(
monkeypatch: pytest.MonkeyPatch,
) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "Strix", parent_id=None)
parked = asyncio.Event()
calls = []
async def cycle(*_args: Any, **kwargs: Any) -> Any:
assert kwargs["interactive"] is True
calls.append(kwargs["initial_input"])
await coordinator.set_status("root", "waiting" if len(calls) == 1 else "completed")
parked.set()
return MagicMock(final_output={"scan_completed": len(calls) == 2})
monkeypatch.setattr(execution, "_run_until_lifecycle", cycle)
task = asyncio.create_task(
execution.run_agent_loop(
agent=MagicMock(),
initial_input="task",
run_config=MagicMock(),
context={"agent_id": "root", "parent_id": None},
max_turns=5,
coordinator=coordinator,
agent_id="root",
interactive=True,
return_on_completion=True,
)
)
try:
await asyncio.wait_for(parked.wait(), timeout=2)
assert not task.done()
assert await coordinator.send("root", {"from": "user", "content": "continue"})
result = await asyncio.wait_for(task, timeout=2)
assert result.final_output == {"scan_completed": True}
assert calls == ["task", []]
assert not await coordinator.send("root", {"from": "user", "content": "too late"})
finally:
task.cancel()
await asyncio.gather(task, return_exceptions=True)
@pytest.mark.asyncio
async def test_cancelled_interactive_assessment_does_not_return_as_completed(
monkeypatch: pytest.MonkeyPatch,
) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "Strix", parent_id=None)
waiting = asyncio.Event()
async def cycle(*_args: Any, **_kwargs: Any) -> None:
await coordinator.set_status("root", "waiting")
async def wait(*_args: Any, **_kwargs: Any) -> bool:
waiting.set()
await asyncio.Event().wait()
return True
monkeypatch.setattr(execution, "_run_until_lifecycle", cycle)
monkeypatch.setattr(coordinator, "wait_for_message", wait)
task = asyncio.create_task(
execution.run_agent_loop(
agent=MagicMock(),
initial_input="task",
run_config=MagicMock(),
context={"agent_id": "root", "parent_id": None},
max_turns=5,
coordinator=coordinator,
agent_id="root",
interactive=True,
return_on_completion=True,
)
)
await asyncio.wait_for(waiting.wait(), timeout=2)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
@pytest.mark.asyncio
async def test_default_interactive_completion_still_parks_for_messages(
monkeypatch: pytest.MonkeyPatch,
) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "Strix", parent_id=None)
waiting = asyncio.Event()
result = MagicMock(final_output={"scan_completed": True})
async def cycle(*_args: Any, **_kwargs: Any) -> Any:
await coordinator.set_status("root", "completed")
return result
async def wait(*_args: Any, **_kwargs: Any) -> bool:
waiting.set()
await asyncio.Event().wait()
return True
monkeypatch.setattr(execution, "_run_until_lifecycle", cycle)
monkeypatch.setattr(coordinator, "wait_for_message", wait)
task = asyncio.create_task(
execution.run_agent_loop(
agent=MagicMock(),
initial_input="task",
run_config=MagicMock(),
context={"agent_id": "root", "parent_id": None},
max_turns=5,
coordinator=coordinator,
agent_id="root",
interactive=True,
)
)
await asyncio.wait_for(waiting.wait(), timeout=2)
assert not task.done()
assert (await coordinator.reachability("root"))[0]
task.cancel()
assert await task is result
def _scripted_cycle(
coordinator: AgentCoordinator,
agent_id: str,

View file

@ -0,0 +1,71 @@
"""Assessment completion must not claim that pending fixes are ready."""
from __future__ import annotations
import argparse
import importlib
from io import StringIO
from pathlib import Path
from types import SimpleNamespace
from typing import Any
import pytest
from rich.console import Console
from strix.report import state
from strix.tools.finish.tool import _do_finish
@pytest.mark.parametrize("pending", [False, True])
def test_finish_result_distinguishes_assessment_from_fix_completion(
monkeypatch: pytest.MonkeyPatch, pending: bool
) -> None:
report = SimpleNamespace(
defer_completion=pending,
vulnerability_reports=[],
update_scan_final_fields=lambda **_: None,
)
monkeypatch.setattr(state, "get_global_report_state", lambda: report)
result = _do_finish(
parent_id=None,
executive_summary="Summary",
methodology="Method",
technical_analysis="Analysis",
recommendations="Recommendations",
agent_graph={},
)
assert result["scan_completed"] is True
assert bool(result.get("fixes_pending")) is pending
if pending:
assert "still being prepared" in result["message"]
def test_final_summary_shows_branches_and_unavailable_fixes(
monkeypatch: pytest.MonkeyPatch,
) -> None:
main: Any = importlib.import_module("strix.interface.main")
report = SimpleNamespace(
run_record={"status": "completed"},
scan_results={
"fix_branches": [
{"title": "Export", "branch": "strix/fix-export", "source_path": "/workspace/repo"}
],
"fix_branch_errors": [{"title": "Import", "error": "Reviewer rejected the fix"}],
},
)
output = StringIO()
monkeypatch.setattr(state, "get_global_report_state", lambda: report)
monkeypatch.setattr(main, "Console", lambda: Console(file=output, width=120, color_system=None))
monkeypatch.setattr(main, "build_final_stats_text", lambda _: main.Text())
main.display_completion_message(
argparse.Namespace(
targets_info=[{"original": "/workspace/repo"}],
run_name="scan",
non_interactive=True,
),
Path("/workspace/results"),
)
assert "strix/fix-export" in output.getvalue()
assert "/workspace/repo" in output.getvalue()
assert "Fix unavailable" in output.getvalue()
assert "Reviewer rejected the fix" in output.getvalue()

View file

@ -958,6 +958,40 @@ async def test_agent_state_sync_projects_completed_report() -> None:
assert runtime.controller.scan_state == "completed"
@pytest.mark.asyncio
async def test_fix_result_is_visible_in_the_live_assessment_transcript() -> None:
runtime = GoTuiRuntime(args())
await runtime.coordinator.register("root", "Strix", parent_id=None)
await runtime._sync_agent_state()
runtime._show_fix_result("Prepared fix branch: strix/fix-export in /workspace/repo")
runtime._show_fix_result("Fix unavailable: independent review rejected it", "error")
transcript = str(runtime.live_view.events_for_agent("root"))
assert "strix/fix-export" in transcript
assert "/workspace/repo" in transcript
assert "independent review rejected it" in transcript
@pytest.mark.asyncio
async def test_agent_state_sync_waits_for_fixes_before_showing_completion() -> None:
runtime = GoTuiRuntime(args())
runtime.report_state = cast(
"Any",
SimpleNamespace(
run_record={"status": "running"},
scan_results={"scan_completed": True},
defer_completion=True,
),
)
await runtime.coordinator.register("root", "Strix", parent_id=None)
await runtime.coordinator.set_status("root", "completed")
await runtime._sync_agent_state()
assert runtime.controller.scan_state == "preparing_fixes"
assert runtime.controller.error is None
runtime.report_state.run_record["status"] = "completed"
await runtime._sync_agent_state()
assert runtime.controller.scan_state == "completed"
@pytest.mark.asyncio
async def test_agent_state_sync_does_not_mask_root_failure_with_completed_report() -> None:
runtime = GoTuiRuntime(args())

View file

@ -103,8 +103,11 @@ 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)]
)
async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
monkeypatch: pytest.MonkeyPatch, tmp_path: Path, interactive: bool, local_branches: bool
) -> None:
_wire_runner(monkeypatch, tmp_path)
events: list[str] = []
@ -122,8 +125,9 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(
events.append("save")
class Fixes:
def __init__(self, **_: Any) -> None:
pass
def __init__(self, **options: Any) -> None:
assert options["publish_local_branches"] is local_branches
assert (options["sink"] is not None) is not local_branches
def start(self, *_: Any) -> None:
events.append("fixes listening")
@ -138,7 +142,12 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(
async def assessment(_: Any) -> None:
events.append("assessment published")
async def root(**_: Any) -> types.SimpleNamespace:
async def platform_fix_sink(*_: Any) -> bool:
return True
async def root(**kwargs: Any) -> types.SimpleNamespace:
assert kwargs["interactive"] is interactive
assert kwargs["return_on_completion"] is True
return types.SimpleNamespace(final_output={"scan_completed": True})
async def cleanup(*_: Any) -> None:
@ -149,19 +158,26 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(
monkeypatch.setattr(runner, "run_agent_loop", root)
monkeypatch.setattr(session_manager, "cleanup", cleanup)
await runner.run_strix_scan(
scan_config={"targets": [], "scan_mode": "deep"},
scan_config={
"targets": [],
"scan_mode": "deep",
"local_fix_branches_enabled": local_branches,
},
scan_id="scan",
image="image",
local_sources=[{"source_path": str(tmp_path)}],
assessment_sink=assessment,
fix_sink=None if local_branches else platform_fix_sink,
interactive=interactive,
)
assert events.index("assessment published") < events.index("fixes finished")
assert events.index("fixes finished") < events.index("sandbox deleted")
@pytest.mark.asyncio
@pytest.mark.parametrize("interactive", [False, True])
async def test_disabled_one_click_fixes_skips_fix_runtime(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
monkeypatch: pytest.MonkeyPatch, tmp_path: Path, interactive: bool
) -> None:
_wire_runner(monkeypatch, tmp_path)
events: list[str] = []
@ -185,7 +201,8 @@ async def test_disabled_one_click_fixes_skips_fix_runtime(
async def assessment(_: Any) -> None:
events.append("assessment published")
async def root(**_: Any) -> types.SimpleNamespace:
async def root(**kwargs: Any) -> types.SimpleNamespace:
assert kwargs["return_on_completion"] is False
return types.SimpleNamespace(final_output={"scan_completed": True})
monkeypatch.setattr(runner, "get_global_report_state", State)
@ -201,6 +218,7 @@ async def test_disabled_one_click_fixes_skips_fix_runtime(
image="image",
local_sources=[{"source_path": str(tmp_path)}],
assessment_sink=assessment,
interactive=interactive,
)
assert events == ["assessment published", "save"]

View file

@ -32,7 +32,7 @@ from tests.test_fix_reliability import LocalSandbox, existing_suite
from tests.test_fix_runtime import _git, _request, _workspace
def setup(tmp_path):
def setup(tmp_path, *, interactive=False):
source, _ = _workspace(tmp_path)
commit = existing_suite(source)
parent = LocalSandbox(tmp_path / "sandbox")
@ -62,8 +62,7 @@ def setup(tmp_path):
coordinator=coordinator,
agents_db_path=tmp_path / "agents.db",
sessions_to_close=sessions,
interactive=False,
**kwargs,
**{"interactive": interactive, **kwargs},
)
async def spawn(**kwargs):
@ -146,8 +145,11 @@ async def test_native_parallel_fixes_deliver_patches_and_preserve_scan_source(
@pytest.mark.asyncio
async def test_ready_fix_creates_local_branch_without_changing_checkout(tmp_path, monkeypatch):
fixes, _, source, _, _, context, sessions = setup(tmp_path)
@pytest.mark.parametrize("interactive", [False, True])
async def test_ready_fix_creates_local_branch_without_changing_checkout(
tmp_path, monkeypatch, interactive
):
fixes, _, source, _, _, context, sessions = setup(tmp_path, interactive=interactive)
fixes.publish_local_branches = True
original_head = _git(source, "rev-parse", "HEAD")
original_branch = _git(source, "branch", "--show-current")
@ -162,7 +164,7 @@ async def test_ready_fix_creates_local_branch_without_changing_checkout(tmp_path
)
assert (await delegate(context))["success"]
branches, errors = await fixes.wait()
branches, errors = await asyncio.wait_for(fixes.wait(), timeout=20)
assert not errors
assert len(branches) == 1
@ -208,7 +210,7 @@ async def test_blocked_native_child_has_no_patch(tmp_path, monkeypatch):
scan_module,
"_run_config",
lambda env: RunConfig(
model=ScriptedModel([*patch(), finish("blocked")]),
model=ScriptedModel([*patch(), finish("blocked", "Required database is unavailable")]),
sandbox=SandboxRunConfig(session=env.session),
tracing_disabled=True,
),
@ -216,7 +218,9 @@ async def test_blocked_native_child_has_no_patch(tmp_path, monkeypatch):
assert (await delegate(context))["success"]
branches, errors = await fixes.wait()
assert branches == []
assert errors == []
assert len(errors) == 1
assert errors[0]["finding_id"] == "finding"
assert "Required database is unavailable" in errors[0]["error"]
assert fixes.records["finding"]["status"] == "stopped"
assert not list((tmp_path / "state/fixes").glob("*/prepared-fix.zip"))
for session in sessions:
@ -224,7 +228,54 @@ async def test_blocked_native_child_has_no_patch(tmp_path, monkeypatch):
@pytest.mark.asyncio
async def test_revised_candidate_does_not_publish_stale_ready_artifact(tmp_path):
async def test_cancelled_verification_is_reported_without_publishing_a_branch(
tmp_path, monkeypatch
):
fixes, _, source, _, _, context, sessions = setup(tmp_path, interactive=True)
fixes.publish_local_branches = True
reviewing = asyncio.Event()
async def verify(*_args):
reviewing.set()
await asyncio.Event().wait()
monkeypatch.setattr(scan_module, "finish_native_fix", verify)
monkeypatch.setattr(
scan_module,
"_run_config",
lambda env: RunConfig(
model=ScriptedModel([finish("blocked")]),
sandbox=SandboxRunConfig(session=env.session),
tracing_disabled=True,
),
)
assert (await delegate(context))["success"]
await asyncio.wait_for(reviewing.wait(), timeout=5)
await fixes.close()
branches, errors = await fixes.wait()
assert branches == []
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"))
for session in sessions:
session.close()
@pytest.mark.asyncio
async def test_dispatch_failure_is_reported_before_an_agent_exists(tmp_path, monkeypatch):
fixes, _, _, _, _, context, _ = setup(tmp_path)
fixes.publish_local_branches = True
fixes._parent_ctx = context.context
monkeypatch.setattr(fixes, "spawn", AsyncMock(side_effect=RuntimeError("Source unavailable")))
await fixes._dispatch("finding")
branches, errors = await fixes.wait()
assert branches == []
assert errors[0]["error"] == "Source unavailable"
@pytest.mark.asyncio
async def test_revised_candidate_does_not_publish_stale_ready_artifact(tmp_path, monkeypatch):
fixes, report, source, _, _, _, _ = setup(tmp_path)
fixes.publish_local_branches = True
original_digest = fixes._finding("finding")[1].digest()
@ -237,6 +288,8 @@ async def test_revised_candidate_does_not_publish_stale_ready_artifact(tmp_path)
"artifact": str(artifact),
}
report["fix_candidate"]["security_invariant"] = "Revised attack"
# Check publication of the old artifact without starting the new attempt.
monkeypatch.setattr(fixes, "_reconcile", AsyncMock())
branches, errors = await fixes.wait()
@ -326,6 +379,7 @@ async def test_finding_changed_before_delivery_discards_reviewed_patch(
)
monkeypatch.setattr(scan_module, "finish_native_fix", finish_preparation)
monkeypatch.setattr(scan_module, "_run_config", lambda _env: RunConfig(tracing_disabled=True))
fixes.sink = AsyncMock(return_value=True)
await fixes.spawn(
"finding", spawn, parent_ctx=context.context, name="Fix", task="Repair", skills=[]