"""Canonical, type-checked graph-program abstract syntax trees.""" from __future__ import annotations from dataclasses import dataclass from typing import TypeAlias from strata.modeling.compose.types import GraphOperator, SemanticType, signature class ProgramTypeError(TypeError): pass @dataclass(frozen=True, slots=True) class AnchorRef: """Episode-local copied anchor; identities never enter the operator vocabulary.""" index: int value_type: SemanticType def __post_init__(self) -> None: if self.index < 0: raise ValueError("anchor index must be non-negative") @property def output_type(self) -> SemanticType: return self.value_type @property def depth(self) -> int: return 0 def canonical(self) -> tuple[str, int, str]: return ("anchor", self.index, self.value_type.value) @dataclass(frozen=True, slots=True) class Apply: operator: GraphOperator argument: "ProgramNode" def __post_init__(self) -> None: operator = GraphOperator(self.operator) object.__setattr__(self, "operator", operator) expected = signature(operator).input_type actual = self.argument.output_type if actual is not expected: raise ProgramTypeError( f"{operator.value} requires {expected.value}, got {actual.value}" ) @property def output_type(self) -> SemanticType: return signature(self.operator).output_type @property def depth(self) -> int: return self.argument.depth + 1 def canonical(self) -> tuple[str, tuple]: return (self.operator.value, self.argument.canonical()) ProgramNode: TypeAlias = AnchorRef | Apply def apply_chain(anchor: AnchorRef, operators: tuple[GraphOperator, ...]) -> ProgramNode: program: ProgramNode = anchor for operator in operators: program = Apply(operator, program) return program def operator_chain(program: ProgramNode) -> tuple[GraphOperator, ...]: operators: list[GraphOperator] = [] node = program while isinstance(node, Apply): operators.append(node.operator) node = node.argument operators.reverse() return tuple(operators) def anchor_ref(program: ProgramNode) -> AnchorRef: node = program while isinstance(node, Apply): node = node.argument return node def canonicalize(program: ProgramNode) -> tuple: """Return an identity-free structural key for equality and hashing.""" return program.canonical() __all__ = [ "AnchorRef", "Apply", "ProgramNode", "ProgramTypeError", "anchor_ref", "apply_chain", "canonicalize", "operator_chain", ]