From b71ed13f4d9de1a09b14e72c345767406f520e7d Mon Sep 17 00:00:00 2001 From: Jonathan Singer Date: Wed, 30 Sep 2026 17:07:08 -0400 Subject: [PATCH] Use native child delegation for finding fixes and strengthen completion guidance --- docs/fix-preparation.md | 11 +- strix/agents/prompts/fix.jinja | 36 +- strix/agents/prompts/fix_workspace.jinja | 6 +- strix/agents/prompts/system_prompt.jinja | 22 +- strix/core/execution.py | 21 +- strix/core/runner.py | 26 +- strix/fix/prepare.py | 1 + strix/fix/runtime.py | 203 ++++++---- strix/fix/scan.py | 472 +++++++++++++---------- strix/interface/fix_cli.py | 2 + strix/report/state.py | 16 +- strix/tools/agents_graph/tools.py | 5 + tests/test_fix_cli.py | 20 + tests/test_fix_preparation.py | 13 + tests/test_runner_teardown.py | 5 +- tests/test_scan_fixes.py | 219 +++++++++-- 16 files changed, 706 insertions(+), 372 deletions(-) diff --git a/docs/fix-preparation.md b/docs/fix-preparation.md index 32edb07ac..0b1899879 100644 --- a/docs/fix-preparation.md +++ b/docs/fix-preparation.md @@ -1,9 +1,9 @@ # Fixing findings during a scan - The assessment investigates and validates an issue, then saves its vulnerability report. -- A confirmed report with an actionable source-backed candidate starts a Fix agent immediately. Unconfirmed reports, duplicate reports, and explicit candidate blockers do not start one. +- After saving a confirmed, source-backed report, the reporting agent calls `create_agent(fix_finding_id=report_id, ...)`. The standard child spawner starts the Fix agent; saving a report alone does not launch work. Unconfirmed reports, duplicate work, and explicit blockers return an actionable tool result. - Each finding gets a Git worktree in the scan's existing sandbox. Fixes run concurrently; the original checkout remains available for assessment and attack chaining. -- One native Strix agent implements the complete fix, adds a regression test, runs it and relevant existing customer unit tests, runs applicable build/lint/type checks, and reviews the change. Test selection and recovery belong to the agent. +- One native Strix child implements the complete fix, retains a regression test, runs it and relevant existing customer unit tests, and runs applicable build/lint/type checks. Before finishing it checks alternate paths to the same attack and affected legitimate callers. Broader suites need a reason; dismissing a relevant failure as pre-existing needs a comparison with the unchanged revision. Test selection and recovery belong to the agent. - The agent calls `agent_finish(outcome="done")` or `agent_finish(outcome="blocked")`. 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 controller binds completion to the final source checkpoint. 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. @@ -12,15 +12,16 @@ ## Implementation - `strix/tools/reporting/tool.py`: persists the finding and its confirmed/unconfirmed validation status. -- `strix/report/state.py`: notifies the scan only after successful finding persistence. -- `strix/fix/scan.py`: deduplicates tasks, creates independent worktrees, preserves turn counts, and joins/cancels tasks during scan teardown. +- `strix/tools/agents_graph/tools.py`: accepts a saved `fix_finding_id` in the existing `create_agent` tool. +- `strix/core/execution.py`: registers, runs, and completes Fix children through the normal child lifecycle. +- `strix/fix/scan.py`: supplies the finding/worktree, deduplicates requests, preserves turn counts, exports completion, and cleans up. It does not schedule fixes from reporting callbacks. - `strix/fix/session.py`: borrows the scan sandbox with a worktree-specific filesystem root and process ownership. Cleanup never stops another agent's processes. - `strix/agents/prompts/fix.jinja`: the single Fix assignment; shared workspace guidance is in `fix_workspace.jinja`. - `strix/fix/runtime.py`: uses `build_strix_agent`, `run_agent_loop`, native tools, persisted sessions, and usage hooks. There is no separate reviewer or custom conversation loop. - `strix/fix/prepare.py`: checks source identity and the completed patch, then exports successful artifacts. - Pro supplies progress/result callbacks. The app registers the inline attempt, stores successful artifacts privately, and creates draft PRs using its existing repository integration. Neither starts another fix sandbox. -The `single_agent` result contract records the Fix agent's completion, command history, final file manifest, and source digest. It does not claim independent verification. Commands include diagnostic failures and superseded attempts; the agent's final summary explains which tests passed and any optional follow-ups. +The `single_agent` result contract exposes the agent's limitations in both `completion.gaps` and top-level `gaps` for compatible readers. It records the Fix agent's completion, command history, final file manifest, and source digest. It does not claim independent verification. Commands include diagnostic failures and superseded attempts; the agent's final summary explains which tests passed and any optional follow-ups. ## Standalone OSS command diff --git a/strix/agents/prompts/fix.jinja b/strix/agents/prompts/fix.jinja index afcd50d16..70c961ebd 100644 --- a/strix/agents/prompts/fix.jinja +++ b/strix/agents/prompts/fix.jinja @@ -1,22 +1,28 @@ -Fix the confirmed vulnerability in your assigned worktree and validate the result. -Use the supplied finding, evidence, and available scan setup details. +Fix the confirmed vulnerability in your assigned worktree. Use the supplied finding, +evidence, suggested edits, and scan setup context. Make the smallest complete fix +that follows repository conventions and preserves legitimate behavior. -Make the smallest complete change that follows repository conventions and preserves -legitimate behavior. Add a focused regression test exercising real application -behavior. Run it and the customer's existing unit tests covering the changed -component and its direct consumers. Both are required. Run applicable lint, -typecheck, or build checks. +Add a regression test exercising the affected application behavior. Do not mock +away the security control being tested. Keep the regression in the delivered patch. +Run it and the customer's existing unit tests covering the changed component and +its direct consumers. Both are required. Run applicable build, lint, or type checks. +Expand to broader suites only when shared behavior or targeted results justify it; +explain why. Find commands in documentation, scripts, and nearby tests; read only +relevant configuration. -Inspect the final change for remaining paths to the reported attack. Correct -problems and rerun affected tests. Keep unrelated hardening as follow-up work. +Before finishing, check that the reported attack is blocked, alternate paths to +that same attack are covered, and affected legitimate callers still work. Trace +relevant callers and entry points. Correct problems and rerun affected tests on the +final code. Keep unrelated hardening outside this fix. -You have at most 300 turns total. Use documented setup and targeted recovery; do -not repeat failed approaches without new evidence. If required tests are missing -or cannot run or pass, report blocked. Incomplete fixes are not delivered. +You have at most 300 turns total. Use documented setup and targeted recovery. +If required tests are missing or cannot run or pass, report blocked. Incomplete +fixes are not delivered. Call agent_finish with outcome done only when the fix is complete and required -checks pass, or blocked when you cannot finish. Summarize the change, test paths, -commands, actual results, and limitations. Put optional follow-ups in -final_recommendations. +checks pass, or blocked when you cannot finish. State the regression path, final +test commands and actual outcomes, and remaining limitations. Distinguish final +checks from failed earlier attempts. Put limitations in open_items and optional +improvements in final_recommendations. {% include "fix_workspace.jinja" %} diff --git a/strix/agents/prompts/fix_workspace.jinja b/strix/agents/prompts/fix_workspace.jinja index 514a1786b..cd0ad7327 100644 --- a/strix/agents/prompts/fix_workspace.jinja +++ b/strix/agents/prompts/fix_workspace.jinja @@ -5,8 +5,10 @@ and test setup. Install needed dependencies, but avoid turning unrelated infrastructure failures into another development project. Attempt a targeted recovery; if still blocked, stop and explain what is needed. -Investigate unrelated test failures only enough to establish whether they occur -without the fix, then document them. Do not repair the repository's entire test +Before dismissing a relevant test failure as pre-existing, reproduce it on the +unchanged revision with the same test and comparable setup in a temporary copy +inside your worktree; preserve the assessment checkout. If you cannot establish +that baseline, report the uncertainty. Do not repair the repository's entire test environment. If required validation remains blocked, stop and report the blocker rather than claiming approval. diff --git a/strix/agents/prompts/system_prompt.jinja b/strix/agents/prompts/system_prompt.jinja index 290e04bd7..ad373962b 100644 --- a/strix/agents/prompts/system_prompt.jinja +++ b/strix/agents/prompts/system_prompt.jinja @@ -119,7 +119,7 @@ WHITE-BOX TESTING (code provided): - 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. - 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. +- Treat the inline changes as a candidate, not as a completed fix. The Fix child you delegate 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): @@ -207,7 +207,7 @@ VALIDATION REQUIREMENTS: - 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 - 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 treat the candidate as a prepared fix. After saving a confirmed source-backed report, delegate implementation and testing to a Fix child through create_agent. - 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. @@ -328,7 +328,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 an initial fix candidate. A dedicated Fix agent starts automatically after a confirmed report and implements/tests the candidate. +3. **WHITE-BOX**: Discovery → Validation → Reporting with an initial fix candidate. After saving a confirmed report, the reporting agent calls create_agent with fix_finding_id set to the returned report ID. 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 @@ -372,15 +372,23 @@ If valid → Spawns "Auth Reporting Agent" (creates the vulnerability report with the initial fix candidate: code_locations fix_before/fix_after + fix_pr_body) ↓ -STOP - a dedicated Fix agent starts automatically for the confirmed report. -It fixes and tests in its own worktree while the assessment continues. +Reporting agent calls create_agent(fix_finding_id=, + name="Auth Fix Agent", task="Fix the reported issue and validate it") +The Fix child implements and tests in its own worktree while assessment continues. ``` CONFIRMED FINDINGS AND FIXES: - Set validation_status="confirmed" only after validation establishes the issue. Use "unconfirmed" for source concerns with unresolved evidence gaps. -- Reporting supplies the finding and initial fix suggestion; the automatically - started Fix agent owns implementation, regression testing and customer unit tests. +- After successfully saving a confirmed, source-backed finding with an actionable + fix candidate, call create_agent with fix_finding_id set to its returned report ID. + Pass reproduction details and working environment/test setup. Check that creation + succeeded before finishing the reporting task; resolve a tool error or report the + blocker to the parent. A duplicate or failed report is not a new fix assignment. +- The Fix child owns implementation, regression testing and customer unit tests. + After revising a candidate, explicitly request its updated Fix child; old work + cannot approve a changed or withdrawn finding. Do not wait for fixes to finish + before completing the assessment. - Keep assessment code unchanged for continued testing and attack chaining. - Finish the assessment when security work is complete. Fix agents may still run; finish_scan publishes the report and the runtime handles their eventual cleanup. diff --git a/strix/core/execution.py b/strix/core/execution.py index 64a9aca94..53e56deed 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -6,7 +6,7 @@ import asyncio import contextlib import logging import uuid -from collections.abc import Callable +from collections.abc import Awaitable, Callable from functools import cache from typing import TYPE_CHECKING, Any, cast @@ -356,12 +356,15 @@ async def spawn_child_agent( parent_history: list[Any], event_sink: StreamEventSink | None = None, hooks: RunHooks[dict[str, Any]] | None = None, + on_complete: Callable[[Any, Any], Awaitable[None]] | None = None, + child_id: str | None = None, ) -> dict[str, Any]: parent_id = parent_ctx.get("agent_id") if not isinstance(parent_id, str): raise TypeError("Parent agent_id missing from context") - child_id = uuid.uuid4().hex[:8] + resuming = child_id is not None + child_id = child_id or uuid.uuid4().hex[:8] child_agent = factory(name=name, skills=skills) await coordinator.register( child_id, @@ -384,7 +387,9 @@ async def spawn_child_agent( name=name, parent_id=parent_id, task=task, - initial_input=child_initial_input( + initial_input=[] + if resuming + else child_initial_input( name=name, child_id=child_id, parent_id=parent_id, @@ -393,6 +398,7 @@ async def spawn_child_agent( ), event_sink=event_sink, hooks=hooks, + on_complete=on_complete, ) return { @@ -1070,6 +1076,7 @@ async def _start_child_runner( start_parked: bool = False, event_sink: StreamEventSink | None = None, hooks: RunHooks[dict[str, Any]] | None = None, + on_complete: Callable[[Any, Any], Awaitable[None]] | None = None, ) -> None: session = open_agent_session(child_id, agents_db_path) sessions_to_close.append(session) @@ -1086,8 +1093,9 @@ async def _start_child_runner( # ``_run_cycle``. Swallow it here so the detached task does not surface a # spurious "Task exception was never retrieved" warning. The root agent # hits the same limit on its next call and tears the scan down. + result = None try: - await run_agent_loop( + result = await run_agent_loop( agent=child_agent, initial_input=initial_input, run_config=run_config, @@ -1106,6 +1114,11 @@ async def _start_child_runner( except SubagentBudgetReservedError: logger.info("child %s stopped at the sub-agent budget reserve", child_id) finally: + if on_complete is not None: + try: + await on_complete(result, session) + except Exception: + logger.exception("child %s completion delivery failed", child_id) if not coordinator.is_shutting_down: await _notify_parent_on_exit(coordinator, child_id) diff --git a/strix/core/runner.py b/strix/core/runner.py index aaef7cecc..ae4e0a7c3 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -499,17 +499,15 @@ async def run_strix_scan( fixes = ScanFixes( session=sandbox_session, coordinator=coordinator, - parent_id=root_id, scan_id=scan_id, state_dir=state_dir, local_sources=local_sources, hooks=hooks, + report_state=report_state, event_sink=event_sink, sink=fix_sink, ) - report_state.fix_finding_callback = fixes.notify - for finding in report_state.get_existing_vulnerabilities(): - fixes.notify(finding) + report_state.defer_completion = True child_agent_builder = make_child_factory( scan_mode=scan_mode, @@ -521,8 +519,8 @@ async def run_strix_scan( system_prompt_context=scope_context, ) - async def spawn_child_agent(**kwargs: Any) -> dict[str, Any]: - return await start_child_agent( + async def native_child(**kwargs: Any) -> dict[str, Any]: + options = dict( # noqa: C408 - merge native defaults with task-specific options coordinator=coordinator, factory=child_agent_builder, agents_db_path=agents_db, @@ -532,8 +530,18 @@ async def run_strix_scan( interactive=interactive, event_sink=event_sink, hooks=hooks, - **kwargs, ) + return await start_child_agent(**{**options, **kwargs}) + + async def spawn_child_agent(**kwargs: Any) -> dict[str, Any]: + finding_id = kwargs.pop("fix_finding_id", None) + if finding_id is not None: + if fixes is None: + raise ValueError( # noqa: TRY301 - actionable tool error + "Fix agents require repository source in a non-interactive scan." + ) + return await fixes.spawn(finding_id, native_child, **kwargs) + return await native_child(**kwargs) context: dict[str, Any] = { "coordinator": coordinator, @@ -553,6 +561,8 @@ async def run_strix_scan( await coordinator.attach_runtime(root_id, session=root_session) if is_resume: + if fixes is not None: + await fixes.restore(native_child, context) await respawn_subagents( coordinator=coordinator, factory=child_agent_builder, @@ -679,7 +689,7 @@ async def run_strix_scan( await fixes.close() report_state = get_global_report_state() if report_state is not None: - report_state.fix_finding_callback = None + report_state.defer_completion = False configure_spill_writer(None) # Settle descendants before closing sessions: on a clean finish a child # can still be mid-turn, and closing its session underneath it crashes it. diff --git a/strix/fix/prepare.py b/strix/fix/prepare.py index 14fd914d7..b1228a056 100644 --- a/strix/fix/prepare.py +++ b/strix/fix/prepare.py @@ -240,6 +240,7 @@ async def prepare_fix( # noqa: PLR0911 candidate=context.candidate, candidate_digest=context.candidate.digest(), completion=completion, + gaps=list(completion.gaps) if completion else [], prepared_source_digest=( completion.source_digest if completion and state is PreparationState.READY else None ), diff --git a/strix/fix/runtime.py b/strix/fix/runtime.py index 9ad03f8af..3d4026a92 100644 --- a/strix/fix/runtime.py +++ b/strix/fix/runtime.py @@ -1,8 +1,4 @@ -"""Repair, run native tests, independently review, and retain a draft PR artifact. - -All agent file tools and commands use one persistent sandbox checkout. The host -checkout is an artifact mirror only, updated at repair and review checkpoints. -""" +"""Fix-agent configuration, execution evidence, and successful patch export.""" from __future__ import annotations @@ -445,6 +441,34 @@ class _Completion: recommendations: list[str] = field(default_factory=list[str]) +def build_fix_agent(*, name: str = "Fix agent", workspace_root: str) -> Any: + settings = load_settings() + agent = build_strix_agent( + name=name, + is_root=False, + base_tools=[think], + instructions_override=render_fix_prompt(workspace_root=workspace_root), + chat_completions_tools=uses_chat_completions_tool_schema( + settings.llm.model or "", settings + ), + strict_tool_schemas=supports_strict_tool_schemas(settings.llm.model or ""), + ) + # Same lifecycle implementation; omit scan-only coverage/reporting guidance. + agent.tools = [ + replace( + tool, + description=( + "Finish this assignment with result_summary and outcome: done or blocked. " + "Summarize actual test results, blockers and optional follow-ups." + ), + ) + if isinstance(tool, FunctionTool) and tool.name == "agent_finish" + else tool + for tool in agent.tools + ] + return agent + + class _FixAgent: """A task adapter around the standard Strix agent, session and lifecycle.""" @@ -456,31 +480,7 @@ class _FixAgent: self.session = open_agent_session( self.agent_id, environment.workspace.parent / "fix-agents.db" ) - settings = load_settings() - self.agent = build_strix_agent( - name="Fix agent", - is_root=False, - base_tools=[think], - instructions_override=render_fix_prompt(workspace_root=environment.sandbox_workspace), - chat_completions_tools=uses_chat_completions_tool_schema( - settings.llm.model or "", settings - ), - strict_tool_schemas=supports_strict_tool_schemas(settings.llm.model or ""), - ) - # Same lifecycle implementation; omit scan-only coverage/reporting guidance. - self.agent.tools = [ - replace( - tool, - description=( - "Finish this assignment with result_summary and outcome: " - + ", ".join(self.outcomes) - + ". Summarize actual test results, blockers and optional follow-ups." - ), - ) - if isinstance(tool, FunctionTool) and tool.name == "agent_finish" - else tool - for tool in self.agent.tools - ] + self.agent = build_fix_agent(workspace_root=environment.sandbox_workspace) self.context = { "coordinator": environment.coordinator, "agent_id": self.agent_id, @@ -613,6 +613,52 @@ async def _create_command_sandbox( return cast("BaseSandboxSession", bundle["session"]) +async def build_fix_artifact( + root: Path, + environment: _RuntimeEnvironment, + artifact_path: Path | None, + session: Any, +) -> tuple[list[FileManifestEntry], str, str | None]: + manifest, summary, _ = await build_git_manifest(root) + if artifact_path is None: + return manifest, summary, None + destination = artifact_path.resolve() + patch_output = await build_git_patch(root, manifest) + with ( + open_secret_file(destination) as stream, + zipfile.ZipFile(stream, mode="w", compression=zipfile.ZIP_DEFLATED) as archive, + ): + archive.writestr( + "manifest.json", + json.dumps( + [entry.model_dump(mode="json") for entry in manifest], + indent=2, + ), + ) + archive.writestr("changes.patch", patch_output) + archive.writestr( + "execution.json", + json.dumps([c.model_dump(mode="json") for c in environment.repair_checks], indent=2), + ) + archive.writestr( + "agent-sessions.json", + json.dumps( + { + "repair": await session.get_items(), + } + ), + ) + tools_path = environment.workspace.parent / "fix-tool-results.jsonl" + if tools_path.exists(): + archive.write(tools_path, "tool-results.jsonl") + for entry in manifest: + if entry.operation == "delete": + continue + source = environment.resolve(entry.path) + archive.write(source, f"files/{entry.path}") + return manifest, summary, str(destination) + + async def run_fix_preparation( request: FixPreparationRequestV1, workspace: Path, @@ -659,50 +705,6 @@ async def run_fix_preparation( await environment.initialize() return matches - async def build_artifact( - root: Path, - ) -> tuple[list[FileManifestEntry], str, str | None]: - manifest, summary, _ = await build_git_manifest(root) - if artifact_path is None: - return manifest, summary, None - destination = artifact_path.resolve() - patch_output = await build_git_patch(root, manifest) - with ( - open_secret_file(destination) as stream, - zipfile.ZipFile(stream, mode="w", compression=zipfile.ZIP_DEFLATED) as archive, - ): - archive.writestr( - "manifest.json", - json.dumps( - [entry.model_dump(mode="json") for entry in manifest], - indent=2, - ), - ) - archive.writestr("changes.patch", patch_output) - archive.writestr( - "execution.json", - json.dumps( - [c.model_dump(mode="json") for c in environment.repair_checks], indent=2 - ), - ) - archive.writestr( - "agent-sessions.json", - json.dumps( - { - "repair": await repair.session.get_items(), - } - ), - ) - tools_path = environment.workspace.parent / "fix-tool-results.jsonl" - if tools_path.exists(): - archive.write(tools_path, "tool-results.jsonl") - for entry in manifest: - if entry.operation == "delete": - continue - source = environment.resolve(entry.path) - archive.write(source, f"files/{entry.path}") - return manifest, summary, str(destination) - environment.max_repair_turns = request.repair_turn_limit environment.max_budget_usd = request.max_budget_usd environment.cancelled = cancelled @@ -714,7 +716,9 @@ async def run_fix_preparation( environment.workspace, repair=repair, evidence_reader=environment.current_checks, - manifest_builder=build_artifact, + manifest_builder=lambda root: build_fix_artifact( + root, environment, artifact_path, repair.session + ), source_verifier=verify_source, cancelled=cancelled, ) @@ -723,6 +727,55 @@ async def run_fix_preparation( await repair.close() +async def finish_native_fix( + request: FixPreparationRequestV1, + environment: _RuntimeEnvironment, + hooks: _FixHooks, + result: Any, + session: Any, + artifact_path: Path, +) -> FixPreparationResultV1: + raw = getattr(result, "final_output", None) + if isinstance(raw, str): + try: + raw = json.loads(raw) + except json.JSONDecodeError: + raw = None + raw = raw if isinstance(raw, dict) else {} + complete = raw.get("agent_completed") and raw.get("outcome") == "done" + completion = RepairOutcome( + status=RepairStatus.COMPLETE if complete else RepairStatus.BLOCKED, + summary=raw.get("summary") + or "The Fix agent stopped without completing required validation.", + gaps=raw.get("open_items") or [], + notes=raw.get("recommendations") or [], + turns_used=hooks.turns, + command_results=environment.repair_checks, + source_digest=hooks.completion_digest, + ) + + async def completed(_context: PreparationContext, _checks: list[CheckResult]) -> RepairOutcome: + return completion + + async def source_matches(_context: PreparationContext) -> bool: + # The mirror was cloned at this revision before the child started. + return bool( + request.candidate.source_identity + 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}) + + async def run_isolated_fix_preparation( request: FixPreparationRequestV1, workspace: Path, diff --git a/strix/fix/scan.py b/strix/fix/scan.py index 5430c84f7..3d3316fda 100644 --- a/strix/fix/scan.py +++ b/strix/fix/scan.py @@ -1,4 +1,4 @@ -"""Run finding-scoped Fix tasks in the live scan sandbox.""" +"""Finding worktrees and delivery for Fix children spawned through create_agent.""" from __future__ import annotations @@ -10,23 +10,28 @@ import json import logging import shutil import subprocess +import time from collections.abc import Awaitable, Callable +from copy import deepcopy from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import Any, cast from strix.fix.contracts import FixCandidateV1, FixPreparationRequestV1, FixPreparationResultV1 -from strix.fix.runtime import _RuntimeEnvironment, run_fix_preparation +from strix.fix.prepare import PreparationContext +from strix.fix.runtime import ( + _finding_assignment, + _FixHooks, + _run_config, + _RuntimeEnvironment, + _untrusted_prompt_data, + build_fix_agent, + finish_native_fix, +) from strix.fix.session import WorktreeSession from strix.fix.workspace import git_metadata_archive from strix.utils.secret_files import open_secret_file -if TYPE_CHECKING: - from agents.sandbox.session import BaseSandboxSession - - from strix.core.agents import AgentCoordinator - from strix.core.hooks import ReportUsageHooks - logger = logging.getLogger(__name__) FixSink = Callable[ [str, dict[str, Any], FixPreparationResultV1 | None, Path | None], Awaitable[bool | None] @@ -37,24 +42,22 @@ class ScanFixes: def __init__( self, *, - session: BaseSandboxSession, - coordinator: AgentCoordinator, - parent_id: str, + session: Any, + coordinator: Any, scan_id: str, state_dir: Path, local_sources: list[dict[str, Any]], - hooks: ReportUsageHooks, + hooks: Any, + report_state: Any, event_sink: Any = None, sink: FixSink | None = None, ) -> None: - self.session, self.coordinator, self.parent_id = session, coordinator, parent_id + self.session, self.coordinator = session, coordinator self.scan_id, self.directory = scan_id, state_dir / "fixes" self.directory.mkdir(parents=True, exist_ok=True, mode=0o700) self.directory.chmod(0o700) self.path = self.directory / "tasks.json" - self.records: dict[str, Any] = ( - json.loads(self.path.read_text()) if self.path.exists() else {} - ) + self.records = json.loads(self.path.read_text()) if self.path.exists() else {} self.sources = [ Path(s["source_path"]).resolve() for s in local_sources @@ -68,11 +71,12 @@ class ScanFixes: if s.get("source_path") } self.hooks, self.event_sink, self.sink = hooks, event_sink, sink - self.tasks: dict[str, asyncio.Task[None]] = {} - self.loop = asyncio.get_running_loop() + self.report_state = report_state + self.tasks: dict[str, asyncio.Task[Any]] = {} self.closed = False self.base = f"/workspace/.strix-fixes/{hashlib.sha256(scan_id.encode()).hexdigest()[:16]}" self._source_lock = asyncio.Lock() + self._finding_locks: dict[str, asyncio.Lock] = {} self._staged: set[str] = set() def _save(self) -> None: @@ -81,55 +85,34 @@ class ScanFixes: stream.write(json.dumps(self.records).encode()) temporary.replace(self.path) - def notify(self, report: dict[str, Any]) -> None: - # Reporting callbacks can execute on a worker thread. - self.loop.call_soon_threadsafe(self._schedule, report) - - def _schedule(self, report: dict[str, Any]) -> None: - if self.closed: - return - finding_id = str(report["id"]) - try: - candidate = FixCandidateV1.model_validate(report.get("fix_candidate")) - eligible = ( - not report.get("deletion") - and report.get("validation_status") == "confirmed" - and candidate.finding is not None - and candidate.finding.validation_status == "confirmed" - and not candidate.blocker - and bool(candidate.draft_edits) - and candidate.source_identity is not None - ) - except ValueError: - eligible = False - candidate = None - digest = candidate.digest() if eligible and candidate else None - previous = self.records.get(finding_id, {}) - if ( - digest - and previous.get("digest") == digest - and (finding_id in self.tasks or previous.get("status") != "running") - ): - return - old_task = self.tasks.get(finding_id) - if old_task and not old_task.done(): - old_task.cancel() - if not digest or not candidate: - if previous: - previous["status"] = "obsolete" - previous.pop("artifact", None) - self._save() - return - used = int(previous.get("turns", 0)) - if used >= 300: - return - resume = previous.get("digest") == digest and previous.get("status") == "running" - self.records[finding_id] = {"digest": digest, "turns": used, "status": "running"} - self._save() - self.tasks[finding_id] = asyncio.create_task( - self._run(report, candidate, resume=resume, previous=old_task), - name=f"fix-{finding_id}", + def _finding(self, finding_id: str) -> tuple[dict[str, Any], FixCandidateV1]: + report = next( + ( + r + for r in self.report_state.get_existing_vulnerabilities() + if str(r.get("id")) == finding_id + ), + None, ) + if report is None: + raise ValueError("Save the vulnerability report before requesting its Fix agent.") + report = deepcopy(report) + candidate = FixCandidateV1.model_validate(report.get("fix_candidate")) + if ( + report.get("validation_status") not in {None, "confirmed"} + or not candidate.finding + or candidate.finding.validation_status != "confirmed" + ): + raise ValueError("Only confirmed findings can start a Fix agent.") + if candidate.blocker or not candidate.draft_edits or not candidate.source_identity: + raise ValueError("The finding needs an unblocked source-backed fix candidate.") + return report, candidate + + def _current(self, finding_id: str, digest: str) -> bool: + try: + return self._finding(finding_id)[1].digest() == digest + except ValueError: + return False async def _emit( self, @@ -138,9 +121,224 @@ class ScanFixes: result: FixPreparationResultV1 | None = None, artifact: Path | None = None, ) -> bool: - if self.sink is None: - return True - return await self.sink(stage, report, result, artifact) is not False + return self.sink is None or await self.sink(stage, report, result, artifact) is not False + + async def spawn(self, finding_id: str, spawn: Any, **kwargs: Any) -> dict[str, Any]: + async with self._finding_locks.setdefault(finding_id, asyncio.Lock()): + return await self._spawn(finding_id, spawn, **kwargs) + + async def _spawn(self, finding_id: str, spawn: Any, **kwargs: Any) -> dict[str, Any]: # noqa: PLR0915 + if self.closed: + raise ValueError("The scan is no longer accepting Fix agents.") + report, candidate = self._finding(finding_id) + assert candidate.source_identity is not None + digest = candidate.digest() + previous = self.records.get(finding_id, {}) + running = self.tasks.get(finding_id) + same = previous.get("digest") == digest + if same and previous.get("agent_id") and (running or previous.get("status") != "running"): + return { + "success": True, + "agent_id": previous["agent_id"], + "status": previous["status"], + "message": "This finding already has a Fix agent.", + } + if running and not running.done(): + running.cancel() + await asyncio.gather(running, return_exceptions=True) + used = int(previous.get("turns", 0)) + if used >= 300: + raise ValueError("This finding has exhausted its 300-turn Fix allowance.") + if len(self.sources) != 1: + raise ValueError("Fix requires one identified Git source checkout.") + resume_id = ( + previous.get("agent_id") if same and previous.get("status") == "running" else None + ) + key = hashlib.sha256(f"{finding_id}:{digest}".encode()).hexdigest()[:24] + directory, root = self.directory / key, f"{self.base}/worktrees/{key}" + artifact = directory / "prepared-fix.zip" + borrowed, base, started = None, None, False + source = self.sources[0] + try: + if not await self._emit("started", report): + raise ValueError( # noqa: TRY301 - resource cleanup must surround setup + "The app declined this fix registration; check the current finding and attempt." + ) + started = True + directory.mkdir(parents=True, exist_ok=True) + mirror = directory / "source" + if not mirror.exists(): + await asyncio.to_thread( + _clone_revision, source, mirror, candidate.source_identity.value + ) + base = await self._stage_base(source) + exists = await self.session.exec("test", "-d", root, shell=False, timeout=30) + if resume_id and exists.exit_code: + raise ValueError( # noqa: TRY301 - setup cleanup boundary + "The previous Fix worktree is unavailable; cannot safely resume." + ) + if exists.exit_code: + await self._exec( + "git", + "-C", + base, + "-c", + "core.hooksPath=/dev/null", + "worktree", + "add", + "--detach", + root, + candidate.source_identity.value, + ) + borrowed = WorktreeSession(self.session, root, f"fix-{key}") + request = FixPreparationRequestV1( + scan_id=self.scan_id, + finding_id=finding_id, + candidate=candidate, + network_allowed=True, + max_agent_turns=300, + ) + record = { + "digest": digest, + "turns": used, + "status": "running", + "parent_id": kwargs["parent_ctx"]["agent_id"], + "name": kwargs["name"], + "task": kwargs["task"], + } + self.records[finding_id] = record + + def turns_used(turns: int) -> None: + record["turns"] = turns + self._save() + + environment = _RuntimeEnvironment( + workspace=mirror, + sandbox_session=borrowed, + sandbox_workspace=root, + base_commit=candidate.source_identity.value, + initialized=True, + execution_id=f"fix-{key}", + coordinator=self.coordinator, + parent_id=record["parent_id"], + turns_used=used, + turn_sink=turns_used, + scan_hooks=self.hooks, + event_sink=self.event_sink, + network_allowed=True, + cancelled=lambda: not self._current(finding_id, digest), + ) + hooks = _FixHooks(environment) + started_at = time.monotonic() + + async def finished(result: Any, session: Any) -> None: + prepared = None + try: + prepared = await finish_native_fix( + request, environment, hooks, result, session, artifact + ) + prepared.elapsed_seconds = time.monotonic() - started_at + await self._emit( + "finished", + report, + prepared, + artifact if prepared.state == "ready" else None, + ) + record["status"] = "done" if prepared.state == "ready" else "stopped" + if prepared.state == "ready": + record["artifact"] = str(artifact) + except Exception: + record["status"] = "stopped" + logger.exception("Fix completion delivery failed for %s", finding_id) + with contextlib.suppress(Exception): + await self._emit("finished", report) + raise + finally: + self._save() + await self._cleanup(borrowed, base, root, directory) + if prepared is None or prepared.state != "ready": + artifact.unlink(missing_ok=True) + + parent_ctx = { + **kwargs["parent_ctx"], + "sandbox_session": borrowed, + "completion_outcomes": ["done", "blocked"], + } + assignment = _untrusted_prompt_data( + { + "finding": _finding_assignment(PreparationContext(request, mirror, candidate)), + "scan_context": { + "assessment_source": self.source_roots[source], + "reproduction": report.get("poc_script_code", ""), + }, + } + ) + spawned = await spawn( + **{ + **kwargs, + "parent_ctx": parent_ctx, + "task": kwargs["task"] + "\n\n" + assignment, + "skills": ["fix_task"], + "factory": lambda **kw: build_fix_agent(name=kw["name"], workspace_root=root), + "run_config": _run_config(environment), + "hooks": hooks, + "max_turns": 300 - used, + "on_complete": finished, + "child_id": resume_id, + }, + ) + record["agent_id"] = spawned["agent_id"] + self.tasks[finding_id] = self.coordinator.runtimes[spawned["agent_id"]].task + self._save() + return cast("dict[str, Any]", spawned) + except BaseException: + if started: + with contextlib.suppress(Exception): + await self._emit("finished", report) + await self._cleanup(borrowed, base, root, directory) + if finding_id in self.records: + self.records[finding_id]["status"] = "stopped" + self._save() + raise + + async def _cleanup(self, borrowed: Any, base: str | None, root: str, directory: Path) -> None: + if borrowed: + with contextlib.suppress(Exception): + await borrowed.pty_terminate_all() + if base: + with contextlib.suppress(Exception): + await self._exec("git", "-C", base, "worktree", "remove", "--force", root) + shutil.rmtree(directory / "source", ignore_errors=True) + + async def restore(self, spawn: Any, parent_ctx: dict[str, Any]) -> None: + # Restore only children explicitly created before interruption, never new findings. + for finding_id, record in list(self.records.items()): + if record.get("status") != "running" or not record.get("agent_id"): + continue + try: + await self.spawn( + finding_id, + spawn, + parent_ctx={**parent_ctx, "agent_id": record["parent_id"]}, + name=record["name"], + task=record["task"], + skills=[], + parent_history=[], + ) + except Exception: + logger.exception("Could not restore Fix child for %s", finding_id) + await self.coordinator.set_status(record["agent_id"], "stopped") + + async def wait(self) -> None: + self.closed = True + await asyncio.gather(*self.tasks.values(), return_exceptions=True) + + async def close(self) -> None: + self.closed = True + for task in self.tasks.values(): + if not task.done(): + task.cancel() + await asyncio.gather(*self.tasks.values(), return_exceptions=True) async def _stage_base(self, source: Path) -> str: key = hashlib.sha256(str(source).encode()).hexdigest()[:16] @@ -167,140 +365,6 @@ class ScanFixes: if result.exit_code: raise RuntimeError(f"Workspace command failed: {argv[0]}") - async def _run( # noqa: PLR0912, PLR0915 - self, - report: dict[str, Any], - candidate: FixCandidateV1, - *, - resume: bool, - previous: asyncio.Task[None] | None, - ) -> None: - finding_id, digest = str(report["id"]), candidate.digest() - key = hashlib.sha256(f"{finding_id}:{digest}".encode()).hexdigest()[:24] - agent_id = f"fix-{key}" - root = f"{self.base}/worktrees/{key}" - directory = self.directory / key - artifact = directory / "prepared-fix.zip" - borrowed: WorktreeSession | None = None - base: str | None = None - result: FixPreparationResultV1 | None = None - started = False - try: - if previous: - await asyncio.gather(previous, return_exceptions=True) - if not await self._emit("started", report): - return - started = True - if len(self.sources) != 1 or candidate.source_identity is None: - raise ValueError("Fix requires one identified Git source checkout") # noqa: TRY301 - source = self.sources[0] - commit = candidate.source_identity.value - directory.mkdir(parents=True, exist_ok=True) - mirror = directory / "source" - if not mirror.exists(): - await asyncio.to_thread(_clone_revision, source, mirror, commit) - base = await self._stage_base(source) - exists = await self.session.exec("test", "-d", root, shell=False, timeout=30) - if resume and exists.exit_code: - raise RuntimeError( # noqa: TRY301 - "The previous fix workspace is unavailable; no automatic restart" - ) - if exists.exit_code: - await self._exec( - "git", - "-C", - base, - "-c", - "core.hooksPath=/dev/null", - "worktree", - "add", - "--detach", - root, - commit, - ) - borrowed = WorktreeSession(self.session, root, agent_id) - - def turns_used(turns: int) -> None: - # One cumulative allowance per finding, including candidate revisions. - self.records[finding_id]["turns"] = turns - self._save() - - request = FixPreparationRequestV1( - scan_id=self.scan_id, - finding_id=finding_id, - candidate=candidate, - network_allowed=True, - max_agent_turns=300, - ) - environment = _RuntimeEnvironment( - workspace=mirror, - sandbox_session=borrowed, - sandbox_workspace=root, - base_commit=commit, - initialized=True, - execution_id=agent_id, - coordinator=self.coordinator, - parent_id=self.parent_id, - turns_used=self.records[finding_id]["turns"], - turn_sink=turns_used, - scan_hooks=self.hooks, - event_sink=self.event_sink, - resume=resume, - scan_context={ - "assessment_source": self.source_roots[source], - "reproduction": report.get("poc_script_code", ""), - }, - ) - result = await run_fix_preparation( - request, - mirror, - sandbox_session=borrowed, - runtime_environment=environment, - artifact_path=artifact, - ) - await self._emit( - "finished", report, result, artifact if result.state == "ready" else None - ) - except asyncio.CancelledError: - raise - except Exception: - logger.exception("Fix task %s stopped", finding_id) - finally: - if started and result is None: - with contextlib.suppress(Exception): - await self._emit("finished", report) - if borrowed: - with contextlib.suppress(Exception): - await borrowed.pty_terminate_all() - if base: - with contextlib.suppress(Exception): - await self._exec("git", "-C", base, "worktree", "remove", "--force", root) - with contextlib.suppress(Exception): - await self.coordinator.set_status(agent_id, "completed") - record = self.records[finding_id] - if record.get("digest") == digest and record.get("status") != "obsolete": - record["status"] = "done" if result and result.state == "ready" else "stopped" - if record["status"] == "done": - record["artifact"] = str(artifact) - self._save() - # Successful artifacts survive; unfinished source never becomes a deliverable. - shutil.rmtree(directory / "source", ignore_errors=True) - if result is None or result.state != "ready": - artifact.unlink(missing_ok=True) - - async def wait(self) -> None: - await asyncio.sleep(0) # Drain report callbacks queued by worker threads. - while pending := [task for task in self.tasks.values() if not task.done()]: - await asyncio.gather(*pending, return_exceptions=True) - self.closed = True - - async def close(self) -> None: - self.closed = True - for task in self.tasks.values(): - if not task.done(): - task.cancel() - await asyncio.gather(*self.tasks.values(), return_exceptions=True) - def _clone_revision(source: Path, mirror: Path, commit: str) -> None: subprocess.run( # noqa: S603 diff --git a/strix/interface/fix_cli.py b/strix/interface/fix_cli.py index ddf4f0de2..81714fbc3 100644 --- a/strix/interface/fix_cli.py +++ b/strix/interface/fix_cli.py @@ -126,6 +126,8 @@ def _summary(result: FixPreparationResultV1) -> str: for check in result.checks ) gaps = result.gaps.copy() + if result.completion: + gaps.extend(result.completion.gaps) if result.verifier: gaps.extend(result.verifier.gaps) if result.blocker: diff --git a/strix/report/state.py b/strix/report/state.py index 22df96a57..4ef975be6 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -239,7 +239,7 @@ class ReportState: self._saved_vuln_ids: set[str] = set() self.caido_url: str | None = None - self.fix_finding_callback: Callable[[dict[str, Any]], None] | None = None + self.defer_completion = False self.vulnerability_found_callback: Callable[[dict[str, Any]], None] | None = None self.vulnerability_updated_callback: Callable[[dict[str, Any]], None] | None = None self.vulnerability_deleted_callback: Callable[[dict[str, Any]], None] | None = None @@ -446,7 +446,6 @@ class ReportState: scarf.finding(severity, cwe=cwe, is_cve=bool(cve)) self.save_run_data() - self._notify_fix(report) return report_id def _deleted_vulnerability_reports(self) -> list[dict[str, Any]]: @@ -563,7 +562,6 @@ class ReportState: ) self.save_run_data() - self._notify_fix(report) return report def delete_vulnerability_report( @@ -645,17 +643,9 @@ class ReportState: except OSError: logger.exception("could not remove %s", md_path) - self._notify_fix({**report, "deletion": entry}) logger.info("Deleted vulnerability report %s - %s", report_id, report.get("title")) return report - def _notify_fix(self, report: dict[str, Any]) -> None: - if self.fix_finding_callback: - try: - self.fix_finding_callback(dict(report)) - except Exception: - logger.exception("Could not schedule fix for %s", report.get("id")) - def get_existing_vulnerabilities(self) -> list[dict[str, Any]]: return list(self.vulnerability_reports) @@ -755,8 +745,8 @@ class ReportState: logger.info("Updated scan final fields") self.run_record["assessment_completed_at"] = datetime.now(UTC).isoformat() - self.save_run_data(mark_complete=self.fix_finding_callback is None) - if self.fix_finding_callback is None: + self.save_run_data(mark_complete=not self.defer_completion) + if not self.defer_completion: posthog.end(self, exit_reason="finished_by_tool") scarf.end(self, exit_reason="finished_by_tool") diff --git a/strix/tools/agents_graph/tools.py b/strix/tools/agents_graph/tools.py index 763870947..ffd6d6ef5 100644 --- a/strix/tools/agents_graph/tools.py +++ b/strix/tools/agents_graph/tools.py @@ -498,6 +498,7 @@ async def create_agent( task: str, inherit_context: bool = True, skills: list[str] | None = None, + fix_finding_id: str | None = None, ) -> str: """Spawn a specialist child agent to run in parallel. @@ -543,6 +544,9 @@ async def create_agent( when starting a clean-slate task. skills: List of skill names (e.g. ``["xss", "sql_injection"]``). Max 5; prefer 1-3. + fix_finding_id: Saved report ID to fix. Starts a Fix child in its own + worktree, with the finding context, required tests and a 300-turn cap. + Use after successfully reporting a confirmed, source-backed issue. """ inner = _ctx(ctx) coordinator = coordinator_from_context(inner) @@ -582,6 +586,7 @@ async def create_agent( task=task, skills=skill_list, parent_history=parent_history, + **({"fix_finding_id": fix_finding_id} if fix_finding_id else {}), ) except Exception as e: logger.exception("create_agent: scan runner failed to spawn child '%s'", name) diff --git a/tests/test_fix_cli.py b/tests/test_fix_cli.py index b1f5ef1bd..eb0c513bb 100644 --- a/tests/test_fix_cli.py +++ b/tests/test_fix_cli.py @@ -330,3 +330,23 @@ def test_summary_includes_followups_from_repair_and_reviewer() -> None: summary = fix_cli._summary(result) assert "re-run the nightly suite" in summary assert "rotate the leaked token" in summary + + +def test_summary_keeps_nested_completion_limitations_for_stored_results(): + candidate = _request("a" * 40).candidate + result = FixPreparationResultV1( + state=PreparationState.READY, + stop_reason="Completed", + source_identity=candidate.source_identity, + candidate=candidate, + candidate_digest=candidate.digest(), + completion=RepairOutcome( + status=RepairStatus.COMPLETE, + summary="Implemented and tested", + gaps=["External integration was not exercised."], + ), + ) + text = fix_cli._summary(result) + assert "External integration was not exercised." in text + result.gaps = list(result.completion.gaps) + assert fix_cli._summary(result).count("External integration was not exercised.") == 1 diff --git a/tests/test_fix_preparation.py b/tests/test_fix_preparation.py index 3fa8e5686..b56447431 100644 --- a/tests/test_fix_preparation.py +++ b/tests/test_fix_preparation.py @@ -426,3 +426,16 @@ def test_new_command_metadata_does_not_change_existing_finding_digest(tmp_path: json.dumps(payload, sort_keys=True, separators=(",", ":")).encode() ).hexdigest() assert candidate.digest() == previous + + +@pytest.mark.asyncio +async def test_completion_limitations_are_preserved_at_top_level(tmp_path): + workspace, commit = _workspace(tmp_path) + + async def agent(context, checks): + completion = await _noop_repair(context, checks) + return completion.model_copy(update={"gaps": ["External integration was not exercised."]}) + + result = await prepare_fix(_request(_candidate(commit)), workspace, repair=agent) + assert result.state is PreparationState.READY + assert result.gaps == result.completion.gaps == ["External integration was not exercised."] diff --git a/tests/test_runner_teardown.py b/tests/test_runner_teardown.py index f95eea3d7..77a115fac 100644 --- a/tests/test_runner_teardown.py +++ b/tests/test_runner_teardown.py @@ -99,7 +99,7 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(monkey events = [] class State: - fix_finding_callback = None + defer_completion = False def __init__(self): self.scan_results = {"scan_completed": True} @@ -114,9 +114,6 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(monkey def __init__(self, **_): pass - def notify(self, _): - pass - async def wait(self): events.append("fixes finished") diff --git a/tests/test_scan_fixes.py b/tests/test_scan_fixes.py index 1facdd02b..187ed10db 100644 --- a/tests/test_scan_fixes.py +++ b/tests/test_scan_fixes.py @@ -1,21 +1,27 @@ -"""Exercise live-scan worktree isolation with native tools and scripted inference.""" +"""Fix delegation through the native create_agent tool, child loop and worktrees.""" from __future__ import annotations import asyncio +import json from pathlib import Path +from types import SimpleNamespace from unittest.mock import AsyncMock import pytest from agents import RunConfig from agents.sandbox import SandboxRunConfig +from agents.tool_context import ToolContext from strix.core.agents import AgentCoordinator +from strix.core.execution import spawn_child_agent from strix.core.hooks import ReportUsageHooks from strix.fix import FindingContext -from strix.fix import runtime as fix_runtime +from strix.fix import scan as scan_module from strix.fix.scan import ScanFixes from strix.fix.session import WorktreeSession +from strix.report.state import ReportState +from strix.tools.agents_graph.tools import create_agent from tests.test_fix_completion import ScriptedModel, finish, patch, suite_commands from tests.test_fix_reliability import LocalSandbox, existing_suite from tests.test_fix_runtime import _git, _request, _workspace @@ -26,16 +32,6 @@ def setup(tmp_path): commit = existing_suite(source) parent = LocalSandbox(tmp_path / "sandbox") coordinator = AgentCoordinator() - fixes = ScanFixes( - session=parent, - coordinator=coordinator, - parent_id="root", - scan_id="scan", - state_dir=tmp_path / "state", - local_sources=[{"source_path": str(source)}], - hooks=ReportUsageHooks(model="test", max_turns=1000), - ) - fixes.base = str(tmp_path / "sandbox" / "fixes") candidate = _request(commit).candidate candidate.finding = FindingContext(validation_status="confirmed", title="Unsafe result") report = { @@ -43,13 +39,70 @@ def setup(tmp_path): "validation_status": "confirmed", "fix_candidate": candidate.model_dump(mode="json"), } - return fixes, report, source, parent + reports = [report] + fixes = ScanFixes( + session=parent, + coordinator=coordinator, + scan_id="scan", + state_dir=tmp_path / "state", + local_sources=[{"source_path": str(source)}], + hooks=ReportUsageHooks(model="test", max_turns=1000), + report_state=SimpleNamespace(get_existing_vulnerabilities=lambda: reports), + ) + fixes.base = str(tmp_path / "sandbox" / "fixes") + sessions = [] + + async def native(**kwargs): + return await spawn_child_agent( + coordinator=coordinator, + agents_db_path=tmp_path / "agents.db", + sessions_to_close=sessions, + interactive=False, + **kwargs, + ) + + async def spawn(**kwargs): + finding_id = kwargs.pop("fix_finding_id") + return await fixes.spawn(finding_id, native, **kwargs) + + context = ToolContext( + tool_name="create_agent", + tool_call_id="spawn-test", + tool_arguments="{}", + context={ + "coordinator": coordinator, + "agent_id": "reporter", + "parent_id": "root", + "spawn_child_agent": spawn, + }, + ) + return fixes, report, source, parent, reports, context, sessions + + +async def delegate(context, finding_id="finding"): + return json.loads( + await create_agent.on_invoke_tool( + context, + json.dumps( + { + "name": "Fix agent", + "task": "Fix the saved issue and run regression and customer tests.", + "skills": [], + "inherit_context": False, + "fix_finding_id": finding_id, + } + ), + ) + ) @pytest.mark.asyncio -async def test_parallel_fixes_use_worktrees_without_modifying_scan_source(tmp_path, monkeypatch): - fixes, report, source, parent = setup(tmp_path) - models = {} +async def test_native_parallel_fixes_deliver_patches_and_preserve_scan_source( + tmp_path, monkeypatch +): + fixes, report, source, parent, reports, context, sessions = setup(tmp_path) + models, stages = {}, [] + reports.append({**report, "id": "another-finding"}) def config(env): model = models.setdefault( @@ -59,38 +112,79 @@ async def test_parallel_fixes_use_worktrees_without_modifying_scan_source(tmp_pa model=model, sandbox=SandboxRunConfig(session=env.session), tracing_disabled=True ) - monkeypatch.setattr(fix_runtime, "_run_config", config) - fixes.notify(report) - fixes.notify({**report, "id": "another-finding"}) + async def sink(stage, report, result, artifact): + stages.append((stage, report["id"], result, artifact)) + return True + + monkeypatch.setattr(scan_module, "_run_config", config) + fixes.sink = sink + first, second = await asyncio.gather(delegate(context), delegate(context, "another-finding")) + assert first["success"] and second["success"], (first, second) + duplicate = await delegate(context) + assert duplicate["agent_id"] == first["agent_id"] await fixes.wait() assert len(models) == 2 assert all(record["status"] == "done" for record in fixes.records.values()), fixes.records assert all(Path(record["artifact"]).exists() for record in fixes.records.values()) + assert all( + fixes.coordinator.parent_of[result["agent_id"]] == "reporter" for result in [first, second] + ) + assert len([s for s in stages if s[0] == "finished" and s[2].state == "ready"]) == 2 assert _git(source, "status", "--porcelain") == "" assert "unsafe" in (source / "app.py").read_text() assert parent.state.manifest.root == str(tmp_path / "sandbox") assert not list((tmp_path / "sandbox/fixes/worktrees").glob("*/app.py")) assert not list((tmp_path / "state/fixes").glob("*/source")) + for session in sessions: + session.close() @pytest.mark.asyncio -async def test_unconfirmed_duplicate_and_exhausted_findings_do_not_start_agent(tmp_path): - fixes, report, _, _ = setup(tmp_path) - fixes._run = AsyncMock() - fixes.notify({**report, "validation_status": "unconfirmed"}) - await asyncio.sleep(0) - fixes._run.assert_not_called() - fixes.notify(report) - fixes.notify(report) +async def test_delegation_errors_reach_reporting_agent_before_any_model_call(tmp_path): + fixes, report, _, _, _, context, _ = setup(tmp_path) + missing = await delegate(context, "unknown") + assert not missing["success"] and "Save the vulnerability" in missing["error"] + report["validation_status"] = "unconfirmed" + assert "confirmed" in (await delegate(context))["error"] + report["validation_status"] = "confirmed" + fixes.records["finding"] = {"turns": 300} + assert "300-turn" in (await delegate(context))["error"] + fixes.records.clear() + fixes.sink = AsyncMock(side_effect=RuntimeError("Missing callback configuration")) + result = await delegate(context) + assert not result["success"] and "Missing callback configuration" in result["error"] + assert not fixes.tasks + + +@pytest.mark.asyncio +async def test_blocked_native_child_has_no_patch(tmp_path, monkeypatch): + fixes, _, _, _, _, context, sessions = setup(tmp_path) + monkeypatch.setattr( + scan_module, + "_run_config", + lambda env: RunConfig( + model=ScriptedModel([*patch(), finish("blocked")]), + sandbox=SandboxRunConfig(session=env.session), + tracing_disabled=True, + ), + ) + assert (await delegate(context))["success"] await fixes.wait() - assert fixes._run.await_count == 1 - fixes.closed = False - fixes.records["finding"]["turns"] = 300 - fixes.records["finding"]["status"] = "running" - fixes.tasks.clear() - fixes.notify(report) - await fixes.wait() - assert fixes._run.await_count == 1 + 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_finding_revision_invalidates_active_completion(tmp_path): + fixes, report, _, _, _, _, _ = setup(tmp_path) + digest = fixes._finding("finding")[1].digest() + assert fixes._current("finding", digest) + report["fix_candidate"]["security_invariant"] = "Revised attack" + assert not fixes._current("finding", digest) + fixes.report_state.get_existing_vulnerabilities().clear() + assert not fixes._current("finding", digest) @pytest.mark.asyncio @@ -104,3 +198,58 @@ async def test_worktree_process_cleanup_never_terminates_parent_sessions(tmp_pat assert parent.exec.call_args.args[3] == "fix-one" with pytest.raises(ValueError, match="another agent"): await child.pty_write_stdin(session_id=123, chars="kill") + + +@pytest.mark.asyncio +async def test_native_child_keeps_cumulative_turn_cap_and_does_not_export_partial_patch( + tmp_path, monkeypatch +): + fixes, _, _, _, _, context, sessions = setup(tmp_path) + fixes.records["finding"] = {"digest": "older-candidate", "turns": 299, "status": "done"} + model = ScriptedModel([*patch(), finish("done")]) + monkeypatch.setattr( + scan_module, + "_run_config", + lambda env: RunConfig( + model=model, + sandbox=SandboxRunConfig(session=env.session), + tracing_disabled=True, + ), + ) + assert (await delegate(context))["success"] + await fixes.wait() + assert fixes.records["finding"]["turns"] == 300 + 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_saving_report_does_not_launch_until_reporting_agent_delegates(tmp_path, monkeypatch): + fixes, report, _, _, _, context, sessions = setup(tmp_path) + state = ReportState("native-handoff") + state._run_dir = tmp_path / "report" + fixes.report_state = state + report_id = state.add_vulnerability_report( + title="Unsafe result", + severity="high", + validation_status="confirmed", + fix_candidate=report["fix_candidate"], + ) + assert not fixes.tasks + monkeypatch.setattr( + scan_module, + "_run_config", + lambda env: RunConfig( + model=ScriptedModel([finish("blocked")]), + sandbox=SandboxRunConfig(session=env.session), + tracing_disabled=True, + ), + ) + spawned = await delegate(context, report_id) + assert spawned["success"], spawned + await fixes.wait() + assert fixes.coordinator.parent_of[spawned["agent_id"]] == "reporter" + for session in sessions: + session.close()