use openvm_stark_backend::{
interaction::InteractionBuilder,
p3_air::AirBuilder,
p3_field::{Field, PrimeCharacteristicRing},
};
use crate::{
var_range::{VariableRangeCheckerBus, VariableRangeCheckerChip},
SubAir, TraceSubRowGenerator,
};
#[cfg(test)]
pub mod tests;
pub use super::assert_less_than::LessThanAuxCols;
#[repr(C)]
#[derive(Clone, Copy, Debug, Default)]
pub struct IsLessThanIo<T> {
pub x: T,
pub y: T,
pub out: T,
pub count: T,
}
impl<T> IsLessThanIo<T> {
pub fn new(x: impl Into<T>, y: impl Into<T>, out: impl Into<T>, count: impl Into<T>) -> Self {
Self {
x: x.into(),
y: y.into(),
out: out.into(),
count: count.into(),
}
}
}
#[derive(Copy, Clone, Debug)]
pub struct IsLtSubAir {
pub bus: VariableRangeCheckerBus,
pub max_bits: usize,
pub decomp_limbs: usize,
}
impl IsLtSubAir {
pub fn new(bus: VariableRangeCheckerBus, max_bits: usize) -> Self {
assert!(max_bits <= 29); let decomp_limbs = max_bits.div_ceil(bus.range_max_bits);
Self {
bus,
max_bits,
decomp_limbs,
}
}
pub fn range_max_bits(&self) -> usize {
self.bus.range_max_bits
}
pub fn when_transition(self) -> IsLtWhenTransitionAir {
IsLtWhenTransitionAir(self)
}
#[inline(always)]
pub(crate) fn eval_without_range_checks<AB: AirBuilder<Var: Copy>>(
&self,
builder: &mut AB,
y_minus_x: impl Into<AB::Expr>,
out: impl Into<AB::Expr>,
condition: impl Into<AB::Expr>,
lower_decomp: &[AB::Var],
) {
assert_eq!(lower_decomp.len(), self.decomp_limbs);
let intermed_val = y_minus_x.into() + AB::Expr::from_usize((1 << self.max_bits) - 1);
let lower = lower_decomp
.iter()
.enumerate()
.fold(AB::Expr::ZERO, |acc, (i, &val)| {
acc + val * AB::Expr::from_usize(1 << (i * self.range_max_bits()))
});
let out = out.into();
let check_val = lower + out.clone() * AB::Expr::from_usize(1 << self.max_bits);
builder.when(condition).assert_eq(intermed_val, check_val);
builder.assert_bool(out);
}
#[inline(always)]
pub(crate) fn eval_range_checks<AB: InteractionBuilder>(
&self,
builder: &mut AB,
lower_decomp: &[AB::Var],
count: impl Into<AB::Expr>,
) {
let count = count.into();
let mut bits_remaining = self.max_bits;
for limb in lower_decomp {
let range_bits = bits_remaining.min(self.bus.range_max_bits);
self.bus
.range_check(*limb, range_bits)
.eval(builder, count.clone());
bits_remaining = bits_remaining.saturating_sub(self.bus.range_max_bits);
}
}
}
impl<AB: InteractionBuilder> SubAir<AB> for IsLtSubAir {
type AirContext<'a>
= (IsLessThanIo<AB::Expr>, &'a [AB::Var])
where
AB::Expr: 'a,
AB::Var: 'a,
AB: 'a;
fn eval<'a>(
&'a self,
builder: &'a mut AB,
(io, lower_decomp): (IsLessThanIo<AB::Expr>, &'a [AB::Var]),
) where
AB::Var: 'a,
AB::Expr: 'a,
{
self.eval_range_checks(builder, lower_decomp, io.count.clone());
self.eval_without_range_checks(builder, io.y - io.x, io.out, io.count, lower_decomp);
}
}
#[derive(Clone, Copy, Debug)]
pub struct IsLtWhenTransitionAir(pub IsLtSubAir);
impl<AB: InteractionBuilder> SubAir<AB> for IsLtWhenTransitionAir {
type AirContext<'a>
= (IsLessThanIo<AB::Expr>, &'a [AB::Var])
where
AB::Expr: 'a,
AB::Var: 'a,
AB: 'a;
fn eval<'a>(
&'a self,
builder: &'a mut AB,
(io, lower_decomp): (IsLessThanIo<AB::Expr>, &'a [AB::Var]),
) where
AB::Var: 'a,
AB::Expr: 'a,
{
self.0
.eval_range_checks(builder, lower_decomp, io.count.clone());
self.0.eval_without_range_checks(
&mut builder.when_transition(),
io.y - io.x,
io.out,
io.count,
lower_decomp,
);
}
}
impl<F: Field> TraceSubRowGenerator<F> for IsLtSubAir {
type TraceContext<'a> = (&'a VariableRangeCheckerChip, u32, u32);
type ColsMut<'a> = (&'a mut [F], &'a mut F);
#[inline(always)]
fn generate_subrow<'a>(
&'a self,
(range_checker, x, y): (&'a VariableRangeCheckerChip, u32, u32),
(lower_decomp, out): (&'a mut [F], &'a mut F),
) {
debug_assert_eq!(lower_decomp.len(), self.decomp_limbs);
debug_assert!(
x < (1 << self.max_bits),
"{x} has more than {} bits",
self.max_bits
);
debug_assert!(
y < (1 << self.max_bits),
"{y} has more than {} bits",
self.max_bits
);
*out = F::from_bool(x < y);
let check_less_than = (1 << self.max_bits) + y - x - 1;
let lower_u32 = check_less_than & ((1 << self.max_bits) - 1);
range_checker.decompose(lower_u32, self.max_bits, lower_decomp);
}
}