diff --git a/strix/core/execution.py b/strix/core/execution.py index 15097b6cf..e308ffa5f 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -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: diff --git a/strix/core/runner.py b/strix/core/runner.py index 7c85199e3..924091b8d 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -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, ) diff --git a/strix/fix/scan.py b/strix/fix/scan.py index c107b08a1..9c44d002b 100644 --- a/strix/fix/scan.py +++ b/strix/fix/scan.py @@ -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 diff --git a/strix/interface/main.py b/strix/interface/main.py index ce4546fda..b16fdda48 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -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") diff --git a/strix/interface/tui/backend/live_view.py b/strix/interface/tui/backend/live_view.py index 549f8bddf..3e40530c2 100644 --- a/strix/interface/tui/backend/live_view.py +++ b/strix/interface/tui/backend/live_view.py @@ -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, diff --git a/strix/interface/tui/internal/app/model_test.go b/strix/interface/tui/internal/app/model_test.go index 4a5c45642..c667a0bbb 100644 --- a/strix/interface/tui/internal/app/model_test.go +++ b/strix/interface/tui/internal/app/model_test.go @@ -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 diff --git a/strix/interface/tui/internal/app/view.go b/strix/interface/tui/internal/app/view.go index 363bd6bd4..450fdb9dc 100644 --- a/strix/interface/tui/internal/app/view.go +++ b/strix/interface/tui/internal/app/view.go @@ -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)) } diff --git a/strix/interface/tui/internal/render/registry.go b/strix/interface/tui/internal/render/registry.go index 9b7f4f35a..2c695c2e5 100644 --- a/strix/interface/tui/internal/render/registry.go +++ b/strix/interface/tui/internal/render/registry.go @@ -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": diff --git a/strix/interface/tui/internal/render/render_test.go b/strix/interface/tui/internal/render/render_test.go index e14169fd8..ba8012830 100644 --- a/strix/interface/tui/internal/render/render_test.go +++ b/strix/interface/tui/internal/render/render_test.go @@ -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") diff --git a/strix/interface/tui/internal/render/scan.go b/strix/interface/tui/internal/render/scan.go index d1fc75c5b..ac83c6ae7 100644 --- a/strix/interface/tui/internal/render/scan.go +++ b/strix/interface/tui/internal/render/scan.go @@ -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) diff --git a/strix/interface/tui/runtime.py b/strix/interface/tui/runtime.py index c33e96d82..30c28b19d 100644 --- a/strix/interface/tui/runtime.py +++ b/strix/interface/tui/runtime.py @@ -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 diff --git a/strix/tools/finish/tool.py b/strix/tools/finish/tool.py index 29d903542..59a0ff1bc 100644 --- a/strix/tools/finish/tool.py +++ b/strix/tools/finish/tool.py @@ -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 diff --git a/tests/test_execution.py b/tests/test_execution.py index 1020879cf..381e36e0b 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -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, diff --git a/tests/test_fix_progress.py b/tests/test_fix_progress.py new file mode 100644 index 000000000..c84e4da9c --- /dev/null +++ b/tests/test_fix_progress.py @@ -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() diff --git a/tests/test_go_tui_runtime.py b/tests/test_go_tui_runtime.py index 18a0256a9..d4a10aba7 100644 --- a/tests/test_go_tui_runtime.py +++ b/tests/test_go_tui_runtime.py @@ -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()) diff --git a/tests/test_runner_teardown.py b/tests/test_runner_teardown.py index 758e8cf03..63df28b33 100644 --- a/tests/test_runner_teardown.py +++ b/tests/test_runner_teardown.py @@ -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"] diff --git a/tests/test_scan_fixes.py b/tests/test_scan_fixes.py index fb55aa68e..615e24d5d 100644 --- a/tests/test_scan_fixes.py +++ b/tests/test_scan_fixes.py @@ -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=[]