summaryrefslogtreecommitdiff
path: root/tests/test_paper_pipeline.py
diff options
context:
space:
mode:
authorAnonymous Authors <anonymous@invalid.example>2026-07-25 13:10:52 -0500
committerAnonymous Authors <anonymous@invalid.example>2026-07-25 13:10:52 -0500
commit84fab096b3a2500755fea3f538933dbed0e8c72c (patch)
tree15a9d2016bf235d3bfba5035b305bc0b7e996a0a /tests/test_paper_pipeline.py
parent6de74d103926d9090f056aeebe7be393ec381ea1 (diff)
Enforce typed responses for pipeline stagesxiang
Diffstat (limited to 'tests/test_paper_pipeline.py')
-rw-r--r--tests/test_paper_pipeline.py65
1 files changed, 65 insertions, 0 deletions
diff --git a/tests/test_paper_pipeline.py b/tests/test_paper_pipeline.py
index 7f1de46..44e9081 100644
--- a/tests/test_paper_pipeline.py
+++ b/tests/test_paper_pipeline.py
@@ -4,6 +4,8 @@ import asyncio
import json
from gap_pipeline.clients import ScriptedClient
+from gap_pipeline.kernel_models import MethodPlan, ProofDAG, ReplacementPlan
+from gap_pipeline.kernel_prompts import diffusion_user
from gap_pipeline.paper_pipeline import PaperKernelPipeline, PaperPipelineConfig
from gap_pipeline.prompts import JUDGE_SYSTEM_PROMPT
from gap_pipeline.store import RunStore
@@ -190,3 +192,66 @@ def test_five_stage_repair_uses_prior_bundle_and_appendix_judge(
assert first_judge_call["system_prompt"] == JUDGE_SYSTEM_PROMPT
assert "METHOD-LABEL SEQUENCE (abstract plan):" in first_judge_call["user_prompt"]
assert "SOURCE PROOF DAG:" not in first_judge_call["user_prompt"]
+
+
+def test_stage4_retries_identical_prompt_when_required_field_is_missing(
+ tmp_path,
+ item,
+) -> None:
+ invalid = _diffused()
+ for node in invalid["nodes"]:
+ node.pop("justification")
+
+ request_prefix = f"{item.item_id}.stage4.diffusion.v01"
+ proposer = ScriptedClient({request_prefix: [invalid, _diffused()]})
+ pipeline = PaperKernelPipeline(
+ proposer=proposer,
+ judges=[ScriptedClient({}) for _ in range(5)],
+ store=RunStore(tmp_path / "run", item.item_id),
+ config=PaperPipelineConfig(
+ proposer_model="scripted",
+ judge_model="scripted",
+ ),
+ )
+ dag = ProofDAG.model_validate(_dag())
+ methods = MethodPlan.model_validate(_methods())
+ replacements = ReplacementPlan.model_validate(_replacement())
+
+ result = asyncio.run(
+ pipeline.diffuse_dag(
+ item,
+ dag,
+ methods,
+ replacements,
+ version=1,
+ )
+ )
+
+ assert all(node.justification for node in result.nodes)
+ assert proposer.calls == [
+ request_prefix,
+ f"{request_prefix}.validation_retry01",
+ ]
+
+ calls_dir = tmp_path / "run" / "items" / item.item_id / "calls"
+ first_call = json.loads((calls_dir / f"{request_prefix}.json").read_text())
+ retry_call = json.loads(
+ (
+ calls_dir / f"{request_prefix}.validation_retry01.json"
+ ).read_text()
+ )
+ expected_prompt = diffusion_user(item, dag, methods, replacements)
+ assert first_call["user_prompt"] == expected_prompt
+ assert retry_call["user_prompt"] == expected_prompt
+
+ stage = json.loads(
+ (
+ tmp_path
+ / "run"
+ / "items"
+ / item.item_id
+ / "stages"
+ / "04_diffused_proof_v01.json"
+ ).read_text()
+ )
+ assert stage["request_id"] == f"{request_prefix}.validation_retry01"