feat: allow scans to skip automatic fixes

This commit is contained in:
yoni 2026-10-01 07:06:52 +00:00
parent bba4aa24dc
commit 60d4ce15e5
2 changed files with 73 additions and 14 deletions

View file

@ -500,7 +500,13 @@ async def run_strix_scan(
)
report_state = get_global_report_state()
if report_state is not None and local_sources and not interactive:
one_click_fixes_enabled = scan_config.get("one_click_fixes_enabled", True) is not False
if (
one_click_fixes_enabled
and report_state is not None
and local_sources
and not interactive
):
fixes = ScanFixes(
session=sandbox_session,
coordinator=coordinator,

View file

@ -2,7 +2,7 @@ from __future__ import annotations
import asyncio
import types
from typing import Any
from typing import TYPE_CHECKING, Any
import pytest
from agents import ModelSettings
@ -15,6 +15,10 @@ from strix.runtime import session_manager
from tests.test_fix_reliability import LocalSandbox
if TYPE_CHECKING:
from pathlib import Path
def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None:
monkeypatch.setattr(runner, "run_dir_for", lambda _scan_id: tmp_path)
monkeypatch.setattr(runner, "runtime_state_dir", lambda _run_dir: tmp_path)
@ -99,42 +103,44 @@ async def test_a_live_child_is_settled_before_sessions_close(
@pytest.mark.asyncio
async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(monkeypatch, tmp_path):
async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
_wire_runner(monkeypatch, tmp_path)
events = []
events: list[str] = []
class State:
defer_completion = False
def __init__(self):
def __init__(self) -> None:
self.scan_results = {"scan_completed": True}
def get_existing_vulnerabilities(self):
def get_existing_vulnerabilities(self) -> list[Any]:
return []
def save_run_data(self, **_):
def save_run_data(self, **_: Any) -> None:
events.append("save")
class Fixes:
def __init__(self, **_):
def __init__(self, **_: Any) -> None:
pass
def start(self, *_):
def start(self, *_: Any) -> None:
events.append("fixes listening")
async def wait(self):
async def wait(self) -> None:
events.append("fixes finished")
async def close(self):
async def close(self) -> None:
events.append("fixes closed")
async def assessment(_):
async def assessment(_: Any) -> None:
events.append("assessment published")
async def root(**_):
async def root(**_: Any) -> types.SimpleNamespace:
return types.SimpleNamespace(final_output={"scan_completed": True})
async def cleanup(*_):
async def cleanup(*_: Any) -> None:
events.append("sandbox deleted")
monkeypatch.setattr(runner, "get_global_report_state", State)
@ -150,3 +156,50 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(monkey
)
assert events.index("assessment published") < events.index("fixes finished")
assert events.index("fixes finished") < events.index("sandbox deleted")
@pytest.mark.asyncio
async def test_disabled_one_click_fixes_skips_fix_runtime(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
_wire_runner(monkeypatch, tmp_path)
events: list[str] = []
class State:
defer_completion = False
def __init__(self) -> None:
self.scan_results = {"scan_completed": True}
def get_existing_vulnerabilities(self) -> list[Any]:
return []
def save_run_data(self, **_: Any) -> None:
events.append("save")
class Fixes:
def __init__(self, **_: Any) -> None:
raise AssertionError("ScanFixes must not be constructed")
async def assessment(_: Any) -> None:
events.append("assessment published")
async def root(**_: Any) -> types.SimpleNamespace:
return types.SimpleNamespace(final_output={"scan_completed": True})
monkeypatch.setattr(runner, "get_global_report_state", State)
monkeypatch.setattr(runner, "ScanFixes", Fixes)
monkeypatch.setattr(runner, "run_agent_loop", root)
await runner.run_strix_scan(
scan_config={
"targets": [],
"scan_mode": "deep",
"one_click_fixes_enabled": False,
},
scan_id="scan",
image="image",
local_sources=[{"source_path": str(tmp_path)}],
assessment_sink=assessment,
)
assert events == ["assessment published", "save"]