"""Literal five-stage GAP kernel pipeline matching the manuscript operations.""" from __future__ import annotations import asyncio from pydantic import BaseModel, ConfigDict, model_validator from .clients import JsonLLM from .kernel_models import ( CandidateBundle, DiffusedProof, JudgeVerdict, KernelRunResult, MethodPlan, ProofDAG, RenderedVariant, ReplacementPlan, VerificationIteration, ) from .kernel_prompts import ( DAG_SYSTEM, DIFFUSION_SYSTEM, JUDGE_SYSTEM, METHOD_SYSTEM, RENDER_SYSTEM, REPLACEMENT_SYSTEM, dag_user, diffusion_user, judge_user, method_user, render_user, replacement_user, ) from .models import CanonicalItem, ModelCallRecord from .store import RunStore, sha256_payload class PaperPipelineConfig(BaseModel): model_config = ConfigDict(extra="forbid") protocol_name: str = "gap-literal-five-stage-J5-K2-T15" proposer_model: str judge_model: str judge_count: int = 5 streak_length: int = 2 max_iterations: int = 15 @model_validator(mode="after") def enforce_protocol(self) -> "PaperPipelineConfig": if (self.judge_count, self.streak_length, self.max_iterations) != (5, 2, 15): raise ValueError("the GAP protocol is fixed at J=5, K=2, T=15") return self class PaperKernelPipeline: """Execute all five paper stages as typed, separately auditable calls.""" def __init__( self, *, proposer: JsonLLM, judges: list[JsonLLM], store: RunStore, config: PaperPipelineConfig, ) -> None: if len(judges) != config.judge_count: raise ValueError( f"expected {config.judge_count} judges, received {len(judges)}" ) self.proposer = proposer self.judges = judges self.store = store self.config = config async def _call( self, client: JsonLLM, *, request_id: str, system_prompt: str, user_prompt: str, ) -> dict: response = await client.generate_json( system_prompt=system_prompt, user_prompt=user_prompt, request_id=request_id, ) self.store.write_call( ModelCallRecord( request_id=request_id, model=response.model, system_prompt=system_prompt, user_prompt=user_prompt, response_data=response.data, raw_text=response.raw_text, provider_response_id=response.response_id, usage=response.usage, ) ) return response.data 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), ) ) 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) return methods async def generate_replacements( self, item: CanonicalItem, dag: ProofDAG, methods: MethodPlan, *, version: int, previous_replacements: ReplacementPlan | None = None, 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) self.store.write_stage( f"03_replacement_v{version:02d}", replacements, request_id=request_id, ) return replacements async def diffuse_dag( self, item: CanonicalItem, dag: ProofDAG, methods: MethodPlan, replacements: ReplacementPlan, *, version: int, previous_diffused: DiffusedProof | None = None, 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) self.store.write_stage( f"04_diffused_proof_v{version:02d}", diffused, request_id=request_id, ) return diffused async def render_variant( self, dag: ProofDAG, replacements: ReplacementPlan, diffused: DiffusedProof, *, version: int, previous_variant: RenderedVariant | None = None, 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) self.store.write_stage( f"05_rendered_variant_v{version:02d}", variant, request_id=request_id, ) return variant async def build_bundle( self, item: CanonicalItem, dag: ProofDAG, methods: MethodPlan, *, version: int, previous_bundle: CandidateBundle | None = None, feedback: str = "", ) -> CandidateBundle: replacements = await self.generate_replacements( item, dag, methods, version=version, previous_replacements=( previous_bundle.replacement_plan if previous_bundle is not None else None ), feedback=feedback, ) diffused = await self.diffuse_dag( item, dag, methods, replacements, version=version, previous_diffused=( previous_bundle.diffused_proof if previous_bundle is not None else None ), feedback=feedback, ) variant = await self.render_variant( dag, replacements, diffused, version=version, previous_variant=( previous_bundle.variant if previous_bundle is not None else None ), feedback=feedback, ) bundle = CandidateBundle( replacement_plan=replacements, diffused_proof=diffused, variant=variant, ) if ( previous_bundle is not None and sha256_payload(bundle) == sha256_payload(previous_bundle) ): raise ValueError("repair pass returned an unchanged candidate bundle") return bundle async def _judge_once( self, judge: JsonLLM, *, item: CanonicalItem, dag: ProofDAG, methods: MethodPlan, bundle: CandidateBundle, iteration: int, 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, ), ) ) return verdict.validate_coverage(dag) @staticmethod def _feedback(verdicts: list[JudgeVerdict]) -> str: rows = [] for index, verdict in enumerate(verdicts, start=1): if verdict.verdict == "reject": rows.append( f"judge {index}: {verdict.blocking_issues}; " f"suggested repair: {verdict.patch_suggestion}" ) return "\n".join(rows) async def verify( self, item: CanonicalItem, dag: ProofDAG, methods: MethodPlan, initial_bundle: CandidateBundle, ) -> KernelRunResult: bundle = initial_bundle iterations: list[VerificationIteration] = [] pass_streak = 0 streak_sha: str | None = None repaired_from_previous = False for iteration in range(1, self.config.max_iterations + 1): bundle_sha = sha256_payload(bundle) candidate_sha = sha256_payload(bundle.variant) verdicts = list( await asyncio.gather( *( self._judge_once( judge, item=item, dag=dag, methods=methods, bundle=bundle, iteration=iteration, judge_id=judge_id, ) for judge_id, judge in enumerate(self.judges, start=1) ) ) ) unanimous = all(verdict.verdict == "accept" for verdict in verdicts) if unanimous: if streak_sha not in {None, bundle_sha}: raise AssertionError("pass streak crossed provenance versions") streak_sha = bundle_sha pass_streak += 1 else: pass_streak = 0 streak_sha = None record = VerificationIteration( iteration=iteration, bundle_sha256=bundle_sha, candidate_sha256=candidate_sha, verdicts=verdicts, unanimous=unanimous, pass_streak_after=pass_streak, repaired_from_previous=repaired_from_previous, ) iterations.append(record) self.store.write_iteration(iteration, record) if pass_streak == self.config.streak_length: result = KernelRunResult( item_id=item.item_id, status="accepted", proof_dag=dag, method_plan=methods, accepted_replacement_plan=bundle.replacement_plan, accepted_diffused_proof=bundle.diffused_proof, accepted_candidate=bundle.variant, iterations=iterations, accepted_candidate_sha256=candidate_sha, accepted_bundle_sha256=bundle_sha, ) self.store.write_final(result) return result if not unanimous and iteration < self.config.max_iterations: bundle = await self.build_bundle( item, dag, methods, version=iteration + 1, previous_bundle=bundle, feedback=self._feedback(verdicts), ) repaired_from_previous = True else: repaired_from_previous = False result = KernelRunResult( item_id=item.item_id, status="rejected", proof_dag=dag, method_plan=methods, iterations=iterations, rejection_reason="no two consecutive unanimous rounds within T=15", ) self.store.write_final(result) return result async def run(self, item: CanonicalItem) -> KernelRunResult: self.store.write_input(item) self.store.write_config(self.config) dag = await self.construct_dag(item) methods = await self.summarize_methods(dag) bundle = await self.build_bundle(item, dag, methods, version=1) return await self.verify(item, dag, methods, bundle)