openvm_circuit_primitives/range/
mod.rs1use core::mem::size_of;
8use std::{
9 borrow::{Borrow, BorrowMut},
10 sync::atomic::AtomicU32,
11};
12
13use openvm_circuit_primitives_derive::AlignedBorrow;
14use openvm_stark_backend::{
15 interaction::InteractionBuilder,
16 p3_air::{Air, BaseAir, PairBuilder},
17 p3_field::Field,
18 p3_matrix::{dense::RowMajorMatrix, Matrix},
19 BaseAirWithPublicValues, PartitionedBaseAir,
20};
21
22use crate::{ColumnsAir, StructReflection, StructReflectionHelper};
23
24mod bus;
25
26#[cfg(test)]
27pub mod tests;
28
29pub use bus::*;
30
31#[derive(Default, AlignedBorrow, StructReflection, Copy, Clone)]
32#[repr(C)]
33pub struct RangeCols<T> {
34 pub mult: T,
36}
37
38#[derive(Default, AlignedBorrow, StructReflection, Copy, Clone)]
39#[repr(C)]
40pub struct RangePreprocessedCols<T> {
41 pub counter: T,
43}
44
45pub const NUM_RANGE_COLS: usize = size_of::<RangeCols<u8>>();
46pub const NUM_RANGE_PREPROCESSED_COLS: usize = size_of::<RangePreprocessedCols<u8>>();
47
48#[derive(Clone, Copy, Debug, derive_new::new, ColumnsAir)]
49#[columns_via(RangeCols<u8>)]
50pub struct RangeCheckerAir {
51 pub bus: RangeCheckBus,
52}
53
54impl RangeCheckerAir {
55 pub fn range_max(&self) -> u32 {
56 self.bus.range_max
57 }
58}
59
60impl<F: Field> BaseAirWithPublicValues<F> for RangeCheckerAir {}
61impl<F: Field> PartitionedBaseAir<F> for RangeCheckerAir {}
62impl<F: Field> BaseAir<F> for RangeCheckerAir {
63 fn width(&self) -> usize {
64 NUM_RANGE_COLS
65 }
66
67 fn preprocessed_trace(&self) -> Option<RowMajorMatrix<F>> {
68 let column = (0..self.range_max()).map(F::from_u32).collect();
70 Some(RowMajorMatrix::new_col(column))
71 }
72}
73
74impl<AB: InteractionBuilder + PairBuilder> Air<AB> for RangeCheckerAir {
75 fn eval(&self, builder: &mut AB) {
76 let preprocessed = builder.preprocessed();
77 let prep_local = preprocessed
78 .row_slice(0)
79 .expect("window should have two elements");
80 let prep_local: &RangePreprocessedCols<AB::Var> = (*prep_local).borrow();
81 let main = builder.main();
82 let local = main.row_slice(0).expect("window should have two elements");
83 let local: &RangeCols<AB::Var> = (*local).borrow();
84 self.bus
86 .receive(prep_local.counter)
87 .eval(builder, local.mult);
88 }
89}
90
91pub struct RangeCheckerChip {
92 pub air: RangeCheckerAir,
93 count: Vec<AtomicU32>,
95}
96
97impl RangeCheckerChip {
98 pub fn new(bus: RangeCheckBus) -> Self {
99 let mut count = vec![];
100 for _ in 0..bus.range_max {
101 count.push(AtomicU32::new(0));
102 }
103
104 Self {
105 air: RangeCheckerAir::new(bus),
106 count,
107 }
108 }
109
110 pub fn bus(&self) -> RangeCheckBus {
111 self.air.bus
112 }
113
114 pub fn range_max(&self) -> u32 {
115 self.air.range_max()
116 }
117
118 pub fn add_count(&self, val: u32) {
119 let val_atomic = &self.count[val as usize];
121 val_atomic.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
122 }
123
124 pub fn generate_trace<F: Field>(&self) -> RowMajorMatrix<F> {
125 let mut rows = F::zero_vec(self.air.range_max() as usize * NUM_RANGE_COLS);
126 for (n, row) in rows.chunks_exact_mut(NUM_RANGE_COLS).enumerate() {
127 let cols: &mut RangeCols<F> = (*row).borrow_mut();
128 cols.mult = F::from_u32(self.count[n].swap(0, std::sync::atomic::Ordering::Relaxed));
130 }
131 RowMajorMatrix::new(rows, NUM_RANGE_COLS)
132 }
133}