"""Prompt-faithful kernel generation and the J=5, K=2, T=15 loop.""" from __future__ import annotations import asyncio from pydantic import BaseModel, ConfigDict, model_validator from .clients import JsonLLM from .models import ( CanonicalItem, IterationRecord, JudgeVerdict, KernelCandidate, KernelPlan, KernelRunResult, ModelCallRecord, ProofPlanDAG, ) from .prompts import ( FIX_SYSTEM_PROMPT, JUDGE_SYSTEM_PROMPT, KERNEL_GENERATE_SYSTEM, KERNEL_PLAN_SYSTEM, fix_user, judge_user, kernel_generate_user, kernel_plan_user, ) from .store import RunStore, sha256_payload class PipelineConfig(BaseModel): model_config = ConfigDict(extra="forbid") protocol_name: str = "gap-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) -> "PipelineConfig": 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 KernelPipeline: def __init__( self, *, proposer: JsonLLM, judges: list[JsonLLM], store: RunStore, config: PipelineConfig, ) -> 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 extract_plan( self, item: CanonicalItem, ) -> tuple[KernelPlan, ProofPlanDAG]: request_id = f"{item.item_id}.plan" payload = await self._call( self.proposer, request_id=request_id, system_prompt=KERNEL_PLAN_SYSTEM, user_prompt=kernel_plan_user(item), ) plan = KernelPlan.model_validate(payload) dag = ProofPlanDAG.from_plan(plan) self.store.write_stage( "01_proof_dag", dag, request_id=request_id, ) self.store.write_stage( "02_method_plan", {"method_labels": plan.core_steps}, request_id=request_id, ) self.store.write_stage( "03_mutable_slots", { key: value.model_dump(mode="json") for key, value in plan.mutable_slots.items() }, request_id=request_id, ) return plan, dag async def generate_candidate( self, item: CanonicalItem, plan: KernelPlan, ) -> KernelCandidate: request_id = f"{item.item_id}.candidate" payload = await self._call( self.proposer, request_id=request_id, system_prompt=KERNEL_GENERATE_SYSTEM, user_prompt=kernel_generate_user(item, plan), ) candidate = KernelCandidate.model_validate(payload) self.store.write_stage( "04_regenerated_proof", {"solution": candidate.solution}, request_id=request_id, ) self.store.write_stage( "05_variant_question", {"question": candidate.question}, request_id=request_id, ) return candidate async def _judge_once( self, judge: JsonLLM, *, item: CanonicalItem, plan: KernelPlan, dag: ProofPlanDAG, candidate: KernelCandidate, iteration: int, judge_id: int, ) -> JudgeVerdict: request_id = f"{item.item_id}.verify.t{iteration:02d}.j{judge_id}" payload = await self._call( judge, request_id=request_id, system_prompt=JUDGE_SYSTEM_PROMPT, user_prompt=judge_user(item, plan, dag, candidate), ) return JudgeVerdict.model_validate(payload).validate_coverage(dag) async def _repair( self, *, item: CanonicalItem, candidate: KernelCandidate, verdicts: list[JudgeVerdict], iteration: int, ) -> KernelCandidate: problem_issues: list[str] = [] solution_issues: list[str] = [] for verdict in verdicts: if verdict.verdict == "reject": problem_issues.append(verdict.blocking_issues) solution_issues.append( verdict.patch_suggestion or verdict.blocking_issues ) request_id = f"{item.item_id}.repair.after_t{iteration:02d}" payload = await self._call( self.proposer, request_id=request_id, system_prompt=FIX_SYSTEM_PROMPT, user_prompt=fix_user( item, candidate, problem_issues="; ".join(problem_issues), solution_issues="; ".join(solution_issues), ), ) repaired = KernelCandidate( question=str(payload["corrected_question"]), solution=str(payload["corrected_solution"]), ) self.store.write_stage( f"repair_after_{iteration:02d}", repaired, request_id=request_id, ) return repaired async def verify_candidate( self, *, item: CanonicalItem, plan: KernelPlan, dag: ProofPlanDAG, candidate: KernelCandidate, ) -> KernelRunResult: dag.validate_plan(plan) current = candidate iterations: list[IterationRecord] = [] pass_streak = 0 streak_sha: str | None = None for iteration in range(1, self.config.max_iterations + 1): current_sha = sha256_payload(current) verdicts = list( await asyncio.gather( *( self._judge_once( judge, item=item, plan=plan, dag=dag, candidate=current, 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, current_sha}: raise AssertionError("pass streak crossed candidate versions") streak_sha = current_sha pass_streak += 1 else: pass_streak = 0 streak_sha = None record = IterationRecord( iteration=iteration, candidate_sha256=current_sha, verdicts=verdicts, unanimous=unanimous, pass_streak_after=pass_streak, ) 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", plan=plan, proof_dag=dag, accepted_candidate=current, iterations=iterations, accepted_candidate_sha256=current_sha, ) self.store.write_final(result) return result if not unanimous and iteration < self.config.max_iterations: current = await self._repair( item=item, candidate=current, verdicts=verdicts, iteration=iteration, ) result = KernelRunResult( item_id=item.item_id, status="rejected", plan=plan, proof_dag=dag, 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) plan, dag = await self.extract_plan(item) candidate = await self.generate_candidate(item, plan) return await self.verify_candidate( item=item, plan=plan, dag=dag, candidate=candidate, )