summaryrefslogtreecommitdiff
path: root/src/gap_pipeline/kernel_models.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/gap_pipeline/kernel_models.py')
-rw-r--r--src/gap_pipeline/kernel_models.py54
1 files changed, 37 insertions, 17 deletions
diff --git a/src/gap_pipeline/kernel_models.py b/src/gap_pipeline/kernel_models.py
index 32abf39..f835473 100644
--- a/src/gap_pipeline/kernel_models.py
+++ b/src/gap_pipeline/kernel_models.py
@@ -38,6 +38,11 @@ class ProofDAG(StrictModel):
def node_ids(self) -> list[str]:
return [node.node_id for node in self.nodes]
+ def leaf_node_ids(self) -> list[str]:
+ """Return source leaves: nodes with no prerequisite dependencies."""
+
+ return [node.node_id for node in self.nodes if not node.dependencies]
+
@model_validator(mode="after")
def validate_graph(self) -> "ProofDAG":
node_ids = self.node_ids()
@@ -113,6 +118,8 @@ class ReplacementChange(StrictModel):
]
if any(not value.strip() for value in text_fields):
raise ValueError("replacement fields must be non-empty")
+ if not re.fullmatch(r"slot[1-9][0-9]*", self.slot_id):
+ raise ValueError("replacement slot IDs must be slot1, slot2, ...")
if self.original_value.strip() == self.replacement_value.strip():
raise ValueError("replacement must differ from the original value")
return self
@@ -140,6 +147,34 @@ class ReplacementPlan(StrictModel):
}
if unknown:
raise ValueError(f"replacement plan references unknown nodes {sorted(unknown)}")
+ leaves = set(dag.leaf_node_ids())
+ non_leaf = {
+ change.source_node_id
+ for change in self.changes
+ if change.source_node_id not in leaves
+ }
+ if non_leaf:
+ raise ValueError(
+ "replacement plan must target source leaf nodes; "
+ f"received {sorted(non_leaf)}"
+ )
+ return self
+
+ def validate_repair_of(
+ self,
+ previous: "ReplacementPlan",
+ ) -> "ReplacementPlan":
+ current_targets = [
+ (change.slot_id, change.source_node_id) for change in self.changes
+ ]
+ previous_targets = [
+ (change.slot_id, change.source_node_id)
+ for change in previous.changes
+ ]
+ if current_targets != previous_targets:
+ raise ValueError(
+ "repair must preserve replacement slot IDs and source leaf nodes"
+ )
return self
@@ -234,13 +269,11 @@ class CandidateBundle(StrictModel):
class JudgeVerdict(StrictModel):
verdict: Literal["accept", "reject"]
step_by_step_check: str
- replacement_check: str
blocking_issues: str = ""
patch_suggestion: str = ""
@field_validator(
"step_by_step_check",
- "replacement_check",
"blocking_issues",
"patch_suggestion",
mode="before",
@@ -279,7 +312,6 @@ class JudgeVerdict(StrictModel):
def validate_coverage(
self,
dag: ProofDAG,
- replacement_plan: ReplacementPlan,
) -> "JudgeVerdict":
missing_nodes = [
node_id
@@ -290,20 +322,8 @@ class JudgeVerdict(StrictModel):
)
is None
]
- missing_slots = [
- change.slot_id
- for change in replacement_plan.changes
- if re.search(
- rf"(?<![A-Za-z0-9_]){re.escape(change.slot_id)}(?![A-Za-z0-9_])",
- self.replacement_check,
- )
- is None
- ]
- if missing_nodes or missing_slots:
- raise ValueError(
- "judge coverage incomplete: "
- f"nodes={missing_nodes}, slots={missing_slots}"
- )
+ if missing_nodes:
+ raise ValueError(f"judge coverage incomplete: nodes={missing_nodes}")
return self