diff options
Diffstat (limited to 'tests/test_pipeline.py')
| -rw-r--r-- | tests/test_pipeline.py | 236 |
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) |
