feat: publish verified fixes as local branches

This commit is contained in:
yoni 2026-10-01 15:16:20 +00:00
parent 3763a677f9
commit 1fa7d211b4
10 changed files with 364 additions and 7 deletions

View file

@ -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

View file

@ -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:

View file

@ -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)

View file

@ -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()

View file

@ -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)

View file

@ -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", []),

View file

@ -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):

View file

@ -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:

View file

@ -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")

View file

@ -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)