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#[derive(Clone, Debug, Serialize, Deserialize)]
13pub enum BaseLeaf<F> {
14 Variable(SymbolicVariable<F>),
16
17 IsFirstRow,
19
20 IsLastRow,
22
23 IsTransition,
25
26 Constant(F),
28}
29
30pub 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 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 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 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 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 BaseEntry::Public => builder.public_values()[v.index].into(),
160 BaseEntry::Periodic => builder.periodic_values()[v.index].into(),
163 },
164 BaseLeaf::IsFirstRow => builder.is_first_row(),
166 BaseLeaf::IsLastRow => builder.is_last_row(),
167 BaseLeaf::IsTransition => builder.is_transition_window(2),
168 BaseLeaf::Constant(c) => AB::Expr::from(*c),
170 },
171 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
184impl<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 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 let constant = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(5)));
319 assert_eq!(constant.poly_degree(N, &[]), 0);
320
321 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 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 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 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 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 let expr: SymbolicExpression<BabyBear> = Default::default();
518
519 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 assert!(matches!(
530 SymbolicExpression::<BabyBear>::ZERO,
531 SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ZERO
532 ));
533 assert!(matches!(
535 SymbolicExpression::<BabyBear>::ONE,
536 SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::ONE
537 ));
538 assert!(matches!(
540 SymbolicExpression::<BabyBear>::TWO,
541 SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if c == BabyBear::TWO
542 ));
543 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 let var = SymbolicVariable::<BabyBear>::new(BaseEntry::Main { offset: 0 }, 3);
554 let expr: SymbolicExpression<BabyBear> = var.into();
556 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 let field_val = BabyBear::new(42);
570 let expr: SymbolicExpression<BabyBear> = field_val.into();
571 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 let prime_subfield_val = <BabyBear as PrimeCharacteristicRing>::PrimeSubfield::new(7);
582 let expr = SymbolicExpression::<BabyBear>::from_prime_subfield(prime_subfield_val);
583 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 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 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 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 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 let result = a - b;
631
632 match result {
634 SymbolicExpr::Sub {
635 x,
636 y,
637 degree_multiple,
638 } => {
639 assert_eq!(degree_multiple, 1);
641
642 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 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 let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
664 BaseEntry::Main { offset: 0 },
665 0,
666 )));
667
668 let result = -var;
670
671 match result {
673 SymbolicExpr::Neg { x, degree_multiple } => {
674 assert_eq!(degree_multiple, 1);
676
677 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 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 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 let constant = SymbolicExpression::Leaf(BaseLeaf::Constant(BabyBear::new(5)));
714
715 let var = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::<BabyBear>::new(
717 BaseEntry::Main { offset: 0 },
718 0,
719 )));
720
721 let result = constant + var;
723
724 match result {
725 SymbolicExpr::Add {
726 x,
727 y,
728 degree_multiple,
729 } => {
730 assert_eq!(degree_multiple, 1);
732
733 assert!(matches!(
735 x.as_ref(),
736 SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if *c == BabyBear::new(5)
737 ));
738
739 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 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 let ab = a * b;
768 assert_eq!(ab.degree_multiple(), 2);
769
770 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 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 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 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 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 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 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 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 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 let result = var.clone() + zero.clone();
885 assert_eq!(result.degree_multiple(), 1);
886
887 let result = var.clone() - zero.clone();
889 assert_eq!(result.degree_multiple(), 1);
890
891 let result = zero.clone() - var.clone();
893 assert_eq!(result.degree_multiple(), 1);
894
895 let result = var.clone() * one;
897 assert_eq!(result.degree_multiple(), 1);
898
899 let result = var * zero;
901 assert_eq!(result.degree_multiple(), 0);
902 }
903
904 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 fn test_builder() -> ResolveTestBuilder {
968 ResolveTestBuilder {
969 main: RowMajorMatrix::new(
970 vec![
971 BabyBear::new(10),
972 BabyBear::new(20), BabyBear::new(30),
974 BabyBear::new(40), ],
976 2, ),
978 public_values: vec![BabyBear::new(99)],
979 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 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 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 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 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 let add = col0.clone() + col1.clone();
1047 assert_eq!(add.resolve(&b), BabyBear::new(30));
1048
1049 let sub = col0.clone() - col1.clone();
1051 assert_eq!(sub.resolve(&b), BabyBear::new(10) - BabyBear::new(20));
1052
1053 let mul = col0.clone() * col1;
1055 assert_eq!(mul.resolve(&b), BabyBear::new(200));
1056
1057 let neg = -col0;
1059 assert_eq!(neg.resolve(&b), -BabyBear::new(10));
1060 }
1061
1062 #[test]
1063 fn resolve_periodic_columns() {
1064 let b = test_builder();
1073
1074 let p0 =
1076 SymbolicExpression::from(SymbolicVariable::<BabyBear>::new(BaseEntry::Periodic, 0));
1077 assert_eq!(p0.resolve(&b), BabyBear::new(7));
1078
1079 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 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 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 assert_eq!(decoded.resolve(&b), expr.resolve(&b));
1134 assert_eq!(serde_json::to_string(&decoded).unwrap(), json);
1136 assert_eq!(decoded.degree_multiple(), expr.degree_multiple());
1137 }
1138}