"""A deterministic, fail-closed forward computation-graph trace validator.""" from __future__ import annotations from dataclasses import dataclass import hashlib import json import math import sys from typing import NoReturn MAX_NODES = 256 MAX_TEXT_LENGTH = 240 def _fail(message: str) -> NoReturn: raise ValueError(message) def _text(value: object, field: str) -> str: if type(value) is not str or not value.strip() or len(value) > MAX_TEXT_LENGTH: _fail(f"{field} must be non-empty bounded text") return value def _number(value: object, field: str) -> float: if type(value) not in (int, float): _fail(f"{field} must be a concrete number") result = float(value) if not math.isfinite(result): _fail(f"{field} must be finite") if result != 0.0 and abs(result) < sys.float_info.min: _fail(f"{field} is below the supported numeric resolution") return result def _digest(payload: object) -> str: data = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True) return "sha256:" + hashlib.sha256(data.encode("utf-8")).hexdigest() @dataclass(frozen=True) class Node: node_id: str operator: str inputs: tuple[str, ...] = () constant: float | None = None def __post_init__(self) -> None: object.__setattr__(self, "inputs", tuple(self.inputs)) _text(self.node_id, "node_id") _text(self.operator, "operator") if self.operator not in {"input", "constant", "add", "multiply", "relu"}: _fail("operator is not supported by graph-contract-v1") if any(type(item) is not str or not item for item in self.inputs): _fail("inputs must contain concrete non-empty node IDs") arity = {"input": 0, "constant": 0, "add": 2, "multiply": 2, "relu": 1}[self.operator] if len(self.inputs) != arity: _fail(f"{self.operator} requires arity {arity}") if self.operator == "constant": _number(self.constant, "constant") elif self.constant is not None: _fail("only constant nodes may carry a constant value") @dataclass(frozen=True) class GraphContract: graph_version: str scope_id: str model_revision: str operator_contract_version: str owner: str derived_absolute_tolerance: float nodes: tuple[Node, ...] topological_order: tuple[str, ...] outputs: tuple[str, ...] def __post_init__(self) -> None: object.__setattr__(self, "nodes", tuple(self.nodes)) object.__setattr__(self, "topological_order", tuple(self.topological_order)) object.__setattr__(self, "outputs", tuple(self.outputs)) for field in ("graph_version", "scope_id", "model_revision", "operator_contract_version", "owner"): _text(getattr(self, field), field) tolerance = _number(self.derived_absolute_tolerance, "derived_absolute_tolerance") if not 0.0 <= tolerance <= 1e-6: _fail("derived_absolute_tolerance must be between zero and 1e-6") if not 1 <= len(self.nodes) <= MAX_NODES: _fail(f"nodes must contain between 1 and {MAX_NODES} records") if any(type(node) is not Node for node in self.nodes): _fail("nodes must contain concrete Node records") for node in self.nodes: node.__post_init__() ids = tuple(node.node_id for node in self.nodes) if len(ids) != len(set(ids)): _fail("node IDs must be unique") if any(type(node_id) is not str or not node_id for node_id in self.topological_order): _fail("topological_order must contain concrete non-empty node IDs") if len(self.topological_order) != len(set(self.topological_order)): _fail("topological_order must not contain duplicates") if set(self.topological_order) != set(ids): _fail("topological_order must contain every declared node exactly once") if any(type(node_id) is not str or not node_id for node_id in self.outputs): _fail("outputs must contain concrete non-empty node IDs") if not self.outputs or len(self.outputs) != len(set(self.outputs)): _fail("outputs must be a non-empty unique sequence") if not set(self.outputs).issubset(ids): _fail("every output must name a declared node") positions = {node_id: index for index, node_id in enumerate(self.topological_order)} node_by_id = {node.node_id: node for node in self.nodes} for node_id in self.topological_order: for dependency in node_by_id[node_id].inputs: if dependency not in positions: _fail("an edge references an undeclared node") if positions[dependency] >= positions[node_id]: _fail("topological_order violates an edge or contains a cycle") reachable: set[str] = set() stack = list(self.outputs) while stack: current = stack.pop() if current in reachable: continue reachable.add(current) stack.extend(node_by_id[current].inputs) if reachable != set(ids): _fail("every declared node must contribute to a declared output") @dataclass(frozen=True) class ValueBinding: node_id: str value: float def __post_init__(self) -> None: _text(self.node_id, "binding.node_id") _number(self.value, "binding.value") @dataclass(frozen=True) class TraceEvidence: trace_id: str graph_version: str operator_contract_version: str scope_id: str model_revision: str input_source_version: str observed_at: str owner: str input_bindings: tuple[ValueBinding, ...] claimed_trace: tuple[ValueBinding, ...] def __post_init__(self) -> None: object.__setattr__(self, "input_bindings", tuple(self.input_bindings)) object.__setattr__(self, "claimed_trace", tuple(self.claimed_trace)) for field in ("trace_id", "graph_version", "operator_contract_version", "scope_id", "model_revision", "input_source_version", "observed_at", "owner"): _text(getattr(self, field), field) for name, bindings in (("input_bindings", self.input_bindings), ("claimed_trace", self.claimed_trace)): if not bindings or len(bindings) > MAX_NODES: _fail(f"{name} has an invalid bounded workload") if any(type(binding) is not ValueBinding for binding in bindings): _fail(f"{name} must contain concrete ValueBinding records") for binding in bindings: binding.__post_init__() ids = [binding.node_id for binding in bindings] if len(ids) != len(set(ids)): _fail(f"{name} contains duplicate node IDs") @dataclass(frozen=True) class TraceAudit: trace_content_id: str node_count: int output_values: tuple[ValueBinding, ...] decision: str def validate_forward_trace(contract: GraphContract, evidence: TraceEvidence) -> TraceAudit: if type(contract) is not GraphContract or type(evidence) is not TraceEvidence: _fail("validator requires concrete GraphContract and TraceEvidence records") contract.__post_init__() evidence.__post_init__() for field in ("graph_version", "operator_contract_version", "scope_id", "model_revision", "owner"): if getattr(evidence, field) != getattr(contract, field): _fail(f"trace {field} does not match the graph contract") node_by_id = {node.node_id: node for node in contract.nodes} input_nodes = {node.node_id for node in contract.nodes if node.operator == "input"} supplied = {binding.node_id: _number(binding.value, "input value") for binding in evidence.input_bindings} if set(supplied) != input_nodes: _fail("input bindings must cover exactly the graph input nodes") claimed = {binding.node_id: _number(binding.value, "claimed value") for binding in evidence.claimed_trace} if tuple(binding.node_id for binding in evidence.claimed_trace) != contract.topological_order: _fail("claimed trace must follow the complete declared topological order") values: dict[str, float] = {} for node_id in contract.topological_order: node = node_by_id[node_id] if node.operator == "input": value = supplied[node_id] elif node.operator == "constant": value = _number(node.constant, "constant") elif node.operator == "add": value = math.fsum(values[parent] for parent in node.inputs) elif node.operator == "multiply": value = math.prod(values[parent] for parent in node.inputs) else: value = max(0.0, values[node.inputs[0]]) if not math.isfinite(value): _fail(f"node {node_id} overflowed the supported numeric range") if value != 0.0 and abs(value) < sys.float_info.min: _fail(f"node {node_id} is below the supported numeric resolution") if ( node.operator == "multiply" and all(values[parent] != 0.0 for parent in node.inputs) and value == 0.0 ): _fail(f"node {node_id} underflowed the supported numeric range") values[node_id] = value if node_id not in claimed: _fail(f"claimed trace does not reproduce node {node_id}") if node.operator in {"input", "constant"}: matches = claimed[node_id] == value else: matches = abs(claimed[node_id] - value) <= contract.derived_absolute_tolerance if not matches: _fail(f"claimed trace does not reproduce node {node_id}") payload = { "contract": { "graph": contract.graph_version, "scope": contract.scope_id, "model": contract.model_revision, "operators": contract.operator_contract_version, "owner": contract.owner, "derived_absolute_tolerance": contract.derived_absolute_tolerance, "nodes": [(n.node_id, n.operator, list(n.inputs), n.constant) for n in contract.nodes], "order": list(contract.topological_order), "outputs": list(contract.outputs), }, "evidence": { "trace": evidence.trace_id, "graph": evidence.graph_version, "scope": evidence.scope_id, "operators": evidence.operator_contract_version, "model": evidence.model_revision, "source": evidence.input_source_version, "observed_at": evidence.observed_at, "owner": evidence.owner, "inputs": [(b.node_id, b.value) for b in evidence.input_bindings], "claimed": [(b.node_id, b.value) for b in evidence.claimed_trace], }, } outputs = tuple(ValueBinding(node_id, values[node_id]) for node_id in contract.outputs) return TraceAudit(_digest(payload), len(values), outputs, "PASS") ILLUSTRATIVE_GRAPH = GraphContract( graph_version="renewal-score-graph-v1", scope_id="illustrative-renewal-score", model_revision="renewal-forward-r3", operator_contract_version="graph-contract-v1", owner="retention-ml", derived_absolute_tolerance=1e-12, nodes=( Node("usage", "input"), Node("weight", "constant", constant=2.0), Node("bias", "constant", constant=-0.5), Node("scaled", "multiply", ("usage", "weight")), Node("margin", "add", ("scaled", "bias")), Node("score", "relu", ("margin",)), ), topological_order=("usage", "weight", "bias", "scaled", "margin", "score"), outputs=("score",), ) ILLUSTRATIVE_TRACE = TraceEvidence( trace_id="trace-0042", graph_version="renewal-score-graph-v1", operator_contract_version="graph-contract-v1", scope_id="illustrative-renewal-score", model_revision="renewal-forward-r3", input_source_version="feature-snapshot-2026-08-25", observed_at="2026-08-25T09:00:00Z", owner="retention-ml", input_bindings=(ValueBinding("usage", 0.8),), claimed_trace=( ValueBinding("usage", 0.8), ValueBinding("weight", 2.0), ValueBinding("bias", -0.5), ValueBinding("scaled", 1.6), ValueBinding("margin", 1.1), ValueBinding("score", 1.1), ), ) def format_example() -> str: audit = validate_forward_trace(ILLUSTRATIVE_GRAPH, ILLUSTRATIVE_TRACE) return "\n".join(( "example=illustrative_only", f"graph_version={ILLUSTRATIVE_GRAPH.graph_version}", f"trace_id={audit.trace_content_id}", f"nodes={audit.node_count}", f"score={audit.output_values[0].value:.3f}", f"decision={audit.decision}", )) if __name__ == "__main__": print(format_example())