diff options
| author | Anonymous Authors <anonymous@invalid.example> | 2026-07-24 14:03:51 -0500 |
|---|---|---|
| committer | Anonymous Authors <anonymous@invalid.example> | 2026-07-24 14:03:51 -0500 |
| commit | 708f2af9c6985e9cb5cd53e434a7d3b8dfa2b4ac (patch) | |
| tree | f4aa3d2240a621e9057ddad2a4d8ec7820db5cde /src/gap_pipeline/models.py | |
| parent | d4eb26780a8a8c70ca75812af0c3c6a295c0797c (diff) | |
Diffstat (limited to 'src/gap_pipeline/models.py')
| -rw-r--r-- | src/gap_pipeline/models.py | 126 |
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: |
