mirror of
https://github.com/usestrix/strix.git
synced 2026-10-02 02:13:43 +00:00
fix: complete interactive autofix scans before cleanup
This commit is contained in:
parent
1fa7d211b4
commit
f14468565a
17 changed files with 467 additions and 30 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
71
tests/test_fix_progress.py
Normal file
71
tests/test_fix_progress.py
Normal 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()
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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=[]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue