Skip to main content

p3_air/symbolic/
expression.rs

1use p3_field::{Algebra, ExtensionField, Field, InjectiveMonomial};
2use serde::{Deserialize, Serialize};
3
4use crate::symbolic::variable::{BaseEntry, SymbolicVariable};
5use crate::symbolic::{SymLeaf, SymbolicExpr};
6use crate::{AirBuilder, WindowAccess};
7
8/// Leaf nodes for base-field symbolic expressions.
9///
10/// These represent the atomic building blocks of AIR constraint expressions:
11/// trace column references, selectors, and field constants.
12#[derive(Clone, Debug, Serialize, Deserialize)]
13pub enum BaseLeaf<F> {
14    /// A reference to a trace column or public input.
15    Variable(SymbolicVariable<F>),
16
17    /// Selector evaluating to a non-zero value only on the first row.
18    IsFirstRow,
19
20    /// Selector evaluating to a non-zero value only on the last row.
21    IsLastRow,
22
23    /// Selector evaluating to zero only on the last row.
24    IsTransition,
25
26    /// A constant field element.
27    Constant(F),
28}
29
30/// A symbolic expression tree for base-field AIR constraints.
31///
32/// This is a type alias for the generic [`SymbolicExpr`] parameterized with
33/// base-field [`BaseLeaf`] nodes.
34pub type SymbolicExpression<F> = SymbolicExpr<BaseLeaf<F>>;
35
36impl<F: Field> SymLeaf for BaseLeaf<F> {
37    type F = F;
38
39    const ZERO: Self = Self::Constant(F::ZERO);
40    const ONE: Self = Self::Constant(F::ONE);
41    const TWO: Self = Self::Constant(F::TWO);
42    const NEG_ONE: Self = Self::Constant(F::NEG_ONE);
43
44    fn degree_multiple(&self) -> usize {
45        match self {
46            Self::Variable(v) => v.degree_multiple(),
47            Self::IsFirstRow | Self::IsLastRow => 1,
48            Self::IsTransition | Self::Constant(_) => 0,
49        }
50    }
51
52    fn degree_multiple_with_transition(&self, transition_degree: usize) -> usize {
53        match self {
54            Self::IsTransition => transition_degree,
55            _ => self.degree_multiple(),
56        }
57    }
58
59    fn poly_degree(&self, trace_len: usize, periodic_periods: &[usize]) -> usize {
60        match self {
61            Self::Variable(v) => v.poly_degree(trace_len, periodic_periods),
62            // Boundary selectors are non-zero at a single row, so they are degree-`(N - 1)`
63            // polynomials, while the transition selector only needs to vanish on the last
64            // row and so is linear.
65            Self::IsFirstRow | Self::IsLastRow => trace_len.saturating_sub(1),
66            Self::IsTransition => 1,
67            Self::Constant(_) => 0,
68        }
69    }
70
71    fn as_const(&self) -> Option<&F> {
72        match self {
73            Self::Constant(c) => Some(c),
74            _ => None,
75        }
76    }
77
78    fn from_const(c: F) -> Self {
79        Self::Constant(c)
80    }
81}
82
83impl<F: Field, EF: ExtensionField<F>> From<SymbolicVariable<F>> for SymbolicExpression<EF> {
84    fn from(var: SymbolicVariable<F>) -> Self {
85        Self::Leaf(BaseLeaf::Variable(SymbolicVariable::new(
86            var.entry, var.index,
87        )))
88    }
89}
90
91impl<F: Field, EF: ExtensionField<F>> From<F> for SymbolicExpression<EF> {
92    fn from(f: F) -> Self {
93        Self::Leaf(BaseLeaf::Constant(f.into()))
94    }
95}
96
97impl<F: Field> SymbolicExpression<F> {
98    /// Evaluate this symbolic expression against a concrete [`AirBuilder`].
99    ///
100    /// # Overview
101    ///
102    /// - Walk the expression tree top-down.
103    /// - Replace each leaf with the builder's concrete value.
104    /// - Recurse into arithmetic nodes; combine in the builder's algebra.
105    ///
106    /// # Algorithm
107    ///
108    /// ```text
109    ///     leaf   → builder lookup (main / preprocessed / public / periodic / selector / constant)
110    ///     x + y  → resolve(x) + resolve(y)
111    ///     x - y  → resolve(x) - resolve(y)
112    ///     x * y  → resolve(x) * resolve(y)
113    ///     -x     → -resolve(x)
114    /// ```
115    ///
116    /// # Panics
117    ///
118    /// - Row offset other than 0 or 1.
119    /// - Column index out of bounds.
120    pub fn resolve<AB>(&self, builder: &AB) -> AB::Expr
121    where
122        AB: AirBuilder<F = F>,
123    {
124        match self {
125            Self::Leaf(leaf) => match leaf {
126                BaseLeaf::Variable(v) => match v.entry {
127                    // Main trace: offset 0 = current row, offset 1 = next row.
128                    // Symbolic builders only emit two-row windows.
129                    BaseEntry::Main { offset } => {
130                        let main = builder.main();
131                        match offset {
132                            0 => main
133                                .current(v.index)
134                                .expect("main column index out of bounds")
135                                .into(),
136                            1 => main
137                                .next(v.index)
138                                .expect("main column index out of bounds")
139                                .into(),
140                            _ => panic!("expressions cannot span more than two rows"),
141                        }
142                    }
143                    // Preprocessed trace: same shape, commitment-free trace.
144                    BaseEntry::Preprocessed { offset } => {
145                        let prep = builder.preprocessed();
146                        match offset {
147                            0 => prep
148                                .current(v.index)
149                                .expect("preprocessed column index out of bounds")
150                                .into(),
151                            1 => prep
152                                .next(v.index)
153                                .expect("preprocessed column index out of bounds")
154                                .into(),
155                            _ => panic!("expressions cannot span more than two rows"),
156                        }
157                    }
158                    // Public input: direct slice lookup.
159                    BaseEntry::Public => builder.public_values()[v.index].into(),
160                    // Periodic column at the current row.
161                    // Empty default slice → out-of-bounds panic on stray emissions.
162                    BaseEntry::Periodic => builder.periodic_values()[v.index].into(),
163                },
164                // Boundary and transition selectors come straight from the builder.
165                BaseLeaf::IsFirstRow => builder.is_first_row(),
166                BaseLeaf::IsLastRow => builder.is_last_row(),
167                BaseLeaf::IsTransition => builder.is_transition_window(2),
168                // Lift the field constant into the builder's expression algebra.
169                BaseLeaf::Constant(c) => AB::Expr::from(*c),
170            },
171            // Arithmetic: recurse on operands, combine in the builder's algebra.
172            Self::Add { x, y, .. } => x.resolve(builder) + y.resolve(builder),
173            Self::Sub { x, y, .. } => x.resolve(builder) - y.resolve(builder),
174            Self::Neg { x, .. } => -x.resolve(builder),
175            Self::Mul { x, y, .. } => x.resolve(builder) * y.resolve(builder),
176        }
177    }
178}
179
180impl<F: Field> Algebra<F> for SymbolicExpression<F> {}
181
182impl<F: Field> Algebra<SymbolicVariable<F>> for SymbolicExpression<F> {}
183
184// Note we cannot implement PermutationMonomial due to the degree_multiple part which makes
185// operations non invertible.
186impl<F: Field + InjectiveMonomial<N>, const N: u64> InjectiveMonomial<N> for SymbolicExpression<F> {}
187
188#[cfg(test)]
189mod tests {
190    use alloc::sync::Arc;
191    use alloc::vec;
192    use alloc::vec::Vec;
193
194    use p3_baby_bear::BabyBear;
195    use p3_field::PrimeCharacteristicRing;
196    use p3_matrix::dense::RowMajorMatrix;
197
198    use super::*;
199    use crate::symbolic::BaseEntry;
200
201    #[test]
202    fn test_symbolic_expression_degree_multiple() {
203        let constant_expr =
204            SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::Constant(BabyBear::new(5)));
205        assert_eq!(
206            constant_expr.degree_multiple(),
207            0,
208            "Constant should have degree 0"
209        );
210
211        let variable_expr = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::new(
212            BaseEntry::Main { offset: 0 },
213            1,
214        )));
215        assert_eq!(
216            variable_expr.degree_multiple(),
217            1,
218            "Main variable should have degree 1"
219        );
220
221        let preprocessed_var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::new(
222            BaseEntry::Preprocessed { offset: 0 },
223            2,
224        )));
225        assert_eq!(
226            preprocessed_var.degree_multiple(),
227            1,
228            "Preprocessed variable should have degree 1"
229        );
230
231        let public_var = SymbolicExpression::Leaf(BaseLeaf::Variable(
232            SymbolicVariable::<BabyBear>::new(BaseEntry::Public, 4),
233        ));
234        assert_eq!(
235            public_var.degree_multiple(),
236            0,
237            "Public variable should have degree 0"
238        );
239
240        let is_first_row = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsFirstRow);
241        assert_eq!(
242            is_first_row.degree_multiple(),
243            1,
244            "IsFirstRow should have degree 1"
245        );
246
247        let is_last_row = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsLastRow);
248        assert_eq!(
249            is_last_row.degree_multiple(),
250            1,
251            "IsLastRow should have degree 1"
252        );
253
254        let is_transition = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsTransition);
255        assert_eq!(
256            is_transition.degree_multiple(),
257            0,
258            "IsTransition should have degree 0"
259        );
260
261        let add_expr = SymbolicExpr::<BaseLeaf<BabyBear>>::Add {
262            x: Arc::new(variable_expr.clone()),
263            y: Arc::new(preprocessed_var.clone()),
264            degree_multiple: 1,
265        };
266        assert_eq!(
267            add_expr.degree_multiple(),
268            1,
269            "Addition should take max degree of inputs"
270        );
271
272        let sub_expr = SymbolicExpr::<BaseLeaf<BabyBear>>::Sub {
273            x: Arc::new(variable_expr.clone()),
274            y: Arc::new(preprocessed_var.clone()),
275            degree_multiple: 1,
276        };
277        assert_eq!(
278            sub_expr.degree_multiple(),
279            1,
280            "Subtraction should take max degree of inputs"
281        );
282
283        let neg_expr = SymbolicExpr::<BaseLeaf<BabyBear>>::Neg {
284            x: Arc::new(variable_expr.clone()),
285            degree_multiple: 1,
286        };
287        assert_eq!(
288            neg_expr.degree_multiple(),
289            1,
290            "Negation should keep the degree"
291        );
292
293        let mul_expr = SymbolicExpr::<BaseLeaf<BabyBear>>::Mul {
294            x: Arc::new(variable_expr),
295            y: Arc::new(preprocessed_var),
296            degree_multiple: 2,
297        };
298        assert_eq!(
299            mul_expr.degree_multiple(),
300            2,
301            "Multiplication should sum degrees"
302        );
303    }
304
305    #[test]
306    fn test_symbolic_expression_poly_degree() {
307        const N: usize = 8;
308
309        // The transition selector is linear, unlike the boundary selectors which are
310        // degree-`(N - 1)` polynomials.
311        let is_transition = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsTransition);
312        assert_eq!(is_transition.poly_degree(N, &[]), 1);
313
314        let is_first_row = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsFirstRow);
315        assert_eq!(is_first_row.poly_degree(N, &[]), N - 1);
316
317        // Constants contribute nothing.
318        let constant = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(5)));
319        assert_eq!(constant.poly_degree(N, &[]), 0);
320
321        // `is_transition * main` is degree `1 + (N - 1) = N`.
322        let main = SymbolicExpression::<BabyBear>::from(SymbolicVariable::new(
323            BaseEntry::Main { offset: 0 },
324            0,
325        ));
326        let guarded = is_transition * main.clone();
327        assert_eq!(guarded.poly_degree(N, &[]), N);
328
329        // Products of periodic columns sum their reduced degrees: two period-2
330        // columns give `(N - N/2) + (N - N/2) = N`, versus `2(N - 1)` for two
331        // regular columns.
332        let p0 =
333            SymbolicExpression::<BabyBear>::from(SymbolicVariable::new(BaseEntry::Periodic, 0));
334        let p1 =
335            SymbolicExpression::<BabyBear>::from(SymbolicVariable::new(BaseEntry::Periodic, 1));
336        let periodic_product = p0 * p1;
337        assert_eq!(periodic_product.poly_degree(N, &[2, 2]), N);
338
339        // Sums take the max degree of their operands.
340        let sum = main + SymbolicExpression::Leaf(BaseLeaf::IsTransition);
341        assert_eq!(sum.poly_degree(N, &[]), N - 1);
342    }
343
344    #[test]
345    fn poly_degree_handles_shared_dag_in_linear_time() {
346        // Repeated squaring builds a DAG of depth `d` whose flattened tree has
347        // `2^d` leaves but only `O(d)` distinct nodes. `poly_degree` must run in
348        // time proportional to the distinct nodes; without memoization this test
349        // would take `O(2^d)` and never finish.
350        const DEPTH: usize = 30;
351        const N: usize = 1 << 10;
352
353        let mut expr =
354            SymbolicExpression::<BabyBear>::from(SymbolicVariable::new(BaseEntry::Periodic, 0));
355        for _ in 0..DEPTH {
356            expr = expr.clone() * expr.clone();
357        }
358
359        // The period-2 column has degree `N - N/2 = N/2`; each squaring doubles it.
360        assert_eq!(expr.poly_degree(N, &[2]), (N / 2) << DEPTH);
361        assert_eq!(expr.degree_multiple_with_transition(1), 1 << DEPTH);
362    }
363
364    #[test]
365    fn degree_multiple_counts_every_transition_factor() {
366        let transition = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsTransition);
367        let main =
368            SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 0 }, 0));
369        let expr = transition.clone().cube() * (main.cube() - transition);
370        assert_eq!(expr.degree_multiple_with_transition(0), 3);
371        assert_eq!(expr.degree_multiple_with_transition(1), 6);
372    }
373
374    #[test]
375    fn test_addition_of_constants() {
376        let a = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(3)));
377        let b = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(4)));
378        let result = a + b;
379        match result {
380            SymbolicExpr::Leaf(BaseLeaf::Constant(val)) => assert_eq!(val, BabyBear::new(7)),
381            _ => panic!("Addition of constants did not simplify correctly"),
382        }
383    }
384
385    #[test]
386    fn test_subtraction_of_constants() {
387        let a = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(10)));
388        let b = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(4)));
389        let result = a - b;
390        match result {
391            SymbolicExpr::Leaf(BaseLeaf::Constant(val)) => assert_eq!(val, BabyBear::new(6)),
392            _ => panic!("Subtraction of constants did not simplify correctly"),
393        }
394    }
395
396    #[test]
397    fn test_negation() {
398        let a = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(7)));
399        let result = -a;
400        match result {
401            SymbolicExpr::Leaf(BaseLeaf::Constant(val)) => {
402                assert_eq!(val, BabyBear::NEG_ONE * BabyBear::new(7));
403            }
404            _ => panic!("Negation did not work correctly"),
405        }
406    }
407
408    #[test]
409    fn test_multiplication_of_constants() {
410        let a = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(3)));
411        let b = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(5)));
412        let result = a * b;
413        match result {
414            SymbolicExpr::Leaf(BaseLeaf::Constant(val)) => assert_eq!(val, BabyBear::new(15)),
415            _ => panic!("Multiplication of constants did not simplify correctly"),
416        }
417    }
418
419    #[test]
420    fn test_degree_multiple_for_addition() {
421        let a = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
422            BaseEntry::Main { offset: 0 },
423            1,
424        )));
425        let b = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
426            BaseEntry::Main { offset: 0 },
427            2,
428        )));
429        let result = a + b;
430        match result {
431            SymbolicExpr::Add {
432                degree_multiple,
433                x,
434                y,
435            } => {
436                assert_eq!(degree_multiple, 1);
437                assert!(
438                    matches!(&*x, SymbolicExpr::Leaf(BaseLeaf::Variable(v)) if v.index == 1 && matches!(v.entry, BaseEntry::Main { offset: 0 }))
439                );
440                assert!(
441                    matches!(&*y, SymbolicExpr::Leaf(BaseLeaf::Variable(v)) if v.index == 2 && matches!(v.entry, BaseEntry::Main { offset: 0 }))
442                );
443            }
444            _ => panic!("Addition did not create an Add expression"),
445        }
446    }
447
448    #[test]
449    fn test_degree_multiple_for_multiplication() {
450        let a = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
451            BaseEntry::Main { offset: 0 },
452            1,
453        )));
454        let b = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
455            BaseEntry::Main { offset: 0 },
456            2,
457        )));
458        let result = a * b;
459
460        match result {
461            SymbolicExpr::Mul {
462                degree_multiple,
463                x,
464                y,
465            } => {
466                assert_eq!(degree_multiple, 2, "Multiplication should sum degrees");
467
468                assert!(
469                    matches!(&*x, SymbolicExpr::Leaf(BaseLeaf::Variable(v))
470                        if v.index == 1 && matches!(v.entry, BaseEntry::Main { offset: 0 })
471                    ),
472                    "Left operand should match `a`"
473                );
474
475                assert!(
476                    matches!(&*y, SymbolicExpr::Leaf(BaseLeaf::Variable(v))
477                        if v.index == 2 && matches!(v.entry, BaseEntry::Main { offset: 0 })
478                    ),
479                    "Right operand should match `b`"
480                );
481            }
482            _ => panic!("Multiplication did not create a `Mul` expression"),
483        }
484    }
485
486    #[test]
487    fn test_sum_operator() {
488        let expressions = vec![
489            SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(2))),
490            SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(3))),
491            SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(5))),
492        ];
493        let result: SymbolicExpression<BabyBear> = expressions.into_iter().sum();
494        match result {
495            SymbolicExpr::Leaf(BaseLeaf::Constant(val)) => assert_eq!(val, BabyBear::new(10)),
496            _ => panic!("Sum did not produce correct result"),
497        }
498    }
499
500    #[test]
501    fn test_product_operator() {
502        let expressions = vec![
503            SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(2))),
504            SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(3))),
505            SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(4))),
506        ];
507        let result: SymbolicExpression<BabyBear> = expressions.into_iter().product();
508        match result {
509            SymbolicExpr::Leaf(BaseLeaf::Constant(val)) => assert_eq!(val, BabyBear::new(24)),
510            _ => panic!("Product did not produce correct result"),
511        }
512    }
513
514    #[test]
515    fn test_default_is_zero() {
516        // Default should produce ZERO constant.
517        let expr: SymbolicExpression<BabyBear> = Default::default();
518
519        // Verify it matches the zero constant.
520        assert!(matches!(
521            expr,
522            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ZERO
523        ));
524    }
525
526    #[test]
527    fn test_ring_constants() {
528        // ZERO is a Constant variant wrapping the field's zero element.
529        assert!(matches!(
530            SymbolicExpression::<BabyBear>::ZERO,
531            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ZERO
532        ));
533        // ONE is a Constant variant wrapping the field's one element.
534        assert!(matches!(
535            SymbolicExpression::<BabyBear>::ONE,
536            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ONE
537        ));
538        // TWO is a Constant variant wrapping the field's two element.
539        assert!(matches!(
540            SymbolicExpression::<BabyBear>::TWO,
541            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::TWO
542        ));
543        // NEG_ONE is a Constant variant wrapping the field's -1 element.
544        assert!(matches!(
545            SymbolicExpression::<BabyBear>::NEG_ONE,
546            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::NEG_ONE
547        ));
548    }
549
550    #[test]
551    fn test_from_symbolic_variable() {
552        // Create a main trace variable at column index 3.
553        let var = SymbolicVariable::<BabyBear>::new(BaseEntry::Main { offset: 0 }, 3);
554        // Convert to expression.
555        let expr: SymbolicExpression<BabyBear> = var.into();
556        // Verify the variable is preserved with correct entry and index.
557        match expr {
558            SymbolicExpr::Leaf(BaseLeaf::Variable(v)) => {
559                assert!(matches!(v.entry, BaseEntry::Main { offset: 0 }));
560                assert_eq!(v.index, 3);
561            }
562            _ => panic!("Expected Variable variant"),
563        }
564    }
565
566    #[test]
567    fn test_from_field_element() {
568        // Convert a field element directly to expression.
569        let field_val = BabyBear::new(42);
570        let expr: SymbolicExpression<BabyBear> = field_val.into();
571        // Verify it becomes a Constant with the same value.
572        assert!(matches!(
573            expr,
574            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == field_val
575        ));
576    }
577
578    #[test]
579    fn test_from_prime_subfield() {
580        // Create expression from prime subfield element.
581        let prime_subfield_val = <BabyBear as PrimeCharacteristicRing>::PrimeSubfield::new(7);
582        let expr = SymbolicExpression::<BabyBear>::from_prime_subfield(prime_subfield_val);
583        // Verify it produces a constant with the converted value.
584        assert!(matches!(
585            expr,
586            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::new(7)
587        ));
588    }
589
590    #[test]
591    fn test_assign_operators() {
592        // Test AddAssign with constants (should simplify).
593        let mut expr = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(5)));
594        expr += SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(3)));
595        assert!(matches!(
596            expr,
597            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::new(8)
598        ));
599
600        // Test SubAssign with constants (should simplify).
601        let mut expr = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(10)));
602        expr -= SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(4)));
603        assert!(matches!(
604            expr,
605            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::new(6)
606        ));
607
608        // Test MulAssign with constants (should simplify).
609        let mut expr = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(6)));
610        expr *= SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(7)));
611        assert!(matches!(
612            expr,
613            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::new(42)
614        ));
615    }
616
617    #[test]
618    fn test_subtraction_creates_sub_node() {
619        // Create two trace variables.
620        let a = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
621            BaseEntry::Main { offset: 0 },
622            0,
623        )));
624        let b = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
625            BaseEntry::Main { offset: 0 },
626            1,
627        )));
628
629        // Subtract them.
630        let result = a - b;
631
632        // Should create Sub node (not simplified).
633        match result {
634            SymbolicExpr::Sub {
635                x,
636                y,
637                degree_multiple,
638            } => {
639                // Both operands have degree 1, so max is 1.
640                assert_eq!(degree_multiple, 1);
641
642                // Verify left operand is main trace variable at index 0, offset 0.
643                assert!(matches!(
644                    x.as_ref(),
645                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
646                        if v.index == 0 && matches!(v.entry, BaseEntry::Main { offset: 0 })
647                ));
648
649                // Verify right operand is main trace variable at index 1, offset 0.
650                assert!(matches!(
651                    y.as_ref(),
652                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
653                        if v.index == 1 && matches!(v.entry, BaseEntry::Main { offset: 0 })
654                ));
655            }
656            _ => panic!("Expected Sub variant"),
657        }
658    }
659
660    #[test]
661    fn test_negation_creates_neg_node() {
662        // Create a trace variable.
663        let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
664            BaseEntry::Main { offset: 0 },
665            0,
666        )));
667
668        // Negate it.
669        let result = -var;
670
671        // Should create Neg node (not simplified).
672        match result {
673            SymbolicExpr::Neg { x, degree_multiple } => {
674                // Degree is preserved from operand.
675                assert_eq!(degree_multiple, 1);
676
677                // Verify operand is main trace variable at index 0, offset 0.
678                assert!(matches!(
679                    x.as_ref(),
680                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
681                        if v.index == 0 && matches!(v.entry, BaseEntry::Main { offset: 0 })
682                ));
683            }
684            _ => panic!("Expected Neg variant"),
685        }
686    }
687
688    #[test]
689    fn test_empty_sum_returns_zero() {
690        // Sum of empty iterator should be additive identity.
691        let empty: Vec<SymbolicExpression<BabyBear>> = vec![];
692        let result: SymbolicExpression<BabyBear> = empty.into_iter().sum();
693        assert!(matches!(
694            result,
695            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ZERO
696        ));
697    }
698
699    #[test]
700    fn test_empty_product_returns_one() {
701        // Product of empty iterator should be multiplicative identity.
702        let empty: Vec<SymbolicExpression<BabyBear>> = vec![];
703        let result: SymbolicExpression<BabyBear> = empty.into_iter().product();
704        assert!(matches!(
705            result,
706            SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ONE
707        ));
708    }
709
710    #[test]
711    fn test_mixed_degree_addition() {
712        // Constant has degree 0.
713        let constant = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(5)));
714
715        // Variable has degree 1.
716        let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
717            BaseEntry::Main { offset: 0 },
718            0,
719        )));
720
721        // Add them: max(0, 1) = 1.
722        let result = constant + var;
723
724        match result {
725            SymbolicExpr::Add {
726                x,
727                y,
728                degree_multiple,
729            } => {
730                // Degree is max(0, 1) = 1.
731                assert_eq!(degree_multiple, 1);
732
733                // Verify left operand is the constant 5.
734                assert!(matches!(
735                    x.as_ref(),
736                    SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if *c == BabyBear::new(5)
737                ));
738
739                // Verify right operand is main trace variable at index 0, offset 0.
740                assert!(matches!(
741                    y.as_ref(),
742                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
743                        if v.index == 0 && matches!(v.entry, BaseEntry::Main { offset: 0 })
744                ));
745            }
746            _ => panic!("Expected Add variant"),
747        }
748    }
749
750    #[test]
751    fn test_chained_multiplication_degree() {
752        // Create three variables, each with degree 1.
753        let a = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
754            BaseEntry::Main { offset: 0 },
755            0,
756        )));
757        let b = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
758            BaseEntry::Main { offset: 0 },
759            1,
760        )));
761        let c = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
762            BaseEntry::Main { offset: 0 },
763            2,
764        )));
765
766        // a * b has degree 1 + 1 = 2.
767        let ab = a * b;
768        assert_eq!(ab.degree_multiple(), 2);
769
770        // (a * b) * c has degree 2 + 1 = 3.
771        let abc = ab * c;
772        assert_eq!(abc.degree_multiple(), 3);
773    }
774
775    #[test]
776    fn test_add_zero_identity_folding() {
777        let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
778            BaseEntry::Main { offset: 0 },
779            0,
780        )));
781        let zero = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::Constant(BabyBear::ZERO));
782
783        // x + 0 should return x, not create an Add node.
784        let result = var.clone() + zero.clone();
785        assert!(
786            matches!(result, SymbolicExpr::Leaf(BaseLeaf::Variable(_))),
787            "x + 0 should fold to x"
788        );
789
790        // 0 + x should return x, not create an Add node.
791        let result = zero + var;
792        assert!(
793            matches!(result, SymbolicExpr::Leaf(BaseLeaf::Variable(_))),
794            "0 + x should fold to x"
795        );
796    }
797
798    #[test]
799    fn test_sub_zero_identity_folding() {
800        let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
801            BaseEntry::Main { offset: 0 },
802            0,
803        )));
804        let zero = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::Constant(BabyBear::ZERO));
805
806        // x - 0 should return x, not create a Sub node.
807        let result = var.clone() - zero.clone();
808        assert!(
809            matches!(result, SymbolicExpr::Leaf(BaseLeaf::Variable(_))),
810            "x - 0 should fold to x"
811        );
812
813        // 0 - x should return -x, not create a Sub node.
814        let result = zero - var;
815        match result {
816            SymbolicExpr::Neg { x, degree_multiple } => {
817                assert_eq!(degree_multiple, 1);
818                assert!(matches!(
819                    x.as_ref(),
820                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
821                        if v.index == 0 && v.entry == BaseEntry::Main { offset: 0 }
822                ));
823            }
824            _ => panic!("0 - x should fold to Neg(x)"),
825        }
826    }
827
828    #[test]
829    fn test_mul_zero_identity_folding() {
830        let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
831            BaseEntry::Main { offset: 0 },
832            0,
833        )));
834        let zero = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::Constant(BabyBear::ZERO));
835
836        // x * 0 should return Constant(0), not create a Mul node.
837        let result = var.clone() * zero.clone();
838        assert!(
839            matches!(result, SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ZERO),
840            "x * 0 should fold to 0"
841        );
842
843        // 0 * x should return Constant(0), not create a Mul node.
844        let result = zero * var;
845        assert!(
846            matches!(result, SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ZERO),
847            "0 * x should fold to 0"
848        );
849    }
850
851    #[test]
852    fn test_mul_one_identity_folding() {
853        let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
854            BaseEntry::Main { offset: 0 },
855            0,
856        )));
857        let one = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::Constant(BabyBear::ONE));
858
859        // x * 1 should return x, not create a Mul node.
860        let result = var.clone() * one.clone();
861        assert!(
862            matches!(result, SymbolicExpr::Leaf(BaseLeaf::Variable(_))),
863            "x * 1 should fold to x"
864        );
865
866        // 1 * x should return x, not create a Mul node.
867        let result = one * var;
868        assert!(
869            matches!(result, SymbolicExpr::Leaf(BaseLeaf::Variable(_))),
870            "1 * x should fold to x"
871        );
872    }
873
874    #[test]
875    fn test_identity_folding_preserves_degree() {
876        let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
877            BaseEntry::Main { offset: 0 },
878            0,
879        )));
880        let zero = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::Constant(BabyBear::ZERO));
881        let one = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::Constant(BabyBear::ONE));
882
883        // x + 0 should preserve degree of x.
884        let result = var.clone() + zero.clone();
885        assert_eq!(result.degree_multiple(), 1);
886
887        // x - 0 should preserve degree of x.
888        let result = var.clone() - zero.clone();
889        assert_eq!(result.degree_multiple(), 1);
890
891        // 0 - x should preserve degree of x.
892        let result = zero.clone() - var.clone();
893        assert_eq!(result.degree_multiple(), 1);
894
895        // x * 1 should preserve degree of x.
896        let result = var.clone() * one;
897        assert_eq!(result.degree_multiple(), 1);
898
899        // x * 0 should have degree 0 (constant).
900        let result = var * zero;
901        assert_eq!(result.degree_multiple(), 0);
902    }
903
904    /// Minimal builder used to drive symbolic-expression resolution.
905    ///
906    /// Carries:
907    /// - a 2-row main trace,
908    /// - a public-value slice,
909    /// - precomputed selector values for the current row,
910    /// - a periodic-column row evaluated at the current step.
911    struct ResolveTestBuilder {
912        main: RowMajorMatrix<BabyBear>,
913        public_values: Vec<BabyBear>,
914        periodic_row: Vec<BabyBear>,
915        is_first: BabyBear,
916        is_last: BabyBear,
917        is_transition: BabyBear,
918    }
919
920    impl AirBuilder for ResolveTestBuilder {
921        type F = BabyBear;
922        type Expr = BabyBear;
923        type Var = BabyBear;
924        type PreprocessedWindow = RowMajorMatrix<BabyBear>;
925        type MainWindow = RowMajorMatrix<BabyBear>;
926        type PublicVar = BabyBear;
927        type PeriodicVar = BabyBear;
928
929        fn main(&self) -> Self::MainWindow {
930            self.main.clone()
931        }
932
933        fn preprocessed(&self) -> &Self::PreprocessedWindow {
934            unimplemented!("no preprocessed columns in test builder")
935        }
936
937        fn is_first_row(&self) -> Self::Expr {
938            self.is_first
939        }
940
941        fn is_last_row(&self) -> Self::Expr {
942            self.is_last
943        }
944
945        fn is_transition(&self) -> Self::Expr {
946            self.is_transition
947        }
948
949        fn assert_zero<I: Into<Self::Expr>>(&mut self, _: I) {}
950
951        fn public_values(&self) -> &[Self::PublicVar] {
952            &self.public_values
953        }
954
955        fn periodic_values(&self) -> &[Self::PeriodicVar] {
956            &self.periodic_row
957        }
958    }
959
960    /// 2-row × 2-column trace, plus a 2-cell periodic row at the current step:
961    ///
962    /// ```text
963    ///     main row 0 (current): [10, 20]
964    ///     main row 1 (next):    [30, 40]
965    ///     periodic_row (curr):  [7, 13]
966    /// ```
967    fn test_builder() -> ResolveTestBuilder {
968        ResolveTestBuilder {
969            main: RowMajorMatrix::new(
970                vec![
971                    BabyBear::new(10),
972                    BabyBear::new(20), // current row
973                    BabyBear::new(30),
974                    BabyBear::new(40), // next row
975                ],
976                2, // width
977            ),
978            public_values: vec![BabyBear::new(99)],
979            // Two periodic columns at the current row.
980            // Distinct primes so any cross-stream mix-up is visible.
981            periodic_row: vec![BabyBear::new(7), BabyBear::new(13)],
982            is_first: BabyBear::ONE,
983            is_last: BabyBear::ZERO,
984            is_transition: BabyBear::ONE,
985        }
986    }
987
988    #[test]
989    fn resolve_main_current_row() {
990        let b = test_builder();
991        // Main column 0, offset 0 → current row value 10.
992        let expr =
993            SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 0 }, 0));
994        assert_eq!(expr.resolve(&b), BabyBear::new(10));
995    }
996
997    #[test]
998    fn resolve_main_next_row() {
999        let b = test_builder();
1000        // Main column 1, offset 1 → next row value 40.
1001        let expr =
1002            SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 1 }, 1));
1003        assert_eq!(expr.resolve(&b), BabyBear::new(40));
1004    }
1005
1006    #[test]
1007    fn resolve_public_value() {
1008        let b = test_builder();
1009        // Public value at index 0 → 99.
1010        let expr = SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Public, 0));
1011        assert_eq!(expr.resolve(&b), BabyBear::new(99));
1012    }
1013
1014    #[test]
1015    fn resolve_constant() {
1016        let b = test_builder();
1017        let expr = SymbolicExpression::<BabyBear>::from(BabyBear::new(42));
1018        assert_eq!(expr.resolve(&b), BabyBear::new(42));
1019    }
1020
1021    #[test]
1022    fn resolve_selectors() {
1023        let b = test_builder();
1024
1025        let first = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsFirstRow);
1026        assert_eq!(first.resolve(&b), BabyBear::ONE, "is_first_row = 1");
1027
1028        let last = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsLastRow);
1029        assert_eq!(last.resolve(&b), BabyBear::ZERO, "is_last_row = 0");
1030
1031        let trans = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsTransition);
1032        assert_eq!(trans.resolve(&b), BabyBear::ONE, "is_transition = 1");
1033    }
1034
1035    #[test]
1036    fn resolve_arithmetic() {
1037        let b = test_builder();
1038
1039        // col0_curr = 10, col1_curr = 20.
1040        let col0 =
1041            SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 0 }, 0));
1042        let col1 =
1043            SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 0 }, 1));
1044
1045        // 10 + 20 = 30.
1046        let add = col0.clone() + col1.clone();
1047        assert_eq!(add.resolve(&b), BabyBear::new(30));
1048
1049        // 10 - 20 = -10 (mod p).
1050        let sub = col0.clone() - col1.clone();
1051        assert_eq!(sub.resolve(&b), BabyBear::new(10) - BabyBear::new(20));
1052
1053        // 10 * 20 = 200.
1054        let mul = col0.clone() * col1;
1055        assert_eq!(mul.resolve(&b), BabyBear::new(200));
1056
1057        // -10 (mod p).
1058        let neg = -col0;
1059        assert_eq!(neg.resolve(&b), -BabyBear::new(10));
1060    }
1061
1062    #[test]
1063    fn resolve_periodic_columns() {
1064        // Invariant: a periodic leaf reads from the builder's
1065        // periodic-value slice, in declared column order.
1066        //
1067        // Fixture:
1068        //
1069        //     periodic row (current step) : [7, 13]
1070        //     index 0 →  7
1071        //     index 1 → 13
1072        let b = test_builder();
1073
1074        // Column 0 → 7.
1075        let p0 =
1076            SymbolicExpression::from(SymbolicVariable::<BabyBear>::new(BaseEntry::Periodic, 0));
1077        assert_eq!(p0.resolve(&b), BabyBear::new(7));
1078
1079        // Column 1 → 13.
1080        let p1 =
1081            SymbolicExpression::from(SymbolicVariable::<BabyBear>::new(BaseEntry::Periodic, 1));
1082        assert_eq!(p1.resolve(&b), BabyBear::new(13));
1083    }
1084
1085    #[test]
1086    fn resolve_periodic_combines_with_arithmetic() {
1087        // Invariant: periodic leaves compose under the same algebra
1088        // as main, public, and preprocessed leaves.
1089        //
1090        // Fixture:
1091        //
1092        //     periodic row : [7, 13]
1093        //     main row 0   : [10, 20]
1094        //
1095        //     expression   : main[0] * periodic[0] + periodic[1]
1096        //                  = 10 * 7 + 13
1097        //                  = 83
1098        let b = test_builder();
1099
1100        let col0 =
1101            SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 0 }, 0));
1102        let p0 =
1103            SymbolicExpression::from(SymbolicVariable::<BabyBear>::new(BaseEntry::Periodic, 0));
1104        let p1 =
1105            SymbolicExpression::from(SymbolicVariable::<BabyBear>::new(BaseEntry::Periodic, 1));
1106
1107        let expr = col0 * p0 + p1;
1108        assert_eq!(expr.resolve(&b), BabyBear::new(83));
1109    }
1110
1111    #[test]
1112    fn serde_round_trip_preserves_resolution() {
1113        // A constraint mixing every leaf kind, both row offsets, and all node kinds:
1114        //   main[0]·main_next[1] - public[0] + periodic[0]·is_transition - constant
1115        let b = test_builder();
1116        let main_cur =
1117            SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 0 }, 0));
1118        let main_next =
1119            SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 1 }, 1));
1120        let public =
1121            SymbolicExpression::from(SymbolicVariable::<BabyBear>::new(BaseEntry::Public, 0));
1122        let periodic =
1123            SymbolicExpression::from(SymbolicVariable::<BabyBear>::new(BaseEntry::Periodic, 0));
1124        let transition = SymbolicExpression::<BabyBear>::Leaf(BaseLeaf::IsTransition);
1125
1126        let expr = main_cur * main_next - public + periodic * transition
1127            - SymbolicExpression::from(BabyBear::new(5));
1128
1129        let json = serde_json::to_string(&expr).unwrap();
1130        let decoded: SymbolicExpression<BabyBear> = serde_json::from_str(&json).unwrap();
1131
1132        // Semantic equality: both trees resolve to the same value.
1133        assert_eq!(decoded.resolve(&b), expr.resolve(&b));
1134        // Structural equality: the decoded tree re-serializes identically.
1135        assert_eq!(serde_json::to_string(&decoded).unwrap(), json);
1136        assert_eq!(decoded.degree_multiple(), expr.degree_multiple());
1137    }
1138}