From 60d4ce15e5778eec0f5522775c159d70dcc80d73 Mon Sep 17 00:00:00 2001 From: yoni Date: Thu, 1 Oct 2026 07:06:52 +0000 Subject: [PATCH] feat: allow scans to skip automatic fixes --- strix/core/runner.py | 8 +++- tests/test_runner_teardown.py | 79 +++++++++++++++++++++++++++++------ 2 files changed, 73 insertions(+), 14 deletions(-) diff --git a/strix/core/runner.py b/strix/core/runner.py index dcc48e9fe..d9e835315 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -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, diff --git a/tests/test_runner_teardown.py b/tests/test_runner_teardown.py index 7d45dcfea..1e88c4559 100644 --- a/tests/test_runner_teardown.py +++ b/tests/test_runner_teardown.py @@ -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"]