Skip to main content

laddu_expr/
expression.rs

1use std::{
2    collections::{HashMap, HashSet},
3    fmt,
4    hash::Hash,
5    sync::{Arc, OnceLock},
6};
7
8use num::complex::Complex64;
9use serde::{Deserialize, Serialize};
10
11use crate::{
12    ExprGraphError, ExprShapeError, ParamError, ParamResult,
13    parameters::{InitialSpec, ParamState, Parameter, ParameterUpdate},
14};
15
16/// Stable identifier for a node in a serialized [`ExprGraph`].
17#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
18pub struct ExprId(u64);
19
20impl ExprId {
21    /// Creates an identifier from a zero-based node index.
22    pub fn from_index(index: usize) -> Self {
23        Self(index as u64)
24    }
25
26    /// Returns the zero-based node index.
27    pub fn index(self) -> usize {
28        self.0 as usize
29    }
30}
31
32/// Runtime value category produced by an expression node.
33#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
34pub enum ValueKind {
35    /// A real scalar.
36    Real,
37    /// A complex scalar.
38    Complex,
39    /// A vector with a fixed number of elements.
40    Vector {
41        /// Number of vector elements.
42        len: usize,
43    },
44    /// A matrix with fixed dimensions.
45    Matrix {
46        /// Number of rows.
47        rows: usize,
48        /// Number of columns.
49        cols: usize,
50    },
51}
52
53/// Statically known scalar number category.
54#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
55pub enum NumberClass {
56    /// No narrower category is known.
57    Unknown,
58    /// The value is real.
59    Real,
60    /// The value is purely imaginary.
61    Imaginary,
62    /// The value may have real and imaginary components.
63    Complex,
64}
65
66/// Context-free value semantics inferred for one expression node.
67#[derive(Copy, Clone, Debug, PartialEq, Eq)]
68pub struct ExprNodeSemantics {
69    /// Runtime value kind.
70    pub value_kind: ValueKind,
71    /// Known relationship between real and imaginary components.
72    pub number_class: NumberClass,
73}
74
75fn add_number_class(lhs: NumberClass, rhs: NumberClass) -> NumberClass {
76    use NumberClass::{Complex, Imaginary, Real, Unknown};
77    match (lhs, rhs) {
78        (Real, Real) => Real,
79        (Imaginary, Imaginary) => Imaginary,
80        (Complex, _) | (_, Complex) => Complex,
81        (Unknown, _) | (_, Unknown) => Unknown,
82        _ => Complex,
83    }
84}
85
86fn mul_number_class(lhs: NumberClass, rhs: NumberClass) -> NumberClass {
87    use NumberClass::{Complex, Imaginary, Real, Unknown};
88    match (lhs, rhs) {
89        (Real, Real) | (Imaginary, Imaginary) => Real,
90        (Real, Imaginary) | (Imaginary, Real) => Imaginary,
91        (Complex, _) | (_, Complex) => Complex,
92        (Unknown, _) | (_, Unknown) => Unknown,
93    }
94}
95
96/// Intrinsic source of a node's evaluation dependencies.
97#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
98pub enum ExprDependencyKind {
99    /// The node is a compile-time constant.
100    Constant,
101    /// The node directly reads a parameter definition.
102    Parameter,
103    /// The node directly reads event data.
104    Event,
105    /// The node inherits the union of its children's dependencies.
106    Children,
107}
108
109/// Structural shape of an expression.
110#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
111pub enum ExprShape {
112    /// A scalar expression.
113    Scalar,
114    /// A vector expression.
115    Vector {
116        /// Number of vector elements.
117        len: usize,
118    },
119    /// A matrix expression.
120    Matrix {
121        /// Number of rows.
122        rows: usize,
123        /// Number of columns.
124        cols: usize,
125    },
126}
127
128impl fmt::Display for ExprShape {
129    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
130        match self {
131            Self::Scalar => write!(f, "scalar"),
132            Self::Vector { len } => write!(f, "vector[{len}]"),
133            Self::Matrix { rows, cols } => write!(f, "matrix[{rows}x{cols}]"),
134        }
135    }
136}
137
138/// Converts a component selector into a zero-based index.
139pub trait ComponentIndex {
140    /// Returns the selected zero-based component index.
141    fn component_index(self) -> usize;
142}
143
144impl ComponentIndex for usize {
145    fn component_index(self) -> usize {
146        self
147    }
148}
149
150impl ComponentIndex for i32 {
151    fn component_index(self) -> usize {
152        usize::try_from(self).expect("component index must be nonnegative")
153    }
154}
155
156#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
157/// A named component of a four-momentum in `(E, px, py, pz)` order.
158pub enum P4Component {
159    /// Energy.
160    E,
161    /// Momentum in the x direction.
162    Px,
163    /// Momentum in the y direction.
164    Py,
165    /// Momentum in the z direction.
166    Pz,
167}
168
169impl P4Component {
170    /// Return the lowercase event-column suffix for this component.
171    pub fn label(self) -> &'static str {
172        match self {
173            Self::E => "e",
174            Self::Px => "px",
175            Self::Py => "py",
176            Self::Pz => "pz",
177        }
178    }
179
180    /// Return the component's position in `(E, px, py, pz)` order.
181    pub fn index(self) -> usize {
182        match self {
183            Self::E => 0,
184            Self::Px => 1,
185            Self::Py => 2,
186            Self::Pz => 3,
187        }
188    }
189}
190
191/// Unary operation in an expression graph.
192#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
193pub enum UnaryOp {
194    /// Arithmetic negation.
195    Neg,
196    /// Real part.
197    Real,
198    /// Imaginary part.
199    Imag,
200    /// Complex conjugate.
201    Conj,
202    /// Squared complex norm.
203    NormSqr,
204    /// Principal square root.
205    Sqrt,
206    /// Exponential.
207    Exp,
208    /// Sine.
209    Sin,
210    /// Cosine.
211    Cos,
212    /// Natural logarithm.
213    Log,
214    /// Integer power.
215    PowI(i32),
216}
217
218impl UnaryOp {
219    /// Applies this operation to a scalar complex value.
220    pub fn evaluate(&self, value: Complex64) -> Complex64 {
221        match self {
222            Self::Neg => -value,
223            Self::Real => Complex64::from(value.re),
224            Self::Imag => Complex64::from(value.im),
225            Self::Conj => value.conj(),
226            Self::NormSqr => Complex64::from(value.norm_sqr()),
227            Self::Sqrt => value.sqrt(),
228            Self::Exp => value.exp(),
229            Self::Sin => value.sin(),
230            Self::Cos => value.cos(),
231            Self::Log => value.ln(),
232            Self::PowI(power) => value.powi(*power),
233        }
234    }
235}
236
237/// Binary operation in an expression graph.
238#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
239pub enum BinaryOp {
240    /// Addition.
241    Add,
242    /// Subtraction.
243    Sub,
244    /// Multiplication.
245    Mul,
246    /// Division.
247    Div,
248    /// Two-argument arctangent of the real parts.
249    Atan2,
250}
251
252impl BinaryOp {
253    /// Applies this operation to two scalar complex values.
254    pub fn evaluate(&self, a: Complex64, b: Complex64) -> Complex64 {
255        match self {
256            Self::Add => a + b,
257            Self::Sub => a - b,
258            Self::Mul => a * b,
259            Self::Div => a / b,
260            Self::Atan2 => Complex64::from(a.re.atan2(b.re)),
261        }
262    }
263}
264
265/// Serialized node in a topologically ordered [`ExprGraph`].
266#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
267pub enum ExprNode {
268    /// A real constant.
269    RealConst(f64),
270    /// A complex constant.
271    ComplexConst(Complex64),
272    /// A scalar fit parameter.
273    ScalarParam(Parameter),
274    /// A named scalar event column.
275    EventScalar(Arc<str>),
276    /// One component of a named event four-momentum.
277    EventP4Component {
278        /// Event column base name.
279        name: Arc<str>,
280        /// Requested four-momentum component.
281        component: P4Component,
282    },
283    /// A unary operation.
284    Unary {
285        /// Operation to apply.
286        op: UnaryOp,
287        /// Input node.
288        input: ExprId,
289    },
290    /// A binary operation.
291    Binary {
292        /// Operation to apply.
293        op: BinaryOp,
294        /// Left operand.
295        lhs: ExprId,
296        /// Right operand.
297        rhs: ExprId,
298    },
299    /// A sum of zero or more terms.
300    NaryAdd {
301        /// Term nodes.
302        terms: Vec<ExprId>,
303    },
304    /// A product of zero or more factors.
305    NaryMul {
306        /// Factor nodes.
307        factors: Vec<ExprId>,
308    },
309    /// A complex scalar assembled from real and imaginary expressions.
310    Complex {
311        /// Real component.
312        re: ExprId,
313        /// Imaginary component.
314        im: ExprId,
315    },
316    /// A vector assembled from scalar elements.
317    Vector {
318        /// Scalar element nodes.
319        elements: Vec<ExprId>,
320    },
321    /// A row-major matrix assembled from scalar elements.
322    Matrix {
323        /// Number of rows.
324        rows: usize,
325        /// Number of columns.
326        cols: usize,
327        /// Row-major scalar elements.
328        elements: Vec<ExprId>,
329    },
330    /// A vector component selection.
331    Component {
332        /// Vector input.
333        input: ExprId,
334        /// Zero-based component index.
335        index: usize,
336    },
337    /// A matrix element selection.
338    MatrixElement {
339        /// Matrix input.
340        input: ExprId,
341        /// Zero-based row index.
342        row: usize,
343        /// Zero-based column index.
344        col: usize,
345    },
346    /// Matrix-matrix multiplication.
347    MatMul {
348        /// Left matrix.
349        lhs: ExprId,
350        /// Right matrix.
351        rhs: ExprId,
352    },
353    /// Matrix-vector multiplication.
354    MatVec {
355        /// Matrix operand.
356        matrix: ExprId,
357        /// Vector operand.
358        vector: ExprId,
359    },
360    /// Vector dot product.
361    Dot {
362        /// Left vector.
363        lhs: ExprId,
364        /// Right vector.
365        rhs: ExprId,
366    },
367    /// Solution of a linear system.
368    Solve {
369        /// Coefficient matrix.
370        matrix: ExprId,
371        /// Right-hand-side vector or matrix.
372        rhs: ExprId,
373    },
374}
375
376/// Bit-exact structural identity for a scalar parameter definition.
377///
378/// Equality includes state, initial-value policy, bounds, periodicity, scale,
379/// and user-facing labels. Floating-point values are compared by their bit
380/// patterns, so signed zero and distinct NaN payloads remain distinct.
381#[doc(hidden)]
382#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
383pub struct ParameterStructuralKey {
384    name: Arc<str>,
385    state: ParameterStateStructuralKey,
386    initial: InitialStructuralKey,
387    bounds: (Option<u64>, Option<u64>),
388    periodic: bool,
389    scale: Option<u64>,
390    unit: Option<Arc<str>>,
391    latex: Option<Arc<str>>,
392    description: Option<Arc<str>>,
393}
394
395#[doc(hidden)]
396#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
397pub enum ParameterStateStructuralKey {
398    Free,
399    Fixed(u64),
400}
401
402#[doc(hidden)]
403#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
404pub enum InitialStructuralKey {
405    Default,
406    Value(u64),
407    Uniform { min: u64, max: u64 },
408}
409
410/// Bit-exact, metadata-free structural identity for an expression node.
411///
412/// The key includes the node variant, semantic payload, child identifiers, and
413/// complete parameter definitions. It deliberately excludes [`ExprMetadata`].
414/// Its ordering is deterministic but its representation and hash values are an
415/// internal workspace contract, not a stable serialized or persisted format.
416#[doc(hidden)]
417#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
418pub enum ExprNodeStructuralKey {
419    RealConst(u64),
420    ComplexConst {
421        re: u64,
422        im: u64,
423    },
424    ScalarParam(ParameterStructuralKey),
425    EventScalar(Arc<str>),
426    EventP4Component {
427        name: Arc<str>,
428        component: P4Component,
429    },
430    Unary {
431        op: UnaryOp,
432        input: ExprId,
433    },
434    Binary {
435        op: BinaryOp,
436        lhs: ExprId,
437        rhs: ExprId,
438    },
439    NaryAdd {
440        terms: Vec<ExprId>,
441    },
442    NaryMul {
443        factors: Vec<ExprId>,
444    },
445    Complex {
446        re: ExprId,
447        im: ExprId,
448    },
449    Vector {
450        elements: Vec<ExprId>,
451    },
452    Matrix {
453        rows: usize,
454        cols: usize,
455        elements: Vec<ExprId>,
456    },
457    Component {
458        input: ExprId,
459        index: usize,
460    },
461    MatrixElement {
462        input: ExprId,
463        row: usize,
464        col: usize,
465    },
466    MatMul {
467        lhs: ExprId,
468        rhs: ExprId,
469    },
470    MatVec {
471        matrix: ExprId,
472        vector: ExprId,
473    },
474    Dot {
475        lhs: ExprId,
476        rhs: ExprId,
477    },
478    Solve {
479        matrix: ExprId,
480        rhs: ExprId,
481    },
482}
483
484impl From<&Parameter> for ParameterStructuralKey {
485    fn from(parameter: &Parameter) -> Self {
486        let state = match parameter.state() {
487            ParamState::Free => ParameterStateStructuralKey::Free,
488            ParamState::Fixed(value) => ParameterStateStructuralKey::Fixed(value.to_bits()),
489        };
490        let initial = match parameter.initial_spec() {
491            InitialSpec::Default => InitialStructuralKey::Default,
492            InitialSpec::Value(value) => InitialStructuralKey::Value(value.to_bits()),
493            InitialSpec::Uniform { min, max } => InitialStructuralKey::Uniform {
494                min: min.to_bits(),
495                max: max.to_bits(),
496            },
497        };
498        Self {
499            name: Arc::from(parameter.name()),
500            state,
501            initial,
502            bounds: (
503                parameter.bounds_spec().min.map(f64::to_bits),
504                parameter.bounds_spec().max.map(f64::to_bits),
505            ),
506            periodic: parameter.is_periodic(),
507            scale: parameter.scale().map(f64::to_bits),
508            unit: parameter.unit_label().map(Arc::from),
509            latex: parameter.latex_label().map(Arc::from),
510            description: parameter.description_text().map(Arc::from),
511        }
512    }
513}
514
515impl From<Complex64> for ExprNode {
516    fn from(value: Complex64) -> Self {
517        if value.im == 0.0 {
518            Self::RealConst(value.re)
519        } else {
520            Self::ComplexConst(value)
521        }
522    }
523}
524
525impl ExprNode {
526    /// Infers this node's context-free value semantics from the already
527    /// computed semantics of earlier nodes in the expression graph.
528    pub fn semantics(&self, children: &[ExprNodeSemantics]) -> ExprNodeSemantics {
529        ExprNodeSemantics {
530            value_kind: self.infer_value_kind(children),
531            number_class: self.infer_number_class(children),
532        }
533    }
534
535    fn infer_value_kind(&self, children: &[ExprNodeSemantics]) -> ValueKind {
536        match self {
537            Self::RealConst(_) | Self::ScalarParam(_) => ValueKind::Real,
538            Self::ComplexConst(value) => {
539                if value.im == 0.0 {
540                    ValueKind::Real
541                } else {
542                    ValueKind::Complex
543                }
544            }
545            Self::EventScalar(_) | Self::EventP4Component { .. } => ValueKind::Real,
546            Self::Unary { op, input } => match op {
547                UnaryOp::Real | UnaryOp::Imag | UnaryOp::NormSqr => ValueKind::Real,
548                UnaryOp::Neg
549                | UnaryOp::Conj
550                | UnaryOp::Sqrt
551                | UnaryOp::Exp
552                | UnaryOp::Sin
553                | UnaryOp::Cos
554                | UnaryOp::Log
555                | UnaryOp::PowI(_) => children[input.index()].value_kind,
556            },
557            Self::Binary { op, lhs, rhs } => {
558                if *op == BinaryOp::Atan2 {
559                    return ValueKind::Real;
560                }
561                if children[lhs.index()].value_kind == ValueKind::Real
562                    && children[rhs.index()].value_kind == ValueKind::Real
563                {
564                    ValueKind::Real
565                } else {
566                    ValueKind::Complex
567                }
568            }
569            Self::NaryAdd { terms } => {
570                if terms
571                    .iter()
572                    .all(|id| children[id.index()].value_kind == ValueKind::Real)
573                {
574                    ValueKind::Real
575                } else {
576                    ValueKind::Complex
577                }
578            }
579            Self::NaryMul { factors } => {
580                if factors
581                    .iter()
582                    .all(|id| children[id.index()].value_kind == ValueKind::Real)
583                {
584                    ValueKind::Real
585                } else {
586                    ValueKind::Complex
587                }
588            }
589            Self::Complex { .. } => ValueKind::Complex,
590            Self::Vector { elements } => ValueKind::Vector {
591                len: elements.len(),
592            },
593            Self::Matrix { rows, cols, .. } => ValueKind::Matrix {
594                rows: *rows,
595                cols: *cols,
596            },
597            Self::Component { input, .. } => match children[input.index()].value_kind {
598                ValueKind::Vector { .. } => ValueKind::Complex,
599                kind => kind,
600            },
601            Self::MatrixElement { .. } | Self::Dot { .. } => ValueKind::Complex,
602            Self::MatMul { lhs, rhs } => {
603                let ValueKind::Matrix { rows, .. } = children[lhs.index()].value_kind else {
604                    return ValueKind::Complex;
605                };
606                let ValueKind::Matrix { cols, .. } = children[rhs.index()].value_kind else {
607                    return ValueKind::Complex;
608                };
609                ValueKind::Matrix { rows, cols }
610            }
611            Self::MatVec { matrix, .. } => {
612                let ValueKind::Matrix { rows, .. } = children[matrix.index()].value_kind else {
613                    return ValueKind::Complex;
614                };
615                ValueKind::Vector { len: rows }
616            }
617            Self::Solve { rhs, .. } => children[rhs.index()].value_kind,
618        }
619    }
620
621    fn infer_number_class(&self, children: &[ExprNodeSemantics]) -> NumberClass {
622        match self {
623            Self::RealConst(_) | Self::ScalarParam(_) => NumberClass::Real,
624            Self::ComplexConst(value) => match (value.re == 0.0, value.im == 0.0) {
625                (_, true) => NumberClass::Real,
626                (true, false) => NumberClass::Imaginary,
627                (false, false) => NumberClass::Complex,
628            },
629            Self::EventScalar(_) | Self::EventP4Component { .. } => NumberClass::Real,
630            Self::Unary { op, input } => match op {
631                UnaryOp::Neg | UnaryOp::Conj => children[input.index()].number_class,
632                UnaryOp::Real | UnaryOp::Imag | UnaryOp::NormSqr => NumberClass::Real,
633                UnaryOp::Exp | UnaryOp::Sin | UnaryOp::Cos | UnaryOp::PowI(_) => {
634                    let input = children[input.index()].number_class;
635                    if input == NumberClass::Real {
636                        NumberClass::Real
637                    } else {
638                        NumberClass::Unknown
639                    }
640                }
641                UnaryOp::Sqrt | UnaryOp::Log => NumberClass::Unknown,
642            },
643            Self::Binary { op, lhs, rhs } => {
644                let lhs = children[lhs.index()].number_class;
645                let rhs = children[rhs.index()].number_class;
646                match op {
647                    BinaryOp::Add | BinaryOp::Sub => add_number_class(lhs, rhs),
648                    BinaryOp::Mul | BinaryOp::Div => mul_number_class(lhs, rhs),
649                    BinaryOp::Atan2 => NumberClass::Real,
650                }
651            }
652            Self::NaryAdd { terms } => {
653                let mut classes = terms.iter().map(|id| children[id.index()].number_class);
654                let Some(first) = classes.next() else {
655                    return NumberClass::Real;
656                };
657                classes.fold(first, add_number_class)
658            }
659            Self::NaryMul { factors } => {
660                let mut classes = factors.iter().map(|id| children[id.index()].number_class);
661                let Some(first) = classes.next() else {
662                    return NumberClass::Real;
663                };
664                classes.fold(first, mul_number_class)
665            }
666            Self::Complex { .. } => NumberClass::Complex,
667            Self::Vector { .. }
668            | Self::Matrix { .. }
669            | Self::Component { .. }
670            | Self::MatrixElement { .. }
671            | Self::MatMul { .. }
672            | Self::MatVec { .. }
673            | Self::Dot { .. }
674            | Self::Solve { .. } => NumberClass::Unknown,
675        }
676    }
677
678    /// Returns the intrinsic source of this node's evaluation dependencies.
679    pub fn dependency_kind(&self) -> ExprDependencyKind {
680        match self {
681            Self::RealConst(_) | Self::ComplexConst(_) => ExprDependencyKind::Constant,
682            Self::ScalarParam(_) => ExprDependencyKind::Parameter,
683            Self::EventScalar(_) | Self::EventP4Component { .. } => ExprDependencyKind::Event,
684            _ => ExprDependencyKind::Children,
685        }
686    }
687
688    /// Returns this node's bit-exact, metadata-free structural identity.
689    #[doc(hidden)]
690    pub fn structural_key(&self) -> ExprNodeStructuralKey {
691        match self {
692            Self::RealConst(value) => ExprNodeStructuralKey::RealConst(value.to_bits()),
693            Self::ComplexConst(value) => ExprNodeStructuralKey::ComplexConst {
694                re: value.re.to_bits(),
695                im: value.im.to_bits(),
696            },
697            Self::ScalarParam(parameter) => {
698                ExprNodeStructuralKey::ScalarParam(ParameterStructuralKey::from(parameter))
699            }
700            Self::EventScalar(name) => ExprNodeStructuralKey::EventScalar(Arc::clone(name)),
701            Self::EventP4Component { name, component } => ExprNodeStructuralKey::EventP4Component {
702                name: Arc::clone(name),
703                component: *component,
704            },
705            Self::Unary { op, input } => ExprNodeStructuralKey::Unary {
706                op: *op,
707                input: *input,
708            },
709            Self::Binary { op, lhs, rhs } => ExprNodeStructuralKey::Binary {
710                op: *op,
711                lhs: *lhs,
712                rhs: *rhs,
713            },
714            Self::NaryAdd { terms } => ExprNodeStructuralKey::NaryAdd {
715                terms: terms.clone(),
716            },
717            Self::NaryMul { factors } => ExprNodeStructuralKey::NaryMul {
718                factors: factors.clone(),
719            },
720            Self::Complex { re, im } => ExprNodeStructuralKey::Complex { re: *re, im: *im },
721            Self::Vector { elements } => ExprNodeStructuralKey::Vector {
722                elements: elements.clone(),
723            },
724            Self::Matrix {
725                rows,
726                cols,
727                elements,
728            } => ExprNodeStructuralKey::Matrix {
729                rows: *rows,
730                cols: *cols,
731                elements: elements.clone(),
732            },
733            Self::Component { input, index } => ExprNodeStructuralKey::Component {
734                input: *input,
735                index: *index,
736            },
737            Self::MatrixElement { input, row, col } => ExprNodeStructuralKey::MatrixElement {
738                input: *input,
739                row: *row,
740                col: *col,
741            },
742            Self::MatMul { lhs, rhs } => ExprNodeStructuralKey::MatMul {
743                lhs: *lhs,
744                rhs: *rhs,
745            },
746            Self::MatVec { matrix, vector } => ExprNodeStructuralKey::MatVec {
747                matrix: *matrix,
748                vector: *vector,
749            },
750            Self::Dot { lhs, rhs } => ExprNodeStructuralKey::Dot {
751                lhs: *lhs,
752                rhs: *rhs,
753            },
754            Self::Solve { matrix, rhs } => ExprNodeStructuralKey::Solve {
755                matrix: *matrix,
756                rhs: *rhs,
757            },
758        }
759    }
760
761    /// Creates the most compact constant-node representation for `value`.
762    pub fn from_folded_const(value: Complex64) -> Self {
763        if value.im == 0.0 && value.im.is_sign_positive() {
764            Self::RealConst(value.re)
765        } else {
766            Self::ComplexConst(value)
767        }
768    }
769
770    /// Returns the node's scalar constant value, if it is a constant.
771    pub fn const_value(&self) -> Option<Complex64> {
772        match self {
773            ExprNode::RealConst(value) => Some(Complex64::from(*value)),
774            ExprNode::ComplexConst(value) => Some(*value),
775            _ => None,
776        }
777    }
778
779    /// Returns whether `node` is the scalar constant zero.
780    pub fn is_zero(node: &ExprNode) -> bool {
781        node.const_value()
782            .is_some_and(|value| value == Complex64::ZERO)
783    }
784
785    /// Returns whether `node` is the scalar constant one.
786    pub fn is_one(node: &ExprNode) -> bool {
787        node.const_value()
788            .is_some_and(|value| value == Complex64::ONE)
789    }
790
791    /// Iterates over this node's direct dependencies in semantic operand order.
792    ///
793    /// The iterator borrows the node and does not allocate. Binary operands are
794    /// returned left-to-right, and vector, matrix, sum, and product children
795    /// retain their stored order.
796    pub fn children(&self) -> impl ExactSizeIterator<Item = ExprId> + DoubleEndedIterator + '_ {
797        (0..self.child_count()).map(|index| self.child_at(index))
798    }
799
800    /// Returns the identifiers of this node's direct dependencies.
801    ///
802    /// This compatibility helper collects [`Self::children`]. Prefer the
803    /// borrowed iterator when an owned vector is not required.
804    pub fn child_ids(&self) -> Vec<ExprId> {
805        self.children().collect()
806    }
807
808    /// Returns a copy of this node with each direct dependency transformed.
809    ///
810    /// Children are passed to `map` in the same semantic order as
811    /// [`Self::children`]. Non-child fields are preserved exactly.
812    pub fn map_children(&self, mut map: impl FnMut(ExprId) -> ExprId) -> Self {
813        match self {
814            Self::RealConst(_)
815            | Self::ComplexConst(_)
816            | Self::ScalarParam(_)
817            | Self::EventScalar(_)
818            | Self::EventP4Component { .. } => self.clone(),
819            Self::Unary { op, input } => Self::Unary {
820                op: *op,
821                input: map(*input),
822            },
823            Self::Binary { op, lhs, rhs } => Self::Binary {
824                op: *op,
825                lhs: map(*lhs),
826                rhs: map(*rhs),
827            },
828            Self::NaryAdd { terms } => Self::NaryAdd {
829                terms: terms.iter().copied().map(&mut map).collect(),
830            },
831            Self::NaryMul { factors } => Self::NaryMul {
832                factors: factors.iter().copied().map(&mut map).collect(),
833            },
834            Self::Complex { re, im } => Self::Complex {
835                re: map(*re),
836                im: map(*im),
837            },
838            Self::Vector { elements } => Self::Vector {
839                elements: elements.iter().copied().map(&mut map).collect(),
840            },
841            Self::Matrix {
842                rows,
843                cols,
844                elements,
845            } => Self::Matrix {
846                rows: *rows,
847                cols: *cols,
848                elements: elements.iter().copied().map(&mut map).collect(),
849            },
850            Self::Component { input, index } => Self::Component {
851                input: map(*input),
852                index: *index,
853            },
854            Self::MatrixElement { input, row, col } => Self::MatrixElement {
855                input: map(*input),
856                row: *row,
857                col: *col,
858            },
859            Self::MatMul { lhs, rhs } => Self::MatMul {
860                lhs: map(*lhs),
861                rhs: map(*rhs),
862            },
863            Self::MatVec { matrix, vector } => Self::MatVec {
864                matrix: map(*matrix),
865                vector: map(*vector),
866            },
867            Self::Dot { lhs, rhs } => Self::Dot {
868                lhs: map(*lhs),
869                rhs: map(*rhs),
870            },
871            Self::Solve { matrix, rhs } => Self::Solve {
872                matrix: map(*matrix),
873                rhs: map(*rhs),
874            },
875        }
876    }
877
878    fn child_count(&self) -> usize {
879        match self {
880            Self::RealConst(_)
881            | Self::ComplexConst(_)
882            | Self::ScalarParam(_)
883            | Self::EventScalar(_)
884            | Self::EventP4Component { .. } => 0,
885            Self::Unary { .. } | Self::Component { .. } | Self::MatrixElement { .. } => 1,
886            Self::Binary { .. }
887            | Self::Complex { .. }
888            | Self::MatMul { .. }
889            | Self::MatVec { .. }
890            | Self::Dot { .. }
891            | Self::Solve { .. } => 2,
892            Self::NaryAdd { terms } => terms.len(),
893            Self::NaryMul { factors } => factors.len(),
894            Self::Vector { elements } | Self::Matrix { elements, .. } => elements.len(),
895        }
896    }
897
898    fn child_at(&self, index: usize) -> ExprId {
899        match self {
900            Self::Unary { input, .. }
901            | Self::Component { input, .. }
902            | Self::MatrixElement { input, .. } => *input,
903            Self::Binary { lhs, rhs, .. }
904            | Self::Complex { re: lhs, im: rhs }
905            | Self::MatMul { lhs, rhs }
906            | Self::Dot { lhs, rhs } => [*lhs, *rhs][index],
907            Self::MatVec { matrix, vector } => [*matrix, *vector][index],
908            Self::Solve { matrix, rhs } => [*matrix, *rhs][index],
909            Self::NaryAdd { terms } => terms[index],
910            Self::NaryMul { factors } => factors[index],
911            Self::Vector { elements } | Self::Matrix { elements, .. } => elements[index],
912            Self::RealConst(_)
913            | Self::ComplexConst(_)
914            | Self::ScalarParam(_)
915            | Self::EventScalar(_)
916            | Self::EventP4Component { .. } => unreachable!("leaf node has no children"),
917        }
918    }
919}
920
921/// Broad origin category recorded in [`ExprMetadata`].
922#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
923pub enum ExprSourceKind {
924    /// Constant literal.
925    Const,
926    /// Fit parameter.
927    Param,
928    /// Event data.
929    Event,
930    /// Unary operation.
931    Unary,
932    /// Binary or n-ary operation.
933    Binary,
934    /// Complex-number construction.
935    Complex,
936    /// Vector construction or selection.
937    Vector,
938    /// Matrix construction or selection.
939    Matrix,
940    /// Linear-algebra operation.
941    LinearAlgebra,
942}
943
944/// User-facing annotations and origin information for an expression node.
945#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
946pub struct ExprMetadata {
947    source: ExprSourceKind,
948    name: Option<Arc<str>>,
949    tags: Vec<Arc<str>>,
950}
951
952impl ExprMetadata {
953    /// Creates metadata for the given source category.
954    pub fn new(source: ExprSourceKind) -> Self {
955        Self {
956            source,
957            name: None,
958            tags: Vec::new(),
959        }
960    }
961
962    /// Returns the node's source category.
963    pub fn source(&self) -> ExprSourceKind {
964        self.source
965    }
966
967    /// Returns the optional user-assigned name.
968    pub fn name(&self) -> Option<&str> {
969        self.name.as_deref()
970    }
971
972    /// Returns the user-assigned tags.
973    pub fn tags(&self) -> &[Arc<str>] {
974        &self.tags
975    }
976
977    /// Returns whether the metadata contains `tag`.
978    pub fn has_tag(&self, tag: &str) -> bool {
979        self.tags.iter().any(|candidate| candidate.as_ref() == tag)
980    }
981}
982
983/// Shareable symbolic expression represented internally as a directed acyclic graph.
984///
985/// Adding, subtracting, or multiplying a vector or matrix by a scalar acts
986/// elementwise in either operand order. Dividing a vector or matrix by a
987/// scalar is also elementwise; scalar divided by a tensor is unsupported.
988#[derive(Clone, Debug)]
989pub struct Expr {
990    node: Arc<DagNode>,
991}
992
993#[derive(Clone, Debug)]
994struct DagNode {
995    kind: DagNodeKind,
996    metadata: ExprMetadata,
997    shape: OnceLock<Result<ExprShape, ExprShapeError>>,
998}
999
1000#[derive(Clone, Debug)]
1001enum DagNodeKind {
1002    RealConst(f64),
1003    ComplexConst(Complex64),
1004    ScalarParam(Parameter),
1005    EventScalar(Arc<str>),
1006    EventP4Component {
1007        name: Arc<str>,
1008        component: P4Component,
1009    },
1010    Unary {
1011        op: UnaryOp,
1012        input: Expr,
1013    },
1014    Binary {
1015        op: BinaryOp,
1016        lhs: Expr,
1017        rhs: Expr,
1018    },
1019    TensorScalar {
1020        op: BinaryOp,
1021        tensor: Expr,
1022        scalar: Expr,
1023        scalar_on_left: bool,
1024    },
1025    Complex {
1026        re: Expr,
1027        im: Expr,
1028    },
1029    Vector {
1030        elements: Vec<Expr>,
1031    },
1032    Matrix {
1033        rows: usize,
1034        cols: usize,
1035        elements: Vec<Expr>,
1036    },
1037    Component {
1038        input: Expr,
1039        index: usize,
1040    },
1041    MatrixElement {
1042        input: Expr,
1043        row: usize,
1044        col: usize,
1045    },
1046    MatMul {
1047        lhs: Expr,
1048        rhs: Expr,
1049    },
1050    MatVec {
1051        matrix: Expr,
1052        vector: Expr,
1053    },
1054    Dot {
1055        lhs: Expr,
1056        rhs: Expr,
1057    },
1058    Solve {
1059        matrix: Expr,
1060        rhs: Expr,
1061    },
1062}
1063
1064impl DagNodeKind {
1065    fn child_count(&self) -> usize {
1066        match self {
1067            Self::RealConst(_)
1068            | Self::ComplexConst(_)
1069            | Self::ScalarParam(_)
1070            | Self::EventScalar(_)
1071            | Self::EventP4Component { .. } => 0,
1072            Self::Unary { .. } | Self::Component { .. } | Self::MatrixElement { .. } => 1,
1073            Self::Binary { .. }
1074            | Self::TensorScalar { .. }
1075            | Self::Complex { .. }
1076            | Self::MatMul { .. }
1077            | Self::MatVec { .. }
1078            | Self::Dot { .. }
1079            | Self::Solve { .. } => 2,
1080            Self::Vector { elements } | Self::Matrix { elements, .. } => elements.len(),
1081        }
1082    }
1083
1084    fn child_at(&self, index: usize) -> &Expr {
1085        match self {
1086            Self::Unary { input, .. }
1087            | Self::Component { input, .. }
1088            | Self::MatrixElement { input, .. } => input,
1089            Self::Binary { lhs, rhs, .. } | Self::MatMul { lhs, rhs } | Self::Dot { lhs, rhs } => {
1090                [lhs, rhs][index]
1091            }
1092            Self::TensorScalar { tensor, scalar, .. } => [tensor, scalar][index],
1093            Self::Complex { re, im } => [re, im][index],
1094            Self::MatVec { matrix, vector } => [matrix, vector][index],
1095            Self::Solve { matrix, rhs } => [matrix, rhs][index],
1096            Self::Vector { elements } | Self::Matrix { elements, .. } => &elements[index],
1097            Self::RealConst(_)
1098            | Self::ComplexConst(_)
1099            | Self::ScalarParam(_)
1100            | Self::EventScalar(_)
1101            | Self::EventP4Component { .. } => unreachable!("leaf nodes have no children"),
1102        }
1103    }
1104
1105    fn map_children(&self, mut map: impl FnMut(&Expr) -> Expr) -> Self {
1106        match self {
1107            Self::RealConst(value) => Self::RealConst(*value),
1108            Self::ComplexConst(value) => Self::ComplexConst(*value),
1109            Self::ScalarParam(parameter) => Self::ScalarParam(parameter.clone()),
1110            Self::EventScalar(name) => Self::EventScalar(Arc::clone(name)),
1111            Self::EventP4Component { name, component } => Self::EventP4Component {
1112                name: Arc::clone(name),
1113                component: *component,
1114            },
1115            Self::Unary { op, input } => Self::Unary {
1116                op: *op,
1117                input: map(input),
1118            },
1119            Self::Binary { op, lhs, rhs } => Self::Binary {
1120                op: *op,
1121                lhs: map(lhs),
1122                rhs: map(rhs),
1123            },
1124            Self::TensorScalar {
1125                op,
1126                tensor,
1127                scalar,
1128                scalar_on_left,
1129            } => Self::TensorScalar {
1130                op: *op,
1131                tensor: map(tensor),
1132                scalar: map(scalar),
1133                scalar_on_left: *scalar_on_left,
1134            },
1135            Self::Complex { re, im } => Self::Complex {
1136                re: map(re),
1137                im: map(im),
1138            },
1139            Self::Vector { elements } => Self::Vector {
1140                elements: elements.iter().map(&mut map).collect(),
1141            },
1142            Self::Matrix {
1143                rows,
1144                cols,
1145                elements,
1146            } => Self::Matrix {
1147                rows: *rows,
1148                cols: *cols,
1149                elements: elements.iter().map(&mut map).collect(),
1150            },
1151            Self::Component { input, index } => Self::Component {
1152                input: map(input),
1153                index: *index,
1154            },
1155            Self::MatrixElement { input, row, col } => Self::MatrixElement {
1156                input: map(input),
1157                row: *row,
1158                col: *col,
1159            },
1160            Self::MatMul { lhs, rhs } => Self::MatMul {
1161                lhs: map(lhs),
1162                rhs: map(rhs),
1163            },
1164            Self::MatVec { matrix, vector } => Self::MatVec {
1165                matrix: map(matrix),
1166                vector: map(vector),
1167            },
1168            Self::Dot { lhs, rhs } => Self::Dot {
1169                lhs: map(lhs),
1170                rhs: map(rhs),
1171            },
1172            Self::Solve { matrix, rhs } => Self::Solve {
1173                matrix: map(matrix),
1174                rhs: map(rhs),
1175            },
1176        }
1177    }
1178}
1179
1180impl Expr {
1181    fn new(kind: DagNodeKind) -> Self {
1182        let source = source_kind(&kind);
1183        Self {
1184            node: Arc::new(DagNode {
1185                kind,
1186                metadata: ExprMetadata::new(source),
1187                shape: OnceLock::new(),
1188            }),
1189        }
1190    }
1191
1192    /// Assigns a display name to the expression root.
1193    pub fn named(self, name: impl Into<Arc<str>>) -> Self {
1194        self.with_metadata(|metadata| metadata.name = Some(name.into()))
1195    }
1196
1197    /// Adds a tag to the expression root.
1198    pub fn tagged(self, tag: impl Into<Arc<str>>) -> Self {
1199        let tag = tag.into();
1200        self.with_metadata(|metadata| {
1201            if !metadata.tags.iter().any(|existing| existing == &tag) {
1202                metadata.tags.push(tag);
1203            }
1204        })
1205    }
1206
1207    /// Adds each supplied tag to the expression root.
1208    pub fn tagged_with(self, tags: impl IntoIterator<Item = impl Into<Arc<str>>>) -> Self {
1209        tags.into_iter().fold(self, Self::tagged)
1210    }
1211
1212    /// Replace tagged components that do not match any requested tag with zero.
1213    ///
1214    /// Untagged nodes remain active, while a matching tagged node retains its complete subtree.
1215    pub fn project_tags<'a>(&self, tags: impl IntoIterator<Item = &'a str>) -> Self {
1216        let tags: Vec<_> = tags.into_iter().collect();
1217        enum Frame {
1218            Visit(Expr),
1219            Rebuild(Expr, usize),
1220        }
1221
1222        let mut projected = Vec::new();
1223        let mut stack = vec![Frame::Visit(self.clone())];
1224        while let Some(frame) = stack.pop() {
1225            match frame {
1226                Frame::Visit(expr) => {
1227                    let has_tags = !expr.node.metadata.tags.is_empty();
1228                    if has_tags || expr.node.kind.child_count() == 0 {
1229                        projected.push(
1230                            if !has_tags
1231                                || expr
1232                                    .node
1233                                    .metadata
1234                                    .tags
1235                                    .iter()
1236                                    .any(|candidate| tags.contains(&candidate.as_ref()))
1237                            {
1238                                expr.clone()
1239                            } else {
1240                                expr.zero_like()
1241                            },
1242                        );
1243                        continue;
1244                    }
1245                    let child_count = expr.node.kind.child_count();
1246                    stack.push(Frame::Rebuild(expr.clone(), child_count));
1247                    for index in (0..child_count).rev() {
1248                        stack.push(Frame::Visit(expr.node.kind.child_at(index).clone()));
1249                    }
1250                }
1251                Frame::Rebuild(expr, child_count) => {
1252                    let child_start = projected.len() - child_count;
1253                    let children = projected.split_off(child_start);
1254                    let mut children = children.into_iter();
1255                    let kind = expr.node.kind.map_children(|original| {
1256                        children.next().unwrap_or_else(|| original.clone())
1257                    });
1258                    projected.push(
1259                        Expr::new(kind)
1260                            .with_metadata(|metadata| *metadata = expr.node.metadata.clone()),
1261                    );
1262                }
1263            }
1264        }
1265        projected.pop().unwrap_or_else(|| self.clone())
1266    }
1267
1268    fn zero_like(&self) -> Self {
1269        match self
1270            .shape()
1271            .expect("valid expression shapes are cached eagerly")
1272        {
1273            ExprShape::Scalar => Expr::from(0.0),
1274            ExprShape::Vector { len } => vector((0..len).map(|_| Expr::from(0.0))),
1275            ExprShape::Matrix { rows, cols } => {
1276                matrix_from_flat(rows, cols, (0..rows * cols).map(|_| Expr::from(0.0)))
1277                    .expect("zero matrix dimensions match")
1278            }
1279        }
1280    }
1281
1282    /// Returns an expression for the real part.
1283    pub fn real(&self) -> Self {
1284        unary(UnaryOp::Real, self)
1285    }
1286
1287    /// Returns an expression for the imaginary part.
1288    pub fn imag(&self) -> Self {
1289        unary(UnaryOp::Imag, self)
1290    }
1291
1292    /// Returns the complex conjugate, applied elementwise to vectors and matrices.
1293    ///
1294    /// # Panics
1295    /// Panics if an internally validated matrix cannot be rebuilt with its cached dimensions.
1296    pub fn conj(&self) -> Self {
1297        match self.shape() {
1298            Ok(ExprShape::Vector { len }) => {
1299                vector((0..len).map(|index| self.component(index).conj()))
1300            }
1301            Ok(ExprShape::Matrix { rows, cols }) => matrix_from_flat(
1302                rows,
1303                cols,
1304                (0..rows)
1305                    .flat_map(|row| (0..cols).map(move |col| self.matrix_element(row, col).conj())),
1306            )
1307            .expect("conjugating a valid matrix preserves its dimensions"),
1308            _ => unary(UnaryOp::Conj, self),
1309        }
1310    }
1311
1312    /// Returns an expression for the squared complex norm.
1313    pub fn norm_sqr(&self) -> Self {
1314        unary(UnaryOp::NormSqr, self)
1315    }
1316
1317    /// Returns an expression for the principal square root.
1318    pub fn sqrt(&self) -> Self {
1319        unary(UnaryOp::Sqrt, self)
1320    }
1321
1322    /// Returns an expression for the exponential.
1323    pub fn exp(&self) -> Self {
1324        unary(UnaryOp::Exp, self)
1325    }
1326
1327    /// Returns an expression for the sine.
1328    pub fn sin(&self) -> Self {
1329        unary(UnaryOp::Sin, self)
1330    }
1331
1332    /// Returns an expression for the cosine.
1333    pub fn cos(&self) -> Self {
1334        unary(UnaryOp::Cos, self)
1335    }
1336
1337    /// Returns an expression for the principal arccosine.
1338    pub fn acos(&self) -> Self {
1339        atan2((Expr::from(1.0) - self.powi(2)).sqrt(), self)
1340    }
1341
1342    /// Returns an expression for the natural logarithm.
1343    pub fn log(&self) -> Self {
1344        unary(UnaryOp::Log, self)
1345    }
1346
1347    /// Returns an expression raised to an integer power.
1348    pub fn powi(&self, power: i32) -> Self {
1349        unary(UnaryOp::PowI(power), self)
1350    }
1351
1352    /// Selects a component from a vector-valued expression.
1353    pub fn component(&self, index: impl ComponentIndex) -> Self {
1354        Expr::new(DagNodeKind::Component {
1355            input: self.clone(),
1356            index: index.component_index(),
1357        })
1358    }
1359
1360    /// Selects an element from a matrix-valued expression.
1361    pub fn matrix_element(&self, row: usize, col: usize) -> Self {
1362        Expr::new(DagNodeKind::MatrixElement {
1363            input: self.clone(),
1364            row,
1365            col,
1366        })
1367    }
1368
1369    /// Serializes the shareable expression DAG into a topologically ordered graph.
1370    pub fn to_graph(&self) -> ExprGraph {
1371        if !self.contains_tensor_scalar() {
1372            return GraphBuilder::new().build(self);
1373        }
1374        let (root, mut nodes) = self.lower_tensor_scalars();
1375        let graph = GraphBuilder::new().build(&root);
1376        drop(root);
1377        // Parents are released before children so a deep scalar subtree is
1378        // not recursively destroyed when one tensor operation is lowered.
1379        while let Some(node) = nodes.pop() {
1380            drop(node);
1381        }
1382        graph
1383    }
1384
1385    fn contains_tensor_scalar(&self) -> bool {
1386        let mut seen = HashSet::new();
1387        let mut stack = vec![self];
1388        while let Some(expr) = stack.pop() {
1389            if !seen.insert(Arc::as_ptr(&expr.node) as usize) {
1390                continue;
1391            }
1392            if matches!(expr.node.kind, DagNodeKind::TensorScalar { .. }) {
1393                return true;
1394            }
1395            for index in 0..expr.node.kind.child_count() {
1396                stack.push(expr.node.kind.child_at(index));
1397            }
1398        }
1399        false
1400    }
1401
1402    fn lower_tensor_scalars(&self) -> (Self, Vec<Self>) {
1403        let mut lowered = HashMap::<usize, Expr>::new();
1404        let mut order = Vec::new();
1405        let mut stack = vec![(self.clone(), false)];
1406        while let Some((expr, visited)) = stack.pop() {
1407            let key = Arc::as_ptr(&expr.node) as usize;
1408            if lowered.contains_key(&key) {
1409                continue;
1410            }
1411            if visited {
1412                let kind = expr
1413                    .node
1414                    .kind
1415                    .map_children(|child| lowered[&(Arc::as_ptr(&child.node) as usize)].clone());
1416                let result = match kind {
1417                    DagNodeKind::TensorScalar {
1418                        op,
1419                        tensor,
1420                        scalar,
1421                        scalar_on_left,
1422                    } => lower_tensor_scalar(op, tensor, scalar, scalar_on_left),
1423                    other => Expr::new(other),
1424                }
1425                .with_metadata(|metadata| *metadata = expr.node.metadata.clone());
1426                order.push(result.clone());
1427                lowered.insert(key, result);
1428            } else {
1429                stack.push((expr.clone(), true));
1430                for index in (0..expr.node.kind.child_count()).rev() {
1431                    stack.push((expr.node.kind.child_at(index).clone(), false));
1432                }
1433            }
1434        }
1435        (lowered[&(Arc::as_ptr(&self.node) as usize)].clone(), order)
1436    }
1437
1438    /// Rebuilds a shareable expression DAG from its serialized graph form.
1439    ///
1440    /// # Errors
1441    ///
1442    /// Returns [`ExprGraphError`] when the graph is empty, its root or a child
1443    /// identifier is invalid, its metadata length does not match its node
1444    /// count, or its nodes are not topologically ordered.
1445    pub fn from_graph(graph: ExprGraph) -> Result<Self, ExprGraphError> {
1446        let ExprGraph {
1447            root,
1448            nodes,
1449            metadata,
1450        } = graph;
1451        let graph = ExprGraph::from_parts(root, nodes, metadata)?;
1452        let mut expressions: Vec<Expr> = Vec::with_capacity(graph.nodes.len());
1453        for (index, node) in graph.nodes.iter().enumerate() {
1454            let child = |id: ExprId| expressions[id.index()].clone();
1455            let expression = match node {
1456                ExprNode::RealConst(value) => Expr::new(DagNodeKind::RealConst(*value)),
1457                ExprNode::ComplexConst(value) => Expr::new(DagNodeKind::ComplexConst(*value)),
1458                ExprNode::ScalarParam(parameter) => {
1459                    Expr::new(DagNodeKind::ScalarParam(parameter.clone()))
1460                }
1461                ExprNode::EventScalar(name) => {
1462                    Expr::new(DagNodeKind::EventScalar(Arc::clone(name)))
1463                }
1464                ExprNode::EventP4Component { name, component } => {
1465                    Expr::new(DagNodeKind::EventP4Component {
1466                        name: Arc::clone(name),
1467                        component: *component,
1468                    })
1469                }
1470                ExprNode::Unary { op, input } => Expr::new(DagNodeKind::Unary {
1471                    op: *op,
1472                    input: child(*input),
1473                }),
1474                ExprNode::Binary { op, lhs, rhs } => Expr::new(DagNodeKind::Binary {
1475                    op: *op,
1476                    lhs: child(*lhs),
1477                    rhs: child(*rhs),
1478                }),
1479                ExprNode::NaryAdd { terms } => terms
1480                    .iter()
1481                    .map(|id| child(*id))
1482                    .reduce(|lhs, rhs| binary(BinaryOp::Add, &lhs, &rhs))
1483                    .unwrap_or_else(|| Expr::from(0.0)),
1484                ExprNode::NaryMul { factors } => factors
1485                    .iter()
1486                    .map(|id| child(*id))
1487                    .reduce(|lhs, rhs| binary(BinaryOp::Mul, &lhs, &rhs))
1488                    .unwrap_or_else(|| Expr::from(1.0)),
1489                ExprNode::Complex { re, im } => Expr::new(DagNodeKind::Complex {
1490                    re: child(*re),
1491                    im: child(*im),
1492                }),
1493                ExprNode::Vector { elements } => Expr::new(DagNodeKind::Vector {
1494                    elements: elements.iter().map(|id| child(*id)).collect(),
1495                }),
1496                ExprNode::Matrix {
1497                    rows,
1498                    cols,
1499                    elements,
1500                } => Expr::new(DagNodeKind::Matrix {
1501                    rows: *rows,
1502                    cols: *cols,
1503                    elements: elements.iter().map(|id| child(*id)).collect(),
1504                }),
1505                ExprNode::Component { input, index } => Expr::new(DagNodeKind::Component {
1506                    input: child(*input),
1507                    index: *index,
1508                }),
1509                ExprNode::MatrixElement { input, row, col } => {
1510                    Expr::new(DagNodeKind::MatrixElement {
1511                        input: child(*input),
1512                        row: *row,
1513                        col: *col,
1514                    })
1515                }
1516                ExprNode::MatMul { lhs, rhs } => Expr::new(DagNodeKind::MatMul {
1517                    lhs: child(*lhs),
1518                    rhs: child(*rhs),
1519                }),
1520                ExprNode::MatVec { matrix, vector } => Expr::new(DagNodeKind::MatVec {
1521                    matrix: child(*matrix),
1522                    vector: child(*vector),
1523                }),
1524                ExprNode::Dot { lhs, rhs } => Expr::new(DagNodeKind::Dot {
1525                    lhs: child(*lhs),
1526                    rhs: child(*rhs),
1527                }),
1528                ExprNode::Solve { matrix, rhs } => Expr::new(DagNodeKind::Solve {
1529                    matrix: child(*matrix),
1530                    rhs: child(*rhs),
1531                }),
1532            };
1533            let mut dag = (*expression.node).clone();
1534            dag.metadata = graph.metadata[index].clone();
1535            expressions.push(Expr {
1536                node: Arc::new(dag),
1537            });
1538        }
1539        Ok(expressions[graph.root.index()].clone())
1540    }
1541
1542    /// Determines and validates the expression's structural shape.
1543    ///
1544    /// # Errors
1545    ///
1546    /// Returns [`ExprShapeError`] when this expression contains an operation
1547    /// whose operand shapes are incompatible.
1548    pub fn shape(&self) -> Result<ExprShape, ExprShapeError> {
1549        self.node
1550            .shape
1551            .get_or_init(|| self.node.kind.shape())
1552            .clone()
1553    }
1554
1555    fn with_metadata(self, f: impl FnOnce(&mut ExprMetadata)) -> Self {
1556        let mut node = (*self.node).clone();
1557        f(&mut node.metadata);
1558        Self {
1559            node: Arc::new(node),
1560        }
1561    }
1562}
1563
1564impl Serialize for Expr {
1565    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
1566    where
1567        S: serde::Serializer,
1568    {
1569        self.to_graph().serialize(serializer)
1570    }
1571}
1572
1573impl<'de> Deserialize<'de> for Expr {
1574    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
1575    where
1576        D: serde::Deserializer<'de>,
1577    {
1578        Expr::from_graph(ExprGraph::deserialize(deserializer)?).map_err(serde::de::Error::custom)
1579    }
1580}
1581
1582impl DagNodeKind {
1583    fn shape(&self) -> Result<ExprShape, ExprShapeError> {
1584        match self {
1585            Self::RealConst(_)
1586            | Self::ComplexConst(_)
1587            | Self::ScalarParam(_)
1588            | Self::EventScalar(_)
1589            | Self::EventP4Component { .. } => Ok(ExprShape::Scalar),
1590            Self::Unary { input, .. } => {
1591                input.expect_shape("unary operation", ExprShape::Scalar)?;
1592                Ok(ExprShape::Scalar)
1593            }
1594            Self::Binary { lhs, rhs, .. } => {
1595                lhs.expect_shape("binary operation", ExprShape::Scalar)?;
1596                rhs.expect_shape("binary operation", ExprShape::Scalar)?;
1597                Ok(ExprShape::Scalar)
1598            }
1599            Self::TensorScalar {
1600                op,
1601                tensor,
1602                scalar,
1603                scalar_on_left,
1604            } => {
1605                scalar.expect_shape("tensor-scalar operation", ExprShape::Scalar)?;
1606                if *scalar_on_left && *op == BinaryOp::Div {
1607                    return Err(ExprShapeError::new(
1608                        "tensor-scalar operation",
1609                        "scalar division by a tensor is unsupported",
1610                    ));
1611                }
1612                match tensor.shape()? {
1613                    shape @ (ExprShape::Vector { .. } | ExprShape::Matrix { .. }) => Ok(shape),
1614                    ExprShape::Scalar => Err(ExprShapeError::new(
1615                        "tensor-scalar operation",
1616                        "expected a vector or matrix operand",
1617                    )),
1618                }
1619            }
1620            Self::Complex { re, im } => {
1621                re.expect_shape("complex constructor", ExprShape::Scalar)?;
1622                im.expect_shape("complex constructor", ExprShape::Scalar)?;
1623                Ok(ExprShape::Scalar)
1624            }
1625            Self::Vector { elements } => {
1626                for element in elements {
1627                    element.expect_shape("vector constructor", ExprShape::Scalar)?;
1628                }
1629                Ok(ExprShape::Vector {
1630                    len: elements.len(),
1631                })
1632            }
1633            Self::Matrix {
1634                rows,
1635                cols,
1636                elements,
1637            } => {
1638                let expected = rows.checked_mul(*cols).ok_or_else(|| {
1639                    ExprShapeError::new("matrix constructor", "row/column product overflowed")
1640                })?;
1641                if elements.len() != expected {
1642                    return Err(ExprShapeError::new(
1643                        "matrix constructor",
1644                        format!(
1645                            "shape {rows}x{cols} requires {expected} elements, got {}",
1646                            elements.len()
1647                        ),
1648                    ));
1649                }
1650                for element in elements {
1651                    element.expect_shape("matrix constructor", ExprShape::Scalar)?;
1652                }
1653                Ok(ExprShape::Matrix {
1654                    rows: *rows,
1655                    cols: *cols,
1656                })
1657            }
1658            Self::Component { input, index } => {
1659                let ExprShape::Vector { len } = input.shape()? else {
1660                    return Err(ExprShapeError::new(
1661                        "component",
1662                        format!("expected vector, got {}", input.shape()?),
1663                    ));
1664                };
1665                if *index >= len {
1666                    return Err(ExprShapeError::new(
1667                        "component",
1668                        format!("index {index} is out of bounds for vector[{len}]"),
1669                    ));
1670                }
1671                Ok(ExprShape::Scalar)
1672            }
1673            Self::MatrixElement { input, row, col } => {
1674                let ExprShape::Matrix { rows, cols } = input.shape()? else {
1675                    return Err(ExprShapeError::new(
1676                        "matrix element",
1677                        format!("expected matrix, got {}", input.shape()?),
1678                    ));
1679                };
1680                if *row >= rows || *col >= cols {
1681                    return Err(ExprShapeError::new(
1682                        "matrix element",
1683                        format!("index ({row}, {col}) is out of bounds for matrix[{rows}x{cols}]"),
1684                    ));
1685                }
1686                Ok(ExprShape::Scalar)
1687            }
1688            Self::MatMul { lhs, rhs } => {
1689                let ExprShape::Matrix {
1690                    rows: lhs_rows,
1691                    cols: lhs_cols,
1692                } = lhs.shape()?
1693                else {
1694                    return Err(ExprShapeError::new(
1695                        "matrix multiplication",
1696                        format!("left input must be a matrix, got {}", lhs.shape()?),
1697                    ));
1698                };
1699                let ExprShape::Matrix {
1700                    rows: rhs_rows,
1701                    cols: rhs_cols,
1702                } = rhs.shape()?
1703                else {
1704                    return Err(ExprShapeError::new(
1705                        "matrix multiplication",
1706                        format!("right input must be a matrix, got {}", rhs.shape()?),
1707                    ));
1708                };
1709                if lhs_cols != rhs_rows {
1710                    return Err(ExprShapeError::new(
1711                        "matrix multiplication",
1712                        format!("cannot multiply {lhs_rows}x{lhs_cols} by {rhs_rows}x{rhs_cols}"),
1713                    ));
1714                }
1715                Ok(ExprShape::Matrix {
1716                    rows: lhs_rows,
1717                    cols: rhs_cols,
1718                })
1719            }
1720            Self::MatVec { matrix, vector } => {
1721                let ExprShape::Matrix { rows, cols } = matrix.shape()? else {
1722                    return Err(ExprShapeError::new(
1723                        "matrix-vector multiplication",
1724                        format!("left input must be a matrix, got {}", matrix.shape()?),
1725                    ));
1726                };
1727                let ExprShape::Vector { len } = vector.shape()? else {
1728                    return Err(ExprShapeError::new(
1729                        "matrix-vector multiplication",
1730                        format!("right input must be a vector, got {}", vector.shape()?),
1731                    ));
1732                };
1733                if cols != len {
1734                    return Err(ExprShapeError::new(
1735                        "matrix-vector multiplication",
1736                        format!("cannot multiply {rows}x{cols} matrix by vector[{len}]"),
1737                    ));
1738                }
1739                Ok(ExprShape::Vector { len: rows })
1740            }
1741            Self::Dot { lhs, rhs } => {
1742                let ExprShape::Vector { len: lhs_len } = lhs.shape()? else {
1743                    return Err(ExprShapeError::new(
1744                        "dot product",
1745                        format!("left input must be a vector, got {}", lhs.shape()?),
1746                    ));
1747                };
1748                let ExprShape::Vector { len: rhs_len } = rhs.shape()? else {
1749                    return Err(ExprShapeError::new(
1750                        "dot product",
1751                        format!("right input must be a vector, got {}", rhs.shape()?),
1752                    ));
1753                };
1754                if lhs_len != rhs_len {
1755                    return Err(ExprShapeError::new(
1756                        "dot product",
1757                        format!("vector lengths differ: {lhs_len} and {rhs_len}"),
1758                    ));
1759                }
1760                Ok(ExprShape::Scalar)
1761            }
1762            Self::Solve { matrix, rhs } => {
1763                let ExprShape::Matrix { rows, cols } = matrix.shape()? else {
1764                    return Err(ExprShapeError::new(
1765                        "linear solve",
1766                        format!("left input must be a matrix, got {}", matrix.shape()?),
1767                    ));
1768                };
1769                let ExprShape::Vector { len } = rhs.shape()? else {
1770                    return Err(ExprShapeError::new(
1771                        "linear solve",
1772                        format!("right input must be a vector, got {}", rhs.shape()?),
1773                    ));
1774                };
1775                if rows != cols || rows != len {
1776                    return Err(ExprShapeError::new(
1777                        "linear solve",
1778                        format!("cannot solve matrix[{rows}x{cols}] against vector[{len}]"),
1779                    ));
1780                }
1781                Ok(ExprShape::Vector { len })
1782            }
1783        }
1784    }
1785}
1786
1787impl Expr {
1788    fn is_tensor_syntax(&self) -> bool {
1789        matches!(
1790            self.node.kind,
1791            DagNodeKind::Vector { .. }
1792                | DagNodeKind::Matrix { .. }
1793                | DagNodeKind::MatMul { .. }
1794                | DagNodeKind::MatVec { .. }
1795                | DagNodeKind::Solve { .. }
1796                | DagNodeKind::TensorScalar { .. }
1797        )
1798    }
1799
1800    fn expect_shape(
1801        &self,
1802        operation: &'static str,
1803        expected: ExprShape,
1804    ) -> Result<(), ExprShapeError> {
1805        let actual = self.shape()?;
1806        if actual != expected {
1807            return Err(ExprShapeError::new(
1808                operation,
1809                format!("expected {expected}, got {actual}"),
1810            ));
1811        }
1812        Ok(())
1813    }
1814}
1815
1816auto_ops::impl_op_ex!(+ |a: &Expr, b: &Expr| -> Expr { binary(BinaryOp::Add, a, b) });
1817
1818auto_ops::impl_op_ex!(+ |a: &Expr, b: &f64| -> Expr { binary(BinaryOp::Add, a, b) });
1819auto_ops::impl_op_ex!(+ |a: &f64, b: &Expr| -> Expr { binary(BinaryOp::Add, a, b) });
1820
1821auto_ops::impl_op_ex!(+ |a: &Expr, b: &Complex64| -> Expr { binary(BinaryOp::Add, a, b) });
1822auto_ops::impl_op_ex!(+ |a: &Complex64, b: &Expr| -> Expr { binary(BinaryOp::Add, a, b) });
1823
1824auto_ops::impl_op_ex!(+ |a: &Expr, b: &Parameter| -> Expr { binary(BinaryOp::Add, a, b) });
1825auto_ops::impl_op_ex!(+ |a: &Parameter, b: &Expr| -> Expr { binary(BinaryOp::Add, a, b) });
1826
1827auto_ops::impl_op_ex!(+ |a: &Parameter, b: &f64| -> Expr { binary(BinaryOp::Add, a, b) });
1828auto_ops::impl_op_ex!(+ |a: &f64, b: &Parameter| -> Expr { binary(BinaryOp::Add, a, b) });
1829
1830auto_ops::impl_op_ex!(+ |a: &Parameter, b: &Complex64| -> Expr { binary(BinaryOp::Add, a, b) });
1831auto_ops::impl_op_ex!(+ |a: &Complex64, b: &Parameter| -> Expr { binary(BinaryOp::Add, a, b) });
1832
1833auto_ops::impl_op_ex!(+ |a: &Parameter, b: &Parameter| -> Expr { binary(BinaryOp::Add, a, b) });
1834
1835auto_ops::impl_op_ex!(-|a: &Expr, b: &Expr| -> Expr { binary(BinaryOp::Sub, a, b) });
1836
1837auto_ops::impl_op_ex!(-|a: &Expr, b: &f64| -> Expr { binary(BinaryOp::Sub, a, b) });
1838
1839auto_ops::impl_op_ex!(-|a: &f64, b: &Expr| -> Expr { binary(BinaryOp::Sub, a, b) });
1840
1841auto_ops::impl_op_ex!(-|a: &Expr, b: &Complex64| -> Expr { binary(BinaryOp::Sub, a, b) });
1842
1843auto_ops::impl_op_ex!(-|a: &Complex64, b: &Expr| -> Expr { binary(BinaryOp::Sub, a, b) });
1844
1845auto_ops::impl_op_ex!(-|a: &Expr, b: &Parameter| -> Expr { binary(BinaryOp::Sub, a, b) });
1846
1847auto_ops::impl_op_ex!(-|a: &Parameter, b: &Expr| -> Expr { binary(BinaryOp::Sub, a, b) });
1848
1849auto_ops::impl_op_ex!(-|a: &f64, b: &Parameter| -> Expr { binary(BinaryOp::Sub, a, b) });
1850
1851auto_ops::impl_op_ex!(-|a: &Parameter, b: &f64| -> Expr { binary(BinaryOp::Sub, a, b) });
1852
1853auto_ops::impl_op_ex!(-|a: &Complex64, b: &Parameter| -> Expr { binary(BinaryOp::Sub, a, b) });
1854
1855auto_ops::impl_op_ex!(-|a: &Parameter, b: &Complex64| -> Expr { binary(BinaryOp::Sub, a, b) });
1856
1857auto_ops::impl_op_ex!(-|a: &Parameter, b: &Parameter| -> Expr { binary(BinaryOp::Sub, a, b) });
1858
1859auto_ops::impl_op_ex!(*|a: &Expr, b: &Expr| -> Expr { binary(BinaryOp::Mul, a, b) });
1860
1861auto_ops::impl_op_ex!(*|a: &Expr, b: &f64| -> Expr { binary(BinaryOp::Mul, a, b) });
1862auto_ops::impl_op_ex!(*|a: &f64, b: &Expr| -> Expr { binary(BinaryOp::Mul, a, b) });
1863
1864auto_ops::impl_op_ex!(*|a: &Expr, b: &Complex64| -> Expr { binary(BinaryOp::Mul, a, b) });
1865auto_ops::impl_op_ex!(*|a: &Complex64, b: &Expr| -> Expr { binary(BinaryOp::Mul, a, b) });
1866
1867auto_ops::impl_op_ex!(*|a: &Expr, b: &Parameter| -> Expr { binary(BinaryOp::Mul, a, b) });
1868auto_ops::impl_op_ex!(*|a: &Parameter, b: &Expr| -> Expr { binary(BinaryOp::Mul, a, b) });
1869
1870auto_ops::impl_op_ex!(*|a: &f64, b: &Parameter| -> Expr { binary(BinaryOp::Mul, a, b) });
1871auto_ops::impl_op_ex!(*|a: &Parameter, b: &f64| -> Expr { binary(BinaryOp::Mul, a, b) });
1872
1873auto_ops::impl_op_ex!(*|a: &Complex64, b: &Parameter| -> Expr { binary(BinaryOp::Mul, a, b) });
1874auto_ops::impl_op_ex!(*|a: &Parameter, b: &Complex64| -> Expr { binary(BinaryOp::Mul, a, b) });
1875
1876auto_ops::impl_op_ex!(*|a: &Parameter, b: &Parameter| -> Expr { binary(BinaryOp::Mul, a, b) });
1877
1878auto_ops::impl_op_ex!(/ |a: &Expr, b: &Expr| -> Expr {
1879    binary(BinaryOp::Div, a, b)
1880});
1881
1882auto_ops::impl_op_ex!(/ |a: &Expr, b: &Complex64| -> Expr { binary(BinaryOp::Div, a, b) });
1883auto_ops::impl_op_ex!(/ |a: &Complex64, b: &Expr| -> Expr { binary(BinaryOp::Div, a, b) });
1884
1885auto_ops::impl_op_ex!(/ |a: &Expr, b: &f64| -> Expr { binary(BinaryOp::Div, a, b) });
1886auto_ops::impl_op_ex!(/ |a: &f64, b: &Expr| -> Expr { binary(BinaryOp::Div, a, b) });
1887
1888auto_ops::impl_op_ex!(/|a: &Expr, b: &Parameter| -> Expr {
1889    binary(BinaryOp::Div, a, b)
1890});
1891auto_ops::impl_op_ex!(/|a: &Parameter, b: &Expr| -> Expr {
1892    binary(BinaryOp::Div, a, b)
1893});
1894
1895auto_ops::impl_op_ex!(/|a: &f64, b: &Parameter| -> Expr {
1896    binary(BinaryOp::Div, a, b)
1897});
1898auto_ops::impl_op_ex!(/|a: &Parameter, b: &f64| -> Expr {
1899    binary(BinaryOp::Div, a, b)
1900});
1901
1902auto_ops::impl_op_ex!(/|a: &Complex64, b: &Parameter| -> Expr {
1903    binary(BinaryOp::Div, a, b)
1904});
1905auto_ops::impl_op_ex!(/|a: &Parameter, b: &Complex64| -> Expr {
1906    binary(BinaryOp::Div, a, b)
1907});
1908
1909auto_ops::impl_op_ex!(/|a: &Parameter, b: &Parameter| -> Expr {
1910    binary(BinaryOp::Div, a, b)
1911});
1912
1913auto_ops::impl_op_ex!(-|a: &Expr| -> Expr { unary(UnaryOp::Neg, a) });
1914auto_ops::impl_op_ex!(-|a: &Parameter| -> Expr { unary(UnaryOp::Neg, a) });
1915
1916auto_ops::impl_op_ex!(+= |a: &mut Expr, b: &Expr| {
1917    *a = binary(BinaryOp::Add, &*a, b);
1918});
1919auto_ops::impl_op_ex!(+= |a: &mut Expr, b: &f64| {
1920    *a = binary(BinaryOp::Add, &*a, b);
1921});
1922auto_ops::impl_op_ex!(+= |a: &mut Expr, b: &Complex64| {
1923    *a = binary(BinaryOp::Add, &*a, b);
1924});
1925auto_ops::impl_op_ex!(+= |a: &mut Expr, b: &Parameter| {
1926    *a = binary(BinaryOp::Add, &*a, b);
1927});
1928
1929auto_ops::impl_op_ex!(-= |a: &mut Expr, b: &Expr| {
1930    *a = binary(BinaryOp::Sub, &*a, b);
1931});
1932auto_ops::impl_op_ex!(-= |a: &mut Expr, b: &f64| {
1933    *a = binary(BinaryOp::Sub, &*a, b);
1934});
1935auto_ops::impl_op_ex!(-= |a: &mut Expr, b: &Complex64| {
1936    *a = binary(BinaryOp::Sub, &*a, b);
1937});
1938auto_ops::impl_op_ex!(-= |a: &mut Expr, b: &Parameter| {
1939    *a = binary(BinaryOp::Sub, &*a, b);
1940});
1941
1942auto_ops::impl_op_ex!(*= |a: &mut Expr, b: &Expr| {
1943    *a = binary(BinaryOp::Mul, &*a, b);
1944});
1945auto_ops::impl_op_ex!(*= |a: &mut Expr, b: &f64| {
1946    *a = binary(BinaryOp::Mul, &*a, b);
1947});
1948auto_ops::impl_op_ex!(*= |a: &mut Expr, b: &Complex64| {
1949    *a = binary(BinaryOp::Mul, &*a, b);
1950});
1951auto_ops::impl_op_ex!(*= |a: &mut Expr, b: &Parameter| {
1952    *a = binary(BinaryOp::Mul, &*a, b);
1953});
1954
1955auto_ops::impl_op_ex!(/= |a: &mut Expr, b: &Expr| {
1956    *a = binary(BinaryOp::Div, &*a, b);
1957});
1958auto_ops::impl_op_ex!(/= |a: &mut Expr, b: &f64| {
1959    *a = binary(BinaryOp::Div, &*a, b);
1960});
1961auto_ops::impl_op_ex!(/= |a: &mut Expr, b: &Complex64| {
1962    *a = binary(BinaryOp::Div, &*a, b);
1963});
1964auto_ops::impl_op_ex!(/= |a: &mut Expr, b: &Parameter| {
1965    *a = binary(BinaryOp::Div, &*a, b);
1966});
1967
1968impl From<f64> for Expr {
1969    fn from(value: f64) -> Self {
1970        Self::new(DagNodeKind::RealConst(value))
1971    }
1972}
1973
1974impl From<&f64> for Expr {
1975    fn from(value: &f64) -> Self {
1976        Self::new(DagNodeKind::RealConst(*value))
1977    }
1978}
1979
1980impl From<Complex64> for Expr {
1981    fn from(value: Complex64) -> Self {
1982        Self::new(DagNodeKind::ComplexConst(value))
1983    }
1984}
1985
1986impl From<&Complex64> for Expr {
1987    fn from(value: &Complex64) -> Self {
1988        Self::new(DagNodeKind::ComplexConst(*value))
1989    }
1990}
1991
1992impl From<&Expr> for Expr {
1993    fn from(value: &Expr) -> Self {
1994        value.clone()
1995    }
1996}
1997
1998impl From<Parameter> for Expr {
1999    fn from(parameter: Parameter) -> Self {
2000        Expr::new(DagNodeKind::ScalarParam(parameter))
2001    }
2002}
2003
2004impl From<&Parameter> for Expr {
2005    fn from(parameter: &Parameter) -> Self {
2006        parameter.clone().into()
2007    }
2008}
2009
2010/// Constructs `cos(phase) + i sin(phase)`.
2011pub fn cis(phase: Expr) -> Expr {
2012    phase.cos() + Complex64::I * phase.sin()
2013}
2014
2015/// Constructs a complex scalar from real and imaginary expressions.
2016pub fn complex(re: impl Into<Expr>, im: impl Into<Expr>) -> Expr {
2017    Expr::new(DagNodeKind::Complex {
2018        re: re.into(),
2019        im: im.into(),
2020    })
2021}
2022
2023/// Constructs a complex scalar from magnitude and phase expressions.
2024pub fn polar_complex(mag: impl Into<Expr>, phase: impl Into<Expr>) -> Expr {
2025    mag.into() * (Complex64::I * phase.into()).exp()
2026}
2027
2028/// References a named scalar column in each event.
2029pub fn event_scalar(name: impl Into<Arc<str>>) -> Expr {
2030    Expr::new(DagNodeKind::EventScalar(name.into()))
2031}
2032
2033/// References one component of a named event four-momentum.
2034pub fn event_p4_component(name: impl Into<Arc<str>>, component: P4Component) -> Expr {
2035    Expr::new(DagNodeKind::EventP4Component {
2036        name: name.into(),
2037        component,
2038    })
2039}
2040
2041/// Constructs the two-argument arctangent `atan2(y, x)`.
2042pub fn atan2(y: impl Into<Expr>, x: impl Into<Expr>) -> Expr {
2043    binary(BinaryOp::Atan2, y, x)
2044}
2045
2046/// Constructs the principal arccosine of `value`.
2047pub fn acos(value: impl Into<Expr>) -> Expr {
2048    value.into().acos()
2049}
2050
2051/// Constructs a vector expression from scalar elements.
2052pub fn vector<E>(elements: impl IntoIterator<Item = E>) -> Expr
2053where
2054    E: Into<Expr>,
2055    Expr: From<E>,
2056{
2057    Expr::new(DagNodeKind::Vector {
2058        elements: elements.into_iter().map(Expr::from).collect(),
2059    })
2060}
2061
2062/// Constructs a row-major matrix expression from a nested array.
2063pub fn matrix<const R: usize, const C: usize, E>(elements: [[E; C]; R]) -> Expr
2064where
2065    E: Into<Expr>,
2066    Expr: From<E>,
2067{
2068    Expr::new(DagNodeKind::Matrix {
2069        rows: R,
2070        cols: C,
2071        elements: elements.into_iter().flatten().map(Expr::from).collect(),
2072    })
2073}
2074
2075/// Constructs a row-major matrix from a flat sequence.
2076///
2077/// # Errors
2078///
2079/// Returns [`ExprShapeError`] when either dimension is zero, the dimension
2080/// product overflows, the element count differs from `rows * cols`, or an
2081/// element is not scalar-valued.
2082pub fn matrix_from_flat<E>(
2083    rows: usize,
2084    cols: usize,
2085    elements: impl IntoIterator<Item = E>,
2086) -> Result<Expr, ExprShapeError>
2087where
2088    E: Into<Expr>,
2089    Expr: From<E>,
2090{
2091    if rows == 0 || cols == 0 {
2092        return Err(ExprShapeError::new(
2093            "matrix constructor",
2094            format!("matrix dimensions must be nonzero, got {rows}x{cols}"),
2095        ));
2096    }
2097    let expected = rows.checked_mul(cols).ok_or_else(|| {
2098        ExprShapeError::new("matrix constructor", "row/column product overflowed")
2099    })?;
2100    let elements = elements.into_iter().map(Expr::from).collect::<Vec<_>>();
2101    if elements.len() != expected {
2102        return Err(ExprShapeError::new(
2103            "matrix constructor",
2104            format!(
2105                "shape {rows}x{cols} requires {expected} elements, got {}",
2106                elements.len()
2107            ),
2108        ));
2109    }
2110    for element in &elements {
2111        element.expect_shape("matrix constructor", ExprShape::Scalar)?;
2112    }
2113    Ok(Expr::new(DagNodeKind::Matrix {
2114        rows,
2115        cols,
2116        elements,
2117    }))
2118}
2119
2120/// Constructs a matrix-matrix multiplication expression.
2121pub fn matmul(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Expr {
2122    Expr::new(DagNodeKind::MatMul {
2123        lhs: lhs.into(),
2124        rhs: rhs.into(),
2125    })
2126}
2127
2128/// Constructs a matrix-vector multiplication expression.
2129pub fn matvec(matrix: impl Into<Expr>, vector: impl Into<Expr>) -> Expr {
2130    Expr::new(DagNodeKind::MatVec {
2131        matrix: matrix.into(),
2132        vector: vector.into(),
2133    })
2134}
2135
2136/// Constructs a vector dot-product expression.
2137pub fn dot(lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Expr {
2138    Expr::new(DagNodeKind::Dot {
2139        lhs: lhs.into(),
2140        rhs: rhs.into(),
2141    })
2142}
2143
2144/// Constructs an expression that solves a linear system.
2145pub fn solve(matrix: impl Into<Expr>, rhs: impl Into<Expr>) -> Expr {
2146    Expr::new(DagNodeKind::Solve {
2147        matrix: matrix.into(),
2148        rhs: rhs.into(),
2149    })
2150}
2151
2152fn unary(op: UnaryOp, expr: impl Into<Expr>) -> Expr {
2153    Expr::new(DagNodeKind::Unary {
2154        op,
2155        input: expr.into(),
2156    })
2157}
2158
2159fn binary(op: BinaryOp, lhs: impl Into<Expr>, rhs: impl Into<Expr>) -> Expr {
2160    let lhs = lhs.into();
2161    let rhs = rhs.into();
2162    match (lhs.is_tensor_syntax(), rhs.is_tensor_syntax()) {
2163        (true, false) if op != BinaryOp::Atan2 => Expr::new(DagNodeKind::TensorScalar {
2164            op,
2165            tensor: lhs,
2166            scalar: rhs,
2167            scalar_on_left: false,
2168        }),
2169        (false, true) if matches!(op, BinaryOp::Add | BinaryOp::Sub | BinaryOp::Mul) => {
2170            Expr::new(DagNodeKind::TensorScalar {
2171                op,
2172                tensor: rhs,
2173                scalar: lhs,
2174                scalar_on_left: true,
2175            })
2176        }
2177        _ => Expr::new(DagNodeKind::Binary { op, lhs, rhs }),
2178    }
2179}
2180
2181fn lower_tensor_scalar(op: BinaryOp, tensor: Expr, scalar: Expr, scalar_on_left: bool) -> Expr {
2182    let element = |value: Expr| {
2183        if scalar_on_left {
2184            binary(op, scalar.clone(), value)
2185        } else {
2186            binary(op, value, scalar.clone())
2187        }
2188    };
2189    let shape = match &tensor.node.kind {
2190        DagNodeKind::Vector { elements } => ExprShape::Vector {
2191            len: elements.len(),
2192        },
2193        DagNodeKind::Matrix { rows, cols, .. } => ExprShape::Matrix {
2194            rows: *rows,
2195            cols: *cols,
2196        },
2197        _ => tensor.shape().expect("tensor-scalar shape was validated"),
2198    };
2199    match shape {
2200        ExprShape::Vector { len } => {
2201            let elements = (0..len).map(|index| {
2202                let value = match &tensor.node.kind {
2203                    DagNodeKind::Vector { elements } => elements[index].clone(),
2204                    _ => tensor.component(index),
2205                };
2206                element(value)
2207            });
2208            vector(elements)
2209        }
2210        ExprShape::Matrix { rows, cols } => {
2211            let elements = (0..rows).flat_map(|row| {
2212                (0..cols).map({
2213                    let element = &element;
2214                    let tensor = &tensor;
2215                    move |col| {
2216                        let value = match &tensor.node.kind {
2217                            DagNodeKind::Matrix { elements, .. } => {
2218                                elements[row * cols + col].clone()
2219                            }
2220                            _ => tensor.matrix_element(row, col),
2221                        };
2222                        element(value)
2223                    }
2224                })
2225            });
2226            matrix_from_flat(rows, cols, elements).expect("tensor-scalar matrix shape is valid")
2227        }
2228        ExprShape::Scalar => unreachable!("tensor-scalar input is a tensor"),
2229    }
2230}
2231
2232/// Topologically ordered, serializable representation of an [`Expr`] DAG.
2233#[derive(Clone, Debug, Serialize, Deserialize)]
2234pub struct ExprGraph {
2235    root: ExprId,
2236    nodes: Vec<ExprNode>,
2237    metadata: Vec<ExprMetadata>,
2238}
2239
2240/// Workspace-internal accumulator for rebuilding expression graphs.
2241///
2242/// Nodes and metadata are emitted together in child-before-parent order. The
2243/// caller owns rewrite policy and chooses the remap key, which may include
2244/// traversal context in addition to the source node identifier.
2245#[doc(hidden)]
2246pub struct ExprGraphRebuilder<K> {
2247    nodes: Vec<ExprNode>,
2248    metadata: Vec<ExprMetadata>,
2249    remapped: HashMap<K, ExprId>,
2250}
2251
2252#[doc(hidden)]
2253impl<K> ExprGraphRebuilder<K>
2254where
2255    K: Eq + Hash,
2256{
2257    /// Creates an empty rebuild accumulator sized for the expected output.
2258    pub fn with_capacity(capacity: usize) -> Self {
2259        Self {
2260            nodes: Vec::with_capacity(capacity),
2261            metadata: Vec::with_capacity(capacity),
2262            remapped: HashMap::with_capacity(capacity),
2263        }
2264    }
2265
2266    /// Returns the emitted identifier associated with `key`, if any.
2267    pub fn remapped(&self, key: &K) -> Option<ExprId> {
2268        self.remapped.get(key).copied()
2269    }
2270
2271    /// Returns the nodes emitted so far in child-before-parent order.
2272    pub fn nodes(&self) -> &[ExprNode] {
2273        &self.nodes
2274    }
2275
2276    /// Returns the metadata emitted so far, aligned with [`Self::nodes`].
2277    pub fn metadata(&self) -> &[ExprMetadata] {
2278        &self.metadata
2279    }
2280
2281    /// Associates `key` with an already emitted node.
2282    ///
2283    /// This supports graph transforms that remove a source node by aliasing it
2284    /// to an existing result.
2285    ///
2286    /// # Panics
2287    ///
2288    /// Panics if `key` was mapped previously or `id` has not been emitted.
2289    pub fn alias(&mut self, key: K, id: ExprId) {
2290        assert!(
2291            !self.remapped.contains_key(&key),
2292            "a rebuild key may only be mapped once"
2293        );
2294        assert!(
2295            id.index() < self.nodes.len(),
2296            "a rebuild alias must reference an emitted node"
2297        );
2298        self.remapped.insert(key, id);
2299    }
2300
2301    /// Emits one node and its aligned metadata without assigning a remap key.
2302    ///
2303    /// This supports replacement fragments whose intermediate nodes do not
2304    /// correspond one-to-one with source nodes.
2305    ///
2306    /// # Panics
2307    ///
2308    /// Panics if the node references a child that has not already been emitted.
2309    pub fn emit_anonymous(&mut self, node: ExprNode, metadata: ExprMetadata) -> ExprId {
2310        let id = ExprId::from_index(self.nodes.len());
2311        assert!(
2312            node.children().all(|child| child.index() < id.index()),
2313            "rebuilt expression children must be emitted before their parent"
2314        );
2315        self.nodes.push(node);
2316        self.metadata.push(metadata);
2317        id
2318    }
2319
2320    /// Emits one node and its aligned metadata after all of its children.
2321    ///
2322    /// # Panics
2323    ///
2324    /// Panics if `key` was mapped previously or if the node references a
2325    /// child that has not already been emitted.
2326    pub fn emit(&mut self, key: K, node: ExprNode, metadata: ExprMetadata) -> ExprId {
2327        assert!(
2328            !self.remapped.contains_key(&key),
2329            "a rebuild key may only be mapped once"
2330        );
2331        let id = self.emit_anonymous(node, metadata);
2332        self.remapped.insert(key, id);
2333        id
2334    }
2335
2336    /// Validates and finishes the rebuilt graph with `root` as its root node.
2337    pub fn finish(self, root: ExprId) -> Result<ExprGraph, ExprGraphError> {
2338        ExprGraph::from_parts(root, self.nodes, self.metadata)
2339    }
2340}
2341
2342impl ExprGraph {
2343    /// Returns the nodes reachable from `roots` in child-before-parent order.
2344    ///
2345    /// The traversal is iterative, visits each node at most once, and preserves
2346    /// semantic child order. This is a workspace-internal traversal primitive
2347    /// for graph consumers that must remain safe for deeply nested graphs.
2348    #[doc(hidden)]
2349    pub fn reachable_post_order(&self, roots: impl IntoIterator<Item = ExprId>) -> Vec<ExprId> {
2350        let roots = roots.into_iter().collect::<Vec<_>>();
2351        let mut visited = HashSet::with_capacity(self.nodes.len());
2352        let mut stack = Vec::new();
2353        let mut order = Vec::new();
2354
2355        for root in roots.into_iter().rev() {
2356            stack.push((root, false));
2357        }
2358        while let Some((id, expanded)) = stack.pop() {
2359            if expanded {
2360                order.push(id);
2361                continue;
2362            }
2363            if self.node(id).is_none() {
2364                continue;
2365            }
2366            if !visited.insert(id) {
2367                continue;
2368            }
2369            stack.push((id, true));
2370            if let Some(node) = self.node(id) {
2371                for child in node.children().rev() {
2372                    stack.push((child, false));
2373                }
2374            }
2375        }
2376        order
2377    }
2378
2379    /// Return a copy of this graph with a batch of parameter updates applied.
2380    ///
2381    /// Updates are keyed by parameter name and every occurrence of a named
2382    /// scalar parameter is updated. The operation is atomic: duplicate or
2383    /// unknown names, invalid patches, conflicts, and invalid final
2384    /// definitions leave the original graph untouched.
2385    ///
2386    /// # Errors
2387    ///
2388    /// Returns [`ParamError`] when an update is duplicated, unknown, invalid,
2389    /// or leaves parameter definitions in conflict.
2390    pub fn with_parameters<N, I>(&self, updates: I) -> ParamResult<Self>
2391    where
2392        N: AsRef<str>,
2393        I: IntoIterator<Item = (N, ParameterUpdate)>,
2394    {
2395        let mut by_name = HashMap::new();
2396        let mut names = Vec::new();
2397        for (name, update) in updates {
2398            let name = name.as_ref().to_owned();
2399            update.validate()?;
2400            if by_name.insert(name.clone(), update).is_some() {
2401                return Err(ParamError::DuplicateName(name));
2402            }
2403            names.push(name);
2404        }
2405
2406        let mut found = HashSet::new();
2407        let mut graph = self.clone();
2408        for node in &mut graph.nodes {
2409            if let ExprNode::ScalarParam(parameter) = node
2410                && let Some(update) = by_name.get(parameter.name())
2411            {
2412                *parameter = parameter.with_update(update)?;
2413                found.insert(parameter.name().to_owned());
2414            }
2415        }
2416        if let Some(name) = names.iter().find(|name| !found.contains(*name)) {
2417            return Err(ParamError::UnknownName(name.clone()));
2418        }
2419
2420        let mut registry = crate::parameters::ParamRegistry::new();
2421        for node in &graph.nodes {
2422            if let ExprNode::ScalarParam(parameter) = node {
2423                registry.register(parameter.clone())?;
2424            }
2425        }
2426        registry.layout()?;
2427        Ok(graph)
2428    }
2429
2430    /// Replaces tagged components that match none of `tags` with zero.
2431    ///
2432    /// Untagged nodes remain active, and a matching tagged node retains its
2433    /// entire subtree.
2434    ///
2435    /// # Panics
2436    ///
2437    /// Panics only if an internal graph-rebuild invariant is violated.
2438    pub fn project_tags<'a>(&self, tags: impl IntoIterator<Item = &'a str>) -> Self {
2439        let tags: Vec<_> = tags.into_iter().collect();
2440        let mut rebuild = ExprGraphRebuilder::with_capacity(self.nodes.len());
2441        let root_key = (self.root, false);
2442        let mut visited = HashSet::with_capacity(self.nodes.len());
2443        let mut stack = vec![(root_key, false)];
2444        while let Some((key @ (old, retain_all), expanded)) = stack.pop() {
2445            if expanded {
2446                let old_metadata = &self.metadata[old.index()];
2447                let matches = old_metadata
2448                    .tags
2449                    .iter()
2450                    .any(|tag| tags.contains(&tag.as_ref()));
2451                let node = if !retain_all && !old_metadata.tags.is_empty() && !matches {
2452                    ExprNode::RealConst(0.0)
2453                } else {
2454                    let retain_children = retain_all || matches;
2455                    self.nodes[old.index()].map_children(|child| {
2456                        rebuild
2457                            .remapped(&(child, retain_children))
2458                            .expect("tag projection emits children before parents")
2459                    })
2460                };
2461                let metadata = if matches || retain_all {
2462                    old_metadata.clone()
2463                } else {
2464                    ExprMetadata::new(old_metadata.source)
2465                };
2466                rebuild.emit(key, node, metadata);
2467                continue;
2468            }
2469            if !visited.insert(key) {
2470                continue;
2471            }
2472            stack.push((key, true));
2473            let old_metadata = &self.metadata[old.index()];
2474            let matches = old_metadata
2475                .tags
2476                .iter()
2477                .any(|tag| tags.contains(&tag.as_ref()));
2478            if retain_all || old_metadata.tags.is_empty() || matches {
2479                let retain_children = retain_all || matches;
2480                for child in self.nodes[old.index()].children().rev() {
2481                    stack.push(((child, retain_children), false));
2482                }
2483            }
2484        }
2485        let root = rebuild
2486            .remapped(&root_key)
2487            .expect("tag projection emits its root");
2488        rebuild
2489            .finish(root)
2490            .expect("tag projection rebuilds a valid expression graph")
2491    }
2492
2493    /// Validates and constructs a graph from its serialized parts.
2494    ///
2495    /// Child nodes must precede their parents, metadata must have one entry per
2496    /// node, and `root` must identify an existing node.
2497    ///
2498    /// # Errors
2499    ///
2500    /// Returns [`ExprGraphError`] when the graph is empty, the metadata and
2501    /// node lengths differ, `root` is invalid, or a child identifier is
2502    /// invalid or does not precede its parent.
2503    pub fn from_parts(
2504        root: ExprId,
2505        nodes: Vec<ExprNode>,
2506        metadata: Vec<ExprMetadata>,
2507    ) -> Result<Self, ExprGraphError> {
2508        if nodes.is_empty() {
2509            return Err(ExprGraphError::Empty);
2510        }
2511        if nodes.len() != metadata.len() {
2512            return Err(ExprGraphError::MetadataLength {
2513                node_len: nodes.len(),
2514                metadata_len: metadata.len(),
2515            });
2516        }
2517        if root.index() >= nodes.len() {
2518            return Err(ExprGraphError::InvalidRoot {
2519                root: root.index(),
2520                node_len: nodes.len(),
2521            });
2522        }
2523        for (index, node) in nodes.iter().enumerate() {
2524            for child in node.children() {
2525                if child.index() >= nodes.len() {
2526                    return Err(ExprGraphError::InvalidChild {
2527                        node: index,
2528                        child: child.index(),
2529                    });
2530                }
2531                if child.index() >= index {
2532                    return Err(ExprGraphError::InvalidChildOrder {
2533                        node: index,
2534                        child: child.index(),
2535                    });
2536                }
2537            }
2538        }
2539        Ok(Self {
2540            root,
2541            nodes,
2542            metadata,
2543        })
2544    }
2545
2546    /// Returns the root node identifier.
2547    pub fn root(&self) -> ExprId {
2548        self.root
2549    }
2550
2551    /// Returns the node identified by `id`, if it exists.
2552    pub fn node(&self, id: ExprId) -> Option<&ExprNode> {
2553        self.nodes.get(id.index())
2554    }
2555
2556    /// Returns all nodes in topological order.
2557    pub fn nodes(&self) -> &[ExprNode] {
2558        &self.nodes
2559    }
2560
2561    /// Returns the metadata associated with `id`, if it exists.
2562    pub fn metadata(&self, id: ExprId) -> Option<&ExprMetadata> {
2563        self.metadata.get(id.index())
2564    }
2565}
2566
2567pub(crate) fn node_children(node: &ExprNode) -> Vec<(String, ExprId)> {
2568    node.children()
2569        .enumerate()
2570        .map(|(index, child)| (node_child_label(node, index), child))
2571        .collect()
2572}
2573
2574fn node_child_label(node: &ExprNode, index: usize) -> String {
2575    match node {
2576        ExprNode::Unary { .. } | ExprNode::Component { .. } | ExprNode::MatrixElement { .. } => {
2577            "input".into()
2578        }
2579        ExprNode::Binary { .. } | ExprNode::MatMul { .. } | ExprNode::Dot { .. } => {
2580            if index == 0 { "lhs" } else { "rhs" }.into()
2581        }
2582        ExprNode::NaryAdd { .. } => format!("term[{index}]"),
2583        ExprNode::NaryMul { .. } => format!("factor[{index}]"),
2584        ExprNode::Complex { .. } => if index == 0 { "re" } else { "im" }.into(),
2585        ExprNode::Vector { .. } => format!("element[{index}]"),
2586        ExprNode::Matrix { cols, .. } => {
2587            format!("element[{},{}]", index / cols, index % cols)
2588        }
2589        ExprNode::MatVec { .. } => if index == 0 { "matrix" } else { "vector" }.into(),
2590        ExprNode::Solve { .. } => if index == 0 { "matrix" } else { "rhs" }.into(),
2591        ExprNode::RealConst(_)
2592        | ExprNode::ComplexConst(_)
2593        | ExprNode::ScalarParam(_)
2594        | ExprNode::EventScalar(_)
2595        | ExprNode::EventP4Component { .. } => unreachable!("leaf nodes have no child labels"),
2596    }
2597}
2598
2599#[derive(Default)]
2600struct GraphBuilder {
2601    nodes: Vec<ExprNode>,
2602    metadata: Vec<ExprMetadata>,
2603    ids: HashMap<usize, ExprId>,
2604}
2605
2606impl GraphBuilder {
2607    fn new() -> Self {
2608        Self::default()
2609    }
2610
2611    fn build(mut self, expr: &Expr) -> ExprGraph {
2612        let mut stack = vec![(expr.clone(), false)];
2613        while let Some((expr, expanded)) = stack.pop() {
2614            let key = Arc::as_ptr(&expr.node) as usize;
2615            if self.ids.contains_key(&key) {
2616                continue;
2617            }
2618            if expanded {
2619                let node = self.lower(&expr.node.kind);
2620                let id = ExprId::from_index(self.nodes.len());
2621                self.nodes.push(node);
2622                self.metadata.push(expr.node.metadata.clone());
2623                self.ids.insert(key, id);
2624                continue;
2625            }
2626            stack.push((expr.clone(), true));
2627            for index in (0..expr.node.kind.child_count()).rev() {
2628                stack.push((expr.node.kind.child_at(index).clone(), false));
2629            }
2630        }
2631        let root = self.id(expr);
2632        ExprGraph {
2633            root,
2634            nodes: self.nodes,
2635            metadata: self.metadata,
2636        }
2637    }
2638
2639    fn id(&self, expr: &Expr) -> ExprId {
2640        let key = Arc::as_ptr(&expr.node) as usize;
2641        self.ids[&key]
2642    }
2643
2644    fn lower(&self, kind: &DagNodeKind) -> ExprNode {
2645        match kind {
2646            DagNodeKind::RealConst(value) => ExprNode::RealConst(*value),
2647            DagNodeKind::ComplexConst(value) => ExprNode::ComplexConst(*value),
2648            DagNodeKind::ScalarParam(parameter) => ExprNode::ScalarParam(parameter.clone()),
2649            DagNodeKind::EventScalar(name) => ExprNode::EventScalar(Arc::clone(name)),
2650            DagNodeKind::EventP4Component { name, component } => ExprNode::EventP4Component {
2651                name: Arc::clone(name),
2652                component: *component,
2653            },
2654            DagNodeKind::Unary { op, input } => {
2655                let input = self.id(input);
2656                ExprNode::Unary { op: *op, input }
2657            }
2658            DagNodeKind::Binary { op, lhs, rhs } => {
2659                let lhs = self.id(lhs);
2660                let rhs = self.id(rhs);
2661                ExprNode::Binary { op: *op, lhs, rhs }
2662            }
2663            DagNodeKind::TensorScalar { .. } => {
2664                unreachable!("tensor-scalar nodes are lowered before graph construction")
2665            }
2666            DagNodeKind::Complex { re, im } => {
2667                let re = self.id(re);
2668                let im = self.id(im);
2669                ExprNode::Complex { re, im }
2670            }
2671            DagNodeKind::Vector { elements } => ExprNode::Vector {
2672                elements: elements.iter().map(|expr| self.id(expr)).collect(),
2673            },
2674            DagNodeKind::Matrix {
2675                rows,
2676                cols,
2677                elements,
2678            } => ExprNode::Matrix {
2679                rows: *rows,
2680                cols: *cols,
2681                elements: elements.iter().map(|expr| self.id(expr)).collect(),
2682            },
2683            DagNodeKind::Component { input, index } => {
2684                let input = self.id(input);
2685                ExprNode::Component {
2686                    input,
2687                    index: *index,
2688                }
2689            }
2690            DagNodeKind::MatrixElement { input, row, col } => {
2691                let input = self.id(input);
2692                ExprNode::MatrixElement {
2693                    input,
2694                    row: *row,
2695                    col: *col,
2696                }
2697            }
2698            DagNodeKind::MatMul { lhs, rhs } => {
2699                let lhs = self.id(lhs);
2700                let rhs = self.id(rhs);
2701                ExprNode::MatMul { lhs, rhs }
2702            }
2703            DagNodeKind::MatVec { matrix, vector } => {
2704                let matrix = self.id(matrix);
2705                let vector = self.id(vector);
2706                ExprNode::MatVec { matrix, vector }
2707            }
2708            DagNodeKind::Dot { lhs, rhs } => {
2709                let lhs = self.id(lhs);
2710                let rhs = self.id(rhs);
2711                ExprNode::Dot { lhs, rhs }
2712            }
2713            DagNodeKind::Solve { matrix, rhs } => {
2714                let matrix = self.id(matrix);
2715                let rhs = self.id(rhs);
2716                ExprNode::Solve { matrix, rhs }
2717            }
2718        }
2719    }
2720}
2721
2722fn source_kind(kind: &DagNodeKind) -> ExprSourceKind {
2723    match kind {
2724        DagNodeKind::RealConst(_) | DagNodeKind::ComplexConst(_) => ExprSourceKind::Const,
2725        DagNodeKind::ScalarParam(_) => ExprSourceKind::Param,
2726        DagNodeKind::EventScalar(_) | DagNodeKind::EventP4Component { .. } => ExprSourceKind::Event,
2727        DagNodeKind::Unary { .. } => ExprSourceKind::Unary,
2728        DagNodeKind::Binary { .. } | DagNodeKind::TensorScalar { .. } => ExprSourceKind::Binary,
2729        DagNodeKind::Complex { .. } => ExprSourceKind::Complex,
2730        DagNodeKind::Vector { .. } | DagNodeKind::Component { .. } | DagNodeKind::Dot { .. } => {
2731            ExprSourceKind::Vector
2732        }
2733        DagNodeKind::Matrix { .. } | DagNodeKind::MatrixElement { .. } => ExprSourceKind::Matrix,
2734        DagNodeKind::MatMul { .. } | DagNodeKind::MatVec { .. } | DagNodeKind::Solve { .. } => {
2735            ExprSourceKind::LinearAlgebra
2736        }
2737    }
2738}
2739
2740#[cfg(test)]
2741mod tests {
2742    use super::*;
2743    use crate::parameter;
2744
2745    #[test]
2746    fn builds_target_syntax_without_layout_or_context() {
2747        let model = (Complex64::I * parameter!("y", initial : 1.0, bounds : (0.0, 2.0))
2748            + parameter!("x"))
2749        .norm_sqr();
2750
2751        let graph = model.to_graph();
2752        assert!(matches!(
2753            graph.node(graph.root()),
2754            Some(ExprNode::Unary {
2755                op: UnaryOp::NormSqr,
2756                ..
2757            })
2758        ));
2759    }
2760
2761    #[test]
2762    fn parameter_nodes_store_specs_but_do_not_make_layouts() {
2763        let graph = Expr::from(parameter!("x", initial: 1.0)).to_graph();
2764        assert!(matches!(
2765            graph.node(graph.root()),
2766            Some(ExprNode::ScalarParam(spec)) if spec.name() == "x"
2767        ));
2768    }
2769
2770    #[test]
2771    fn complex_constructor_builds_expression_node() {
2772        let graph = complex(parameter!("re"), parameter!("im")).to_graph();
2773
2774        assert!(matches!(
2775            graph.node(graph.root()),
2776            Some(ExprNode::Complex { .. })
2777        ));
2778    }
2779
2780    #[test]
2781    fn polar_complex_lowers_to_expression_graph() {
2782        let graph = polar_complex(parameter!("mag"), parameter!("phase")).to_graph();
2783
2784        assert!(graph.nodes().iter().any(|node| matches!(
2785            node,
2786            ExprNode::Unary {
2787                op: UnaryOp::Exp,
2788                ..
2789            }
2790        )));
2791    }
2792
2793    #[test]
2794    fn metadata_survives_graph_construction() {
2795        let graph = event_scalar("mass")
2796            .named("event mass")
2797            .tagged("data")
2798            .tagged("data")
2799            .to_graph();
2800        let metadata = graph.metadata(graph.root()).unwrap();
2801        assert_eq!(metadata.name(), Some("event mass"));
2802        assert_eq!(metadata.tags(), &[Arc::from("data")]);
2803        assert!(metadata.has_tag("data"));
2804    }
2805
2806    #[test]
2807    fn parameter_updates_rewrite_all_same_named_occurrences() {
2808        let expression = Expr::from(parameter!("x")) + Expr::from(parameter!("x"));
2809        let graph = expression
2810            .to_graph()
2811            .with_parameters([(
2812                String::from("x"),
2813                ParameterUpdate {
2814                    state: Some(ParamState::Fixed(2.0)),
2815                    ..Default::default()
2816                },
2817            )])
2818            .unwrap();
2819
2820        assert_eq!(
2821            graph
2822                .nodes()
2823                .iter()
2824                .filter_map(|node| match node {
2825                    ExprNode::ScalarParam(parameter) => Some(parameter),
2826                    _ => None,
2827                })
2828                .count(),
2829            2
2830        );
2831        assert!(graph.nodes().iter().all(|node| {
2832            !matches!(node, ExprNode::ScalarParam(parameter) if parameter.name() == "x" && parameter.is_free())
2833        }));
2834    }
2835
2836    #[test]
2837    fn parameter_updates_reject_duplicates_and_unknown_names_without_mutation() {
2838        let graph = Expr::from(parameter!("x")).to_graph();
2839        let original = graph.to_string();
2840        assert!(matches!(
2841            graph.with_parameters([
2842                ("x", ParameterUpdate::default()),
2843                ("x", ParameterUpdate::default()),
2844            ]),
2845            Err(ParamError::DuplicateName(name)) if name == "x"
2846        ));
2847        assert!(matches!(
2848            graph.with_parameters([("missing", ParameterUpdate::default())]),
2849            Err(ParamError::UnknownName(name)) if name == "missing"
2850        ));
2851        assert_eq!(graph.to_string(), original);
2852    }
2853
2854    #[test]
2855    fn expressions_round_trip_through_serde_with_metadata() {
2856        let expression = ((parameter!("x", initial: 1.0) + 2.0).named("offset")
2857            * event_scalar("mass").tagged("data"))
2858        .tagged("model");
2859        let encoded = serde_json::to_string(&expression).unwrap();
2860        let decoded: Expr = serde_json::from_str(&encoded).unwrap();
2861
2862        assert_eq!(
2863            serde_json::to_value(expression.to_graph()).unwrap(),
2864            serde_json::to_value(decoded.to_graph()).unwrap()
2865        );
2866    }
2867
2868    #[test]
2869    fn display_formats_graph_as_labeled_tree() {
2870        let graph = ((parameter!("x") + 1.0).named("offset") * event_scalar("mass").tagged("data"))
2871            .to_graph();
2872        let display = graph.display_tree().to_string();
2873
2874        assert!(display.starts_with("ExprGraph(root=#"));
2875        assert!(display.contains("Binary(Mul)"));
2876        assert!(display.contains("┣ lhs:"));
2877        assert!(display.contains("┗ rhs:"));
2878        assert!(display.contains("Binary(Add) name=\"offset\""));
2879        assert!(display.contains("ScalarParam(x)"));
2880        assert!(display.contains("RealConst(1)"));
2881        assert!(display.contains("EventScalar(mass) tags=[data]"));
2882    }
2883
2884    #[test]
2885    fn display_formats_graph_as_expression() {
2886        let costheta = Expr::from(parameter!("costheta"));
2887        let phi = event_scalar("phi");
2888        let p = Expr::from(parameter!("p"));
2889        let phase = Expr::from(7.0) * Complex64::I;
2890        let graph =
2891            (((costheta.powi(2) * phi.sin()) - 5.2).norm_sqr() * p.conj() - phase.exp()).to_graph();
2892
2893        assert_eq!(
2894            graph.to_string(),
2895            "|costheta^2 * sin(phi) - 5.2|^2 * conj(p) - exp(7 * i)"
2896        );
2897    }
2898
2899    #[test]
2900    fn display_parenthesizes_when_precedence_requires_it() {
2901        let a = Expr::from(parameter!("a"));
2902        let b = Expr::from(parameter!("b"));
2903        let c = Expr::from(parameter!("c"));
2904
2905        assert_eq!(
2906            (a.clone() * (b.clone() + c.clone())).to_graph().to_string(),
2907            "a * (b + c)"
2908        );
2909        assert_eq!(
2910            (a.clone() - (b.clone() - c.clone())).to_graph().to_string(),
2911            "a - (b - c)"
2912        );
2913        assert_eq!(((a / b) / c).to_graph().to_string(), "a / b / c");
2914    }
2915
2916    #[test]
2917    fn display_rounds_tiny_float_representation_noise() {
2918        let metadata = ExprMetadata::new(ExprSourceKind::Const);
2919        let graph = ExprGraph::from_parts(
2920            ExprId::from_index(2),
2921            vec![
2922                ExprNode::RealConst(2.9999999999999996),
2923                ExprNode::ComplexConst(Complex64::new(0.30000000000000004, 1.9999999999999998)),
2924                ExprNode::Binary {
2925                    op: BinaryOp::Add,
2926                    lhs: ExprId::from_index(0),
2927                    rhs: ExprId::from_index(1),
2928                },
2929            ],
2930            vec![metadata.clone(), metadata.clone(), metadata],
2931        )
2932        .unwrap();
2933
2934        assert_eq!(graph.to_string(), "3 + 0.3 + 2i");
2935        assert!(graph.display_tree().to_string().contains("RealConst(3)"));
2936        assert!(
2937            graph
2938                .display_tree()
2939                .to_string()
2940                .contains("ComplexConst(0.3 + 2i)")
2941        );
2942    }
2943
2944    #[test]
2945    fn display_formats_p4_components_and_atan2() {
2946        let expr = atan2(
2947            event_p4_component("ks1", P4Component::Py),
2948            event_p4_component("ks1", P4Component::Px),
2949        );
2950
2951        assert_eq!(expr.to_graph().to_string(), "atan2(ks1.py, ks1.px)");
2952    }
2953
2954    #[test]
2955    fn graph_from_parts_validates_structure() {
2956        let metadata = ExprMetadata::new(ExprSourceKind::Const);
2957        let graph = ExprGraph::from_parts(
2958            ExprId::from_index(1),
2959            vec![
2960                ExprNode::RealConst(1.0),
2961                ExprNode::Unary {
2962                    op: UnaryOp::Neg,
2963                    input: ExprId::from_index(0),
2964                },
2965            ],
2966            vec![metadata.clone(), metadata.clone()],
2967        )
2968        .unwrap();
2969        assert!(matches!(
2970            graph.node(graph.root()),
2971            Some(ExprNode::Unary {
2972                op: UnaryOp::Neg,
2973                ..
2974            })
2975        ));
2976
2977        let err = ExprGraph::from_parts(
2978            ExprId::from_index(0),
2979            vec![ExprNode::RealConst(1.0)],
2980            Vec::new(),
2981        )
2982        .unwrap_err();
2983        assert!(matches!(err, ExprGraphError::MetadataLength { .. }));
2984
2985        let err = ExprGraph::from_parts(
2986            ExprId::from_index(0),
2987            vec![ExprNode::Unary {
2988                op: UnaryOp::Neg,
2989                input: ExprId::from_index(0),
2990            }],
2991            vec![metadata],
2992        )
2993        .unwrap_err();
2994        assert!(matches!(err, ExprGraphError::InvalidChildOrder { .. }));
2995    }
2996
2997    #[test]
2998    fn graph_preserves_unsimplified_expression_shape() {
2999        let graph = (parameter!("x") + 0.0).to_graph();
3000        assert!(matches!(
3001            graph.node(graph.root()),
3002            Some(ExprNode::Binary {
3003                op: BinaryOp::Add,
3004                ..
3005            })
3006        ));
3007    }
3008
3009    #[test]
3010    fn graph_preserves_written_operand_order_for_commutative_ops() {
3011        let left_param = (parameter!("x") + 1.0).to_graph();
3012        assert!(matches!(
3013            left_param.node(left_param.root()),
3014            Some(ExprNode::Binary {
3015                op: BinaryOp::Add,
3016                lhs,
3017                rhs
3018            }) if matches!(left_param.node(*lhs), Some(ExprNode::ScalarParam(parameter)) if parameter.name() == "x")
3019                && matches!(left_param.node(*rhs), Some(ExprNode::RealConst(1.0)))
3020        ));
3021
3022        let right_param = (1.0 + parameter!("x")).to_graph();
3023        assert!(matches!(
3024            right_param.node(right_param.root()),
3025            Some(ExprNode::Binary {
3026                op: BinaryOp::Add,
3027                lhs,
3028                rhs
3029            }) if matches!(right_param.node(*lhs), Some(ExprNode::RealConst(1.0)))
3030                && matches!(right_param.node(*rhs), Some(ExprNode::ScalarParam(parameter)) if parameter.name() == "x")
3031        ));
3032    }
3033
3034    #[test]
3035    fn represents_kmatrix_style_solve_graph() {
3036        let beta = vector([
3037            complex(parameter!("b0_re"), parameter!("b0_im")),
3038            complex(parameter!("b1_re"), parameter!("b1_im")),
3039        ]);
3040        let a = matrix([
3041            [Complex64::new(1.0, 0.0), Complex64::new(0.0, 1.0)],
3042            [Complex64::new(0.0, -1.0), Complex64::new(1.0, 0.0)],
3043        ]);
3044        let graph = solve(a, beta).component(0).to_graph();
3045
3046        assert!(
3047            graph
3048                .nodes()
3049                .iter()
3050                .any(|node| matches!(node, ExprNode::Solve { .. }))
3051        );
3052    }
3053
3054    #[test]
3055    fn graph_builder_preserves_shared_dag_nodes() {
3056        let shared = event_scalar("x").sin();
3057        let expression = vector((0..1_000).map(|_| shared.clone()));
3058        let graph = expression.to_graph();
3059
3060        assert_eq!(graph.nodes().len(), 3);
3061        let ExprNode::Vector { elements } = graph.node(graph.root()).unwrap() else {
3062            panic!("root should be a vector");
3063        };
3064        assert!(elements.windows(2).all(|pair| pair[0] == pair[1]));
3065    }
3066
3067    #[test]
3068    fn expression_projection_preserves_occurrence_rebuild_behavior() {
3069        let shared = event_scalar("x").sin();
3070        let projected = (shared.clone() + shared).project_tags(["selected"]);
3071        let graph = projected.to_graph();
3072
3073        assert_eq!(
3074            graph
3075                .nodes()
3076                .iter()
3077                .filter(|node| matches!(
3078                    node,
3079                    ExprNode::Unary {
3080                        op: UnaryOp::Sin,
3081                        ..
3082                    }
3083                ))
3084                .count(),
3085            2
3086        );
3087    }
3088
3089    #[test]
3090    fn iterative_construction_and_projection_handle_deep_expressions() {
3091        let mut expression = event_scalar("x");
3092        for _ in 0..10_000 {
3093            expression = expression.sin();
3094        }
3095
3096        let projected = expression.project_tags(["selected"]);
3097        let graph = projected.to_graph();
3098
3099        assert_eq!(graph.nodes().len(), 10_001);
3100        assert_eq!(
3101            graph.reachable_post_order([graph.root()]).len(),
3102            graph.nodes().len()
3103        );
3104
3105        // Deep `Arc` chains also recurse when their final owner is dropped;
3106        // this test targets traversal behavior rather than destructor policy.
3107        std::mem::forget(expression);
3108        std::mem::forget(projected);
3109    }
3110
3111    #[test]
3112    fn tensor_scalar_lowering_handles_a_deep_scalar_element() {
3113        let mut scalar = event_scalar("x");
3114        for _ in 0..10_000 {
3115            scalar = scalar.sin();
3116        }
3117        let expression = vector([scalar.clone()]) * 2.0;
3118        let graph = expression.to_graph();
3119        assert!(
3120            matches!(graph.node(graph.root()), Some(ExprNode::Vector { elements }) if elements.len() == 1)
3121        );
3122        std::mem::forget(scalar);
3123        std::mem::forget(expression);
3124    }
3125
3126    #[test]
3127    fn reachable_post_order_preserves_child_order_and_deduplicates_shared_nodes() {
3128        let shared = event_scalar("x").sin();
3129        let graph = (shared.clone() + shared).to_graph();
3130        let order = graph.reachable_post_order([graph.root()]);
3131
3132        assert_eq!(order.len(), graph.nodes().len());
3133        assert_eq!(order.last(), Some(&graph.root()));
3134        for id in order {
3135            for child in graph.node(id).unwrap().children() {
3136                assert!(child.index() < id.index());
3137            }
3138        }
3139    }
3140
3141    #[test]
3142    fn dynamic_matrices_and_shapes_are_checked_eagerly() {
3143        let dynamic = matrix_from_flat(2, 2, [1.0, 2.0, 3.0, 4.0]).unwrap();
3144        assert_eq!(
3145            dynamic.shape().unwrap(),
3146            ExprShape::Matrix { rows: 2, cols: 2 }
3147        );
3148        assert!(matrix_from_flat(2, 2, [1.0, 2.0, 3.0]).is_err());
3149        assert!(matmul(dynamic, matrix([[1.0, 2.0, 3.0]])).shape().is_err());
3150    }
3151
3152    #[test]
3153    fn assignment_operators_build_binary_expression_nodes() {
3154        let mut expr = Expr::from(parameter!("x"));
3155        expr += parameter!("y");
3156        expr -= 1.0;
3157        expr *= Complex64::I;
3158        expr /= Expr::from(parameter!("z"));
3159
3160        let graph = expr.to_graph();
3161        assert!(matches!(
3162            graph.node(graph.root()),
3163            Some(ExprNode::Binary {
3164                op: BinaryOp::Div,
3165                ..
3166            })
3167        ));
3168        assert_eq!(
3169            graph
3170                .nodes()
3171                .iter()
3172                .filter(|node| matches!(node, ExprNode::Binary { .. }))
3173                .count(),
3174            4
3175        );
3176    }
3177
3178    #[test]
3179    fn assignment_operators_accept_borrowed_rhs_values() {
3180        let y = parameter!("y");
3181        let one = 1.0;
3182        let i = Complex64::I;
3183        let z = Expr::from(parameter!("z"));
3184
3185        let mut expr = Expr::from(parameter!("x"));
3186        expr += &y;
3187        expr -= &one;
3188        expr *= &i;
3189        expr /= &z;
3190
3191        let graph = expr.to_graph();
3192        assert!(matches!(
3193            graph.node(graph.root()),
3194            Some(ExprNode::Binary {
3195                op: BinaryOp::Div,
3196                ..
3197            })
3198        ));
3199    }
3200}