From 84fab096b3a2500755fea3f538933dbed0e8c72c Mon Sep 17 00:00:00 2001 From: Anonymous Authors Date: Sat, 25 Jul 2026 13:10:52 -0500 Subject: Enforce typed responses for pipeline stages --- src/gap_pipeline/clients.py | 17 ++- src/gap_pipeline/paper_pipeline.py | 212 +++++++++++++++++++++++-------------- 2 files changed, 150 insertions(+), 79 deletions(-) (limited to 'src/gap_pipeline') 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: -- cgit v1.2.3