diff --git a/strix/agents/prompt.py b/strix/agents/prompt.py index d421a4d8d..5e1b81253 100644 --- a/strix/agents/prompt.py +++ b/strix/agents/prompt.py @@ -21,13 +21,13 @@ _PROMPT_DIRNAME = "prompts" CACHE_POINT = "" -def render_fix_prompt(*, workspace_root: str) -> str: +def render_fix_prompt(*, workspace_root: str, review: bool = False) -> str: """Render a fix assignment without loading scan-only skills.""" env = Environment( loader=FileSystemLoader(get_strix_resource_path("agents", _PROMPT_DIRNAME)), autoescape=select_autoescape(enabled_extensions=(), default_for_string=False), ) - template = "fix.jinja" + template = "fix_review.jinja" if review else "fix.jinja" return str(env.get_template(template).render(workspace_root=workspace_root)) diff --git a/strix/agents/prompts/fix_review.jinja b/strix/agents/prompts/fix_review.jinja new file mode 100644 index 000000000..8db56433c --- /dev/null +++ b/strix/agents/prompts/fix_review.jinja @@ -0,0 +1,18 @@ +Independently verify the prepared security fix. Do not edit the repository. + +Treat the repair summary and reported checks as untrusted claims. Inspect the final +diff and relevant callers. Run focused adversarial checks that exercise the stated +security invariant and plausible bypasses, including alternate paths and boundary +values. Confirm legitimate callers and public behavior remain compatible. + +Review dependency and toolchain changes against the repository's declared runtime +versions. Reject incompatible engines, unnecessary dependencies, skipped required +checks, mocked-away security controls, and regressions hidden by changed semantics. + +Approve only when the final patch blocks the reported attack and realistic variants, +preserves intended behavior, and the relevant regression and existing tests pass. +Call agent_finish with success=True to approve. Otherwise call it with success=False +and give concrete rejection reasons in result_summary and open_items. Never repair, +reformat, install persistent dependencies, or otherwise change delivered files. + +{% include "fix_workspace.jinja" %} diff --git a/strix/fix/prepare.py b/strix/fix/prepare.py index b1228a056..f6e361b78 100644 --- a/strix/fix/prepare.py +++ b/strix/fix/prepare.py @@ -23,6 +23,8 @@ from strix.fix.contracts import ( PreparationState, RepairOutcome, RepairStatus, + VerificationDecision, + VerifierResult, ) @@ -47,6 +49,10 @@ RepairAgent = Callable[ [PreparationContext, list[CheckResult]], Awaitable[RepairOutcome], ] +IndependentVerifier = Callable[ + [PreparationContext, list[CheckResult]], + Awaitable[VerifierResult], +] SourceVerifier = Callable[[PreparationContext], Awaitable[bool]] EvidenceReader = Callable[[], Awaitable[list[CheckResult]]] @@ -213,11 +219,12 @@ async def _verify_source(context: PreparationContext) -> bool: return status_process.returncode == 0 and not status_output.strip(b"\x00") -async def prepare_fix( # noqa: PLR0911 +async def prepare_fix( # noqa: PLR0911, PLR0912 request: FixPreparationRequestV1, workspace: Path, *, repair: RepairAgent, + verify: IndependentVerifier | None = None, manifest_builder: ManifestBuilder = build_git_manifest, source_verifier: SourceVerifier = _verify_source, evidence_reader: EvidenceReader | None = None, @@ -227,6 +234,8 @@ async def prepare_fix( # noqa: PLR0911 started = time.monotonic() context = PreparationContext(request=request, workspace=workspace, candidate=request.candidate) completion: RepairOutcome | None = None + verifier: VerifierResult | None = None + attempt_history: list[FixPreparationAttempt] = [] async def finish(state: PreparationState, reason: str) -> FixPreparationResultV1: manifest: list[FileManifestEntry] = [] @@ -240,9 +249,14 @@ async def prepare_fix( # noqa: PLR0911 candidate=context.candidate, candidate_digest=context.candidate.digest(), completion=completion, + validation_mode="agent_review" if verifier else "single_agent", gaps=list(completion.gaps) if completion else [], prepared_source_digest=( - completion.source_digest if completion and state is PreparationState.READY else None + verifier.source_digest + if verifier and state is PreparationState.READY + else completion.source_digest + if completion and state is PreparationState.READY + else None ), final_file_manifest=manifest, changed_files=[entry.path for entry in manifest], @@ -251,6 +265,8 @@ async def prepare_fix( # noqa: PLR0911 checks=await evidence_reader() if evidence_reader else (completion.command_results if completion else []), + verifier=verifier, + attempt_history=attempt_history, attempts=1 if completion else 0, elapsed_seconds=time.monotonic() - started, ) @@ -273,7 +289,38 @@ async def prepare_fix( # noqa: PLR0911 return await finish(PreparationState.BLOCKED, "The agent produced no patch.") if completion.source_digest != await workspace_digest(workspace): return await finish(PreparationState.BLOCKED, "Source changed after completion.") - return await finish(PreparationState.READY, completion.summary) + if verify is None: + return await finish(PreparationState.READY, completion.summary) + before_review = await workspace_digest(workspace) + attempt = FixPreparationAttempt( + attempt=1, + repair=completion, + checks=await evidence_reader() + if evidence_reader + else list(completion.command_results), + workspace_digest=before_review, + ) + attempt_history.append(attempt) + context.feedback.append(attempt) + verifier = await verify(context, attempt.checks) + attempt.verifier = verifier + after_review = await workspace_digest(workspace) + if after_review != before_review: + return await finish( + PreparationState.BLOCKED, + "The deliverable changed during independent verification.", + ) + if verifier.decision is not VerificationDecision.VERIFIED: + return await finish(PreparationState.BLOCKED, verifier.summary) + if verifier.source_digest != after_review: + return await finish( + PreparationState.BLOCKED, + "Independent verification did not approve the final source snapshot.", + ) + return await finish( + PreparationState.READY, + "Independent verification approved the draft PR.", + ) except PreparationCancelledError: return await finish(PreparationState.FAILED, "Fix preparation was cancelled.") except TimeoutError: diff --git a/strix/fix/runtime.py b/strix/fix/runtime.py index 925d28971..4899f27fb 100644 --- a/strix/fix/runtime.py +++ b/strix/fix/runtime.py @@ -47,6 +47,8 @@ from strix.fix import ( PreparationContext, RepairOutcome, RepairStatus, + VerificationDecision, + VerifierResult, build_git_manifest, build_git_patch, prepare_fix, @@ -85,11 +87,18 @@ def _output_text(text: str, *, max_chars: int | None = _MAX_TOOL_OUTPUT_CHARS) - class _FixHooks(ReportUsageHooks): """Use Strix usage hooks and retain native tool evidence without deciding test success.""" - def __init__(self, environment: _RuntimeEnvironment) -> None: - self.max_turns = min(environment.max_repair_turns, 300) + def __init__( + self, + environment: _RuntimeEnvironment, + *, + max_turns: int | None = None, + track_environment_turns: bool = True, + ) -> None: + self.max_turns = min(max_turns or environment.max_repair_turns, 300) super().__init__(model=load_settings().llm.model or "", max_turns=self.max_turns) self.environment = environment - self.turns = environment.turns_used + self.track_environment_turns = track_environment_turns + self.turns = environment.turns_used if track_environment_turns else 0 self.completion_digest: str | None = None self._recent_commands: deque[tuple[str, int, str]] = deque(maxlen=_REPEAT_WINDOW) self._repetition_warning = False @@ -123,8 +132,9 @@ class _FixHooks(ReportUsageHooks): self._sync_scan_budget() await super().on_llm_start(context, agent, system_prompt, input_items) self.turns += 1 - self.environment.turns_used = self.turns - if self.environment.turn_sink: + if self.track_environment_turns: + self.environment.turns_used = self.turns + if self.track_environment_turns and self.environment.turn_sink: self.environment.turn_sink(self.turns) if self._repetition_warning: input_items.append( @@ -280,6 +290,7 @@ class _RuntimeEnvironment: base_commit: str = "" validated_digest: str | None = None max_repair_turns: int = 300 + max_review_turns: int = 250 turns_used: int = 0 turn_sink: Callable[[int], None] | None = None scan_hooks: ReportUsageHooks | None = None @@ -292,6 +303,7 @@ class _RuntimeEnvironment: usage: LLMUsageLedger = field(default_factory=LLMUsageLedger) coordinator: AgentCoordinator = field(default_factory=AgentCoordinator) pending_commands: dict[int, dict[str, Any]] = field(default_factory=dict[int, dict[str, Any]]) + run_config_factory: Callable[[], RunConfig] | None = None def record_command(self, result: CheckResult) -> None: self.repair_checks.append(result) @@ -457,13 +469,18 @@ class _Completion: recommendations: list[str] = field(default_factory=list[str]) -def build_fix_agent(*, name: str = "Fix agent", workspace_root: str) -> Any: +def build_fix_agent( + *, + name: str = "Fix agent", + workspace_root: str, + review: bool = False, +) -> Any: settings = load_settings() agent = build_strix_agent( name=name, is_root=False, base_tools=[think, stop_process], - instructions_override=render_fix_prompt(workspace_root=workspace_root), + instructions_override=render_fix_prompt(workspace_root=workspace_root, review=review), chat_completions_tools=uses_chat_completions_tool_schema( settings.llm.model or "", settings ), @@ -474,8 +491,9 @@ def build_fix_agent(*, name: str = "Fix agent", workspace_root: str) -> Any: replace( tool, description=( - "Finish this assignment with result_summary and success=True when complete, " - "or success=False when blocked. " + "Finish this assignment with result_summary and success=True only when " + + ("the patch is independently verified, " if review else "the fix is complete, ") + + "or success=False when it must be rejected or is blocked. " "Summarize actual test results, blockers and optional follow-ups." ), timeout_seconds=180, @@ -499,18 +517,29 @@ def build_fix_agent(*, name: str = "Fix agent", workspace_root: str) -> Any: class _FixAgent: """A task adapter around the standard Strix agent, session and lifecycle.""" - def __init__(self, environment: _RuntimeEnvironment) -> None: + def __init__(self, environment: _RuntimeEnvironment, *, review: bool = False) -> None: self.environment = environment - self.agent_id = environment.execution_id - self.hooks = _FixHooks(environment) + self.review = review + self.agent_id = f"{environment.execution_id}-review" if review else environment.execution_id + self.hooks = _FixHooks( + environment, + max_turns=environment.max_review_turns if review else environment.max_repair_turns, + track_environment_turns=not review, + ) self.session = open_agent_session( self.agent_id, environment.workspace.parent / "fix-agents.db" ) - self.agent = build_fix_agent(workspace_root=environment.sandbox_workspace) + self.agent = build_fix_agent( + name="Independent fix verifier" if review else "Fix agent", + workspace_root=environment.sandbox_workspace, + review=review, + ) self.context = { "coordinator": environment.coordinator, "agent_id": self.agent_id, - "parent_id": environment.parent_id or "fix-standalone", + "parent_id": environment.execution_id + if review + else environment.parent_id or "fix-standalone", "sandbox_session": environment.session, "before_agent_finish": self.hooks.before_finish, "interactive": False, @@ -520,12 +549,17 @@ class _FixAgent: start_turns = self.hooks.turns self.hooks.completion_digest = None env = self.environment + parent_id = env.execution_id if self.review else env.parent_id or "fix-standalone" await env.coordinator.register( self.agent_id, self.agent.name, - env.parent_id or "fix-standalone", + parent_id, skills=["fix_task"], - task="Implement and test the confirmed finding", + task=( + "Independently verify the prepared security fix" + if self.review + else "Implement and test the confirmed finding" + ), ) await env.coordinator.attach_runtime( self.agent_id, @@ -544,8 +578,10 @@ class _FixAgent: ) result = await run_agent_loop( agent=self.agent, - initial_input=[] if env.resume else _untrusted_prompt_data(payload), - run_config=_run_config(env), + initial_input=[] + if env.resume and not self.review + else _untrusted_prompt_data(payload), + run_config=env.run_config_factory() if env.run_config_factory else _run_config(env), context=self.context, max_turns=remaining, coordinator=env.coordinator, @@ -629,6 +665,58 @@ class ManagedRepairAgent(_FixAgent): ) +class ManagedIndependentVerifier(_FixAgent): + def __init__(self, environment: _RuntimeEnvironment) -> None: + super().__init__(environment, review=True) + + async def __call__( + self, context: PreparationContext, checks: list[CheckResult] + ) -> VerifierResult: + manifest, _, _ = await build_git_manifest(context.workspace) + patch = (await build_git_patch(context.workspace, manifest)).decode(errors="replace") + first_command = len(self.environment.repair_checks) + completion = await self.run( + { + "finding": _finding_assignment(context), + "repair": context.feedback[-1].repair.model_dump( + mode="json", exclude={"command_results"} + ) + if context.feedback + else None, + "repository_root": self.environment.sandbox_workspace, + "network_allowed": self.environment.network_allowed, + "diff": patch[:150_000], + "diff_truncated": len(patch) > 150_000, + "changed_files": [entry.model_dump(mode="json") for entry in manifest], + "requested_checks": [ + check.model_dump(mode="json") for check in context.request.checks + ], + "checks": [_command_preview(check, max_chars=2000) for check in checks], + } + ) + extra_checks = self.environment.repair_checks[first_command:] + return VerifierResult( + decision=( + VerificationDecision.VERIFIED + if completion.outcome == "done" + else VerificationDecision.REJECTED + ), + summary=completion.summary, + gaps=completion.open_items, + notes=completion.recommendations, + review_basis=( + "execution" + if any( + check.status is CheckStatus.PASSED and check.exit_code == 0 + for check in extra_checks + ) + else "code_review" + ), + source_digest=self.hooks.completion_digest, + turns_used=completion.turns, + ) + + async def _create_command_sandbox( sandbox_id: str, ) -> BaseSandboxSession: @@ -646,6 +734,7 @@ async def build_fix_artifact( environment: _RuntimeEnvironment, artifact_path: Path | None, session: Any, + review_session: Any | None = None, ) -> tuple[list[FileManifestEntry], str, str | None]: manifest, summary, _ = await build_git_manifest(root) if artifact_path is None: @@ -673,6 +762,11 @@ async def build_fix_artifact( json.dumps( { "repair": await session.get_items(), + **( + {"review": await review_session.get_items()} + if review_session is not None + else {} + ), } ), ) @@ -734,18 +828,21 @@ async def run_fix_preparation( return matches environment.max_repair_turns = request.repair_turn_limit + environment.max_review_turns = request.review_turn_limit environment.max_budget_usd = request.max_budget_usd environment.cancelled = cancelled repair = ManagedRepairAgent(environment) + reviewer = ManagedIndependentVerifier(environment) try: result = await prepare_fix( request, environment.workspace, repair=repair, + verify=reviewer, evidence_reader=environment.current_checks, manifest_builder=lambda root: build_fix_artifact( - root, environment, artifact_path, repair.session + root, environment, artifact_path, repair.session, reviewer.session ), source_verifier=verify_source, cancelled=cancelled, @@ -753,6 +850,7 @@ async def run_fix_preparation( return result.model_copy(update={"cost_usd": environment.usage.total_cost}) finally: await repair.close() + await reviewer.close() async def finish_native_fix( @@ -794,16 +892,24 @@ async def finish_native_fix( and environment.base_commit == request.candidate.source_identity.value ) - finished = await prepare_fix( - request, - environment.workspace, - repair=completed, - evidence_reader=environment.current_checks, - manifest_builder=lambda root: build_fix_artifact(root, environment, artifact_path, session), - source_verifier=source_matches, - cancelled=environment.cancelled, - ) - return finished.model_copy(update={"cost_usd": environment.usage.total_cost}) + environment.max_review_turns = request.review_turn_limit + reviewer = ManagedIndependentVerifier(environment) + try: + finished = await prepare_fix( + request, + environment.workspace, + repair=completed, + verify=reviewer, + evidence_reader=environment.current_checks, + manifest_builder=lambda root: build_fix_artifact( + root, environment, artifact_path, session, reviewer.session + ), + source_verifier=source_matches, + cancelled=environment.cancelled, + ) + return finished.model_copy(update={"cost_usd": environment.usage.total_cost}) + finally: + await reviewer.close() async def run_isolated_fix_preparation( diff --git a/strix/fix/scan.py b/strix/fix/scan.py index 1b09f9a85..4a1278c27 100644 --- a/strix/fix/scan.py +++ b/strix/fix/scan.py @@ -74,6 +74,7 @@ class ScanFixes: self.report_state = report_state self.tasks: dict[str, asyncio.Task[Any]] = {} self.dispatches: set[asyncio.Task[Any]] = set() + self.loop: asyncio.AbstractEventLoop | None = None self.closed = False self.base = f"/workspace/.strix-fixes/{hashlib.sha256(scan_id.encode()).hexdigest()[:16]}" self._source_lock = asyncio.Lock() @@ -81,6 +82,7 @@ class ScanFixes: self._staged: set[str] = set() def start(self, spawn: Any, parent_ctx: dict[str, Any]) -> None: + self.loop = asyncio.get_running_loop() self._native_spawn, self._parent_ctx = spawn, parent_ctx self.report_state.finding_persisted_callback = self.notify for report in self.report_state.get_existing_vulnerabilities(): @@ -89,10 +91,44 @@ class ScanFixes: def notify(self, report: dict[str, Any]) -> None: if self.closed: return - task = asyncio.create_task(self._dispatch(str(report["id"]))) + if self.loop is None: + raise RuntimeError("Start the Fix dispatcher from its owning event loop first.") + finding_id = str(report["id"]) + try: + current = asyncio.get_running_loop() + except RuntimeError: + current = None + if current is self.loop: + self._schedule(finding_id) + else: + self.loop.call_soon_threadsafe(self._schedule, finding_id) + + def _schedule(self, finding_id: str) -> None: + if self.closed: + return + task = asyncio.create_task(self._dispatch(finding_id)) self.dispatches.add(task) task.add_done_callback(self.dispatches.discard) + async def _reconcile(self) -> None: + pending = [] + for report in self.report_state.get_existing_vulnerabilities(): + finding_id = str(report["id"]) + try: + _, candidate = self._finding(finding_id) + except ValueError: + continue + previous = self.records.get(finding_id, {}) + running = self.tasks.get(finding_id) + if previous.get("digest") == candidate.digest() and ( + previous.get("status") in {"done", "stopped", "failed"} + or (running and not running.done()) + ): + continue + pending.append(self._dispatch(finding_id)) + if pending: + await asyncio.gather(*pending, return_exceptions=True) + async def _dispatch(self, finding_id: str) -> None: try: report, _ = self._finding(finding_id) @@ -283,6 +319,7 @@ class ScanFixes: network_allowed=True, cancelled=lambda: not self._current(finding_id, digest), ) + environment.run_config_factory = lambda: _run_config(environment) hooks = _FixHooks(environment) started_at = time.monotonic() @@ -371,8 +408,9 @@ class ScanFixes: shutil.rmtree(directory / "source", ignore_errors=True) async def wait(self) -> None: - self.closed = True await asyncio.gather(*self.dispatches, return_exceptions=True) + await self._reconcile() + self.closed = True await asyncio.gather(*self.tasks.values(), return_exceptions=True) async def close(self) -> None: diff --git a/tests/test_fix_preparation.py b/tests/test_fix_preparation.py index b56447431..41452ad58 100644 --- a/tests/test_fix_preparation.py +++ b/tests/test_fix_preparation.py @@ -1,3 +1,5 @@ +# mypy: allow-untyped-defs, allow-untyped-calls, disable-error-code="union-attr" + """Tests for fix candidate anchoring and repository preparation.""" from __future__ import annotations @@ -26,6 +28,8 @@ from strix.fix.contracts import ( ReproductionSpec, SourceIdentity, SourceIdentityKind, + VerificationDecision, + VerifierResult, candidate_from_legacy_report, ) from strix.fix.locations import AnchorStatus, anchor_location @@ -375,6 +379,81 @@ async def test_agent_tests_are_reused_without_controller_execution(tmp_path): assert "Ran 1 test" in result.checks[0].output +@pytest.mark.asyncio +async def test_ready_requires_independent_verifier_approval(tmp_path): + workspace, commit = _workspace(tmp_path) + + async def verify(context, checks): + assert checks + return VerifierResult( + decision=VerificationDecision.VERIFIED, + summary="Attack variants are blocked and legitimate behavior is preserved.", + review_basis="execution", + source_digest=await workspace_digest(context.workspace), + ) + + result = await prepare_fix( + _request(_candidate(commit)), + workspace, + repair=_noop_repair, + verify=verify, + ) + + assert result.state is PreparationState.READY + assert result.validation_mode == "agent_review" + assert result.verifier.decision is VerificationDecision.VERIFIED + assert result.attempt_history[0].verifier == result.verifier + + +@pytest.mark.asyncio +async def test_rejected_independent_verification_blocks_delivery(tmp_path): + workspace, commit = _workspace(tmp_path) + + async def verify(_context, _checks): + return VerifierResult( + decision=VerificationDecision.REJECTED, + summary="A sibling path still exposes arbitrary repository files.", + gaps=["The security invariant is bypassable."], + review_basis="code_review", + ) + + result = await prepare_fix( + _request(_candidate(commit)), + workspace, + repair=_noop_repair, + verify=verify, + ) + + assert result.state is PreparationState.BLOCKED + assert result.verifier.decision is VerificationDecision.REJECTED + assert not result.final_file_manifest + + +@pytest.mark.asyncio +async def test_verifier_cannot_mutate_the_deliverable(tmp_path): + workspace, commit = _workspace(tmp_path) + + async def verify(context, _checks): + (context.workspace / "app.py").write_text("review mutation\n", encoding="utf-8") + return VerifierResult( + decision=VerificationDecision.VERIFIED, + summary="Approved after changing the patch.", + review_basis="code_review", + source_digest=await workspace_digest(context.workspace), + ) + + result = await prepare_fix( + _request(_candidate(commit)), + workspace, + repair=_noop_repair, + verify=verify, + ) + + assert result.state is PreparationState.BLOCKED + assert "changed during independent verification" in result.stop_reason + assert not result.final_file_manifest + + @pytest.mark.asyncio @pytest.mark.parametrize("stop", ["blocked", "exception", "cancel"]) async def test_failed_fix_never_exports_partial_work(tmp_path, stop): @@ -415,7 +494,6 @@ async def test_completed_agent_cannot_deliver_changed_checkpoint(tmp_path): def test_new_command_metadata_does_not_change_existing_finding_digest(tmp_path: Path) -> None: - _root, commit = _workspace(tmp_path) candidate = _candidate(commit) payload = candidate.model_dump(mode="json") diff --git a/tests/test_scan_fixes.py b/tests/test_scan_fixes.py index b5257d309..2a8cb0006 100644 --- a/tests/test_scan_fixes.py +++ b/tests/test_scan_fixes.py @@ -1,3 +1,5 @@ +# mypy: allow-untyped-defs, allow-untyped-calls, disable-error-code="method-assign,var-annotated" + """Persisted findings launch native Fix children in isolated worktrees.""" from __future__ import annotations @@ -104,7 +106,8 @@ async def test_native_parallel_fixes_deliver_patches_and_preserve_scan_source( def config(env): model = models.setdefault( - env.execution_id, ScriptedModel([*patch(), *suite_commands(), finish("done")]) + env.execution_id, + ScriptedModel([*patch(), *suite_commands(), finish("done"), finish("done")]), ) return RunConfig( model=model, sandbox=SandboxRunConfig(session=env.session), tracing_disabled=True @@ -204,7 +207,7 @@ async def test_native_child_keeps_cumulative_turn_cap_and_does_not_export_partia ): fixes, _, _, _, _, context, sessions = setup(tmp_path) fixes.records["finding"] = {"digest": "older-candidate", "turns": 299, "status": "done"} - model = ScriptedModel([*patch(), finish("done")]) + model = ScriptedModel([*patch(), finish("done"), finish("done")]) monkeypatch.setattr( scan_module, "_run_config", @@ -235,7 +238,7 @@ async def test_seven_saved_findings_start_seven_native_children_without_model_ha scan_module, "_run_config", lambda env: RunConfig( - model=ScriptedModel([*patch(), *suite_commands(), finish("done")]), + model=ScriptedModel([*patch(), *suite_commands(), finish("done"), finish("done")]), sandbox=SandboxRunConfig(session=env.session), tracing_disabled=True, ), @@ -270,6 +273,79 @@ async def test_seven_saved_findings_start_seven_native_children_without_model_ha session.close() +@pytest.mark.asyncio +async def test_worker_thread_persistence_starts_exactly_one_native_child(tmp_path, monkeypatch): + fixes, report, _, _, _, context, sessions = setup(tmp_path) + state = ReportState("threaded-handoff") + state._run_dir = tmp_path / "report" + fixes.report_state = state + 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, + ), + ) + stages = [] + + async def sink(stage, saved, _result, _artifact): + stages.append((stage, saved["id"])) + return True + + fixes.sink = sink + fixes.start(fixes._native_spawn, context.context) + finding_id = await asyncio.to_thread( + state.add_vulnerability_report, + title="Unsafe threaded result", + severity="high", + agent_id="reporter", + validation_status="confirmed", + fix_candidate=report["fix_candidate"], + ) + for saved in state.get_existing_vulnerabilities(): + fixes.notify(saved) + await fixes.wait() + assert fixes.records[finding_id]["status"] == "done" + assert [stage for stage in stages if stage == ("started", finding_id)] == [ + ("started", finding_id) + ] + for session in sessions: + session.close() + + +@pytest.mark.asyncio +async def test_wait_reconciles_persisted_candidate_after_missed_notification(tmp_path, monkeypatch): + fixes, report, _, _, _, context, sessions = setup(tmp_path) + state = ReportState("missed-handoff") + state._run_dir = tmp_path / "report" + fixes.report_state = state + 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, + ), + ) + fixes.start(fixes._native_spawn, context.context) + state.finding_persisted_callback = None + finding_id = state.add_vulnerability_report( + title="Unsafe missed result", + severity="high", + agent_id="reporter", + validation_status="confirmed", + fix_candidate=report["fix_candidate"], + ) + assert finding_id not in fixes.records + await fixes.wait() + assert fixes.records[finding_id]["status"] == "done" + for session in sessions: + session.close() + + @pytest.mark.asyncio async def test_persistence_failure_does_not_launch(tmp_path): fixes, report, _, _, _, context, _ = setup(tmp_path)