diff --git a/docs/fix-preparation.md b/docs/fix-preparation.md index 0b1899879..fff812618 100644 --- a/docs/fix-preparation.md +++ b/docs/fix-preparation.md @@ -1,21 +1,21 @@ # Fixing findings during a scan - The assessment investigates and validates an issue, then saves its vulnerability report. -- 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. +- Saving a confirmed, source-backed report automatically starts a Fix agent through the standard child spawner. No model handoff is required. Duplicate notifications reuse the current job; changed candidates invalidate it and start a replacement. Unconfirmed findings and explicit blockers are not eligible. - 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 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. +- The agent calls `agent_finish(success=True)` or `agent_finish(success=False)`. The controller enforces **300 total model turns per finding**, including resumed execution and candidate revisions. It does not start a fresh agent after exhaustion. +- The finish tool checkpoints source before completing. Packaging errors return to the agent for correction; three identical completion errors stop the job with that reason. Untracked dependency/cache paths stay out of the patch; new source and tests stay in. Only a completed, nonempty patch becomes an artifact. Blocked, interrupted, or capped work produces no deliverable patch. - Assessment completion publishes the security report. Fixes may continue in the same sandbox; execution and sandbox cleanup finish after all Fix tasks stop. Scan cancellation and the shared model budget also stop fix work. - In the hosted app, successful fixes become available for **user-initiated draft PR creation** on the issue. Incomplete patches are not shown. Internal diagnostic logs and terminal status remain available to operators. ## Implementation - `strix/tools/reporting/tool.py`: persists the finding and its confirmed/unconfirmed validation status. -- `strix/tools/agents_graph/tools.py`: accepts a saved `fix_finding_id` in the existing `create_agent` tool. +- `strix/report/state.py`: notifies the fix launcher only after persistence succeeds. - `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/fix/scan.py`: supplies the finding/worktree, deduplicates requests, preserves turn counts, exports completion, and cleans up. It replays persisted findings on resume. Terminal failures are reported truthfully; an explicit retry starts a fresh attempt using the remaining turn allowance. +- `strix/runtime/agent_session.py`: borrows the scan sandbox with a worktree-specific filesystem root and process ownership. All scan agents get a process scope. Use `stop_process` or Ctrl-C on an owned tool session; broad shell kill commands are rejected. This prevents accidental interference, not hostile code escaping an OS security boundary. - `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. diff --git a/strix/agents/factory.py b/strix/agents/factory.py index f3564feac..5b2bf78df 100644 --- a/strix/agents/factory.py +++ b/strix/agents/factory.py @@ -45,6 +45,7 @@ from strix.tools.notes.tools import ( ) from strix.tools.nullish import is_nullish from strix.tools.output_store import bound_and_store, bound_text +from strix.tools.processes import stop_process from strix.tools.proxy.tools import ( list_requests, list_sitemap, @@ -441,6 +442,14 @@ def _wrap_exec_command(tool: FunctionTool) -> FunctionTool: except (json.JSONDecodeError, TypeError): parsed = None if isinstance(parsed, dict): + # Guard against accidental shared-sandbox cleanup, not adversarial code. + command = str(parsed.get("cmd", "")) + if re.search(r"(?:^|[\s;/|&()`])(?:pkill|killall|kill)(?:\s|$)", command): + return ( + "Use stop_process(pid) for your own background process, " + "or Ctrl-C through write_stdin. " + "Shared-sandbox process cleanup is not allowed." + ) if "shell" not in parsed: parsed["shell"] = "bash" _apply_shell_output_cap(parsed) @@ -567,6 +576,7 @@ def _finish_tool_use_behavior( _BASE_TOOLS: tuple[Tool, ...] = ( think, + stop_process, load_skill, create_todo, list_todos, diff --git a/strix/agents/prompts/fix.jinja b/strix/agents/prompts/fix.jinja index 70c961ebd..cd0a3380e 100644 --- a/strix/agents/prompts/fix.jinja +++ b/strix/agents/prompts/fix.jinja @@ -19,8 +19,8 @@ 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. State the regression path, final +Call agent_finish with success=True only when the fix is complete and required +checks pass, or success=False 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. diff --git a/strix/agents/prompts/fix_workspace.jinja b/strix/agents/prompts/fix_workspace.jinja index cd0ad7327..6ede07914 100644 --- a/strix/agents/prompts/fix_workspace.jinja +++ b/strix/agents/prompts/fix_workspace.jinja @@ -1,6 +1,8 @@ Your worktree is inside the scan sandbox; other agents use different directories. Keep edits and test resources in your assigned worktree. Never change the original -assessment checkout or stop another agent's services. Use distinct ports for services you start. Use the repository's documented runtime +assessment checkout or stop another agent's services. Use distinct ports for services you start. +Stop only your own processes using stop_process or Ctrl-C through write_stdin. +Never use pkill/killall or kill a process merely because it occupies a port. Use the repository's documented runtime 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. diff --git a/strix/agents/prompts/system_prompt.jinja b/strix/agents/prompts/system_prompt.jinja index ad373962b..7290fc62b 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. The Fix child you delegate can inspect and modify any required repository file. +- Treat the inline changes as a candidate, not as a completed fix. The automatically started Fix child 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. After saving a confirmed source-backed report, delegate implementation and testing to a Fix child through create_agent. +- Do not treat the candidate as a prepared fix. After a confirmed source-backed report is saved, the runtime starts a Fix child for implementation and testing. Do not create a separate repair child. - 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. After saving a confirmed report, the reporting agent calls create_agent with fix_finding_id set to the returned report ID. +3. **WHITE-BOX**: Discovery → Validation → Reporting with an initial fix candidate. Saving a confirmed source-backed report automatically starts a Fix child. 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,24 +372,24 @@ If valid → Spawns "Auth Reporting Agent" (creates the vulnerability report with the initial fix candidate: code_locations fix_before/fix_after + fix_pr_body) ↓ -Reporting agent calls create_agent(fix_finding_id=, - name="Auth Fix Agent", task="Fix the reported issue and validate it") +The runtime automatically starts a Fix child after the confirmed finding is saved. 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. -- 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. +- Saving a confirmed, source-backed finding with an actionable candidate automatically + starts a dedicated Fix child. Do not create repair children yourself or edit + assessment source to fix vulnerabilities. Report findings and continue assessment. +- The Fix child owns implementation, regression testing and customer unit tests in + its own worktree. Candidate revisions automatically invalidate and replace old + work. Fixes may continue after the security report is published. - Keep assessment code unchanged for continued testing and attack chaining. +- Other agents share this sandbox. Stop only your own processes using stop_process + or Ctrl-C through your own write_stdin session. Never use process-wide cleanup + such as pkill/killall or terminate a service because it occupies a port; choose + another port instead. - 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 53e56deed..c0cd1f865 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -7,6 +7,7 @@ import contextlib import logging import uuid from collections.abc import Awaitable, Callable +from dataclasses import replace from functools import cache from typing import TYPE_CHECKING, Any, cast @@ -35,6 +36,7 @@ from strix.core.sessions import ( ) from strix.llm import request_log from strix.llm.compaction import is_context_overflow, maybe_compact +from strix.runtime.agent_session import AgentSandboxSession if TYPE_CHECKING: @@ -1086,6 +1088,14 @@ async def _start_child_runner( child_ctx["agent_id"] = child_id child_ctx["parent_id"] = parent_id child_ctx["task"] = task + if run_config.sandbox and run_config.sandbox.session: + sandbox = run_config.sandbox.session + # Fix tasks already supply their worktree and process scope. Normal children + # receive their own scope while keeping the shared assessment directory. + if not parent_ctx.get("before_agent_finish"): + sandbox = AgentSandboxSession(sandbox, sandbox.state.manifest.root, child_id) + child_ctx["sandbox_session"] = sandbox + run_config = replace(run_config, sandbox=replace(run_config.sandbox, session=sandbox)) async def _child_loop() -> None: # A budget stop is a clean scan-wide shutdown, not a child failure: the diff --git a/strix/core/runner.py b/strix/core/runner.py index ae4e0a7c3..dcc48e9fe 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -377,11 +377,16 @@ async def run_strix_scan( prompt_cache=settings.llm.prompt_cache, extra_headers=settings.llm.extra_headers, ) + from strix.runtime.agent_session import AgentSandboxSession + + root_sandbox = AgentSandboxSession( + sandbox_session, sandbox_session.state.manifest.root, root_id + ) run_config = RunConfig( model=resolved_model, model_provider=StrixProvider(), model_settings=model_settings, - sandbox=SandboxRunConfig(client=bundle["client"], session=bundle["session"]), + sandbox=SandboxRunConfig(client=bundle["client"], session=root_sandbox), trace_include_sensitive_data=False, # A hallucinated tool name is a recoverable model mistake, not a scan-ending # error: hand it back as a tool result so the agent can correct itself. @@ -533,25 +538,15 @@ async def run_strix_scan( ) 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, - "sandbox_session": bundle["session"], + "sandbox_session": root_sandbox, "caido_client": bundle["caido_client"], "mcp_registry": mcp_registry, "agent_id": root_id, "parent_id": None, "interactive": interactive, - "spawn_child_agent": spawn_child_agent, + "spawn_child_agent": native_child, "scan_targets": build_scan_targets(scan_config), "max_context_images": settings.runtime.max_context_images, } @@ -560,9 +555,9 @@ async def run_strix_scan( sessions_to_close.append(root_session) await coordinator.attach_runtime(root_id, session=root_session) + if fixes is not None: + fixes.start(native_child, context) if is_resume: - if fixes is not None: - await fixes.restore(native_child, context) await respawn_subagents( coordinator=coordinator, factory=child_agent_builder, @@ -690,6 +685,7 @@ async def run_strix_scan( report_state = get_global_report_state() if report_state is not None: report_state.defer_completion = False + report_state.finding_persisted_callback = None 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/runtime.py b/strix/fix/runtime.py index 3d4026a92..925d28971 100644 --- a/strix/fix/runtime.py +++ b/strix/fix/runtime.py @@ -61,6 +61,7 @@ from strix.fix.workspace import ( ) from strix.report.usage import LLMUsageLedger from strix.runtime import session_manager +from strix.tools.processes import stop_process from strix.tools.thinking.tool import think from strix.utils.secret_files import open_secret_file @@ -92,10 +93,26 @@ class _FixHooks(ReportUsageHooks): self.completion_digest: str | None = None self._recent_commands: deque[tuple[str, int, str]] = deque(maxlen=_REPEAT_WINDOW) self._repetition_warning = False + self.completion_error: str | None = None + self.agent_id = environment.execution_id + self._finish_errors: list[str] = [] + + async def before_finish(self, success: bool) -> str | None: + if not success: + return None + try: + await self.environment.checkpoint() + self.completion_digest = self.environment.validated_digest + except Exception as error: # noqa: BLE001 - recover before lifecycle side effects + return f"Could not package this fix: {error}. Correct the workspace and finish again." + return None async def on_llm_start( self, context: Any, agent: Any, system_prompt: Any, input_items: Any ) -> None: + self.agent_id = str(context.context.get("agent_id", self.agent_id)) + if self.completion_error: + raise RuntimeError(self.completion_error) if self.environment.cancelled(): raise PreparationCancelledError limit = self.environment.max_budget_usd @@ -163,7 +180,8 @@ class _FixHooks(ReportUsageHooks): f"[Fix turn budget] {turns_used}/{self.max_turns} total turns used. " f"{urgency} {action} Do not start new investigations. If required validation is " "incomplete, report it honestly; do not claim approval. Call agent_finish with " - "result_summary and outcome. Incomplete fixes will not be delivered." + "result_summary and success=True or success=False. " + "Incomplete fixes will not be delivered." ) async def on_llm_end( @@ -199,14 +217,12 @@ class _FixHooks(ReportUsageHooks): try: completion = json.loads(raw) except json.JSONDecodeError: - # The SDK returns plain-text schema errors to the agent for correction. - return - if not isinstance(completion, dict): - return - completion = cast("dict[str, Any]", completion) - if completion.get("agent_completed") and completion.get("outcome") == "done": - await env.checkpoint() - self.completion_digest = env.validated_digest + completion = None + if not isinstance(completion, dict) or not completion.get("agent_completed"): + error = str(completion.get("error", raw)) if isinstance(completion, dict) else raw + self._finish_errors.append(error) + if self._finish_errors[-3:] == [error] * 3: + self.completion_error = f"Completion failed three times: {error}" return if context.tool_name not in {"exec_command", "write_stdin"}: return @@ -446,7 +462,7 @@ def build_fix_agent(*, name: str = "Fix agent", workspace_root: str) -> Any: agent = build_strix_agent( name=name, is_root=False, - base_tools=[think], + base_tools=[think, stop_process], instructions_override=render_fix_prompt(workspace_root=workspace_root), chat_completions_tools=uses_chat_completions_tool_schema( settings.llm.model or "", settings @@ -458,9 +474,20 @@ def build_fix_agent(*, name: str = "Fix agent", workspace_root: str) -> Any: replace( tool, description=( - "Finish this assignment with result_summary and outcome: done or blocked. " + "Finish this assignment with result_summary and success=True when complete, " + "or success=False when blocked. " "Summarize actual test results, blockers and optional follow-ups." ), + timeout_seconds=180, + params_json_schema={ + **tool.params_json_schema, + "properties": { + k: v for k, v in tool.params_json_schema["properties"].items() if k != "outcome" + }, + "required": [ + k for k in tool.params_json_schema.get("required", []) if k != "outcome" + ], + }, ) if isinstance(tool, FunctionTool) and tool.name == "agent_finish" else tool @@ -475,7 +502,6 @@ class _FixAgent: def __init__(self, environment: _RuntimeEnvironment) -> None: self.environment = environment self.agent_id = environment.execution_id - self.outcomes = ["done", "blocked"] self.hooks = _FixHooks(environment) self.session = open_agent_session( self.agent_id, environment.workspace.parent / "fix-agents.db" @@ -486,7 +512,7 @@ class _FixAgent: "agent_id": self.agent_id, "parent_id": environment.parent_id or "fix-standalone", "sandbox_session": environment.session, - "completion_outcomes": self.outcomes, + "before_agent_finish": self.hooks.before_finish, "interactive": False, } @@ -534,14 +560,9 @@ class _FixAgent: completion = json.loads(completion) if isinstance(completion, dict): completed = cast("dict[str, Any]", completion) - outcome = completed.get("outcome") - if ( - completed.get("agent_completed") - and isinstance(outcome, str) - and outcome in self.outcomes - ): + if completed.get("agent_completed"): return _Completion( - outcome, + "done" if completed.get("task_success") is True else "blocked", str(completed.get("summary", "")), self.hooks.turns - start_turns, open_items=list(completed.get("open_items") or []), @@ -549,10 +570,17 @@ class _FixAgent: ) return _Completion( "blocked", - "The agent stopped without a completion outcome; " - "no incomplete fix will be delivered.", + self.hooks.completion_error + or env.coordinator.errors.get(self.agent_id) + or "The agent stopped before completing; no incomplete fix will be delivered.", self.hooks.turns - start_turns, ) + except RuntimeError: + if not self.hooks.completion_error: + raise + return _Completion( + "blocked", self.hooks.completion_error, self.hooks.turns - start_turns + ) except (MaxTurnsExceeded, BudgetExceededError): return _Completion( "blocked", @@ -742,10 +770,12 @@ async def finish_native_fix( except json.JSONDecodeError: raw = None raw = raw if isinstance(raw, dict) else {} - complete = raw.get("agent_completed") and raw.get("outcome") == "done" + complete = raw.get("agent_completed") and raw.get("task_success") is True + failure = hooks.completion_error or environment.coordinator.errors.get(hooks.agent_id) completion = RepairOutcome( status=RepairStatus.COMPLETE if complete else RepairStatus.BLOCKED, summary=raw.get("summary") + or failure or "The Fix agent stopped without completing required validation.", gaps=raw.get("open_items") or [], notes=raw.get("recommendations") or [], diff --git a/strix/fix/scan.py b/strix/fix/scan.py index 3d3316fda..1b09f9a85 100644 --- a/strix/fix/scan.py +++ b/strix/fix/scan.py @@ -1,4 +1,4 @@ -"""Finding worktrees and delivery for Fix children spawned through create_agent.""" +"""Persisted findings start Fix children through the native agent lifecycle.""" from __future__ import annotations @@ -73,12 +73,58 @@ class ScanFixes: self.hooks, self.event_sink, self.sink = hooks, event_sink, sink self.report_state = report_state self.tasks: dict[str, asyncio.Task[Any]] = {} + self.dispatches: set[asyncio.Task[Any]] = set() 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 start(self, spawn: Any, parent_ctx: dict[str, Any]) -> None: + 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(): + self.notify(report) + + def notify(self, report: dict[str, Any]) -> None: + if self.closed: + return + task = asyncio.create_task(self._dispatch(str(report["id"]))) + self.dispatches.add(task) + task.add_done_callback(self.dispatches.discard) + + async def _dispatch(self, finding_id: str) -> None: + try: + report, _ = self._finding(finding_id) + parent_id = report.get("agent_id") or self._parent_ctx["agent_id"] + await self.spawn( + finding_id, + self._native_spawn, + parent_ctx={**self._parent_ctx, "agent_id": parent_id}, + name=f"Fix: {report.get('title', finding_id)}", + task="Implement and validate the saved finding in your assigned worktree.", + skills=[], + parent_history=[], + ) + except Exception as error: # noqa: BLE001 - report launch failure to the scan + active = self.tasks.get(finding_id) + if active and not active.done(): + active.cancel() + await asyncio.gather(active, return_exceptions=True) + logger.warning("fix.dispatch finding=%s rejected=%s", finding_id, error) + await self.coordinator.send( + self._parent_ctx["agent_id"], + { + "from": "fix-runtime", + "type": "information", + "priority": "normal", + "content": ( + f"Fix for {finding_id} was not started: {error}. " + "Do not create a replacement scan child to repair it." + ), + }, + ) + def _save(self) -> None: temporary = self.path.with_suffix(".tmp") with open_secret_file(temporary) as stream: @@ -97,7 +143,9 @@ class ScanFixes: 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 not report.get("fix_candidate"): + raise ValueError("The finding has no source-backed fix candidate.") + candidate = FixCandidateV1.model_validate(report["fix_candidate"]) if ( report.get("validation_status") not in {None, "confirmed"} or not candidate.finding @@ -128,21 +176,28 @@ class ScanFixes: 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"): + if ( + same + and previous.get("agent_id") + and ((running and not running.done()) or previous.get("status") == "done") + ): return { "success": True, "agent_id": previous["agent_id"], "status": previous["status"], "message": "This finding already has a Fix agent.", } + retry = kwargs.pop("retry", False) + if same and previous.get("status") in {"stopped", "failed"} and not retry: + raise ValueError( + previous.get("reason") or "The previous Fix attempt failed; it was not restarted." + ) if running and not running.done(): running.cancel() await asyncio.gather(running, return_exceptions=True) @@ -160,7 +215,7 @@ class ScanFixes: borrowed, base, started = None, None, False source = self.sources[0] try: - if not await self._emit("started", report): + if not await self._emit("retrying" if retry and same else "started", report): raise ValueError( # noqa: TRY301 - resource cleanup must surround setup "The app declined this fix registration; check the current finding and attempt." ) @@ -238,17 +293,21 @@ class ScanFixes: request, environment, hooks, result, session, artifact ) prepared.elapsed_seconds = time.monotonic() - started_at - await self._emit( + delivered = await self._emit( "finished", report, prepared, artifact if prepared.state == "ready" else None, ) + if not delivered: + raise RuntimeError("The app declined the completed Fix result") # noqa: TRY301 record["status"] = "done" if prepared.state == "ready" else "stopped" + record["reason"] = prepared.stop_reason if prepared.state == "ready": record["artifact"] = str(artifact) - except Exception: + except Exception as error: record["status"] = "stopped" + record["reason"] = str(error) logger.exception("Fix completion delivery failed for %s", finding_id) with contextlib.suppress(Exception): await self._emit("finished", report) @@ -262,7 +321,7 @@ class ScanFixes: parent_ctx = { **kwargs["parent_ctx"], "sandbox_session": borrowed, - "completion_outcomes": ["done", "blocked"], + "before_agent_finish": hooks.before_finish, } assignment = _untrusted_prompt_data( { @@ -291,13 +350,14 @@ class ScanFixes: self.tasks[finding_id] = self.coordinator.runtimes[spawned["agent_id"]].task self._save() return cast("dict[str, Any]", spawned) - except BaseException: + except BaseException as error: 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.records[finding_id]["reason"] = str(error) self._save() raise @@ -310,31 +370,16 @@ class ScanFixes: 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.dispatches, return_exceptions=True) await asyncio.gather(*self.tasks.values(), return_exceptions=True) async def close(self) -> None: self.closed = True + for task in self.dispatches: + task.cancel() + await asyncio.gather(*self.dispatches, return_exceptions=True) for task in self.tasks.values(): if not task.done(): task.cancel() diff --git a/strix/fix/session.py b/strix/fix/session.py index da0b011db..2ce50429b 100644 --- a/strix/fix/session.py +++ b/strix/fix/session.py @@ -1,125 +1,6 @@ -"""A worktree-scoped view of a scan sandbox; it never owns the backend.""" +"""Compatibility name for the worktree view of an agent sandbox.""" -from __future__ import annotations - -import copy -import uuid -from typing import TYPE_CHECKING, Any - -from agents.sandbox.session import BaseSandboxSession +from strix.runtime.agent_session import AgentSandboxSession -if TYPE_CHECKING: - from pathlib import Path - - -class WorktreeSession(BaseSandboxSession): - def __init__(self, parent: BaseSandboxSession, root: str, task_id: str) -> None: - self.parent = parent - self.state = copy.deepcopy(parent.state) - self.state.manifest.root = root - self.state.manifest.entries = {} - self.state.session_id = uuid.uuid5(parent.state.session_id, task_id) - self.task_id = task_id - self.processes: set[int] = set() - - def _command(self, command: tuple[Any, ...]) -> list[str]: - return [ - "env", - f"STRIX_FIX_TASK={self.task_id}", - "sh", - "-c", - 'cd -- "$1"; shift; exec "$@"', - "sh", - self.state.manifest.root, - *map(str, command), - ] - - async def _exec_internal(self, *command: Any, timeout: float | None = None) -> Any: - return await self.parent.exec( - *self._command(command), - shell=False, - timeout=min(timeout or 600, 600), - ) - - def supports_pty(self) -> bool: - return self.parent.supports_pty() - - async def pty_exec_start(self, *command: Any, **kwargs: Any) -> Any: - # Intentional native sandbox shell, never host execution. - prepared = self._prepare_exec_command( # nosec B604 - *command, - shell=kwargs.pop("shell", True), - user=kwargs.pop("user", None), - ) - kwargs["timeout"] = min(kwargs.get("timeout") or 600, 600) - update = await self.parent.pty_exec_start( - *self._command(tuple(prepared)), - shell=False, - **kwargs, - ) - if update.process_id is not None: - self.processes.add(update.process_id) - return update - - async def pty_write_stdin(self, *, session_id: int, **kwargs: Any) -> Any: - if session_id not in self.processes: - raise ValueError("That process belongs to another agent") - return await self.parent.pty_write_stdin(session_id=session_id, **kwargs) - - async def read(self, path: Path, **kwargs: Any) -> Any: - return await self.parent.read(self.normalize_path(path), **kwargs) - - async def write(self, path: Path, data: Any, **kwargs: Any) -> None: - await self.parent.write(self.normalize_path(path, for_write=True), data, **kwargs) - - async def running(self) -> bool: - return await self.parent.running() - - async def hydrate_workspace(self, _data: Any) -> None: - raise RuntimeError("A borrowed worktree cannot restore the scan workspace") - - async def persist_workspace(self) -> Any: - raise RuntimeError("The scan owns workspace persistence") - - async def stop(self) -> None: - await self.pty_terminate_all() - - async def shutdown(self) -> None: - await self.pty_terminate_all() - - async def pty_terminate_all(self) -> None: - # Background services inherit this task marker too. Never call the - # parent session's terminate_all(), which would kill assessment tools. - await self.parent.exec( - "python", - "-c", - _STOP_OWN_PROCESSES, - self.task_id, - shell=False, - timeout=15, - ) - self.processes.clear() - - -_STOP_OWN_PROCESSES = r""" -import os, pathlib, signal, sys, time -marker = b"STRIX_FIX_TASK=" + sys.argv[1].encode() -def owned(): - found = [] - for path in pathlib.Path("/proc").glob("[0-9]*/environ"): - try: - if marker in path.read_bytes().split(b"\0"): - found.append(int(path.parent.name)) - except (OSError, ValueError): - pass - return found -for sig in (signal.SIGTERM, signal.SIGKILL): - for pid in owned(): - try: - os.kill(pid, sig) - except ProcessLookupError: - pass - if sig == signal.SIGTERM: - time.sleep(0.2) -""" +WorktreeSession = AgentSandboxSession diff --git a/strix/fix/workspace.py b/strix/fix/workspace.py index d7186da68..dd0a70826 100644 --- a/strix/fix/workspace.py +++ b/strix/fix/workspace.py @@ -144,6 +144,17 @@ current = set( ) - {""} paths = original | current +# Only untracked environment artifacts are excluded. A tracked dependency, +# generated file, or genuine source symlink remains part of the source delta. +environment_roots = { + "node_modules", ".venv", "venv", "__pycache__", ".pytest_cache", + ".mypy_cache", ".ruff_cache", "coverage", ".coverage", ".nyc_output", +} +current = { + name for name in current + if name in original or not environment_roots.intersection(pathlib.PurePosixPath(name).parts) +} + def safe(name): p = root / name @@ -164,6 +175,10 @@ changed = set( git("diff", "--name-only", "--no-renames", "-z", base).decode().split("\0") ) - {""} changed |= current - original +changed = { + name for name in changed + if name in original or not environment_roots.intersection(pathlib.PurePosixPath(name).parts) +} changed.discard(str(pathlib.Path(sys.argv[3]).relative_to(root))) manifest = [] with tarfile.open(sys.argv[3], "w") as archive: diff --git a/strix/report/state.py b/strix/report/state.py index 4ef975be6..6d22a3700 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -240,6 +240,7 @@ class ReportState: self.caido_url: str | None = None self.defer_completion = False + self.finding_persisted_callback: Callable[[dict[str, Any]], None] | None = None 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,6 +447,7 @@ class ReportState: scarf.finding(severity, cwe=cwe, is_cve=bool(cve)) self.save_run_data() + self.notify_finding_persisted(report) return report_id def _deleted_vulnerability_reports(self) -> list[dict[str, Any]]: @@ -562,6 +564,7 @@ class ReportState: ) self.save_run_data() + self.notify_finding_persisted(report) return report def delete_vulnerability_report( @@ -644,8 +647,14 @@ class ReportState: logger.exception("could not remove %s", md_path) logger.info("Deleted vulnerability report %s - %s", report_id, report.get("title")) + self.notify_finding_persisted(report) return report + def notify_finding_persisted(self, report: dict[str, Any]) -> None: + """Called only after hosted persistence and local report saving succeed.""" + if self.finding_persisted_callback: + self.finding_persisted_callback(report) + def get_existing_vulnerabilities(self) -> list[dict[str, Any]]: return list(self.vulnerability_reports) diff --git a/strix/runtime/agent_session.py b/strix/runtime/agent_session.py new file mode 100644 index 000000000..97520505d --- /dev/null +++ b/strix/runtime/agent_session.py @@ -0,0 +1,145 @@ +"""Per-agent process ownership within a shared sandbox (not a security boundary).""" + +from __future__ import annotations + +import copy +import uuid +from typing import TYPE_CHECKING, Any + +from agents.sandbox.session import BaseSandboxSession + + +if TYPE_CHECKING: + from pathlib import Path + + +class AgentSandboxSession(BaseSandboxSession): + def __init__(self, parent: BaseSandboxSession, root: str, task_id: str) -> None: + while isinstance(parent, AgentSandboxSession): + parent = parent.parent + self.parent: BaseSandboxSession = parent + self.state = copy.deepcopy(parent.state) + self.state.manifest.root = root + self.state.manifest.entries = {} + self.state.session_id = uuid.uuid5(parent.state.session_id, task_id) + self.task_id = task_id + self.processes: set[int] = set() + + async def stop_process(self, pid: int) -> Any: + if pid <= 1: + raise ValueError("A positive process ID greater than 1 is required") + return await self.parent.exec( + "python", + "-c", + _STOP_OWN_PROCESSES, + self.task_id, + str(pid), + shell=False, + timeout=15, + ) + + def _command(self, command: tuple[Any, ...]) -> list[str]: + return [ + "env", + f"STRIX_AGENT_TASK={self.task_id}", + "sh", + "-c", + 'cd -- "$1" || exit; shift; exec "$@"', + "sh", + self.state.manifest.root, + *map(str, command), + ] + + async def _exec_internal(self, *command: Any, timeout: float | None = None) -> Any: + return await self.parent.exec( + *self._command(command), + shell=False, + timeout=min(timeout or 600, 600), + ) + + def supports_pty(self) -> bool: + return self.parent.supports_pty() + + async def pty_exec_start(self, *command: Any, **kwargs: Any) -> Any: + # Intentional native sandbox shell, never host execution. + prepared = self._prepare_exec_command( # nosec B604 + *command, + shell=kwargs.pop("shell", True), + user=kwargs.pop("user", None), + ) + kwargs["timeout"] = min(kwargs.get("timeout") or 600, 600) + update = await self.parent.pty_exec_start( + *self._command(tuple(prepared)), + shell=False, + **kwargs, + ) + if update.process_id is not None: + self.processes.add(update.process_id) + return update + + async def pty_write_stdin(self, *, session_id: int, **kwargs: Any) -> Any: + if session_id not in self.processes: + raise ValueError("That process belongs to another agent") + return await self.parent.pty_write_stdin(session_id=session_id, **kwargs) + + async def read(self, path: Path, **kwargs: Any) -> Any: + return await self.parent.read(self.normalize_path(path), **kwargs) + + async def write(self, path: Path, data: Any, **kwargs: Any) -> None: + await self.parent.write(self.normalize_path(path, for_write=True), data, **kwargs) + + async def running(self) -> bool: + return await self.parent.running() + + async def hydrate_workspace(self, _data: Any) -> None: + raise RuntimeError("A borrowed worktree cannot restore the scan workspace") + + async def persist_workspace(self) -> Any: + raise RuntimeError("The scan owns workspace persistence") + + async def stop(self) -> None: + await self.pty_terminate_all() + + async def shutdown(self) -> None: + await self.pty_terminate_all() + + async def pty_terminate_all(self) -> None: + # Background services inherit this task marker too. Never call the + # parent session's terminate_all(), which would kill assessment tools. + await self.parent.exec( + "python", + "-c", + _STOP_OWN_PROCESSES, + self.task_id, + shell=False, + timeout=15, + ) + self.processes.clear() + + +_STOP_OWN_PROCESSES = r""" +import os, pathlib, signal, sys, time +marker = b"STRIX_AGENT_TASK=" + sys.argv[1].encode() +requested = int(sys.argv[2]) if len(sys.argv) > 2 else None +def owned(): + found = [] + for path in pathlib.Path("/proc").glob("[0-9]*/environ"): + try: + if marker in path.read_bytes().split(b"\0"): + found.append(int(path.parent.name)) + except (OSError, ValueError): + pass + return found +if requested is not None and requested not in owned(): + sys.exit("Process is absent or belongs to another agent; nothing was stopped") +for sig in (signal.SIGTERM, signal.SIGKILL): + for pid in owned(): + if requested is not None and pid != requested: + continue + try: + os.kill(pid, sig) + except ProcessLookupError: + pass + if sig == signal.SIGTERM: + time.sleep(0.2) +""" diff --git a/strix/tools/agents_graph/tools.py b/strix/tools/agents_graph/tools.py index ffd6d6ef5..aa89391b8 100644 --- a/strix/tools/agents_graph/tools.py +++ b/strix/tools/agents_graph/tools.py @@ -498,7 +498,6 @@ 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. @@ -544,9 +543,6 @@ 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) @@ -586,7 +582,6 @@ 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) @@ -700,6 +695,12 @@ async def agent_finish( default=str, ) + before_finish = inner.get("before_agent_finish") + if before_finish is not None: + error = await before_finish(success) + if error: + return json.dumps({"success": False, "error": error}) + filed_reports = _filed_reports_by(me) filed_report_ids = [str(r.get("id")) for r in filed_reports] diff --git a/strix/tools/processes.py b/strix/tools/processes.py new file mode 100644 index 000000000..eba408569 --- /dev/null +++ b/strix/tools/processes.py @@ -0,0 +1,21 @@ +"""Process cleanup that cannot accidentally target another agent's services.""" + +from typing import Any + +from agents import RunContextWrapper, function_tool + +from strix.runtime.agent_session import AgentSandboxSession + + +@function_tool +async def stop_process(ctx: RunContextWrapper[dict[str, Any]], pid: int) -> str: + """Stop a background PID started by this agent. For tool sessions use write_stdin Ctrl-C. + + Args: + pid: OS process ID printed when this agent started the background process. + """ + session = ctx.context.get("sandbox_session") + if not isinstance(session, AgentSandboxSession): + return "Process ownership is unavailable; use Ctrl-C on your own write_stdin session." + result = await session.stop_process(pid) + return str(result.stdout or result.stderr or f"Process stop returned {result.exit_code}") diff --git a/tests/test_agent_factory_shell.py b/tests/test_agent_factory_shell.py index 24949ddae..022e322fb 100644 --- a/tests/test_agent_factory_shell.py +++ b/tests/test_agent_factory_shell.py @@ -132,3 +132,15 @@ def test_specialized_tools_do_not_inherit_scan_or_registered_tools(monkeypatch) assert ( not {"scan_extension", "create_agent", "record_coverage", "finish_scan"} & specialized_names ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "command", ["pkill -f server", "sudo killall node", "kill $(lsof -ti:3007)", "/bin/kill -9 123"] +) +async def test_shared_process_cleanup_is_rejected_before_execution(command): + captured = {} + wrapped = factory._wrap_exec_command(_capturing_exec_tool(captured)) + result = await wrapped.on_invoke_tool(None, json.dumps({"cmd": command})) + assert "stop_process" in result + assert not captured diff --git a/tests/test_fix_completion.py b/tests/test_fix_completion.py index 7922b2453..b585e272b 100644 --- a/tests/test_fix_completion.py +++ b/tests/test_fix_completion.py @@ -45,7 +45,7 @@ def call(name: str, **arguments: Any) -> ResponseFunctionToolCall: def finish(outcome: str, summary: str = "Fix and validation results reviewed.") -> Any: - return call("agent_finish", outcome=outcome, result_summary=summary) + return call("agent_finish", success=outcome == "done", result_summary=summary) def shell(cmd: str) -> Any: @@ -204,11 +204,17 @@ async def test_blocked_or_capped_agent_discards_patch(tmp_path, monkeypatch, end @pytest.mark.asyncio -async def test_native_lifecycle_retries_invalid_outcome(tmp_path, monkeypatch): - model = ScriptedModel([*patch(), finish("approved"), *suite_commands(), finish("done")]) +async def test_native_lifecycle_ignores_legacy_quoted_outcome(tmp_path, monkeypatch): + model = ScriptedModel( + [ + *patch(), + *suite_commands(), + call("agent_finish", outcome='"done"', success=True, result_summary="Complete"), + ] + ) result, _ = await scenario(tmp_path, monkeypatch, model) assert result.state is PreparationState.READY - assert result.completion.turns_used == 6 + assert result.completion.turns_used == 5 @pytest.mark.asyncio @@ -267,3 +273,57 @@ async def test_fix_respects_live_scan_budget_without_double_counting(tmp_path, m ModelResponse(output=[], usage=Usage(requests=1), response_id=None), ) assert len(recorded) == 1 + + +@pytest.mark.asyncio +async def test_dependency_symlink_does_not_hide_new_source_or_regression(tmp_path, monkeypatch): + model = ScriptedModel( + [ + *patch(), + shell("ln -s /tmp node_modules"), + shell("printf 'VALUE = 1\\n' > helper.py"), + *suite_commands(), + finish("done"), + ] + ) + result, _ = await scenario(tmp_path, monkeypatch, model) + assert result.state is PreparationState.READY, result.stop_reason + with zipfile.ZipFile(tmp_path / "prepared.zip") as artifact: + names = artifact.namelist() + assert "files/helper.py" in names and "files/tests/test_security.py" in names + assert not any("node_modules" in name for name in names) + + +@pytest.mark.asyncio +async def test_checkpoint_failure_is_correctable_before_completion(tmp_path, monkeypatch): + model = ScriptedModel( + [ + *patch(), + shell("ln -s app.py helper.py"), + finish("done"), + shell("rm helper.py"), + *suite_commands(), + finish("done"), + ] + ) + result, _ = await scenario(tmp_path, monkeypatch, model) + assert result.state is PreparationState.READY, result.stop_reason + assert any("Could not package this fix" in str(items) for items in model.inputs["repair"]) + + +@pytest.mark.asyncio +async def test_three_identical_completion_errors_stop_with_real_reason(tmp_path, monkeypatch): + model = ScriptedModel( + [ + *patch(), + shell("ln -s app.py helper.py"), + finish("done"), + finish("done"), + finish("done"), + ] + ) + result, _ = await scenario(tmp_path, monkeypatch, model, turns=300) + assert result.state is PreparationState.BLOCKED + assert result.completion.turns_used == 6 + assert "Completion failed three times" in result.completion.summary + assert not (tmp_path / "prepared.zip").exists() diff --git a/tests/test_runner_interrupt.py b/tests/test_runner_interrupt.py index 054f5b7a2..b53ad72da 100644 --- a/tests/test_runner_interrupt.py +++ b/tests/test_runner_interrupt.py @@ -12,6 +12,7 @@ import strix.tools.todo.tools as todo_tools from strix.core import runner from strix.core.agents import AgentCoordinator from strix.runtime import session_manager +from tests.test_fix_reliability import LocalSandbox def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: @@ -40,7 +41,11 @@ def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _state_dir: None) async def _create_or_reuse(*_args: Any, **_kwargs: Any) -> dict[str, Any]: - return {"client": object(), "session": object(), "caido_client": None} + return { + "client": object(), + "session": LocalSandbox(tmp_path / "sandbox"), + "caido_client": None, + } async def _cleanup(*_args: Any, **_kwargs: Any) -> None: return None diff --git a/tests/test_runner_mcp.py b/tests/test_runner_mcp.py index 25a5c888d..903d00e2e 100644 --- a/tests/test_runner_mcp.py +++ b/tests/test_runner_mcp.py @@ -21,6 +21,7 @@ from strix.core import runner from strix.core.agents import AgentCoordinator from strix.runtime import session_manager from strix.tools.mcp import McpConnectionConfig, McpConnectionRequest +from tests.test_fix_reliability import LocalSandbox def _settings() -> Any: @@ -49,7 +50,11 @@ def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _d: None) async def _create_or_reuse(*_a: Any, **_k: Any) -> dict[str, Any]: - return {"client": object(), "session": object(), "caido_client": None} + return { + "client": object(), + "session": LocalSandbox(tmp_path / "sandbox"), + "caido_client": None, + } async def _cleanup(*_a: Any, **_k: Any) -> None: return None diff --git a/tests/test_runner_rate_limit.py b/tests/test_runner_rate_limit.py index 3110ae2cb..d6ee012e8 100644 --- a/tests/test_runner_rate_limit.py +++ b/tests/test_runner_rate_limit.py @@ -16,6 +16,7 @@ import strix.tools.todo.tools as todo_tools from strix.core import runner from strix.core.agents import AgentCoordinator from strix.runtime import session_manager +from tests.test_fix_reliability import LocalSandbox def _make_rate_limit_error() -> RateLimitError: @@ -55,7 +56,11 @@ async def test_persistent_rate_limit_stops_gracefully( monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _state_dir: None) async def _create_or_reuse(*_args: Any, **_kwargs: Any) -> dict[str, Any]: - return {"client": object(), "session": object(), "caido_client": None} + return { + "client": object(), + "session": LocalSandbox(tmp_path / "sandbox"), + "caido_client": None, + } async def _cleanup(*_args: Any, **_kwargs: Any) -> None: return None diff --git a/tests/test_runner_root_prompt.py b/tests/test_runner_root_prompt.py index 31d153a71..52e1f7f72 100644 --- a/tests/test_runner_root_prompt.py +++ b/tests/test_runner_root_prompt.py @@ -25,6 +25,7 @@ from strix.core.agents import AgentCoordinator from strix.core.inputs import make_model_settings from strix.runtime import session_manager from strix.tools.mcp import BearerAuth, McpConnectionConfig, McpConnectionRequest +from tests.test_fix_reliability import LocalSandbox def _make_rate_limit_error() -> RateLimitError: @@ -71,7 +72,11 @@ def _patch_engine_scaffold( monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _state_dir: None) async def _create_or_reuse(*_args: Any, **_kwargs: Any) -> dict[str, Any]: - return {"client": object(), "session": object(), "caido_client": None} + return { + "client": object(), + "session": LocalSandbox(tmp_path / "sandbox"), + "caido_client": None, + } async def _cleanup(*_args: Any, **_kwargs: Any) -> None: return None diff --git a/tests/test_runner_teardown.py b/tests/test_runner_teardown.py index 77a115fac..7d45dcfea 100644 --- a/tests/test_runner_teardown.py +++ b/tests/test_runner_teardown.py @@ -12,6 +12,7 @@ import strix.tools.todo.tools as todo_tools from strix.core import runner from strix.core.agents import AgentCoordinator from strix.runtime import session_manager +from tests.test_fix_reliability import LocalSandbox def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: @@ -28,7 +29,11 @@ def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _d: None) async def _create_or_reuse(*_a: Any, **_k: Any) -> dict[str, Any]: - return {"client": object(), "session": object(), "caido_client": None} + return { + "client": object(), + "session": LocalSandbox(tmp_path / "sandbox"), + "caido_client": None, + } async def _cleanup(*_a: Any, **_k: Any) -> None: return None @@ -114,6 +119,9 @@ async def test_assessment_publishes_before_fixes_end_and_sandbox_teardown(monkey def __init__(self, **_): pass + def start(self, *_): + events.append("fixes listening") + async def wait(self): events.append("fixes finished") diff --git a/tests/test_scan_fixes.py b/tests/test_scan_fixes.py index 187ed10db..b5257d309 100644 --- a/tests/test_scan_fixes.py +++ b/tests/test_scan_fixes.py @@ -1,12 +1,11 @@ -"""Fix delegation through the native create_agent tool, child loop and worktrees.""" +"""Persisted findings launch native Fix children in isolated worktrees.""" from __future__ import annotations import asyncio -import json from pathlib import Path from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, Mock import pytest from agents import RunConfig @@ -21,7 +20,6 @@ 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 @@ -65,6 +63,8 @@ def setup(tmp_path): finding_id = kwargs.pop("fix_finding_id") return await fixes.spawn(finding_id, native, **kwargs) + fixes._native_spawn = native + context = ToolContext( tool_name="create_agent", tool_call_id="spawn-test", @@ -79,21 +79,19 @@ def setup(tmp_path): 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, - } - ), +async def delegate(context, finding_id="finding", **options): + try: + return await context.context["spawn_child_agent"]( + fix_finding_id=finding_id, + parent_ctx=context.context, + name="Fix agent", + task="Fix the saved issue and test it.", + skills=[], + parent_history=[], + **options, ) - ) + except ValueError as error: + return {"success": False, "error": str(error)} @pytest.mark.asyncio @@ -151,8 +149,8 @@ async def test_delegation_errors_reach_reporting_agent_before_any_model_call(tmp 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"] + with pytest.raises(RuntimeError, match="Missing callback configuration"): + await delegate(context) assert not fixes.tasks @@ -226,18 +224,70 @@ async def test_native_child_keeps_cumulative_turn_cap_and_does_not_export_partia @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) +async def test_seven_saved_findings_start_seven_native_children_without_model_handoff( + tmp_path, monkeypatch +): + fixes, report, source, _, _, 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"], + monkeypatch.setattr( + scan_module, + "_run_config", + lambda env: RunConfig( + model=ScriptedModel([*patch(), *suite_commands(), finish("done")]), + sandbox=SandboxRunConfig(session=env.session), + tracing_disabled=True, + ), ) - assert not fixes.tasks + stages = [] + + async def sink(stage, report, _result, _artifact): + stages.append((stage, report["id"])) + return True + + fixes.sink = sink + fixes.start(fixes._native_spawn, context.context) + ids = [ + state.add_vulnerability_report( + title=f"Unsafe result {i}", + severity="high", + agent_id="reporter", + validation_status="confirmed", + fix_candidate=report["fix_candidate"], + ) + for i in range(7) + ] + # Replayed/no-op notifications must not create extra jobs. + for saved in state.get_existing_vulnerabilities(): + fixes.notify(saved) + await fixes.wait() + assert set(fixes.records) == set(ids) + assert all(r["status"] == "done" for r in fixes.records.values()), fixes.records + assert len([s for s in stages if s[0] == "started"]) == 7 + assert _git(source, "status", "--porcelain") == "" + 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) + state = ReportState("failed-persistence") + state._run_dir = tmp_path / "report" + fixes.report_state = state + fixes.start(fixes._native_spawn, context.context) + state.vulnerability_found_callback = Mock(side_effect=RuntimeError("Database rejected finding")) + with pytest.raises(RuntimeError, match="Database rejected"): + state.add_vulnerability_report( + title="Unsafe", severity="high", fix_candidate=report["fix_candidate"] + ) + assert not fixes.dispatches and not fixes.tasks + + +@pytest.mark.asyncio +async def test_explicit_retry_gets_new_agent_with_remaining_turns(tmp_path, monkeypatch): + fixes, _, _, _, _, context, sessions = setup(tmp_path) monkeypatch.setattr( scan_module, "_run_config", @@ -247,9 +297,13 @@ async def test_saving_report_does_not_launch_until_reporting_agent_delegates(tmp tracing_disabled=True, ), ) - spawned = await delegate(context, report_id) - assert spawned["success"], spawned + first = await delegate(context) + await fixes.tasks["finding"] + again = await delegate(context) + assert not again["success"] + second = await delegate(context, retry=True) await fixes.wait() - assert fixes.coordinator.parent_of[spawned["agent_id"]] == "reporter" + assert second["agent_id"] != first["agent_id"] + assert fixes.records["finding"]["turns"] == 2 for session in sessions: session.close()