summaryrefslogtreecommitdiff
path: root/src/gap_pipeline/models.py
diff options
context:
space:
mode:
authorAnonymous Authors <anonymous@invalid.example>2026-07-24 14:03:51 -0500
committerAnonymous Authors <anonymous@invalid.example>2026-07-24 14:03:51 -0500
commit708f2af9c6985e9cb5cd53e434a7d3b8dfa2b4ac (patch)
treef4aa3d2240a621e9057ddad2a4d8ec7820db5cde /src/gap_pipeline/models.py
parentd4eb26780a8a8c70ca75812af0c3c6a295c0797c (diff)
Validate path-structured proof-plan DAGHEADmain
Diffstat (limited to 'src/gap_pipeline/models.py')
-rw-r--r--src/gap_pipeline/models.py126
1 files changed, 126 insertions, 0 deletions
diff --git a/src/gap_pipeline/models.py b/src/gap_pipeline/models.py
index bf80c5b..6fab32f 100644
--- a/src/gap_pipeline/models.py
+++ b/src/gap_pipeline/models.py
@@ -2,6 +2,7 @@
from __future__ import annotations
+import re
from datetime import datetime, timezone
from typing import Any, Literal
@@ -67,6 +68,115 @@ class KernelPlan(StrictModel):
return self
+class ProofPlanNode(StrictModel):
+ node_id: str
+ method_label: str
+ dependencies: list[str] = Field(default_factory=list)
+
+ @model_validator(mode="after")
+ def validate_content(self) -> "ProofPlanNode":
+ if not self.node_id.strip():
+ raise ValueError("proof-plan node ID must be non-empty")
+ if not self.method_label.strip():
+ raise ValueError("proof-plan method label must be non-empty")
+ return self
+
+
+class ProofPlanDAG(StrictModel):
+ """Typed path-DAG induced by Prompt-A's ordered proof-plan steps."""
+
+ nodes: list[ProofPlanNode] = Field(min_length=1)
+ terminal_node_id: str
+ structure: Literal["path"] = "path"
+
+ @classmethod
+ def from_plan(cls, plan: KernelPlan) -> "ProofPlanDAG":
+ nodes = [
+ ProofPlanNode(
+ node_id=f"n{index}",
+ method_label=step,
+ dependencies=[] if index == 1 else [f"n{index - 1}"],
+ )
+ for index, step in enumerate(plan.core_steps, start=1)
+ ]
+ return cls(nodes=nodes, terminal_node_id=nodes[-1].node_id)
+
+ def method_labels_payload(self) -> list[str]:
+ """Keep the judge input a sequence while making node coverage auditable."""
+
+ return [f"{node.node_id}: {node.method_label}" for node in self.nodes]
+
+ def validate_plan(self, plan: KernelPlan) -> "ProofPlanDAG":
+ labels = [node.method_label for node in self.nodes]
+ if labels != plan.core_steps:
+ raise ValueError(
+ "proof-plan DAG labels must equal Prompt-A core_steps in order"
+ )
+ return self
+
+ @model_validator(mode="after")
+ def validate_graph(self) -> "ProofPlanDAG":
+ node_ids = [node.node_id for node in self.nodes]
+ if len(node_ids) != len(set(node_ids)):
+ raise ValueError("proof-plan DAG node IDs must be unique")
+
+ known = set(node_ids)
+ if self.terminal_node_id not in known:
+ raise ValueError("proof-plan DAG terminal node is unknown")
+
+ dependents: dict[str, list[str]] = {node_id: [] for node_id in node_ids}
+ indegree = {node.node_id: len(node.dependencies) for node in self.nodes}
+ for node in self.nodes:
+ if len(node.dependencies) != len(set(node.dependencies)):
+ raise ValueError(f"{node.node_id} has duplicate dependencies")
+ for dependency in node.dependencies:
+ if dependency not in known:
+ raise ValueError(
+ f"{node.node_id} depends on unknown node {dependency}"
+ )
+ if dependency == node.node_id:
+ raise ValueError(f"{node.node_id} cannot depend on itself")
+ dependents[dependency].append(node.node_id)
+
+ ready = [node_id for node_id in node_ids if indegree[node_id] == 0]
+ visited = 0
+ while ready:
+ node_id = ready.pop()
+ visited += 1
+ for dependent in dependents[node_id]:
+ indegree[dependent] -= 1
+ if indegree[dependent] == 0:
+ ready.append(dependent)
+ if visited != len(node_ids):
+ raise ValueError("proof-plan graph contains a cycle")
+
+ by_id = {node.node_id: node for node in self.nodes}
+ ancestors: set[str] = set()
+ pending = [self.terminal_node_id]
+ while pending:
+ node_id = pending.pop()
+ if node_id in ancestors:
+ continue
+ ancestors.add(node_id)
+ pending.extend(by_id[node_id].dependencies)
+ if ancestors != known:
+ raise ValueError("every proof-plan node must lead to the terminal node")
+
+ expected_ids = [f"n{index}" for index in range(1, len(self.nodes) + 1)]
+ if node_ids != expected_ids:
+ raise ValueError("path-DAG nodes must be ordered n1 through nN")
+ for index, node in enumerate(self.nodes):
+ expected_dependencies = [] if index == 0 else [node_ids[index - 1]]
+ if node.dependencies != expected_dependencies:
+ raise ValueError(
+ f"{node.node_id} must depend exactly on "
+ f"{expected_dependencies or 'no prior node'}"
+ )
+ if self.terminal_node_id != node_ids[-1]:
+ raise ValueError("path-DAG terminal must be the final ordered node")
+ return self
+
+
class KernelCandidate(StrictModel):
question: str
solution: str
@@ -92,6 +202,20 @@ class JudgeVerdict(StrictModel):
raise ValueError("reject verdict must identify a blocking issue")
return self
+ def validate_coverage(self, dag: ProofPlanDAG) -> "JudgeVerdict":
+ missing = [
+ node.node_id
+ for node in dag.nodes
+ if re.search(
+ rf"(?<![A-Za-z0-9_]){re.escape(node.node_id)}(?![A-Za-z0-9_])",
+ self.step_by_step_check,
+ )
+ is None
+ ]
+ if missing:
+ raise ValueError(f"judge check does not cover DAG nodes: {missing}")
+ return self
+
class IterationRecord(StrictModel):
iteration: int
@@ -105,6 +229,7 @@ class KernelRunResult(StrictModel):
item_id: str
status: Literal["accepted", "rejected"]
plan: KernelPlan
+ proof_dag: ProofPlanDAG
accepted_candidate: KernelCandidate | None = None
iterations: list[IterationRecord]
rejection_reason: str = ""
@@ -113,6 +238,7 @@ class KernelRunResult(StrictModel):
@model_validator(mode="after")
def validate_terminal_state(self) -> "KernelRunResult":
+ self.proof_dag.validate_plan(self.plan)
if self.status == "accepted" and self.accepted_candidate is None:
raise ValueError("accepted run must include its candidate")
if self.status == "rejected" and not self.rejection_reason: