diff options
| -rw-r--r-- | README.md | 5 | ||||
| -rw-r--r-- | src/gap_pipeline/clients.py | 17 | ||||
| -rw-r--r-- | src/gap_pipeline/paper_pipeline.py | 212 | ||||
| -rw-r--r-- | tests/test_paper_pipeline.py | 65 |
4 files changed, 219 insertions, 80 deletions
@@ -50,7 +50,10 @@ five-stage path is the reproduction entry point. The OpenAI adapter does not send `temperature`; this is compatible with `o3`, whose supported value is its default. See `STAGE_MAP.md` for the exact -paper-to-code map. +paper-to-code map. Each typed stage also requests a strict JSON schema. If a +provider response still fails local structural or provenance validation, the +same stage is retried once with the exact same prompt; both responses remain +separately auditable. `ProofDAG` supports branching dependencies and validates topological order, known dependencies, acyclicity, and terminal connectivity. Replacement plans diff --git a/src/gap_pipeline/clients.py b/src/gap_pipeline/clients.py index 3b3e40f..bc35742 100644 --- a/src/gap_pipeline/clients.py +++ b/src/gap_pipeline/clients.py @@ -8,6 +8,8 @@ from dataclasses import dataclass, field from pathlib import Path from typing import Any, Protocol +from pydantic import BaseModel + @dataclass(frozen=True) class ModelResponse: @@ -27,6 +29,7 @@ class JsonLLM(Protocol): system_prompt: str, user_prompt: str, request_id: str, + response_model: type[BaseModel] | None = None, ) -> ModelResponse: ... @@ -61,14 +64,24 @@ class OpenAIJsonClient: system_prompt: str, user_prompt: str, request_id: str, + response_model: type[BaseModel] | None = None, ) -> ModelResponse: + if response_model is None: + response_format: dict[str, Any] = {"type": "json_object"} + else: + from openai.lib._parsing._completions import ( + type_to_response_format_param, + ) + + response_format = type_to_response_format_param(response_model) + kwargs: dict[str, Any] = { "model": self.model, "messages": [ {"role": "system", "content": system_prompt}, {"role": "user", "content": user_prompt}, ], - "response_format": {"type": "json_object"}, + "response_format": response_format, "max_completion_tokens": self.max_completion_tokens, } if self.seed is not None: @@ -111,6 +124,7 @@ class ReplayClient: system_prompt: str, user_prompt: str, request_id: str, + response_model: type[BaseModel] | None = None, ) -> ModelResponse: if request_id not in self._responses: raise KeyError(f"no replay response for {request_id}") @@ -145,6 +159,7 @@ class ScriptedClient: system_prompt: str, user_prompt: str, request_id: str, + response_model: type[BaseModel] | None = None, ) -> ModelResponse: self.calls.append(request_id) prefixes = sorted( diff --git a/src/gap_pipeline/paper_pipeline.py b/src/gap_pipeline/paper_pipeline.py index c90f3f6..ffe9f15 100644 --- a/src/gap_pipeline/paper_pipeline.py +++ b/src/gap_pipeline/paper_pipeline.py @@ -3,8 +3,10 @@ from __future__ import annotations import asyncio +from collections.abc import Callable +from typing import TypeVar -from pydantic import BaseModel, ConfigDict, model_validator +from pydantic import BaseModel, ConfigDict, ValidationError, model_validator from .clients import JsonLLM from .kernel_models import ( @@ -36,6 +38,9 @@ from .models import CanonicalItem, ModelCallRecord from .store import RunStore, sha256_payload +ValidatedModel = TypeVar("ValidatedModel", bound=BaseModel) + + class PaperPipelineConfig(BaseModel): model_config = ConfigDict(extra="forbid") @@ -80,11 +85,13 @@ class PaperKernelPipeline: request_id: str, system_prompt: str, user_prompt: str, + response_model: type[BaseModel] | None = None, ) -> dict: response = await client.generate_json( system_prompt=system_prompt, user_prompt=user_prompt, request_id=request_id, + response_model=response_model, ) self.store.write_call( ModelCallRecord( @@ -100,30 +107,75 @@ class PaperKernelPipeline: ) return response.data + async def _call_validated( + self, + client: JsonLLM, + *, + request_id: str, + system_prompt: str, + user_prompt: str, + response_model: type[ValidatedModel], + validate: Callable[[ValidatedModel], ValidatedModel] | None = None, + ) -> tuple[ValidatedModel, str]: + """Request a typed response, retrying the identical prompt once if invalid.""" + + last_error: ValidationError | ValueError | None = None + for attempt in range(2): + attempt_request_id = ( + request_id + if attempt == 0 + else f"{request_id}.validation_retry{attempt:02d}" + ) + data = await self._call( + client, + request_id=attempt_request_id, + system_prompt=system_prompt, + user_prompt=user_prompt, + response_model=response_model, + ) + try: + result = response_model.model_validate(data) + if validate is not None: + result = validate(result) + return result, attempt_request_id + except (ValidationError, ValueError) as exc: + last_error = exc + + raise ValueError( + f"{request_id} failed typed validation after two identical-prompt attempts" + ) from last_error + async def construct_dag(self, item: CanonicalItem) -> ProofDAG: request_id = f"{item.item_id}.stage1.dag" - dag = ProofDAG.model_validate( - await self._call( - self.proposer, - request_id=request_id, - system_prompt=DAG_SYSTEM, - user_prompt=dag_user(item), - ) + dag, accepted_request_id = await self._call_validated( + self.proposer, + request_id=request_id, + system_prompt=DAG_SYSTEM, + user_prompt=dag_user(item), + response_model=ProofDAG, + ) + self.store.write_stage( + "01_proof_dag", + dag, + request_id=accepted_request_id, ) - self.store.write_stage("01_proof_dag", dag, request_id=request_id) return dag async def summarize_methods(self, dag: ProofDAG) -> MethodPlan: request_id = f"{self.store.item_id}.stage2.methods" - methods = MethodPlan.model_validate( - await self._call( - self.proposer, - request_id=request_id, - system_prompt=METHOD_SYSTEM, - user_prompt=method_user(dag), - ) - ).validate_against(dag) - self.store.write_stage("02_method_plan", methods, request_id=request_id) + methods, accepted_request_id = await self._call_validated( + self.proposer, + request_id=request_id, + system_prompt=METHOD_SYSTEM, + user_prompt=method_user(dag), + response_model=MethodPlan, + validate=lambda value: value.validate_against(dag), + ) + self.store.write_stage( + "02_method_plan", + methods, + request_id=accepted_request_id, + ) return methods async def generate_replacements( @@ -137,26 +189,30 @@ class PaperKernelPipeline: feedback: str = "", ) -> ReplacementPlan: request_id = f"{item.item_id}.stage3.replacement.v{version:02d}" - replacements = ReplacementPlan.model_validate( - await self._call( - self.proposer, - request_id=request_id, - system_prompt=REPLACEMENT_SYSTEM, - user_prompt=replacement_user( - item, - dag, - methods, - previous_replacements=previous_replacements, - feedback=feedback, - ), - ) - ).validate_against(dag) - if previous_replacements is not None: - replacements.validate_repair_of(previous_replacements) + def validate_replacements(value: ReplacementPlan) -> ReplacementPlan: + value.validate_against(dag) + if previous_replacements is not None: + value.validate_repair_of(previous_replacements) + return value + + replacements, accepted_request_id = await self._call_validated( + self.proposer, + request_id=request_id, + system_prompt=REPLACEMENT_SYSTEM, + user_prompt=replacement_user( + item, + dag, + methods, + previous_replacements=previous_replacements, + feedback=feedback, + ), + response_model=ReplacementPlan, + validate=validate_replacements, + ) self.store.write_stage( f"03_replacement_v{version:02d}", replacements, - request_id=request_id, + request_id=accepted_request_id, ) return replacements @@ -172,25 +228,25 @@ class PaperKernelPipeline: feedback: str = "", ) -> DiffusedProof: request_id = f"{item.item_id}.stage4.diffusion.v{version:02d}" - diffused = DiffusedProof.model_validate( - await self._call( - self.proposer, - request_id=request_id, - system_prompt=DIFFUSION_SYSTEM, - user_prompt=diffusion_user( - item, - dag, - methods, - replacements, - previous_diffused=previous_diffused, - feedback=feedback, - ), - ) - ).validate_against(dag, methods) + diffused, accepted_request_id = await self._call_validated( + self.proposer, + request_id=request_id, + system_prompt=DIFFUSION_SYSTEM, + user_prompt=diffusion_user( + item, + dag, + methods, + replacements, + previous_diffused=previous_diffused, + feedback=feedback, + ), + response_model=DiffusedProof, + validate=lambda value: value.validate_against(dag, methods), + ) self.store.write_stage( f"04_diffused_proof_v{version:02d}", diffused, - request_id=request_id, + request_id=accepted_request_id, ) return diffused @@ -205,23 +261,23 @@ class PaperKernelPipeline: feedback: str = "", ) -> RenderedVariant: request_id = f"{self.store.item_id}.stage5.render.v{version:02d}" - variant = RenderedVariant.model_validate( - await self._call( - self.proposer, - request_id=request_id, - system_prompt=RENDER_SYSTEM, - user_prompt=render_user( - replacements, - diffused, - previous_variant=previous_variant, - feedback=feedback, - ), - ) - ).validate_against(dag, diffused) + variant, accepted_request_id = await self._call_validated( + self.proposer, + request_id=request_id, + system_prompt=RENDER_SYSTEM, + user_prompt=render_user( + replacements, + diffused, + previous_variant=previous_variant, + feedback=feedback, + ), + response_model=RenderedVariant, + validate=lambda value: value.validate_against(dag, diffused), + ) self.store.write_stage( f"05_rendered_variant_v{version:02d}", variant, - request_id=request_id, + request_id=accepted_request_id, ) return variant @@ -294,20 +350,20 @@ class PaperKernelPipeline: judge_id: int, ) -> JudgeVerdict: request_id = f"{item.item_id}.verify.t{iteration:02d}.j{judge_id}" - verdict = JudgeVerdict.model_validate( - await self._call( - judge, - request_id=request_id, - system_prompt=JUDGE_SYSTEM, - user_prompt=judge_user( - item, - methods, - bundle.replacement_plan, - bundle.variant, - ), - ) + verdict, _ = await self._call_validated( + judge, + request_id=request_id, + system_prompt=JUDGE_SYSTEM, + user_prompt=judge_user( + item, + methods, + bundle.replacement_plan, + bundle.variant, + ), + response_model=JudgeVerdict, + validate=lambda value: value.validate_coverage(dag), ) - return verdict.validate_coverage(dag) + return verdict @staticmethod def _feedback(verdicts: list[JudgeVerdict]) -> str: 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" |
