summaryrefslogtreecommitdiff
path: root/tests/test_pipeline.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_pipeline.py')
-rw-r--r--tests/test_pipeline.py236
1 files changed, 236 insertions, 0 deletions
diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py
new file mode 100644
index 0000000..ea1077a
--- /dev/null
+++ b/tests/test_pipeline.py
@@ -0,0 +1,236 @@
+from __future__ import annotations
+
+import asyncio
+import json
+
+from gap_pipeline.clients import ScriptedClient
+from gap_pipeline.pipeline import KernelPipeline, PipelineConfig
+from gap_pipeline.store import RunStore
+
+from conftest import changed_candidate
+
+
+def accept_verdict() -> dict:
+ return {
+ "verdict": "accept",
+ "step_by_step_check": "n1 passes; n2 passes",
+ "blocking_issues": "",
+ "patch_suggestion": "",
+ }
+
+
+def reject_verdict(iteration: int) -> dict:
+ return {
+ "verdict": "reject",
+ "step_by_step_check": f"n1 passes; n2 fails at iteration {iteration}",
+ "blocking_issues": "the terminal algebra needs correction",
+ "patch_suggestion": "correct the terminal algebra without changing the plan",
+ }
+
+
+def proposer_responses(
+ item_id: str,
+ dag_dict: dict,
+ plan_dict: dict,
+ replacement_dict: dict,
+ diffused_dict: dict,
+ candidate_dict: dict,
+) -> dict:
+ return {
+ f"{item_id}.stage1": dag_dict,
+ f"{item_id}.stage2": plan_dict,
+ f"{item_id}.stage3": replacement_dict,
+ f"{item_id}.stage4": diffused_dict,
+ f"{item_id}.stage5": candidate_dict,
+ }
+
+
+def test_two_consecutive_unanimous_passes_accept(
+ tmp_path,
+ item,
+ dag_dict,
+ plan_dict,
+ replacement_dict,
+ diffused_dict,
+ candidate_dict,
+) -> None:
+ proposer = ScriptedClient(
+ proposer_responses(
+ item.item_id,
+ dag_dict,
+ plan_dict,
+ replacement_dict,
+ diffused_dict,
+ candidate_dict,
+ )
+ )
+ judges = [
+ ScriptedClient({f"{item.item_id}.verify": [accept_verdict(), accept_verdict()]})
+ for _ in range(5)
+ ]
+ pipeline = KernelPipeline(
+ proposer=proposer,
+ judges=judges,
+ store=RunStore(tmp_path, item.item_id),
+ config=PipelineConfig(proposer_model="scripted", judge_model="scripted"),
+ )
+ result = asyncio.run(pipeline.run(item))
+ assert result.status == "accepted"
+ assert len(result.iterations) == 2
+ assert [row.pass_streak_after for row in result.iterations] == [1, 2]
+ assert not any(row.repaired for row in result.iterations)
+
+ final_path = tmp_path / "items" / item.item_id / "final.json"
+ final = json.loads(final_path.read_text())
+ assert final["human_audit_selected"] in {True, False}
+ assert len(list((final_path.parent / "calls").glob("*.json"))) == 15
+
+
+def test_rejection_resets_streak_and_repairs(
+ tmp_path,
+ item,
+ dag_dict,
+ plan_dict,
+ replacement_dict,
+ diffused_dict,
+ candidate_dict,
+) -> None:
+ repaired = changed_candidate(candidate_dict, "fixed")
+ responses = proposer_responses(
+ item.item_id,
+ dag_dict,
+ plan_dict,
+ replacement_dict,
+ diffused_dict,
+ candidate_dict,
+ )
+ responses[f"{item.item_id}.repair"] = repaired
+ proposer = ScriptedClient(responses)
+ judges = []
+ for judge_id in range(1, 6):
+ first = reject_verdict(1) if judge_id == 1 else accept_verdict()
+ judges.append(
+ ScriptedClient(
+ {
+ f"{item.item_id}.verify": [
+ first,
+ accept_verdict(),
+ accept_verdict(),
+ ]
+ }
+ )
+ )
+ pipeline = KernelPipeline(
+ proposer=proposer,
+ judges=judges,
+ store=RunStore(tmp_path, item.item_id),
+ config=PipelineConfig(proposer_model="scripted", judge_model="scripted"),
+ )
+ result = asyncio.run(pipeline.run(item))
+ assert result.status == "accepted"
+ assert len(result.iterations) == 3
+ assert result.iterations[0].repaired
+ assert [row.pass_streak_after for row in result.iterations] == [0, 1, 2]
+ assert (
+ result.iterations[0].repaired_candidate_sha256
+ == result.iterations[1].candidate_sha256
+ == result.iterations[2].candidate_sha256
+ )
+
+
+def test_rejection_after_one_pass_resets_streak_on_new_candidate(
+ tmp_path,
+ item,
+ dag_dict,
+ plan_dict,
+ replacement_dict,
+ diffused_dict,
+ candidate_dict,
+) -> None:
+ repaired = changed_candidate(candidate_dict, "after-broken-streak")
+ responses = proposer_responses(
+ item.item_id,
+ dag_dict,
+ plan_dict,
+ replacement_dict,
+ diffused_dict,
+ candidate_dict,
+ )
+ responses[f"{item.item_id}.repair"] = repaired
+ proposer = ScriptedClient(responses)
+ judges = []
+ for judge_id in range(1, 6):
+ second = reject_verdict(2) if judge_id == 3 else accept_verdict()
+ judges.append(
+ ScriptedClient(
+ {
+ f"{item.item_id}.verify": [
+ accept_verdict(),
+ second,
+ accept_verdict(),
+ accept_verdict(),
+ ]
+ }
+ )
+ )
+ pipeline = KernelPipeline(
+ proposer=proposer,
+ judges=judges,
+ store=RunStore(tmp_path, item.item_id),
+ config=PipelineConfig(proposer_model="scripted", judge_model="scripted"),
+ )
+ result = asyncio.run(pipeline.run(item))
+ assert result.status == "accepted"
+ assert [row.pass_streak_after for row in result.iterations] == [1, 0, 1, 2]
+ assert result.iterations[0].candidate_sha256 == result.iterations[1].candidate_sha256
+ assert result.iterations[1].repaired_candidate_sha256 != result.iterations[1].candidate_sha256
+ assert (
+ result.iterations[1].repaired_candidate_sha256
+ == result.iterations[2].candidate_sha256
+ == result.iterations[3].candidate_sha256
+ )
+
+
+def test_fifteen_failed_rounds_reject(
+ tmp_path,
+ item,
+ dag_dict,
+ plan_dict,
+ replacement_dict,
+ diffused_dict,
+ candidate_dict,
+) -> None:
+ responses = proposer_responses(
+ item.item_id,
+ dag_dict,
+ plan_dict,
+ replacement_dict,
+ diffused_dict,
+ candidate_dict,
+ )
+ responses[f"{item.item_id}.repair"] = [
+ changed_candidate(candidate_dict, f"repair-{iteration}")
+ for iteration in range(1, 15)
+ ]
+ proposer = ScriptedClient(responses)
+ judges = [
+ ScriptedClient(
+ {
+ f"{item.item_id}.verify": [
+ reject_verdict(iteration) for iteration in range(1, 16)
+ ]
+ }
+ )
+ for _ in range(5)
+ ]
+ pipeline = KernelPipeline(
+ proposer=proposer,
+ judges=judges,
+ store=RunStore(tmp_path, item.item_id),
+ config=PipelineConfig(proposer_model="scripted", judge_model="scripted"),
+ )
+ result = asyncio.run(pipeline.run(item))
+ assert result.status == "rejected"
+ assert len(result.iterations) == 15
+ assert sum(row.repaired for row in result.iterations) == 14
+ assert all(row.pass_streak_after == 0 for row in result.iterations)