mirror of
https://github.com/usestrix/strix.git
synced 2026-10-02 02:13:43 +00:00
feat: allow scans to skip automatic fixes
This commit is contained in:
parent
bba4aa24dc
commit
60d4ce15e5
2 changed files with 73 additions and 14 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue