Skip to main content

p3_air/symbolic/
mod.rs

1//! Symbolic expression types for AIR constraint representation.
2
3mod builder;
4mod expression;
5pub(crate) mod expression_ext;
6mod flatten;
7mod variable;
8
9use alloc::collections::BTreeMap;
10use alloc::sync::Arc;
11use core::iter::{Product, Sum};
12use core::ops;
13
14pub use builder::*;
15pub use expression::{BaseLeaf, SymbolicExpression};
16pub use expression_ext::{ExtLeaf, SymbolicExpressionExt};
17use p3_field::{Dup, ExtensionField, Field, PrimeCharacteristicRing};
18pub use variable::{BaseEntry, ExtEntry, SymbolicVariable, SymbolicVariableExt};
19
20/// Properties that leaf nodes must provide for the generic expression tree.
21///
22/// Both [`BaseLeaf`] (base-field) and
23/// [`ExtLeaf`] (extension-field) implement this trait,
24/// enabling [`SymbolicExpr`] to handle constant folding, degree tracking, and
25/// arithmetic generically.
26pub trait SymLeaf: Clone + core::fmt::Debug {
27    /// The base field type used for constant folding.
28    type F: Field;
29
30    const ZERO: Self;
31    const ONE: Self;
32    const TWO: Self;
33    const NEG_ONE: Self;
34
35    /// Returns the degree multiple of this leaf.
36    fn degree_multiple(&self) -> usize;
37
38    /// Degree multiple with a domain-specific weight for the transition selector.
39    fn degree_multiple_with_transition(&self, _transition_degree: usize) -> usize {
40        self.degree_multiple()
41    }
42
43    /// Returns the exact polynomial degree of this leaf over a trace of length
44    /// `trace_len`, given the period of each periodic column.
45    fn poly_degree(&self, trace_len: usize, periodic_periods: &[usize]) -> usize;
46
47    /// Try to view this leaf as a base-field constant.
48    fn as_const(&self) -> Option<&Self::F>;
49
50    /// Create a leaf from a base-field constant.
51    fn from_const(c: Self::F) -> Self;
52}
53
54/// A symbolic expression tree, generic over its leaf type `A`.
55///
56/// This enum captures the shared tree structure — Add/Sub/Neg/Mul nodes with
57/// `Arc`-wrapped children and cached degree multiples — used by both base-field
58/// and extension-field symbolic expressions.
59///
60/// Concrete types are provided via type aliases:
61/// - [`SymbolicExpression<F>`] = `SymbolicExpr<BaseLeaf<F>>` (base-field constraints)
62/// - [`SymbolicExpressionExt<F, EF>`] = `SymbolicExpr<ExtLeaf<F, EF>>` (extension-field constraints)
63#[derive(Clone, Debug)]
64pub enum SymbolicExpr<A> {
65    /// A leaf node (variable, constant, selector, or lifted sub-expression).
66    Leaf(A),
67
68    /// Addition of two sub-expressions.
69    Add {
70        x: Arc<Self>,
71        y: Arc<Self>,
72        degree_multiple: usize,
73    },
74
75    /// Subtraction of two sub-expressions.
76    Sub {
77        x: Arc<Self>,
78        y: Arc<Self>,
79        degree_multiple: usize,
80    },
81
82    /// Negation of a sub-expression.
83    Neg {
84        x: Arc<Self>,
85        degree_multiple: usize,
86    },
87
88    /// Multiplication of two sub-expressions.
89    Mul {
90        x: Arc<Self>,
91        y: Arc<Self>,
92        degree_multiple: usize,
93    },
94}
95
96impl<A: SymLeaf> SymbolicExpr<A> {
97    /// Returns the degree multiple of this expression.
98    pub fn degree_multiple(&self) -> usize {
99        match self {
100            Self::Leaf(a) => a.degree_multiple(),
101            Self::Add {
102                degree_multiple, ..
103            }
104            | Self::Sub {
105                degree_multiple, ..
106            }
107            | Self::Neg {
108                degree_multiple, ..
109            }
110            | Self::Mul {
111                degree_multiple, ..
112            } => *degree_multiple,
113        }
114    }
115
116    /// Degree multiple counting each transition-selector factor with the given weight.
117    ///
118    /// Weight zero uses the cached two-adic degree. Weight one also counts Circle
119    /// transition selectors as full trace-space polynomials. Periodic columns
120    /// retain their full degree multiple: Circle's doubling map does not give
121    /// them the reduced degree of a two-adic periodic polynomial.
122    pub fn degree_multiple_with_transition(&self, transition_degree: usize) -> usize {
123        if transition_degree == 0 {
124            return self.degree_multiple();
125        }
126        self.degree_with(&|leaf| leaf.degree_multiple_with_transition(transition_degree))
127    }
128
129    /// Returns the exact polynomial degree of this expression over a trace of
130    /// length `trace_len`, given the period of each periodic column (indexed by
131    /// periodic column index).
132    ///
133    /// This is trace-size aware: it treats the transition selector as the linear
134    /// polynomial it is and accounts for the reduced degree of periodic columns,
135    /// unlike the trace-size-independent [`Self::degree_multiple`].
136    pub fn poly_degree(&self, trace_len: usize, periodic_periods: &[usize]) -> usize {
137        self.degree_with(&|leaf| leaf.poly_degree(trace_len, periodic_periods))
138    }
139
140    fn degree_with(&self, leaf_degree: &impl Fn(&A) -> usize) -> usize {
141        // The expression is a DAG: arithmetic nodes share `Arc` children, so a naive
142        // recursion would revisit shared subtrees exponentially. Memoize on node
143        // identity to keep this linear in the number of distinct nodes.
144        let mut cache: BTreeMap<*const Self, usize> = BTreeMap::new();
145        self.degree_memo(leaf_degree, &mut cache)
146    }
147
148    fn degree_memo(
149        &self,
150        leaf_degree: &impl Fn(&A) -> usize,
151        cache: &mut BTreeMap<*const Self, usize>,
152    ) -> usize {
153        match self {
154            Self::Leaf(a) => leaf_degree(a),
155            Self::Add { x, y, .. } | Self::Sub { x, y, .. } => Self::child_degree(
156                x,
157                leaf_degree,
158                cache,
159            )
160            .max(Self::child_degree(y, leaf_degree, cache)),
161            Self::Neg { x, .. } => Self::child_degree(x, leaf_degree, cache),
162            Self::Mul { x, y, .. } => {
163                Self::child_degree(x, leaf_degree, cache)
164                    + Self::child_degree(y, leaf_degree, cache)
165            }
166        }
167    }
168
169    /// Degree of an `Arc`-shared child, looked up by pointer identity so each
170    /// distinct node is evaluated at most once.
171    fn child_degree(
172        node: &Arc<Self>,
173        leaf_degree: &impl Fn(&A) -> usize,
174        cache: &mut BTreeMap<*const Self, usize>,
175    ) -> usize {
176        let key = Arc::as_ptr(node);
177        if let Some(&degree) = cache.get(&key) {
178            return degree;
179        }
180        let degree = node.degree_memo(leaf_degree, cache);
181        cache.insert(key, degree);
182        degree
183    }
184
185    /// Try to view this expression as a base-field constant.
186    fn as_const(&self) -> Option<&A::F> {
187        match self {
188            Self::Leaf(a) => a.as_const(),
189            _ => None,
190        }
191    }
192
193    /// Addition with constant folding and zero-identity elimination.
194    fn sym_add(self, rhs: Self) -> Self {
195        if let (Some(&a), Some(&b)) = (self.as_const(), rhs.as_const()) {
196            return Self::Leaf(A::from_const(a + b));
197        }
198        if self.as_const().is_some_and(|c| c.is_zero()) {
199            return rhs;
200        }
201        if rhs.as_const().is_some_and(|c| c.is_zero()) {
202            return self;
203        }
204        let dm = self.degree_multiple().max(rhs.degree_multiple());
205        Self::Add {
206            x: Arc::new(self),
207            y: Arc::new(rhs),
208            degree_multiple: dm,
209        }
210    }
211
212    /// Subtraction with constant folding and zero-identity elimination.
213    fn sym_sub(self, rhs: Self) -> Self {
214        if let (Some(&a), Some(&b)) = (self.as_const(), rhs.as_const()) {
215            return Self::Leaf(A::from_const(a - b));
216        }
217        if self.as_const().is_some_and(|c| c.is_zero()) {
218            return rhs.sym_neg();
219        }
220        if rhs.as_const().is_some_and(|c| c.is_zero()) {
221            return self;
222        }
223        let dm = self.degree_multiple().max(rhs.degree_multiple());
224        Self::Sub {
225            x: Arc::new(self),
226            y: Arc::new(rhs),
227            degree_multiple: dm,
228        }
229    }
230
231    /// Negation with constant folding.
232    fn sym_neg(self) -> Self {
233        if let Some(&c) = self.as_const() {
234            return Self::Leaf(A::from_const(-c));
235        }
236        let dm = self.degree_multiple();
237        Self::Neg {
238            x: Arc::new(self),
239            degree_multiple: dm,
240        }
241    }
242
243    /// Multiplication with constant folding, zero-annihilation, and one-identity.
244    fn sym_mul(self, rhs: Self) -> Self {
245        if let (Some(&a), Some(&b)) = (self.as_const(), rhs.as_const()) {
246            return Self::Leaf(A::from_const(a * b));
247        }
248        if self.as_const().is_some_and(|c| c.is_zero())
249            || rhs.as_const().is_some_and(|c| c.is_zero())
250        {
251            return Self::Leaf(A::from_const(A::F::ZERO));
252        }
253        if self.as_const().is_some_and(|c| c.is_one()) {
254            return rhs;
255        }
256        if rhs.as_const().is_some_and(|c| c.is_one()) {
257            return self;
258        }
259        let dm = self.degree_multiple() + rhs.degree_multiple();
260        Self::Mul {
261            x: Arc::new(self),
262            y: Arc::new(rhs),
263            degree_multiple: dm,
264        }
265    }
266}
267
268impl<A: SymLeaf> PrimeCharacteristicRing for SymbolicExpr<A> {
269    type PrimeSubfield = <A::F as PrimeCharacteristicRing>::PrimeSubfield;
270
271    const ZERO: Self = Self::Leaf(A::ZERO);
272    const ONE: Self = Self::Leaf(A::ONE);
273    const TWO: Self = Self::Leaf(A::TWO);
274    const NEG_ONE: Self = Self::Leaf(A::NEG_ONE);
275
276    #[inline]
277    fn from_prime_subfield(f: Self::PrimeSubfield) -> Self {
278        Self::Leaf(A::from_const(A::F::from_prime_subfield(f)))
279    }
280}
281
282impl<A: SymLeaf> Dup for SymbolicExpr<A> {
283    #[inline(always)]
284    fn dup(&self) -> Self {
285        self.clone()
286    }
287}
288
289impl<A: SymLeaf> Default for SymbolicExpr<A> {
290    fn default() -> Self {
291        Self::ZERO
292    }
293}
294
295impl<A: SymLeaf, T: Into<Self>> ops::Add<T> for SymbolicExpr<A> {
296    type Output = Self;
297    fn add(self, rhs: T) -> Self {
298        self.sym_add(rhs.into())
299    }
300}
301
302impl<A: SymLeaf, T: Into<Self>> ops::Sub<T> for SymbolicExpr<A> {
303    type Output = Self;
304    fn sub(self, rhs: T) -> Self {
305        self.sym_sub(rhs.into())
306    }
307}
308
309impl<A: SymLeaf> ops::Neg for SymbolicExpr<A> {
310    type Output = Self;
311    fn neg(self) -> Self {
312        self.sym_neg()
313    }
314}
315
316impl<A: SymLeaf, T: Into<Self>> ops::Mul<T> for SymbolicExpr<A> {
317    type Output = Self;
318    fn mul(self, rhs: T) -> Self {
319        self.sym_mul(rhs.into())
320    }
321}
322
323impl<A: SymLeaf, T: Into<Self>> ops::AddAssign<T> for SymbolicExpr<A> {
324    fn add_assign(&mut self, rhs: T) {
325        *self = self.clone() + rhs.into();
326    }
327}
328
329impl<A: SymLeaf, T: Into<Self>> ops::SubAssign<T> for SymbolicExpr<A> {
330    fn sub_assign(&mut self, rhs: T) {
331        *self = self.clone() - rhs.into();
332    }
333}
334
335impl<A: SymLeaf, T: Into<Self>> ops::MulAssign<T> for SymbolicExpr<A> {
336    fn mul_assign(&mut self, rhs: T) {
337        *self = self.clone() * rhs.into();
338    }
339}
340
341impl<A: SymLeaf, T: Into<Self>> Sum<T> for SymbolicExpr<A> {
342    fn sum<I: Iterator<Item = T>>(iter: I) -> Self {
343        iter.map(Into::into)
344            .reduce(|a, b| a + b)
345            .unwrap_or(Self::ZERO)
346    }
347}
348
349impl<A: SymLeaf, T: Into<Self>> Product<T> for SymbolicExpr<A> {
350    fn product<I: Iterator<Item = T>>(iter: I) -> Self {
351        iter.map(Into::into)
352            .reduce(|a, b| a * b)
353            .unwrap_or(Self::ONE)
354    }
355}
356
357impl<F: Field, T: Into<SymbolicExpression<F>>> ops::Add<T> for SymbolicVariable<F> {
358    type Output = SymbolicExpression<F>;
359    fn add(self, rhs: T) -> Self::Output {
360        Self::Output::from(self) + rhs.into()
361    }
362}
363
364impl<F: Field, T: Into<SymbolicExpression<F>>> ops::Sub<T> for SymbolicVariable<F> {
365    type Output = SymbolicExpression<F>;
366    fn sub(self, rhs: T) -> Self::Output {
367        Self::Output::from(self) - rhs.into()
368    }
369}
370
371impl<F: Field, T: Into<SymbolicExpression<F>>> ops::Mul<T> for SymbolicVariable<F> {
372    type Output = SymbolicExpression<F>;
373    fn mul(self, rhs: T) -> Self::Output {
374        Self::Output::from(self) * rhs.into()
375    }
376}
377
378impl<F: Field, EF: ExtensionField<F>, T: Into<SymbolicExpressionExt<F, EF>>> ops::Add<T>
379    for SymbolicVariableExt<F, EF>
380{
381    type Output = SymbolicExpressionExt<F, EF>;
382    fn add(self, rhs: T) -> Self::Output {
383        Self::Output::from(self) + rhs.into()
384    }
385}
386
387impl<F: Field, EF: ExtensionField<F>, T: Into<SymbolicExpressionExt<F, EF>>> ops::Sub<T>
388    for SymbolicVariableExt<F, EF>
389{
390    type Output = SymbolicExpressionExt<F, EF>;
391    fn sub(self, rhs: T) -> Self::Output {
392        Self::Output::from(self) - rhs.into()
393    }
394}
395
396impl<F: Field, EF: ExtensionField<F>, T: Into<SymbolicExpressionExt<F, EF>>> ops::Mul<T>
397    for SymbolicVariableExt<F, EF>
398{
399    type Output = SymbolicExpressionExt<F, EF>;
400    fn mul(self, rhs: T) -> Self::Output {
401        Self::Output::from(self) * rhs.into()
402    }
403}
404
405#[cfg(test)]
406mod tests {
407    use p3_baby_bear::BabyBear;
408    use p3_field::extension::BinomialExtensionField;
409
410    use super::*;
411    use crate::symbolic::expression::BaseLeaf;
412    use crate::symbolic::expression_ext::ExtLeaf;
413    use crate::symbolic::variable::{BaseEntry, ExtEntry};
414
415    type F = BabyBear;
416    type EF = BinomialExtensionField<BabyBear, 4>;
417
418    #[test]
419    fn symbolic_variable_add_produces_add_node() {
420        // Adding a variable and a non-zero constant creates an addition node.
421        let var = SymbolicVariable::<F>::new(BaseEntry::Main { offset: 0 }, 0);
422        let expr = SymbolicExpression::from(F::new(5));
423        let result = var + expr;
424        match result {
425            SymbolicExpr::Add {
426                x,
427                y,
428                degree_multiple,
429            } => {
430                assert_eq!(degree_multiple, 1);
431                assert!(matches!(
432                    x.as_ref(),
433                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
434                        if v.index == 0 && v.entry == BaseEntry::Main { offset: 0 }
435                ));
436                assert!(matches!(
437                    y.as_ref(),
438                    SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if *c == F::new(5)
439                ));
440            }
441            _ => panic!("Expected an Add node"),
442        }
443    }
444
445    #[test]
446    fn symbolic_variable_sub_produces_sub_node() {
447        // Subtracting two variables creates a subtraction node.
448        let var = SymbolicVariable::<F>::new(BaseEntry::Main { offset: 0 }, 0);
449        let other = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::new(
450            BaseEntry::Main { offset: 0 },
451            1,
452        )));
453        let result = var - other;
454        match result {
455            SymbolicExpr::Sub {
456                x,
457                y,
458                degree_multiple,
459            } => {
460                assert_eq!(degree_multiple, 1);
461                assert!(matches!(
462                    x.as_ref(),
463                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
464                        if v.index == 0 && v.entry == BaseEntry::Main { offset: 0 }
465                ));
466                assert!(matches!(
467                    y.as_ref(),
468                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
469                        if v.index == 1 && v.entry == BaseEntry::Main { offset: 0 }
470                ));
471            }
472            _ => panic!("Expected a Sub node"),
473        }
474    }
475
476    #[test]
477    fn symbolic_variable_mul_produces_mul_node() {
478        // Multiplying two variables creates a multiplication node with summed degree.
479        let var = SymbolicVariable::<F>::new(BaseEntry::Main { offset: 0 }, 0);
480        let other = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::new(
481            BaseEntry::Main { offset: 0 },
482            1,
483        )));
484        let result = var * other;
485        match result {
486            SymbolicExpr::Mul {
487                x,
488                y,
489                degree_multiple,
490            } => {
491                assert_eq!(degree_multiple, 2);
492                assert!(matches!(
493                    x.as_ref(),
494                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
495                        if v.index == 0 && v.entry == BaseEntry::Main { offset: 0 }
496                ));
497                assert!(matches!(
498                    y.as_ref(),
499                    SymbolicExpr::Leaf(BaseLeaf::Variable(v))
500                        if v.index == 1 && v.entry == BaseEntry::Main { offset: 0 }
501                ));
502            }
503            _ => panic!("Expected a Mul node"),
504        }
505    }
506
507    #[test]
508    fn symbolic_variable_ext_add_produces_add_node() {
509        // Adding an extension variable and a non-zero constant creates an addition node.
510        let var = SymbolicVariableExt::<F, EF>::new(ExtEntry::Permutation { offset: 0 }, 0);
511        let expr = SymbolicExpressionExt::<F, EF>::from(F::new(3));
512        let result = var + expr;
513        match result {
514            SymbolicExpr::Add {
515                x,
516                y,
517                degree_multiple,
518            } => {
519                assert_eq!(degree_multiple, 1);
520                assert!(matches!(
521                    x.as_ref(),
522                    SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
523                        if v.index == 0 && v.entry == ExtEntry::Permutation { offset: 0 }
524                ));
525                assert!(matches!(
526                    y.as_ref(),
527                    SymbolicExpr::Leaf(ExtLeaf::Base(SymbolicExpr::Leaf(BaseLeaf::Constant(c))))
528                        if *c == F::new(3)
529                ));
530            }
531            _ => panic!("Expected an Add node"),
532        }
533    }
534
535    #[test]
536    fn symbolic_variable_ext_sub_produces_sub_node() {
537        // Subtracting two extension variables creates a subtraction node.
538        let var = SymbolicVariableExt::<F, EF>::new(ExtEntry::Permutation { offset: 0 }, 0);
539        let other = SymbolicExpressionExt::<F, EF>::from(SymbolicVariableExt::<F, EF>::new(
540            ExtEntry::Permutation { offset: 0 },
541            1,
542        ));
543        let result = var - other;
544        match result {
545            SymbolicExpr::Sub {
546                x,
547                y,
548                degree_multiple,
549            } => {
550                assert_eq!(degree_multiple, 1);
551                assert!(matches!(
552                    x.as_ref(),
553                    SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
554                        if v.index == 0 && v.entry == ExtEntry::Permutation { offset: 0 }
555                ));
556                assert!(matches!(
557                    y.as_ref(),
558                    SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
559                        if v.index == 1 && v.entry == ExtEntry::Permutation { offset: 0 }
560                ));
561            }
562            _ => panic!("Expected a Sub node"),
563        }
564    }
565
566    #[test]
567    fn symbolic_variable_ext_mul_produces_mul_node() {
568        // Multiplying two extension variables creates a multiplication node with summed degree.
569        let var = SymbolicVariableExt::<F, EF>::new(ExtEntry::Permutation { offset: 0 }, 0);
570        let other = SymbolicExpressionExt::<F, EF>::from(SymbolicVariableExt::<F, EF>::new(
571            ExtEntry::Permutation { offset: 0 },
572            1,
573        ));
574        let result = var * other;
575        match result {
576            SymbolicExpr::Mul {
577                x,
578                y,
579                degree_multiple,
580            } => {
581                assert_eq!(degree_multiple, 2);
582                assert!(matches!(
583                    x.as_ref(),
584                    SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
585                        if v.index == 0 && v.entry == ExtEntry::Permutation { offset: 0 }
586                ));
587                assert!(matches!(
588                    y.as_ref(),
589                    SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
590                        if v.index == 1 && v.entry == ExtEntry::Permutation { offset: 0 }
591                ));
592            }
593            _ => panic!("Expected a Mul node"),
594        }
595    }
596}