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