From f23afb438e7216f1c9c5d6a1c9062cc8b622d0b3 Mon Sep 17 00:00:00 2001 From: yoni Date: Fri, 25 Sep 2026 05:16:06 +0000 Subject: [PATCH] Add verified fix preparation engine --- strix/agents/prompts/system_prompt.jinja | 21 +- strix/fix/__init__.py | 52 +++ strix/fix/contracts.py | 278 +++++++++++++++ strix/fix/locations.py | 97 +++++ strix/fix/prepare.py | 435 +++++++++++++++++++++++ strix/interface/utils.py | 4 +- strix/report/sarif.py | 12 +- strix/report/state.py | 12 + strix/report/writer.py | 16 +- strix/tools/reporting/tool.py | 150 ++++++-- tests/test_fix_preparation.py | 384 ++++++++++++++++++++ tests/test_report_writer.py | 2 +- tests/test_sarif.py | 29 +- 13 files changed, 1441 insertions(+), 51 deletions(-) create mode 100644 strix/fix/__init__.py create mode 100644 strix/fix/contracts.py create mode 100644 strix/fix/locations.py create mode 100644 strix/fix/prepare.py create mode 100644 tests/test_fix_preparation.py diff --git a/strix/agents/prompts/system_prompt.jinja b/strix/agents/prompts/system_prompt.jinja index 491b544c3..234a865f4 100644 --- a/strix/agents/prompts/system_prompt.jinja +++ b/strix/agents/prompts/system_prompt.jinja @@ -153,8 +153,9 @@ WHITE-BOX TESTING (code provided): - Local execution, unit/integration testing, patch verification, and HTTP requests against locally started in-scope services are normal authorized white-box validation - If dynamically running the code proves impossible after exhaustive attempts, pivot to comprehensive static analysis. - Try to infer how to run the code based on its structure and content. -- Derive the code fix as PART OF reporting, not as a separate later pass: create_vulnerability_report already requires the concrete patch inline (`code_locations` with verbatim `fix_before`/`fix_after` and `fix_pr_body`), so the reporting agent that analyzes the root cause is the one that produces the fix. Do NOT spawn a downstream agent afterwards to re-derive/re-apply the same patch. -- If you also apply and verify the patch in the repo (edit the file, re-test that the vulnerability is gone), do it in the same agent/turn while the analysis is fresh — right before or as part of filing the report — never as a second re-analysis pass. +- Draft the initial fix candidate when you file the report. Use `code_locations` with verbatim `fix_before` and `fix_after`, plus `fix_pr_body`. +- Treat the inline changes as a candidate, not as a completed fix. A later preparation stage can inspect and modify any required repository file. +- Record checks you ran in `fix_verification`. Do not describe reasoned checks as executed checks. COMBINED MODE (code + deployed target present): - Treat this as static analysis plus dynamic testing simultaneously @@ -240,7 +241,9 @@ VALIDATION REQUIREMENTS: - THREAT MODEL: before you start testing, call `get_threat_model` on the target you were pointed at — it is the scan's shared answer to who the attacker is, where the trust boundaries sit, and what counts as critical here. It is scoped to this scan and nothing carries over from an earlier run, so `found: false` means no agent on this run has derived one yet. Read it instead of re-deriving trust boundaries yourself; where your testing disproves it — a boundary it calls trusted turns out to be attacker-reachable, a role it did not know about, a host or endpoint it never listed — record that with `amend_threat_model` so the agents after you inherit the correction. Amending is not optional politeness: a model nobody corrects turns the first agent's guesses into everyone's assumptions. - Before filing any report, run the counterevidence pass: argue the strongest case AGAINST the finding, record what you found in the `counterevidence` field, set `confidence` honestly (a static-only trace you couldn't execute is at best `medium`), and state what evidence would change the severity. See the counterevidence and severity-calibration knowledge above. - A vulnerability is ONLY considered reported when a reporting agent uses create_vulnerability_report (or create_dependency_report for known-CVE dependency/supply-chain findings) with full details. Mentions in agent_finish, finish_scan, or generic messages are NOT sufficient -- Reporting and fixing are ONE step, not two: when source is available, the reporting agent derives the concrete fix and files it INLINE via create_vulnerability_report (`code_locations` with `fix_before`/`fix_after` + `fix_pr_body`) — the report is not complete without it. Do NOT report first and then spawn a separate downstream agent to re-derive and re-apply the same patch; that just re-does the analysis and wastes tokens. (Do not silently patch a finding WITHOUT filing a report — the report, with its embedded fix, is the deliverable.) +- When source is available, the reporting agent files an initial fix candidate with the report. The candidate uses `code_locations` with `fix_before` and `fix_after`, plus `fix_pr_body`. +- Do not treat the candidate as a prepared fix. The fix preparation stage applies, repairs, tests, and independently verifies it after finding discovery closes. +- Do not silently patch a finding without filing a report. - DEDUPLICATION: The create_vulnerability_report tool uses LLM-based deduplication. If it rejects your report as a duplicate, DO NOT attempt to re-submit the same vulnerability. Accept the rejection and move on to testing other areas. The vulnerability has already been reported by another agent. If your evidence proves more than the finding it matched (a working exploit where that one had only a static trace, a chain that raises the impact), revise that finding with update_vulnerability_report using the duplicate_of id — never re-file it. - HTTP EVIDENCE: a finding you validated through the proxy is not fully filed until `http_exchange_ids` carries the proxy request ids of the exchanges that prove it — the request that triggers the vulnerability plus the baseline/control request it differs from (an unauthenticated success next to the authenticated one, the payload response next to the benign one). Copy the ids exactly as `list_requests`/`view_request` show them, never invent or guess one, and never omit the field to bypass validation. Leave it out only when there is no captured HTTP exchange at all (static-only code findings, dependency CVEs). If you filed before the proving exchanges existed, attach them afterwards with update_vulnerability_report. Without the ids, the finding ships as prose nobody can replay. - REVISING A FINDING: use update_vulnerability_report (report id + the fields you want to replace + update_reason) when you learn something a finding already on file does not carry — you built the PoC after filing it, a chain raised its impact, further testing weakened it, or its counterevidence/remediation/code locations were wrong. Editing a finding needs no duplicate verdict, and it is always better than filing a second report for the same issue. Read the finding first with get_report, and pass only the fields that change. @@ -360,7 +363,7 @@ ROOT AGENT ROLE: 1. **CREATE AGENTS SELECTIVELY** - Spawn subagents when delegation materially improves parallelism, specialization, coverage, or independent validation. Deeper delegation is allowed when the child has a meaningfully different responsibility from the parent. Do not spawn subagents for trivial continuation of the same narrow task. 2. **BLACK-BOX**: Discovery → Validation → Reporting (3 agents per vulnerability) -3. **WHITE-BOX**: Discovery → Validation → Reporting-with-fix (3 agents per vulnerability — the reporting agent derives and files the fix inline; do NOT add a separate fixing agent that re-derives the same patch) +3. **WHITE-BOX**: Discovery → Validation → Reporting with an initial fix candidate. The later preparation phase repairs and verifies the candidate. 4. **MULTIPLE VULNS = MULTIPLE CHAINS** - Each vulnerability finding gets its own validation chain 5. **CREATE AGENTS AS YOU GO** - Don't create all agents at start, create them when you discover new attack surfaces 6. **ONE JOB PER AGENT** - Each agent has ONE specific task only @@ -379,7 +382,7 @@ BLACK-BOX (domain/URL only): WHITE-BOX (source code provided): - Found authentication code issues? → Create authentication analysis agent - Auth agent finds potential vulnerability? → Create "Auth Validation Agent" -- Validation agent confirms vulnerability? → Create "Auth Reporting Agent" that files the report AND its inline fix (`code_locations` + `fix_pr_body`) in one shot — no separate fixing agent +- Validation agent confirms vulnerability? → Create "Auth Reporting Agent" that files the report and an initial fix candidate (`code_locations` + `fix_pr_body`) VULNERABILITY WORKFLOW (MANDATORY FOR EVERY FINDING): @@ -401,10 +404,10 @@ Authentication Code Agent finds weak password validation Spawns "Auth Validation Agent" (proves it's exploitable) ↓ If valid → Spawns "Auth Reporting Agent" (creates the vulnerability report - WITH the fix inline: code_locations fix_before/fix_after + fix_pr_body, - applying/verifying the patch in the same turn if desired) + with the initial fix candidate: code_locations fix_before/fix_after + + fix_pr_body) ↓ -STOP - no separate fixing agent; the fix was derived once, at report time +STOP - the scan preparation phase handles repository-wide repair and verification ``` CRITICAL RULES: @@ -440,7 +443,7 @@ FOCUS PRINCIPLES: REALISTIC TESTING OUTCOMES: - **No Findings**: Agent completes testing but finds no vulnerabilities - **Validation Failed**: Initial finding was false positive, validation agent confirms it's not exploitable -- **Valid Vulnerability**: Validation succeeds, spawns a reporting agent that files the report with the fix inline (white-box) — no separate fixing agent +- **Valid Vulnerability**: Validation succeeds. The reporting agent files the report and an initial fix candidate for white-box findings. PERSISTENCE IS MANDATORY: - Real vulnerabilities take TIME - expect to need 2000+ steps minimum diff --git a/strix/fix/__init__.py b/strix/fix/__init__.py new file mode 100644 index 000000000..c11663424 --- /dev/null +++ b/strix/fix/__init__.py @@ -0,0 +1,52 @@ +"""Verified fix preparation contracts and runtime.""" + +from strix.fix.contracts import ( + CandidateLocation, + CheckResult, + CheckStatus, + CommandSpec, + FileManifestEntry, + FixCandidateV1, + FixEdit, + FixPreparationRequestV1, + FixPreparationResultV1, + PreparationState, + ReproductionSpec, + SourceIdentity, + SourceIdentityKind, + VerificationDecision, + VerifierResult, + candidate_from_legacy_report, +) +from strix.fix.prepare import ( + PreparationCancelledError, + PreparationContext, + PreparationPolicy, + build_git_manifest, + prepare_fix, +) + + +__all__ = [ + "CandidateLocation", + "CheckResult", + "CheckStatus", + "CommandSpec", + "FileManifestEntry", + "FixCandidateV1", + "FixEdit", + "FixPreparationRequestV1", + "FixPreparationResultV1", + "PreparationCancelledError", + "PreparationContext", + "PreparationPolicy", + "PreparationState", + "ReproductionSpec", + "SourceIdentity", + "SourceIdentityKind", + "VerificationDecision", + "VerifierResult", + "build_git_manifest", + "candidate_from_legacy_report", + "prepare_fix", +] diff --git a/strix/fix/contracts.py b/strix/fix/contracts.py new file mode 100644 index 000000000..98a52a78a --- /dev/null +++ b/strix/fix/contracts.py @@ -0,0 +1,278 @@ +"""Versioned contracts for verified fix preparation.""" + +from __future__ import annotations + +import hashlib +import json +from enum import StrEnum +from pathlib import PurePosixPath +from typing import TYPE_CHECKING, Literal, cast + +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + + +if TYPE_CHECKING: + from collections.abc import Mapping + + +class ContractModel(BaseModel): + model_config = ConfigDict(extra="forbid") + + +class SourceIdentityKind(StrEnum): + COMMIT = "commit" + ARCHIVE = "archive" + + +class PreparationState(StrEnum): + PREPARING = "preparing" + READY = "ready" + READY_WITH_GAPS = "ready_with_gaps" + NEEDS_REVIEW = "needs_review" + BLOCKED = "blocked" + FAILED = "failed" + STALE = "stale" + + +class CheckStatus(StrEnum): + PASSED = "passed" + FAILED = "failed" + UNAVAILABLE = "unavailable" + CANCELLED = "cancelled" + + +class VerificationDecision(StrEnum): + VERIFIED = "verified" + REJECTED = "rejected" + INCONCLUSIVE = "inconclusive" + + +class SourceIdentity(ContractModel): + kind: SourceIdentityKind + value: str = Field(min_length=1) + repository: str | None = None + + @field_validator("value") + @classmethod + def validate_value(cls, value: str) -> str: + normalized = value.strip().lower() + if not normalized or any(char not in "0123456789abcdef" for char in normalized): + raise ValueError("source identity must be a hexadecimal digest") + return normalized + + @model_validator(mode="after") + def validate_digest_length(self) -> SourceIdentity: + valid_lengths = {40, 64} if self.kind is SourceIdentityKind.COMMIT else {64} + if len(self.value) not in valid_lengths: + expected = "40 or 64" if self.kind is SourceIdentityKind.COMMIT else "64" + raise ValueError(f"{self.kind} identity must contain {expected} hexadecimal characters") + return self + + +def _validate_relative_path(value: str) -> str: + path = value.strip() + parsed = PurePosixPath(path) + if not path or parsed.is_absolute() or ".." in parsed.parts or path.startswith("./"): + raise ValueError("path must be relative to the repository root") + return path + + +class CandidateLocation(ContractModel): + file: str + start_line: int = Field(ge=1) + end_line: int = Field(ge=1) + snippet: str | None = None + label: str | None = None + + _relative_file = field_validator("file")(_validate_relative_path) + + @model_validator(mode="after") + def validate_range(self) -> CandidateLocation: + if self.end_line < self.start_line: + raise ValueError("end_line must be greater than or equal to start_line") + return self + + +class FixEdit(CandidateLocation): + before: str = Field(min_length=1) + after: str + original_sha256: str | None = None + + +class CommandSpec(ContractModel): + name: str = Field(min_length=1, max_length=120) + argv: list[str] = Field(min_length=1) + required: bool = True + timeout_seconds: int = Field(default=300, ge=1, le=3600) + cwd: str = "." + + _relative_cwd = field_validator("cwd")(_validate_relative_path) + + @field_validator("argv") + @classmethod + def validate_argv(cls, value: list[str]) -> list[str]: + if any(not argument or "\x00" in argument for argument in value): + raise ValueError("command arguments must be non-empty and cannot contain NUL") + return value + + +class ReproductionSpec(ContractModel): + instructions: str = Field(min_length=1) + command: CommandSpec | None = None + + +class ReportedCheck(ContractModel): + name: str + result: str + executed: bool = False + + +class FixCandidateV1(ContractModel): + version: Literal["1"] = "1" + source_identity: SourceIdentity | None = None + security_invariant: str = Field(min_length=1) + finding_locations: list[CandidateLocation] = [] + draft_edits: list[FixEdit] = [] + reproduction: ReproductionSpec | None = None + reported_checks: list[ReportedCheck] = [] + known_gaps: list[str] = [] + + def digest(self) -> str: + payload = json.dumps( + self.model_dump(mode="json"), + sort_keys=True, + separators=(",", ":"), + ).encode() + return hashlib.sha256(payload).hexdigest() + + +class FixPreparationRequestV1(ContractModel): + version: Literal["1"] = "1" + organization_id: str | None = None + scan_id: str + finding_id: str + repository_id: str | None = None + candidate: FixCandidateV1 + checks: list[CommandSpec] = [] + max_repair_attempts: int = Field(default=2, ge=1, le=5) + timeout_seconds: int = Field(default=1800, ge=30, le=14400) + network_allowed: bool = False + credentials_allowed: list[str] = [] + + +class CheckResult(ContractModel): + name: str + argv: list[str] + status: CheckStatus + exit_code: int | None = None + duration_seconds: float = Field(ge=0) + output: str = "" + required: bool = True + + +class VerifierResult(ContractModel): + decision: VerificationDecision + summary: str + security_invariant_closed: bool = False + reproduction_executed: bool = False + reproduction_summary: str | None = None + sibling_paths_reviewed: list[str] = [] + preserved_behaviors: list[str] = [] + gaps: list[str] = [] + + +class FileManifestEntry(ContractModel): + path: str + operation: Literal["add", "modify", "delete"] + original_sha256: str | None = None + resulting_sha256: str | None = None + artifact_ref: str | None = None + + _relative_path = field_validator("path")(_validate_relative_path) + + +class FixPreparationResultV1(ContractModel): + version: Literal["1"] = "1" + state: PreparationState + stop_reason: str + source_identity: SourceIdentity | None + candidate_digest: str + final_file_manifest: list[FileManifestEntry] = [] + artifact_ref: str | None = None + changed_files: list[str] = [] + diff_summary: str = "" + checks: list[CheckResult] = [] + security_reproduction: CheckResult | None = None + verifier: VerifierResult | None = None + gaps: list[str] = [] + attempts: int = Field(default=0, ge=0) + elapsed_seconds: float = Field(default=0, ge=0) + cost_usd: float | None = Field(default=None, ge=0) + + +def candidate_from_legacy_report( + report: Mapping[str, object], + *, + source_identity: SourceIdentity | None = None, +) -> FixCandidateV1 | None: + raw_locations = report.get("code_locations") + if not isinstance(raw_locations, list): + return None + location_values = cast("list[object]", raw_locations) + + locations: list[CandidateLocation] = [] + edits: list[FixEdit] = [] + for raw_value in location_values: + if not isinstance(raw_value, dict): + continue + raw = cast("dict[str, object]", raw_value) + file_path = raw.get("file") + start_line = raw.get("start_line") + end_line = raw.get("end_line") + if not isinstance(file_path, str) or type(start_line) is not int: + continue + if type(end_line) is not int: + end_line = start_line + common = { + "file": file_path, + "start_line": start_line, + "end_line": end_line, + "snippet": raw.get("snippet") if isinstance(raw.get("snippet"), str) else None, + "label": raw.get("label") if isinstance(raw.get("label"), str) else None, + } + try: + locations.append(CandidateLocation.model_validate(common)) + before = raw.get("fix_before") + after = raw.get("fix_after") + if isinstance(before, str) and isinstance(after, str): + edits.append(FixEdit.model_validate({**common, "before": before, "after": after})) + except ValueError: + continue + + if not locations: + return None + + invariant = str( + report.get("remediation_steps") or report.get("technical_analysis") or "" + ).strip() + if not invariant: + invariant = "Resolve the reported security finding without changing legitimate behavior." + + reported_checks: list[ReportedCheck] = [] + verification = str(report.get("fix_verification") or "").strip() + if verification: + reported_checks.append( + ReportedCheck(name="reporting-agent verification", result=verification, executed=False) + ) + + return FixCandidateV1( + source_identity=source_identity, + security_invariant=invariant, + finding_locations=locations, + draft_edits=edits, + reproduction=ReproductionSpec( + instructions=str(report.get("poc_description") or report.get("evidence") or invariant) + ), + reported_checks=reported_checks, + known_gaps=["The reporting-agent verification is not independent."], + ) diff --git a/strix/fix/locations.py b/strix/fix/locations.py new file mode 100644 index 000000000..d5f56c8cc --- /dev/null +++ b/strix/fix/locations.py @@ -0,0 +1,97 @@ +"""Deterministic source anchoring for finding locations and draft edits.""" + +from __future__ import annotations + +import hashlib +from dataclasses import dataclass +from enum import StrEnum +from typing import TYPE_CHECKING + +from strix.fix.contracts import CandidateLocation, FixCandidateV1, FixEdit + + +if TYPE_CHECKING: + from pathlib import Path + + +class AnchorStatus(StrEnum): + UNIQUE = "unique" + MISSING = "missing" + AMBIGUOUS = "ambiguous" + STALE = "stale" + + +@dataclass(frozen=True, slots=True) +class AnchorResult: + status: AnchorStatus + location: CandidateLocation | FixEdit + matches: tuple[int, ...] = () + + +def sha256_text(value: str) -> str: + return hashlib.sha256(value.encode()).hexdigest() + + +def _find_blocks(content: str, block: str) -> tuple[int, ...]: + source_lines = content.splitlines() + block_lines = block.splitlines() + if not block_lines: + return () + width = len(block_lines) + return tuple( + index + 1 + for index in range(len(source_lines) - width + 1) + if source_lines[index : index + width] == block_lines + ) + + +def anchor_location( + root: Path, + location: CandidateLocation | FixEdit, +) -> AnchorResult: + file_path = root / location.file + if not file_path.is_file(): + return AnchorResult(AnchorStatus.MISSING, location) + content = file_path.read_text(encoding="utf-8") + block = location.before if isinstance(location, FixEdit) else location.snippet + if not block: + return AnchorResult(AnchorStatus.UNIQUE, location, (location.start_line,)) + + matches = _find_blocks(content, block) + if not matches: + if ( + isinstance(location, FixEdit) + and location.original_sha256 + and sha256_text(content) != location.original_sha256 + ): + return AnchorResult(AnchorStatus.STALE, location) + return AnchorResult(AnchorStatus.MISSING, location) + if len(matches) > 1: + return AnchorResult(AnchorStatus.AMBIGUOUS, location, matches) + + start_line = next(iter(matches)) + end_line = start_line + len(block.splitlines()) - 1 + anchored = location.model_copy(update={"start_line": start_line, "end_line": end_line}) + return AnchorResult(AnchorStatus.UNIQUE, anchored, matches) + + +def anchor_candidate( + root: Path, + candidate: FixCandidateV1, +) -> tuple[FixCandidateV1, list[AnchorResult]]: + results = [ + *(anchor_location(root, location) for location in candidate.finding_locations), + *(anchor_location(root, edit) for edit in candidate.draft_edits), + ] + finding_count = len(candidate.finding_locations) + anchored_locations = [result.location for result in results[:finding_count]] + anchored_edits = [result.location for result in results[finding_count:]] + return ( + candidate.model_copy( + update={ + "finding_locations": anchored_locations, + "draft_edits": anchored_edits, + } + ), + results, + ) diff --git a/strix/fix/prepare.py b/strix/fix/prepare.py new file mode 100644 index 000000000..4fdf36e72 --- /dev/null +++ b/strix/fix/prepare.py @@ -0,0 +1,435 @@ +"""Repository fix preparation engine.""" + +from __future__ import annotations + +import asyncio +import hashlib +import os +import subprocess +import time +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from pathlib import Path +from typing import Literal + +from strix.fix.contracts import ( + CheckResult, + CheckStatus, + CommandSpec, + FileManifestEntry, + FixCandidateV1, + FixPreparationRequestV1, + FixPreparationResultV1, + PreparationState, + VerificationDecision, + VerifierResult, +) +from strix.fix.locations import AnchorStatus, anchor_candidate + + +class PreparationCancelledError(RuntimeError): + pass + + +CommandRunner = Callable[[Path, CommandSpec], Awaitable[CheckResult]] +ManifestBuilder = Callable[[Path], Awaitable[tuple[list[FileManifestEntry], str, str | None]]] +CancellationCheck = Callable[[], bool] + + +@dataclass(slots=True) +class PreparationPolicy: + max_repair_attempts: int = 2 + timeout_seconds: int = 1800 + max_output_chars: int = 20000 + + +@dataclass(slots=True) +class PreparationContext: + request: FixPreparationRequestV1 + workspace: Path + candidate: FixCandidateV1 + attempt: int = 0 + + +RepairAgent = Callable[ + [PreparationContext, list[CheckResult]], + Awaitable[None], +] +IndependentVerifier = Callable[ + [PreparationContext, list[CheckResult], CheckResult | None], + Awaitable[VerifierResult], +] +SourceVerifier = Callable[[PreparationContext], Awaitable[bool]] + + +async def run_command(workspace: Path, command: CommandSpec) -> CheckResult: + started = time.monotonic() + cwd = (workspace / command.cwd).resolve() + if not cwd.is_relative_to(workspace.resolve()) or not cwd.is_dir(): + return CheckResult( + name=command.name, + argv=command.argv, + status=CheckStatus.UNAVAILABLE, + duration_seconds=time.monotonic() - started, + output="The command working directory is unavailable.", + required=command.required, + ) + process: asyncio.subprocess.Process | None = None + try: + process = await asyncio.create_subprocess_exec( + *command.argv, + cwd=cwd, + env={**os.environ, "PYTHONDONTWRITEBYTECODE": "1"}, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.STDOUT, + ) + output, _ = await asyncio.wait_for(process.communicate(), command.timeout_seconds) + except (FileNotFoundError, PermissionError) as exc: + return CheckResult( + name=command.name, + argv=command.argv, + status=CheckStatus.UNAVAILABLE, + duration_seconds=time.monotonic() - started, + output=str(exc), + required=command.required, + ) + except TimeoutError: + if process is not None: + process.kill() + await process.wait() + return CheckResult( + name=command.name, + argv=command.argv, + status=CheckStatus.FAILED, + duration_seconds=time.monotonic() - started, + output=f"Timed out after {command.timeout_seconds} seconds.", + required=command.required, + ) + return CheckResult( + name=command.name, + argv=command.argv, + status=CheckStatus.PASSED if process.returncode == 0 else CheckStatus.FAILED, + exit_code=process.returncode, + duration_seconds=time.monotonic() - started, + output=output.decode(errors="replace")[-20000:], + required=command.required, + ) + + +def _apply_edits(workspace: Path, candidate: FixCandidateV1) -> None: + by_file: dict[str, list[tuple[int, int, str, str]]] = {} + for edit in candidate.draft_edits: + by_file.setdefault(edit.file, []).append( + (edit.start_line, edit.end_line, edit.before, edit.after) + ) + + for file_path, edits in by_file.items(): + path = workspace / file_path + content = path.read_text(encoding="utf-8") + lines = content.splitlines(keepends=True) + newline = "\r\n" if "\r\n" in content else "\n" + for start, end, before, after in sorted(edits, reverse=True): + original_segment = "".join(lines[start - 1 : end]) + actual = original_segment.rstrip("\r\n") + if actual != before.rstrip("\r\n"): + raise ValueError(f"Draft edit source changed at {file_path}:{start}-{end}") + replacement = after.splitlines(keepends=True) + preserve_newline = original_segment.endswith(("\n", "\r")) or end < len(lines) + if replacement and not replacement[-1].endswith(("\n", "\r")) and preserve_newline: + replacement[-1] += newline + lines[start - 1 : end] = replacement + path.write_text("".join(lines), encoding="utf-8", newline="") + + +async def build_git_manifest( + workspace: Path, +) -> tuple[list[FileManifestEntry], str, str | None]: + process = await asyncio.create_subprocess_exec( + "git", + "status", + "--porcelain=v1", + "-z", + cwd=workspace, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + output, error = await process.communicate() + if process.returncode != 0: + raise RuntimeError(error.decode(errors="replace")) + + entries: list[FileManifestEntry] = [] + changed_files: list[str] = [] + records = [record for record in output.decode(errors="replace").split("\0") if record] + index = 0 + while index < len(records): + record = records[index] + status = record[:2] + path_text = record[3:] + if "R" in status or "C" in status: + index += 1 + path = workspace / path_text + changed_files.append(path_text) + operation: Literal["add", "modify", "delete"] + if status == "??" or "A" in status: + operation = "add" + elif "D" in status: + operation = "delete" + else: + operation = "modify" + resulting = hashlib.sha256(path.read_bytes()).hexdigest() if path.is_file() else None + original: str | None = None + if operation != "add": + original_process = await asyncio.create_subprocess_exec( + "git", + "show", + f"HEAD:{path_text}", + cwd=workspace, + stdout=asyncio.subprocess.PIPE, + stderr=subprocess.DEVNULL, + ) + original_bytes, _ = await original_process.communicate() + if original_process.returncode == 0: + original = hashlib.sha256(original_bytes).hexdigest() + entries.append( + FileManifestEntry( + path=path_text, + operation=operation, + original_sha256=original, + resulting_sha256=resulting, + ) + ) + index += 1 + + summary_process = await asyncio.create_subprocess_exec( + "git", + "diff", + "--stat", + "--", + cwd=workspace, + stdout=asyncio.subprocess.PIPE, + stderr=subprocess.DEVNULL, + ) + summary, _ = await summary_process.communicate() + return entries, summary.decode(errors="replace").strip(), None + + +async def _verify_source(context: PreparationContext) -> bool: + identity = context.candidate.source_identity + if identity is None or identity.kind != "commit": + return identity is not None + process = await asyncio.create_subprocess_exec( + "git", + "rev-parse", + "HEAD", + cwd=context.workspace, + stdout=asyncio.subprocess.PIPE, + stderr=subprocess.DEVNULL, + ) + output, _ = await process.communicate() + return process.returncode == 0 and output.decode().strip().lower() == identity.value + + +def _result( + context: PreparationContext, + *, + state: PreparationState, + reason: str, + started: float, + checks: list[CheckResult] | None = None, + reproduction: CheckResult | None = None, + verifier: VerifierResult | None = None, + gaps: list[str] | None = None, + manifest: list[FileManifestEntry] | None = None, + diff_summary: str = "", + artifact_ref: str | None = None, +) -> FixPreparationResultV1: + return FixPreparationResultV1( + state=state, + stop_reason=reason, + source_identity=context.candidate.source_identity, + candidate_digest=context.candidate.digest(), + final_file_manifest=manifest or [], + artifact_ref=artifact_ref, + changed_files=[entry.path for entry in manifest or []], + diff_summary=diff_summary, + checks=checks or [], + security_reproduction=reproduction, + verifier=verifier, + gaps=gaps or [], + attempts=context.attempt, + elapsed_seconds=time.monotonic() - started, + ) + + +async def prepare_fix( + request: FixPreparationRequestV1, + workspace: Path, + *, + repair: RepairAgent, + verify: IndependentVerifier, + command_runner: CommandRunner = run_command, + manifest_builder: ManifestBuilder = build_git_manifest, + source_verifier: SourceVerifier = _verify_source, + cancelled: CancellationCheck = lambda: False, + policy: PreparationPolicy | None = None, +) -> FixPreparationResultV1: + started = time.monotonic() + resolved_policy = policy or PreparationPolicy( + max_repair_attempts=request.max_repair_attempts, + timeout_seconds=request.timeout_seconds, + ) + context = PreparationContext(request=request, workspace=workspace, candidate=request.candidate) + + async def execute() -> FixPreparationResultV1: # noqa: PLR0911, PLR0912 + if cancelled(): + raise PreparationCancelledError + if not await source_verifier(context): + return _result( + context, + state=PreparationState.STALE, + reason="The workspace does not match the recorded source identity.", + started=started, + ) + + anchored, anchors = anchor_candidate(workspace, context.candidate) + edit_results = anchors[len(context.candidate.finding_locations) :] + if any(result.status is AnchorStatus.STALE for result in edit_results): + return _result( + context, + state=PreparationState.STALE, + reason="A draft edit does not match the recorded source.", + started=started, + ) + if any(result.status is not AnchorStatus.UNIQUE for result in anchors): + gaps = [ + f"{result.location.file}: {result.status}" + for result in anchors + if result.status is not AnchorStatus.UNIQUE + ] + return _result( + context, + state=PreparationState.NEEDS_REVIEW, + reason="One or more candidate locations could not be anchored uniquely.", + gaps=gaps, + started=started, + ) + context.candidate = anchored + _apply_edits(workspace, context.candidate) + + checks: list[CheckResult] = [] + reproduction: CheckResult | None = None + for attempt in range(1, resolved_policy.max_repair_attempts + 1): + context.attempt = attempt + if cancelled(): + raise PreparationCancelledError + await repair(context, checks) + checks = [await command_runner(workspace, check) for check in request.checks] + if context.candidate.reproduction and context.candidate.reproduction.command: + reproduction = await command_runner( + workspace, + context.candidate.reproduction.command, + ) + + failed = any( + result.required and result.status is CheckStatus.FAILED for result in checks + ) + reproduction_failed = ( + reproduction is not None and reproduction.status is CheckStatus.FAILED + ) + if not failed and not reproduction_failed: + break + else: + return _result( + context, + state=PreparationState.FAILED, + reason="Required checks still fail after the repair limit.", + checks=checks, + reproduction=reproduction, + started=started, + ) + + verifier = await verify(context, checks, reproduction) + manifest, summary, artifact_ref = await manifest_builder(workspace) + gaps = [ + result.name + for result in checks + if result.required and result.status is CheckStatus.UNAVAILABLE + ] + if reproduction is None and not verifier.reproduction_executed: + gaps.append("No executable security reproduction was available.") + elif reproduction is not None and reproduction.status is CheckStatus.UNAVAILABLE: + gaps.append(reproduction.name) + gaps.extend(verifier.gaps) + + if ( + verifier.decision is not VerificationDecision.VERIFIED + or not verifier.security_invariant_closed + ): + return _result( + context, + state=PreparationState.NEEDS_REVIEW, + reason=verifier.summary, + checks=checks, + reproduction=reproduction, + verifier=verifier, + gaps=gaps, + manifest=manifest, + diff_summary=summary, + artifact_ref=artifact_ref, + started=started, + ) + if not manifest: + return _result( + context, + state=PreparationState.NEEDS_REVIEW, + reason="The prepared workspace does not contain a source change.", + checks=checks, + reproduction=reproduction, + verifier=verifier, + gaps=[*gaps, "No prepared source change was produced."], + started=started, + ) + if gaps: + return _result( + context, + state=PreparationState.READY_WITH_GAPS, + reason="The fix passed available checks, but required verification gaps remain.", + checks=checks, + reproduction=reproduction, + verifier=verifier, + gaps=gaps, + manifest=manifest, + diff_summary=summary, + artifact_ref=artifact_ref, + started=started, + ) + return _result( + context, + state=PreparationState.READY, + reason="The fix passed required checks and independent verification.", + checks=checks, + reproduction=reproduction, + verifier=verifier, + manifest=manifest, + diff_summary=summary, + artifact_ref=artifact_ref, + started=started, + ) + + try: + async with asyncio.timeout(resolved_policy.timeout_seconds): + return await execute() + except PreparationCancelledError: + return _result( + context, + state=PreparationState.FAILED, + reason="Fix preparation was cancelled.", + started=started, + ) + except TimeoutError: + return _result( + context, + state=PreparationState.FAILED, + reason="Fix preparation exceeded its time limit.", + started=started, + ) diff --git a/strix/interface/utils.py b/strix/interface/utils.py index a61f90a69..76521e3b0 100644 --- a/strix/interface/utils.py +++ b/strix/interface/utils.py @@ -209,7 +209,9 @@ def format_vulnerability_report(report: dict[str, Any]) -> Text: # noqa: PLR091 text.append("\n ") text.append(loc["snippet"], style="dim") if loc.get("fix_before") or loc.get("fix_after"): - text.append("\n Fix:") + preparation = report.get("fix_preparation") + prepared = isinstance(preparation, dict) and preparation.get("state") == "ready" + text.append("\n Prepared fix:" if prepared else "\n Draft fix candidate:") if loc.get("fix_before"): text.append("\n - ", style="dim") text.append(loc["fix_before"], style="dim") diff --git a/strix/report/sarif.py b/strix/report/sarif.py index 821fffbc0..ab97598b8 100644 --- a/strix/report/sarif.py +++ b/strix/report/sarif.py @@ -583,14 +583,14 @@ def _result_properties( def _build_fixes(report: dict[str, Any]) -> list[dict[str, Any]] | None: - """Build SARIF ``fixes`` from a finding's code-location fix pairs. + """Build SARIF ``fixes`` from a prepared finding. - Strix findings carry the suggested change inline on each code - location as ``fix_before`` + ``fix_after``. We map every location - that has both (and a safe repo-relative URI + start line) into a - SARIF ``artifactChange``, replacing the finding's region with the - fixed text. Returns None when no location carries a usable fix pair. + SARIF consumers can apply ``fixes`` automatically. Strix emits them only + after the preparation stage records a ``ready`` result. """ + preparation = report.get("fix_preparation") + if not isinstance(preparation, dict) or preparation.get("state") != "ready": + return None raw_locations = report.get("code_locations") if not isinstance(raw_locations, list): return None diff --git a/strix/report/state.py b/strix/report/state.py index 5d13483e2..adec8c93f 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -93,6 +93,8 @@ UPDATABLE_REPORT_FIELDS = frozenset( "http_exchange_ids", "fix_verification", "fix_pr_body", + "fix_candidate", + "fix_preparation", } ) @@ -357,6 +359,8 @@ class ReportState: fix_pr_body: str | None = None, finding_class: str | None = None, dependency_metadata: dict[str, str] | None = None, + fix_candidate: dict[str, Any] | None = None, + fix_preparation: dict[str, Any] | None = None, agent_id: str | None = None, agent_name: str | None = None, ) -> str: @@ -420,6 +424,10 @@ class ReportState: report["finding_class"] = (finding_class or "dynamic").strip().lower() if dependency_metadata: report["dependency_metadata"] = dependency_metadata + if fix_candidate: + report["fix_candidate"] = fix_candidate + if fix_preparation: + report["fix_preparation"] = fix_preparation if agent_id: report["agent_id"] = agent_id if agent_name: @@ -913,6 +921,9 @@ class ReportState: context["ref"] = f"refs/heads/{branch}" return context + def get_repository_context(self) -> dict[str, Any] | None: + return self._derive_repository_context() + def _sync_llm_usage_record(self) -> None: self.run_record["llm_usage"] = self._build_llm_usage_record() @@ -933,6 +944,7 @@ def openrouter_stream_cost(usage: Any) -> float | None: """ if not isinstance(usage, dict): return None + total = 0.0 cost = usage.get("cost") if isinstance(cost, int | float) and cost > 0: diff --git a/strix/report/writer.py b/strix/report/writer.py index cf858f7d7..6d96f0db6 100644 --- a/strix/report/writer.py +++ b/strix/report/writer.py @@ -332,7 +332,10 @@ def render_vulnerability_md(report: dict[str, Any]) -> str: # noqa: PLR0912, PL lines.extend(f" {ln}" for ln in snippet.splitlines()) lines.append(f" {fence}") if loc.get("fix_before") or loc.get("fix_after"): - lines.append("\n **Suggested Fix:**") + preparation = report.get("fix_preparation") + prepared = isinstance(preparation, dict) and preparation.get("state") == "ready" + label = "Prepared Fix" if prepared else "Draft Fix Candidate" + lines.append(f"\n **{label}:**") lines.append("```diff") if loc.get("fix_before"): lines.extend(f"- {ln}" for ln in str(loc["fix_before"]).splitlines()) @@ -347,10 +350,19 @@ def render_vulnerability_md(report: dict[str, Any]) -> str: # noqa: PLR0912, PL lines.append("") if report.get("fix_verification"): - lines.append("## Fix Verification\n") + lines.append("## Reported Candidate Checks\n") lines.append(str(report["fix_verification"])) lines.append("") + if isinstance(report.get("fix_preparation"), dict): + preparation = report["fix_preparation"] + lines.append("## Fix Preparation\n") + lines.append(f"State: `{preparation.get('state', 'unknown')}`") + if preparation.get("stop_reason"): + lines.append("") + lines.append(str(preparation["stop_reason"])) + lines.append("") + if report.get("assumptions"): lines.append("## Assumptions\n") lines.append(str(report["assumptions"])) diff --git a/strix/tools/reporting/tool.py b/strix/tools/reporting/tool.py index 45f4f6586..ffc6b40fb 100644 --- a/strix/tools/reporting/tool.py +++ b/strix/tools/reporting/tool.py @@ -11,8 +11,8 @@ import asyncio import json import logging import re -from pathlib import PurePosixPath -from typing import TYPE_CHECKING, Any +from pathlib import Path, PurePosixPath +from typing import TYPE_CHECKING, Any, cast from agents import RunContextWrapper, function_tool @@ -21,6 +21,9 @@ from strix.tools.proxy.tools import existing_request_ids if TYPE_CHECKING: + from collections.abc import Mapping + + from strix.fix.contracts import SourceIdentity from strix.report.state import ReportState @@ -342,16 +345,11 @@ def _validate_fix_verification( if str(fix_verification or "").strip(): return [] return [ - "fix_verification is REQUIRED when any code_location carries a 'fix_after' - " - "a suggestion a reviewer can click to apply must be verified first. State, in " - "order: (1) security closure - re-trace the source->sink path through the " - "PATCHED code and say why it is now blocked; (2) bypass review - re-read the " - "diff without your original rationale and name the equivalent sinks, sibling " - "call sites, and alternate malicious input classes you checked; (3) preserved " - "behavior - the legitimate inputs, APIs, and error semantics that still work; " - "(4) how each was checked (executed vs. reasoned), naming any unrun check as " - "an explicit gap. If you cannot make these statements, drop 'fix_after' and " - "leave the location informational." + "fix_verification is REQUIRED when any code_location carries a 'fix_after'. " + "Describe the checks you performed on this draft candidate. Separate executed " + "checks from reasoned checks and list every gap. This statement does not make " + "the candidate ready for automatic application. If the candidate fails a check, " + "revise it or remove 'fix_after'." ] @@ -371,6 +369,100 @@ def _finding_class_of(report: dict[str, Any]) -> str: return "dynamic" +def _fix_source_context( + report_state: ReportState, +) -> tuple[SourceIdentity | None, Path | None]: + from strix.fix.contracts import SourceIdentity, SourceIdentityKind + + raw_context = report_state.get_repository_context() + if raw_context is None: + return None, None + context = cast("dict[str, object]", raw_context) + commit = context.get("commitSha") + if not isinstance(commit, str) or not commit: + return None, None + raw_targets = cast("object", report_state.run_record.get("targets_info")) + targets = cast("list[object]", raw_targets) if isinstance(raw_targets, list) else [] + if len(targets) != 1: + return SourceIdentity(kind=SourceIdentityKind.COMMIT, value=commit), None + raw_target = targets[0] + target = cast("dict[str, object]", raw_target) if isinstance(raw_target, dict) else {} + raw_details = target.get("details") + details = cast("dict[str, object]", raw_details) if isinstance(raw_details, dict) else {} + repo_path = details.get("cloned_repo_path") + repository = context.get("repositoryUri") + return ( + SourceIdentity( + kind=SourceIdentityKind.COMMIT, + value=commit, + repository=repository if isinstance(repository, str) else None, + ), + Path(repo_path) if isinstance(repo_path, str) and repo_path else None, + ) + + +def _build_fix_candidate( + report_state: ReportState, + report_fields: Mapping[str, object], +) -> dict[str, object] | None: + from strix.fix.contracts import candidate_from_legacy_report + from strix.fix.locations import AnchorStatus, anchor_candidate + + source_identity, repo_path = _fix_source_context(report_state) + candidate = candidate_from_legacy_report(report_fields, source_identity=source_identity) + if candidate is None: + return None + if repo_path is not None: + anchored, results = anchor_candidate(repo_path, candidate) + candidate = anchored + gaps = [ + f"{result.location.file}: {result.status}" + for result in results + if result.status is not AnchorStatus.UNIQUE + ] + if gaps: + candidate = candidate.model_copy(update={"known_gaps": [*candidate.known_gaps, *gaps]}) + return cast("dict[str, object]", candidate.model_dump(mode="json")) + + +_FIX_CANDIDATE_FIELDS = frozenset( + { + "code_locations", + "remediation_steps", + "technical_analysis", + "poc_description", + "evidence", + "fix_verification", + } +) + + +def _refresh_fix_candidate( + report_state: ReportState, + report_id: str, + changes: dict[str, Any], +) -> None: + if not _FIX_CANDIDATE_FIELDS.intersection(changes): + return + existing = next( + ( + report + for report in report_state.get_existing_vulnerabilities() + if report.get("id") == report_id + ), + None, + ) + if existing is None: + return + candidate = _build_fix_candidate(report_state, {**existing, **changes}) + if candidate is not None: + changes["fix_candidate"] = candidate + changes["fix_preparation"] = { + "state": "stale", + "stop_reason": "The finding or draft candidate changed after preparation.", + } + + _UPDATE_TEXT_FIELDS = ( "title", "description", @@ -651,6 +743,7 @@ def _do_update( class_error = _fit_revision_to_class(report_state, report_id, changes) if class_error is not None: return class_error + _refresh_fix_candidate(report_state, report_id, changes) try: updated = report_state.update_vulnerability_report( @@ -912,6 +1005,9 @@ async def _do_create( "fix_pr_body": fix_pr_body, "http_exchange_ids": normalized_http_exchange_ids, } + fix_candidate = _build_fix_candidate(report_state, report_fields) + if fix_candidate: + report_fields["fix_candidate"] = fix_candidate dedupe = await check_duplicate(candidate, existing) if dedupe.get("is_duplicate"): @@ -1259,11 +1355,10 @@ async def create_vulnerability_report( ``update_vulnerability_report`` once the proxy responds. Keep IDs out of ``evidence`` and all other report text. - **How ``fix_before`` / ``fix_after`` work**: they're used as - literal GitHub/GitLab PR suggestion blocks. When a reviewer - accepts the suggestion, the platform replaces the **exact - lines from ``start_line`` to ``end_line``** with - ``fix_after``. Therefore: + **How ``fix_before`` / ``fix_after`` work**: they describe an + initial fix candidate. A later preparation stage applies, + tests, repairs, and independently verifies the candidate before + Strix can offer an automatic pull request. Therefore: 1. ``fix_before`` must be a **VERBATIM** copy of the source at those lines — same whitespace, indentation, line @@ -1284,9 +1379,9 @@ async def create_vulnerability_report( before SQL"``). Order primary fix first, supporting changes (imports, config) after. - **Informational vs actionable**: + **Informational vs candidate**: - With ``fix_before`` / ``fix_after``: actionable fix - (renders as a PR suggestion block). + candidate for the preparation stage. - Without them: informational context (e.g. showing the source of tainted data, or a sink that doesn't need direct editing). @@ -1321,11 +1416,10 @@ async def create_vulnerability_report( that aren't part of the fix. - Duplicating the same change across multiple locations. fix_verification: REQUIRED whenever any ``code_locations`` entry - carries a ``fix_after``. A reviewer can apply that - suggestion with one click, so an unverified fix ships - straight into the codebase. Before writing this field, work - the gates **in order** and never trade an earlier one for a - later one: + carries a ``fix_after``. This field records the reporting + agent's checks on the draft candidate. It is not an independent + verification result. Before writing this field, work the gates + **in order**: 1. **Security closure** — re-trace the source → sink path through the *patched* code and state why it is now @@ -1344,11 +1438,9 @@ async def create_vulnerability_report( Then write what you did: the commands you ran and their results, and every gate you could only reason about rather - than execute, marked explicitly as a gap. Do not claim a - gate passed because it looks right. If a gate fails, revise - the patch or drop ``fix_after`` and leave the location - informational — never compensate for a failed security - closure with a smaller diff or extra prose. + than execute, marked explicitly as a gap. Do not claim that a + gate passed because the candidate looks correct. If a gate + fails, revise the candidate or drop ``fix_after``. Also use this field to record the narrowest-complete-change judgement: prefer the smallest repository-native fix that diff --git a/tests/test_fix_preparation.py b/tests/test_fix_preparation.py new file mode 100644 index 000000000..02933cedd --- /dev/null +++ b/tests/test_fix_preparation.py @@ -0,0 +1,384 @@ +"""Tests for fix candidate anchoring and repository preparation.""" + +from __future__ import annotations + +import hashlib +import subprocess +import sys +from typing import TYPE_CHECKING + +import pytest + +from strix.fix.contracts import ( + CandidateLocation, + CheckResult, + CheckStatus, + CommandSpec, + FileManifestEntry, + FixCandidateV1, + FixEdit, + FixPreparationRequestV1, + PreparationState, + ReproductionSpec, + SourceIdentity, + SourceIdentityKind, + VerificationDecision, + VerifierResult, + candidate_from_legacy_report, +) +from strix.fix.locations import AnchorStatus, anchor_location +from strix.fix.prepare import PreparationContext, prepare_fix + + +if TYPE_CHECKING: + from pathlib import Path + + +def _git(workspace: Path, *args: str) -> str: + return subprocess.run( # noqa: S603 + ["/usr/bin/git", *args], + cwd=workspace, + check=True, + capture_output=True, + text=True, + ).stdout.strip() + + +def _workspace(tmp_path: Path) -> tuple[Path, str]: + workspace = tmp_path / "repo" + workspace.mkdir() + _git(workspace, "init") + _git(workspace, "config", "user.email", "test@example.com") + _git(workspace, "config", "user.name", "Test") + (workspace / "app.py").write_text("def result():\n return 'unsafe'\n", encoding="utf-8") + _git(workspace, "add", "app.py") + _git(workspace, "commit", "-m", "initial") + return workspace, _git(workspace, "rev-parse", "HEAD") + + +def _candidate( + commit: str, + *, + reproduction: ReproductionSpec | None = None, +) -> FixCandidateV1: + return FixCandidateV1( + source_identity=SourceIdentity(kind=SourceIdentityKind.COMMIT, value=commit), + security_invariant="Return a safe value.", + finding_locations=[ + CandidateLocation( + file="app.py", + start_line=1, + end_line=2, + snippet="def result():\n return 'unsafe'", + ) + ], + draft_edits=[ + FixEdit( + file="app.py", + start_line=2, + end_line=2, + before=" return 'unsafe'", + after=" return 'safe'", + ) + ], + reproduction=reproduction, + ) + + +def _request(candidate: FixCandidateV1, *, attempts: int = 2) -> FixPreparationRequestV1: + return FixPreparationRequestV1( + scan_id="scan-1", + finding_id="finding-1", + candidate=candidate, + checks=[ + CommandSpec( + name="compile", + argv=[ + sys.executable, + "-c", + "compile(open('app.py', encoding='utf-8').read(), 'app.py', 'exec')", + ], + ) + ], + max_repair_attempts=attempts, + ) + + +async def _noop_repair( + _context: PreparationContext, + _checks: list[CheckResult], +) -> None: + return None + + +async def _verified( + _context: PreparationContext, + _checks: list[CheckResult], + _reproduction: CheckResult | None, +) -> VerifierResult: + return VerifierResult( + decision=VerificationDecision.VERIFIED, + summary="The invariant is closed.", + security_invariant_closed=True, + reproduction_executed=True, + reproduction_summary="The vulnerable input is rejected.", + sibling_paths_reviewed=["app.py"], + preserved_behaviors=["The module compiles."], + ) + + +def test_candidate_from_legacy_report_preserves_draft_and_checks() -> None: + candidate = candidate_from_legacy_report( + { + "remediation_steps": "Reject unsafe input.", + "poc_description": "Call the vulnerable function.", + "fix_verification": "Reasoned about the patched branch.", + "code_locations": [ + { + "file": "src/app.py", + "start_line": 9, + "end_line": 9, + "snippet": "sink(value)", + "fix_before": "sink(value)", + "fix_after": "sink(clean(value))", + } + ], + } + ) + + assert candidate is not None + assert candidate.security_invariant == "Reject unsafe input." + assert candidate.draft_edits[0].after == "sink(clean(value))" + assert candidate.reported_checks[0].executed is False + assert candidate.known_gaps == ["The reporting-agent verification is not independent."] + + +def test_anchor_location_rewrites_invented_line_numbers(tmp_path: Path) -> None: + (tmp_path / "app.py").write_text("first\nsecond\ntarget\nlast\n", encoding="utf-8") + location = CandidateLocation( + file="app.py", + start_line=99, + end_line=99, + snippet="target", + ) + + result = anchor_location(tmp_path, location) + + assert result.status is AnchorStatus.UNIQUE + assert result.location.start_line == 3 + assert result.location.end_line == 3 + + +@pytest.mark.parametrize( + ("content", "expected"), + [ + ("first\nlast\n", AnchorStatus.MISSING), + ("target\nmiddle\ntarget\n", AnchorStatus.AMBIGUOUS), + ], +) +def test_anchor_location_rejects_non_unique_source( + tmp_path: Path, + content: str, + expected: AnchorStatus, +) -> None: + (tmp_path / "app.py").write_text(content, encoding="utf-8") + edit = FixEdit( + file="app.py", + start_line=1, + end_line=1, + before="target", + after="safe", + ) + + assert anchor_location(tmp_path, edit).status is expected + + +def test_anchor_location_detects_stale_file_digest(tmp_path: Path) -> None: + original = "target\n" + (tmp_path / "app.py").write_text("changed\n", encoding="utf-8") + edit = FixEdit( + file="app.py", + start_line=1, + end_line=1, + before="target", + after="safe", + original_sha256=hashlib.sha256(original.encode()).hexdigest(), + ) + + assert anchor_location(tmp_path, edit).status is AnchorStatus.STALE + + +@pytest.mark.asyncio +async def test_prepare_fix_returns_ready_with_manifest(tmp_path: Path) -> None: + workspace, commit = _workspace(tmp_path) + reproduction = ReproductionSpec( + instructions="Confirm that result returns safe.", + command=CommandSpec( + name="security reproduction", + argv=[ + sys.executable, + "-c", + "from app import result; assert result() == 'safe'", + ], + ), + ) + + result = await prepare_fix( + _request(_candidate(commit, reproduction=reproduction)), + workspace, + repair=_noop_repair, + verify=_verified, + ) + + assert result.state is PreparationState.READY + assert result.changed_files == ["app.py"] + assert result.final_file_manifest[0].operation == "modify" + assert result.security_reproduction is not None + assert result.security_reproduction.status is CheckStatus.PASSED + assert (workspace / "app.py").read_text(encoding="utf-8").endswith("return 'safe'\n") + + +@pytest.mark.asyncio +async def test_prepare_fix_retries_failed_checks(tmp_path: Path) -> None: + workspace, commit = _workspace(tmp_path) + calls = 0 + + async def runner(_workspace: Path, command: CommandSpec) -> CheckResult: + nonlocal calls + calls += 1 + return CheckResult( + name=command.name, + argv=command.argv, + status=CheckStatus.FAILED if calls == 1 else CheckStatus.PASSED, + exit_code=1 if calls == 1 else 0, + duration_seconds=0, + required=command.required, + ) + + result = await prepare_fix( + _request(_candidate(commit)), + workspace, + repair=_noop_repair, + verify=_verified, + command_runner=runner, + ) + + assert result.state is PreparationState.READY + assert result.attempts == 2 + assert calls == 2 + + +@pytest.mark.asyncio +async def test_prepare_fix_stops_at_repair_limit(tmp_path: Path) -> None: + workspace, commit = _workspace(tmp_path) + + async def runner(_workspace: Path, command: CommandSpec) -> CheckResult: + return CheckResult( + name=command.name, + argv=command.argv, + status=CheckStatus.FAILED, + exit_code=1, + duration_seconds=0, + required=command.required, + ) + + result = await prepare_fix( + _request(_candidate(commit), attempts=2), + workspace, + repair=_noop_repair, + verify=_verified, + command_runner=runner, + ) + + assert result.state is PreparationState.FAILED + assert result.attempts == 2 + assert "repair limit" in result.stop_reason + + +@pytest.mark.asyncio +async def test_prepare_fix_requires_independent_verifier_approval(tmp_path: Path) -> None: + workspace, commit = _workspace(tmp_path) + + async def rejected( + _context: PreparationContext, + _checks: list[CheckResult], + _reproduction: CheckResult | None, + ) -> VerifierResult: + return VerifierResult( + decision=VerificationDecision.REJECTED, + summary="A sibling path remains vulnerable.", + gaps=["Review the sibling handler."], + ) + + result = await prepare_fix( + _request(_candidate(commit)), + workspace, + repair=_noop_repair, + verify=rejected, + ) + + assert result.state is PreparationState.NEEDS_REVIEW + assert result.verifier is not None + assert result.verifier.decision is VerificationDecision.REJECTED + + +@pytest.mark.asyncio +async def test_prepare_fix_requires_closed_security_invariant(tmp_path: Path) -> None: + workspace, commit = _workspace(tmp_path) + + async def incomplete( + _context: PreparationContext, + _checks: list[CheckResult], + _reproduction: CheckResult | None, + ) -> VerifierResult: + return VerifierResult( + decision=VerificationDecision.VERIFIED, + summary="The local edit works, but the invariant is not closed.", + reproduction_executed=True, + ) + + result = await prepare_fix( + _request(_candidate(commit)), + workspace, + repair=_noop_repair, + verify=incomplete, + ) + + assert result.state is PreparationState.NEEDS_REVIEW + + +@pytest.mark.asyncio +async def test_prepare_fix_requires_a_source_change(tmp_path: Path) -> None: + workspace, commit = _workspace(tmp_path) + + async def empty_manifest( + _workspace: Path, + ) -> tuple[list[FileManifestEntry], str, str | None]: + return [], "No changes.", None + + result = await prepare_fix( + _request(_candidate(commit)), + workspace, + repair=_noop_repair, + verify=_verified, + manifest_builder=empty_manifest, + ) + + assert result.state is PreparationState.NEEDS_REVIEW + assert "source change" in result.stop_reason + + +@pytest.mark.asyncio +async def test_prepare_fix_rejects_wrong_source_commit(tmp_path: Path) -> None: + workspace, _commit = _workspace(tmp_path) + candidate = _candidate("0" * 40) + + result = await prepare_fix( + _request(candidate), + workspace, + repair=_noop_repair, + verify=_verified, + ) + + assert result.state is PreparationState.STALE + assert "source identity" in result.stop_reason diff --git a/tests/test_report_writer.py b/tests/test_report_writer.py index 9b8490101..56c80734b 100644 --- a/tests/test_report_writer.py +++ b/tests/test_report_writer.py @@ -264,4 +264,4 @@ def test_render_vulnerability_md_surfaces_calibration_metadata() -> None: assert "Egress appears filtered at the network layer." in md assert "## Confidence Rationale" in md assert "## What Would Change This Severity" in md - assert "## Fix Verification" in md + assert "## Reported Candidate Checks" in md diff --git a/tests/test_sarif.py b/tests/test_sarif.py index 61ffc39bb..a81bc129f 100644 --- a/tests/test_sarif.py +++ b/tests/test_sarif.py @@ -129,14 +129,15 @@ def test_write_sarif_never_embeds_poc_script(tmp_path: Path) -> None: assert poc["description"] == "Send a crafted request to trigger the sink." -def test_write_sarif_builds_fixes_from_code_location_fix_pairs(tmp_path: Path) -> None: - # A code location carrying fix_before/fix_after must surface as a SARIF - # fix (artifactChange/replacement) so consumers can offer a one-click fix. +def test_write_sarif_builds_fixes_from_ready_preparation(tmp_path: Path) -> None: + # SARIF fixes are automatically applicable, so only a prepared result can + # expose the candidate as an artifactChange. write_sarif( tmp_path, [ _finding( remediation_steps="Use a parameterized query.", + fix_preparation={"state": "ready"}, code_locations=[ { "file": "app.py", @@ -159,6 +160,28 @@ def test_write_sarif_builds_fixes_from_code_location_fix_pairs(tmp_path: Path) - assert replacement["insertedContent"]["text"] == 'query = "SELECT * FROM u WHERE id=%s"' +def test_write_sarif_hides_unprepared_fix_candidate(tmp_path: Path) -> None: + write_sarif( + tmp_path, + [ + _finding( + code_locations=[ + { + "file": "app.py", + "start_line": 4, + "end_line": 4, + "fix_before": "unsafe(value)", + "fix_after": "safe(value)", + } + ], + ) + ], + ) + + result = _read(tmp_path)["runs"][0]["results"][0] + assert "fixes" not in result + + def test_write_sarif_omits_fixes_without_fix_pairs(tmp_path: Path) -> None: write_sarif(tmp_path, [_finding()]) assert "fixes" not in _read(tmp_path)["runs"][0]["results"][0]