use alloc::vec::Vec;
use core::fmt::{Display, Formatter};
use math::{FieldElement, StarkField};
use crate::air::Assertion;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ConstraintDivisor<B: StarkField> {
pub(super) numerator: Vec<(usize, B)>,
pub(super) exemptions: Vec<B>,
}
impl<B: StarkField> ConstraintDivisor<B> {
fn new(numerator: Vec<(usize, B)>, exemptions: Vec<B>) -> Self {
ConstraintDivisor { numerator, exemptions }
}
pub fn from_transition(
constraint_enforcement_domain_size: usize,
num_exemptions: usize,
) -> Self {
let exemptions = (constraint_enforcement_domain_size - num_exemptions
..constraint_enforcement_domain_size)
.map(|step| get_trace_domain_value_at::<B>(constraint_enforcement_domain_size, step))
.collect();
Self::new(vec![(constraint_enforcement_domain_size, B::ONE)], exemptions)
}
pub fn from_assertion<E>(assertion: &Assertion<E>, trace_length: usize) -> Self
where
E: FieldElement<BaseField = B>,
{
let num_steps = assertion.get_num_steps(trace_length);
if assertion.first_step == 0 {
Self::new(vec![(num_steps, B::ONE)], vec![])
} else {
let trace_offset = num_steps * assertion.first_step;
let offset = get_trace_domain_value_at::<B>(trace_length, trace_offset);
Self::new(vec![(num_steps, offset)], vec![])
}
}
pub fn numerator(&self) -> &[(usize, B)] {
&self.numerator
}
pub fn exemptions(&self) -> &[B] {
&self.exemptions
}
pub fn degree(&self) -> usize {
let numerator_degree = self.numerator.iter().fold(0, |degree, term| degree + term.0);
let denominator_degree = self.exemptions.len();
numerator_degree - denominator_degree
}
pub fn evaluate_at<E: FieldElement<BaseField = B>>(&self, x: E) -> E {
let mut numerator = E::ONE;
for (degree, constant) in self.numerator.iter() {
let v = x.exp((*degree as u32).into());
let v = v - E::from(*constant);
numerator *= v;
}
let denominator = self.evaluate_exemptions_at(x);
numerator / denominator
}
#[inline(always)]
pub fn evaluate_exemptions_at<E: FieldElement<BaseField = B>>(&self, x: E) -> E {
self.exemptions.iter().fold(E::ONE, |r, &e| r * (x - E::from(e)))
}
}
impl<B: StarkField> Display for ConstraintDivisor<B> {
fn fmt(&self, f: &mut Formatter) -> core::fmt::Result {
for (degree, offset) in self.numerator.iter() {
write!(f, "(x^{degree} - {offset})")?;
}
if !self.exemptions.is_empty() {
write!(f, " / ")?;
for x in self.exemptions.iter() {
write!(f, "(x - {x})")?;
}
}
Ok(())
}
}
fn get_trace_domain_value_at<B: StarkField>(trace_length: usize, step: usize) -> B {
debug_assert!(step < trace_length, "step must be in the trace domain [0, {trace_length})");
let g = B::get_root_of_unity(trace_length.ilog2());
g.exp((step as u64).into())
}
#[cfg(test)]
mod tests {
use math::{fields::f128::BaseElement, polynom};
use super::*;
#[test]
fn constraint_divisor_degree() {
let div = ConstraintDivisor::new(vec![(4, BaseElement::ONE)], vec![]);
assert_eq!(4, div.degree());
let div = ConstraintDivisor::new(
vec![(4, BaseElement::ONE), (2, BaseElement::new(2)), (3, BaseElement::new(3))],
vec![],
);
assert_eq!(9, div.degree());
let div = ConstraintDivisor::new(
vec![(4, BaseElement::ONE), (2, BaseElement::new(2)), (3, BaseElement::new(3))],
vec![BaseElement::ONE, BaseElement::new(2)],
);
assert_eq!(7, div.degree());
}
#[test]
fn constraint_divisor_evaluation() {
let div = ConstraintDivisor::new(vec![(4, BaseElement::ONE)], vec![]);
assert_eq!(BaseElement::new(15), div.evaluate_at(BaseElement::new(2)));
let div = ConstraintDivisor::new(
vec![(4, BaseElement::ONE), (2, BaseElement::new(2)), (3, BaseElement::new(3))],
vec![],
);
let expected = BaseElement::new(15) * BaseElement::new(2) * BaseElement::new(5);
assert_eq!(expected, div.evaluate_at(BaseElement::new(2)));
let div = ConstraintDivisor::new(
vec![(4, BaseElement::ONE), (2, BaseElement::new(2)), (3, BaseElement::new(3))],
vec![BaseElement::ONE, BaseElement::new(2)],
);
let expected = BaseElement::new(255) * BaseElement::new(14) * BaseElement::new(61)
/ BaseElement::new(6);
assert_eq!(expected, div.evaluate_at(BaseElement::new(4)));
}
#[test]
fn constraint_divisor_equivalence() {
let n = 8_usize;
let g = BaseElement::get_root_of_unity(n.trailing_zeros());
let k = 4_u32;
let j = n as u32 / k;
let assertion = Assertion::periodic(0, 0, j as usize, BaseElement::ONE);
let divisor = ConstraintDivisor::from_assertion(&assertion, n);
let poly = polynom::mul(
&polynom::mul(
&[-BaseElement::ONE, BaseElement::ONE],
&[-g.exp(j.into()), BaseElement::ONE],
),
&polynom::mul(
&[-g.exp((2 * j).into()), BaseElement::ONE],
&[-g.exp((3 * j).into()), BaseElement::ONE],
),
);
for i in 0..n {
let expected = polynom::eval(&poly, g.exp((i as u32).into()));
let actual = divisor.evaluate_at(g.exp((i as u32).into()));
assert_eq!(expected, actual);
if i.is_multiple_of(j as usize) {
assert_eq!(BaseElement::ZERO, actual);
}
}
let offset = 1_u32;
let assertion = Assertion::periodic(0, offset as usize, j as usize, BaseElement::ONE);
let divisor = ConstraintDivisor::from_assertion(&assertion, n);
assert_eq!(ConstraintDivisor::new(vec![(k as usize, g.exp(k.into()))], vec![]), divisor);
let poly = polynom::mul(
&polynom::mul(
&[-g.exp(offset.into()), BaseElement::ONE],
&[-g.exp((offset + j).into()), BaseElement::ONE],
),
&polynom::mul(
&[-g.exp((offset + 2 * j).into()), BaseElement::ONE],
&[-g.exp((offset + 3 * j).into()), BaseElement::ONE],
),
);
for i in 0..n {
let expected = polynom::eval(&poly, g.exp((i as u32).into()));
let actual = divisor.evaluate_at(g.exp((i as u32).into()));
assert_eq!(expected, actual);
if i % (j as usize) == offset as usize {
assert_eq!(BaseElement::ZERO, actual);
}
}
let offset = 3_u32;
let k = 2_u32;
let j = n as u32 / k;
let assertion = Assertion::periodic(0, offset as usize, j as usize, BaseElement::ONE);
let divisor = ConstraintDivisor::from_assertion(&assertion, n);
assert_eq!(
ConstraintDivisor::new(vec![(k as usize, g.exp((offset * k).into()))], vec![]),
divisor
);
let poly = polynom::mul(
&[-g.exp(offset.into()), BaseElement::ONE],
&[-g.exp((offset + j).into()), BaseElement::ONE],
);
for i in 0..n {
let expected = polynom::eval(&poly, g.exp((i as u32).into()));
let actual = divisor.evaluate_at(g.exp((i as u32).into()));
assert_eq!(expected, actual);
if i % (j as usize) == offset as usize {
assert_eq!(BaseElement::ZERO, actual);
}
}
}
}