diff --git a/docs/fix-preparation.md b/docs/fix-preparation.md index cf7182bcd..f4767d3c9 100644 --- a/docs/fix-preparation.md +++ b/docs/fix-preparation.md @@ -7,6 +7,9 @@ - The agent calls `agent_finish(success=True)` or `agent_finish(success=False)`. The controller enforces **300 total model turns per finding**, including resumed execution and candidate revisions. It does not start a fresh agent after exhaustion. - The finish tool checkpoints source before completing. Packaging errors return to the agent for correction; three identical completion errors stop the job with that reason. Untracked dependency/cache paths stay out of the patch; new source and tests stay in. Only a completed, nonempty patch becomes an artifact. Blocked, interrupted, or capped work produces no deliverable patch. - Assessment completion publishes the security report. Fixes may continue in the same sandbox; execution and sandbox cleanup finish after all Fix tasks stop. Scan cancellation and the shared model budget also stop fix work. +- The open-source CLI enables automatic fixes by default for source scans. Use `--no-auto-fix` to disable Fix and verifier agents. +- After final reconciliation, the CLI creates one local branch for each current approved fix. It does not change the active checkout or push a branch. +- The final output and `run.json` list each prepared branch. Blocked, rejected, and superseded fixes do not create branches. - In the hosted app, successful fixes become available for **user-initiated draft PR creation** on the issue. Incomplete patches are not shown. Internal diagnostic logs and terminal status remain available to operators. ## Implementation diff --git a/strix/core/runner.py b/strix/core/runner.py index d9e835315..7c85199e3 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -500,12 +500,21 @@ async def run_strix_scan( ) report_state = get_global_report_state() - one_click_fixes_enabled = scan_config.get("one_click_fixes_enabled", True) is not False + auto_fix_enabled = ( + scan_config.get( + "auto_fix_enabled", + scan_config.get("one_click_fixes_enabled", True), + ) + is not False + ) + local_fix_branches_enabled = ( + scan_config.get("local_fix_branches_enabled") is True and fix_sink is None + ) if ( - one_click_fixes_enabled + auto_fix_enabled and report_state is not None and local_sources - and not interactive + and (not interactive or local_fix_branches_enabled) ): fixes = ScanFixes( session=sandbox_session, @@ -517,6 +526,7 @@ async def run_strix_scan( report_state=report_state, event_sink=event_sink, sink=fix_sink, + publish_local_branches=local_fix_branches_enabled, ) report_state.defer_completion = True @@ -650,7 +660,9 @@ async def run_strix_scan( logger.exception("Could not publish assessment before fix completion") if fixes is not None: report("Assessment complete ยท Fixes in progress") - await fixes.wait() + fix_branches, fix_branch_errors = await fixes.wait() + report_state.scan_results["fix_branches"] = fix_branches + report_state.scan_results["fix_branch_errors"] = fix_branch_errors report_state.save_run_data(mark_complete=True) return result # noqa: TRY300 except BudgetExceededError as exc: diff --git a/strix/fix/scan.py b/strix/fix/scan.py index 9b756bcf1..c107b08a1 100644 --- a/strix/fix/scan.py +++ b/strix/fix/scan.py @@ -8,9 +8,12 @@ import hashlib import io import json import logging +import os +import re import shutil import subprocess import time +import zipfile from collections.abc import Awaitable, Callable from copy import deepcopy from pathlib import Path @@ -51,6 +54,7 @@ class ScanFixes: report_state: Any, event_sink: Any = None, sink: FixSink | None = None, + publish_local_branches: bool = False, ) -> None: self.session, self.coordinator = session, coordinator self.scan_id, self.directory = scan_id, state_dir / "fixes" @@ -71,6 +75,7 @@ class ScanFixes: if s.get("source_path") } self.hooks, self.event_sink, self.sink = hooks, event_sink, sink + self.publish_local_branches = publish_local_branches self.report_state = report_state self.tasks: dict[str, asyncio.Task[Any]] = {} self.dispatches: set[asyncio.Task[Any]] = set() @@ -406,11 +411,62 @@ class ScanFixes: await self._exec("git", "-C", base, "worktree", "remove", "--force", root) shutil.rmtree(directory / "source", ignore_errors=True) - async def wait(self) -> None: + async def wait(self) -> tuple[list[dict[str, str]], list[dict[str, str]]]: await asyncio.gather(*self.dispatches, return_exceptions=True) await self._reconcile() self.closed = True await asyncio.gather(*self.tasks.values(), return_exceptions=True) + if not self.publish_local_branches: + return [], [] + + branches: list[dict[str, str]] = [] + errors: list[dict[str, str]] = [] + for finding_id, record in sorted(self.records.items()): + if record.get("status") != "done" or not record.get("artifact"): + continue + title = finding_id + try: + report, candidate = self._finding(finding_id) + assert candidate.source_identity is not None + if candidate.digest() != record.get("digest"): + continue + title = str( + report.get("title") + or (candidate.finding.title if candidate.finding else "") + or finding_id + ) + branch = await asyncio.to_thread( + _publish_local_branch, + self.sources[0], + Path(record["artifact"]), + finding_id, + title, + candidate.source_identity.value, + candidate.digest(), + ) + record["branch"] = branch + record["source_path"] = str(self.sources[0]) + record.pop("branch_error", None) + branches.append( + { + "finding_id": finding_id, + "title": title, + "branch": branch, + "source_path": str(self.sources[0]), + } + ) + except Exception as error: + record["branch_error"] = str(error) + errors.append( + { + "finding_id": finding_id, + "title": title, + "error": str(error), + } + ) + logger.exception("Could not publish local Fix branch for %s", finding_id) + self._save() + return branches, errors async def close(self) -> None: self.closed = True @@ -485,3 +541,95 @@ def _clone_revision(source: Path, mirror: Path, commit: str) -> None: capture_output=True, timeout=60, ) + + +def _publish_local_branch( + source: Path, + artifact: Path, + finding_id: str, + title: str, + base_commit: str, + candidate_digest: str, +) -> str: + slug = re.sub(r"[^a-z0-9]+", "-", title.lower()).strip("-")[:40] or "finding" + finding = re.sub(r"[^a-zA-Z0-9]+", "", finding_id)[:8].lower() or "finding" + branch = f"strix/fix-{slug}-{finding}-{candidate_digest[:8]}" + workspace = artifact.parent / "branch-source" + shutil.rmtree(workspace, ignore_errors=True) + _clone_revision(source, workspace, base_commit) + patch = workspace.parent / "changes.patch" + try: + with zipfile.ZipFile(artifact) as archive: + patch.write_bytes(archive.read("changes.patch")) + git = shutil.which("git") or "/usr/bin/git" + subprocess.run( # noqa: S603 + [git, "apply", "--index", "--binary", "--", str(patch)], + cwd=workspace, + check=True, + capture_output=True, + timeout=120, + ) + tree = ( + subprocess.check_output( # noqa: S603 + [git, "write-tree"], + cwd=workspace, + timeout=30, + ) + .decode() + .strip() + ) + existing = subprocess.run( # noqa: S603 + [git, "rev-parse", "--verify", f"refs/heads/{branch}^{{tree}}"], + cwd=source, + check=False, + capture_output=True, + text=True, + timeout=30, + ) + if existing.returncode == 0: + parent = subprocess.check_output( # noqa: S603 + [git, "rev-parse", f"refs/heads/{branch}^"], + cwd=source, + text=True, + timeout=30, + ).strip() + if existing.stdout.strip() != tree or parent != base_commit: + raise RuntimeError(f"Local branch {branch} already contains different changes.") + return branch + + subject = " ".join(title.split())[:120] or finding_id + message = f"fix: {subject}\n" + environment = { + **os.environ, + "GIT_AUTHOR_NAME": "Strix", + "GIT_AUTHOR_EMAIL": "fixes@strix.ai", + "GIT_COMMITTER_NAME": "Strix", + "GIT_COMMITTER_EMAIL": "fixes@strix.ai", + } + commit = subprocess.check_output( # noqa: S603 + [git, "commit-tree", tree, "-p", base_commit], + cwd=workspace, + input=message, + text=True, + env=environment, + timeout=30, + ).strip() + subprocess.run( # noqa: S603 + [ + git, + "fetch", + "--no-tags", + "--no-write-fetch-head", + "--", + str(workspace), + f"{commit}:refs/heads/{branch}", + ], + cwd=source, + check=True, + capture_output=True, + timeout=120, + ) + return branch + finally: + patch.unlink(missing_ok=True) + shutil.rmtree(workspace, ignore_errors=True) diff --git a/strix/interface/cli.py b/strix/interface/cli.py index a8992ce49..e61ab8b56 100644 --- a/strix/interface/cli.py +++ b/strix/interface/cli.py @@ -93,6 +93,8 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915 "diff_scope": getattr(args, "diff_scope", {"active": False}), "scan_mode": scan_mode, "non_interactive": bool(getattr(args, "non_interactive", False)), + "auto_fix_enabled": bool(args.auto_fix), + "local_fix_branches_enabled": bool(args.auto_fix), "local_sources": getattr(args, "local_sources", None) or [], "workspace_files": getattr(args, "workspace_files", None) or [], "scope_mode": getattr(args, "scope_mode", "auto"), @@ -257,3 +259,32 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915 console.print(final_report_panel) console.print() + fix_branches = (report_state.scan_results or {}).get("fix_branches") or [] + fix_branch_errors = (report_state.scan_results or {}).get("fix_branch_errors") or [] + if fix_branches: + lines = [ + f"{item['title']}\n {item['branch']}\n {item['source_path']}" + for item in fix_branches + ] + console.print( + Panel( + "\n\n".join(lines), + title="[bold white]Prepared fix branches", + title_align="left", + border_style="#60a5fa", + padding=(1, 2), + ) + ) + console.print() + if fix_branch_errors: + lines = [f"{item['title']}: {item['error']}" for item in fix_branch_errors] + console.print( + Panel( + "\n".join(lines), + title="[bold white]Fix branch errors", + title_align="left", + border_style="#ef4444", + padding=(1, 2), + ) + ) + console.print() diff --git a/strix/interface/cli_args.py b/strix/interface/cli_args.py index 340a5d471..7e729784d 100644 --- a/strix/interface/cli_args.py +++ b/strix/interface/cli_args.py @@ -186,6 +186,16 @@ Strix Cloud: ), ) + parser.add_argument( + "--auto-fix", + action=argparse.BooleanOptionalAction, + default=None, + help=( + "Prepare verified local fix branches for confirmed source findings. " + "Use --no-auto-fix to disable this behavior. Default: enabled." + ), + ) + parser.add_argument( "-m", "--scan-mode", @@ -378,6 +388,8 @@ Strix Cloud: except ValueError as e: parser.error(str(e)) + if args.auto_fix is None: + args.auto_fix = True return args @@ -429,6 +441,8 @@ def _load_resume_state(args: argparse.Namespace, parser: argparse.ArgumentParser if args.instruction is None: args.instruction = state.get("instruction") + if args.auto_fix is None: + args.auto_fix = state.get("auto_fix", True) is not False if not getattr(args, "user_instruction", None): args.user_instruction = state.get("user_instruction") or None args.local_sources = collect_local_sources(args.targets_info) diff --git a/strix/interface/scan_setup.py b/strix/interface/scan_setup.py index ae7caf2fc..52d1f314b 100644 --- a/strix/interface/scan_setup.py +++ b/strix/interface/scan_setup.py @@ -255,6 +255,7 @@ def _persist_run_record(args: argparse.Namespace) -> None: # transcript replays this as the user's opening message. "user_instruction": getattr(args, "user_instruction", None), "non_interactive": args.non_interactive, + "auto_fix": args.auto_fix, "local_sources": getattr(args, "local_sources", []), # Persisted so --resume places the same workspace files again. "workspace_files": getattr(args, "workspace_files", []), diff --git a/strix/interface/tui/runtime.py b/strix/interface/tui/runtime.py index 64c77a7b6..c33e96d82 100644 --- a/strix/interface/tui/runtime.py +++ b/strix/interface/tui/runtime.py @@ -91,6 +91,8 @@ class GoTuiRuntime: "diff_scope": self.args.diff_scope, "scan_mode": self.args.scan_mode, "non_interactive": False, + "auto_fix_enabled": bool(self.args.auto_fix), + "local_fix_branches_enabled": bool(self.args.auto_fix), "local_sources": self.args.local_sources or [], "workspace_files": getattr(self.args, "workspace_files", None) or [], "scope_mode": self.args.scope_mode, @@ -253,6 +255,17 @@ class GoTuiRuntime: mcp_status_sink=self.capture_mcp_status, ) await self._sync_agent_state() + 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']}" + ) + for item in results.get("fix_branch_errors") or []: + self.controller.add_message( + f"Could not create fix branch for {item['title']}: {item['error']}", + "error", + ) if self.controller.scan_state == "running": self.controller.scan_state = "stopped" except (asyncio.CancelledError, BudgetExceededError): diff --git a/tests/test_cli_target_list.py b/tests/test_cli_target_list.py index d20230b16..160a6b370 100644 --- a/tests/test_cli_target_list.py +++ b/tests/test_cli_target_list.py @@ -69,6 +69,29 @@ def test_parse_arguments_combines_target_and_target_list( ] +@pytest.mark.parametrize( + ("extra", "expected"), + [ + ([], True), + (["--auto-fix"], True), + (["--no-auto-fix"], False), + ], +) +def test_parse_arguments_controls_automatic_fixes( + monkeypatch: pytest.MonkeyPatch, extra: list[str], expected: bool +) -> None: + _stub_settings(monkeypatch) + monkeypatch.setattr( + sys, + "argv", + ["strix", "--target", "https://example.com", "--non-interactive", *extra], + ) + + args = cli_main.parse_arguments() + + assert args.auto_fix is expected + + def test_parse_arguments_rejects_resume_with_target_list( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] ) -> None: @@ -128,6 +151,32 @@ def test_resume_restores_a_target_less_workspace_mount( assert args.instruction == "audit the auth flow" +def test_resume_restores_disabled_automatic_fixes( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + work = tmp_path / "project" + work.mkdir() + monkeypatch.chdir(tmp_path) + _write_run_record( + tmp_path / "strix_runs", + "pentest_abcd", + { + "run_name": "pentest_abcd", + "targets_info": [], + "local_sources": [], + "workspace_mount": str(work), + "instruction": "audit the auth flow", + "scan_mode": "deep", + "auto_fix": False, + }, + ) + monkeypatch.setattr(sys, "argv", ["strix", "--resume", "pentest_abcd"]) + + args = cli_main.parse_arguments() + + assert args.auto_fix is False + + def test_resume_revalidates_persisted_workspace_files( tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: diff --git a/tests/test_runner_teardown.py b/tests/test_runner_teardown.py index 1e88c4559..758e8cf03 100644 --- a/tests/test_runner_teardown.py +++ b/tests/test_runner_teardown.py @@ -128,8 +128,9 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown( def start(self, *_: Any) -> None: events.append("fixes listening") - async def wait(self) -> None: + async def wait(self) -> tuple[list[Any], list[Any]]: events.append("fixes finished") + return [], [] async def close(self) -> None: events.append("fixes closed") diff --git a/tests/test_scan_fixes.py b/tests/test_scan_fixes.py index 9ed9d53d7..fb55aa68e 100644 --- a/tests/test_scan_fixes.py +++ b/tests/test_scan_fixes.py @@ -145,6 +145,44 @@ async def test_native_parallel_fixes_deliver_patches_and_preserve_scan_source( session.close() +@pytest.mark.asyncio +async def test_ready_fix_creates_local_branch_without_changing_checkout(tmp_path, monkeypatch): + fixes, _, source, _, _, context, sessions = setup(tmp_path) + fixes.publish_local_branches = True + original_head = _git(source, "rev-parse", "HEAD") + original_branch = _git(source, "branch", "--show-current") + monkeypatch.setattr( + scan_module, + "_run_config", + lambda env: RunConfig( + model=ScriptedModel([*patch(), *suite_commands(), finish("done"), finish("done")]), + sandbox=SandboxRunConfig(session=env.session), + tracing_disabled=True, + ), + ) + + assert (await delegate(context))["success"] + branches, errors = await fixes.wait() + + assert not errors + assert len(branches) == 1 + branch = branches[0]["branch"] + assert branch.startswith("strix/fix-unsafe-result-finding-") + assert fixes.records["finding"]["branch"] == branch + assert _git(source, "branch", "--show-current") == original_branch + assert _git(source, "rev-parse", "HEAD") == original_head + assert _git(source, "rev-parse", f"{branch}^") == original_head + assert _git(source, "status", "--porcelain") == "" + assert "return 'safe'" in _git(source, "show", f"{branch}:app.py") + assert "test_safe" in _git(source, "show", f"{branch}:tests/test_security.py") + repeated, repeated_errors = await fixes.wait() + assert repeated_errors == [] + assert repeated == branches + assert _git(source, "show-ref", "--verify", f"refs/heads/{branch}") + for session in sessions: + session.close() + + @pytest.mark.asyncio async def test_delegation_errors_reach_reporting_agent_before_any_model_call(tmp_path): fixes, report, _, _, _, context, _ = setup(tmp_path) @@ -165,6 +203,7 @@ async def test_delegation_errors_reach_reporting_agent_before_any_model_call(tmp @pytest.mark.asyncio async def test_blocked_native_child_has_no_patch(tmp_path, monkeypatch): fixes, _, _, _, _, context, sessions = setup(tmp_path) + fixes.publish_local_branches = True monkeypatch.setattr( scan_module, "_run_config", @@ -175,13 +214,59 @@ async def test_blocked_native_child_has_no_patch(tmp_path, monkeypatch): ), ) assert (await delegate(context))["success"] - await fixes.wait() + branches, errors = await fixes.wait() + assert branches == [] + assert errors == [] assert fixes.records["finding"]["status"] == "stopped" assert not list((tmp_path / "state/fixes").glob("*/prepared-fix.zip")) for session in sessions: session.close() +@pytest.mark.asyncio +async def test_revised_candidate_does_not_publish_stale_ready_artifact(tmp_path): + fixes, report, source, _, _, _, _ = setup(tmp_path) + fixes.publish_local_branches = True + original_digest = fixes._finding("finding")[1].digest() + artifact = tmp_path / "state/fixes/stale/prepared-fix.zip" + artifact.parent.mkdir(parents=True) + artifact.write_bytes(b"unused") + fixes.records["finding"] = { + "digest": original_digest, + "status": "done", + "artifact": str(artifact), + } + report["fix_candidate"]["security_invariant"] = "Revised attack" + + branches, errors = await fixes.wait() + + assert branches == [] + assert errors == [] + assert _git(source, "branch", "--list", "strix/fix-*") == "" + + +@pytest.mark.asyncio +async def test_local_branch_failure_is_persisted_and_returned(tmp_path): + fixes, _, _, _, _, _, _ = setup(tmp_path) + fixes.publish_local_branches = True + digest = fixes._finding("finding")[1].digest() + artifact = tmp_path / "state/fixes/broken/prepared-fix.zip" + artifact.parent.mkdir(parents=True) + artifact.write_bytes(b"not a zip") + fixes.records["finding"] = { + "digest": digest, + "status": "done", + "artifact": str(artifact), + } + + branches, errors = await fixes.wait() + + assert branches == [] + assert len(errors) == 1 + assert errors[0]["finding_id"] == "finding" + assert fixes.records["finding"]["branch_error"] == errors[0]["error"] + + @pytest.mark.asyncio async def test_finding_revision_invalidates_active_completion(tmp_path): fixes, report, _, _, _, _, _ = setup(tmp_path)