use alloc::{
collections::{BTreeMap, BTreeSet},
vec::Vec,
};
use math::{ExtensionOf, FieldElement};
use super::{AirContext, Assertion, ConstraintDivisor};
mod constraint;
pub use constraint::BoundaryConstraint;
mod constraint_group;
pub use constraint_group::BoundaryConstraintGroup;
#[cfg(test)]
mod tests;
pub struct BoundaryConstraints<E: FieldElement> {
main_constraints: Vec<BoundaryConstraintGroup<E::BaseField, E>>,
aux_constraints: Vec<BoundaryConstraintGroup<E, E>>,
}
impl<E: FieldElement> BoundaryConstraints<E> {
pub fn new(
context: &AirContext<E::BaseField>,
main_assertions: Vec<Assertion<E::BaseField>>,
aux_assertions: Vec<Assertion<E>>,
composition_coefficients: &[E],
) -> Self {
assert_eq!(
main_assertions.len(),
context.num_main_assertions,
"expected {} assertions against main trace segment, but received {}",
context.num_main_assertions,
main_assertions.len(),
);
assert_eq!(
aux_assertions.len(),
context.num_aux_assertions,
"expected {} assertions against the auxiliary trace segment, but received {}",
context.num_aux_assertions,
aux_assertions.len(),
);
assert_eq!(
context.num_assertions(),
composition_coefficients.len(),
"number of assertions must match the number of composition coefficient tuples"
);
let trace_length = context.trace_info.length();
let main_trace_width = context.trace_info.main_trace_width();
let aux_trace_width = context.trace_info.aux_segment_width();
let main_assertions = prepare_assertions(main_assertions, main_trace_width, trace_length);
let aux_assertions = prepare_assertions(aux_assertions, aux_trace_width, trace_length);
let inv_g = context.trace_domain_generator.inv();
let mut twiddle_map = BTreeMap::new();
let (main_composition_coefficients, aux_composition_coefficients) =
composition_coefficients.split_at(main_assertions.len());
let main_constraints = group_constraints(
main_assertions,
context,
main_composition_coefficients,
inv_g,
&mut twiddle_map,
);
let aux_constraints = group_constraints(
aux_assertions,
context,
aux_composition_coefficients,
inv_g,
&mut twiddle_map,
);
Self { main_constraints, aux_constraints }
}
pub fn main_constraints(&self) -> &[BoundaryConstraintGroup<E::BaseField, E>] {
&self.main_constraints
}
pub fn aux_constraints(&self) -> &[BoundaryConstraintGroup<E, E>] {
&self.aux_constraints
}
}
fn group_constraints<F, E>(
assertions: Vec<Assertion<F>>,
context: &AirContext<F::BaseField>,
composition_coefficients: &[E],
inv_g: F::BaseField,
twiddle_map: &mut BTreeMap<usize, Vec<F::BaseField>>,
) -> Vec<BoundaryConstraintGroup<F, E>>
where
F: FieldElement,
E: FieldElement<BaseField = F::BaseField> + ExtensionOf<F>,
{
let mut groups = BTreeMap::new();
for (assertion, &cc) in assertions.into_iter().zip(composition_coefficients) {
let key = (assertion.stride(), assertion.first_step());
let group = groups.entry(key).or_insert_with(|| {
BoundaryConstraintGroup::new(ConstraintDivisor::from_assertion(
&assertion,
context.trace_len(),
))
});
group.add(assertion, inv_g, twiddle_map, cc);
}
groups.into_iter().map(|e| e.1).collect::<Vec<_>>()
}
fn prepare_assertions<E: FieldElement>(
assertions: Vec<Assertion<E>>,
trace_width: usize,
trace_length: usize,
) -> Vec<Assertion<E>> {
let mut result = BTreeSet::<Assertion<E>>::new();
for assertion in assertions.into_iter() {
assertion.validate_trace_width(trace_width).unwrap_or_else(|err| {
panic!("assertion {assertion} is invalid: {err}");
});
assertion.validate_trace_length(trace_length).unwrap_or_else(|err| {
panic!("assertion {assertion} is invalid: {err}");
});
for a in result.iter().filter(|a| a.column == assertion.column) {
assert!(
!a.overlaps_with(&assertion),
"assertion {assertion} overlaps with assertion {a}"
);
}
result.insert(assertion);
}
result.into_iter().collect()
}