1use std::iter;
8
9use binius_compute::Allocator;
10use binius_core::word::Word;
11use binius_field::{BinaryField, PackedField};
12use binius_ip_prover::{
13 channel::IPProverChannel,
14 sumcheck::{
15 MleToSumCheckDecorator, PaddedSumcheckDecorator, batch::batch_prove_and_write_evals,
16 common::MleCheckProver, mle_store::MleStore, multilinear_eval::MultilinearEvalEvaluator,
17 round_evaluator::SharedMleCheckProver,
18 },
19};
20use binius_math::inner_product::inner_product;
21use binius_verifier::protocols::rerand::{OperandClaims, RerandOutput, SUMCHECK_DEGREE};
22use either::Either;
23
24use crate::fold_word::BitAxisFolder;
25
26pub struct OperandWitness<'a, F> {
28 pub words: Vec<&'a [Word]>,
30 pub claims: OperandClaims<'a, F>,
32}
33
34pub fn prove<F, P, Channel, A>(
55 bitand: impl MleCheckProver<F>,
56 bitand_claim: F,
57 lagrange: &[F],
58 operands: &[OperandWitness<'_, F>],
59 channel: &mut Channel,
60 alloc: &A,
61) -> RerandOutput<F>
62where
63 F: BinaryField,
64 P: PackedField<Scalar = F>,
65 Channel: IPProverChannel<F>,
66 A: Allocator,
67{
68 let _scope = tracing::debug_span!("BitAnd batched sumcheck").entered();
69
70 let folder = BitAxisFolder::new(lagrange);
73 let operand_provers = operands.iter().map(|operand| {
74 let OperandClaims { point, columns } = &operand.claims;
75 assert_eq!(operand.words.len(), columns.len());
76
77 let mut store = MleStore::<A, P>::new(point.len(), alloc);
78 let evaluators = iter::zip(&operand.words, columns)
79 .map(|(words, per_bit_claims)| {
80 let col = store.push_owned(folder.fold::<P, _>(alloc, words));
81 let claim = inner_product(per_bit_claims.iter().copied(), lagrange.iter().copied());
82 (claim, MultilinearEvalEvaluator::new(col))
83 })
84 .collect::<Vec<_>>();
85 let claims = evaluators.iter().map(|&(claim, _)| claim).collect();
86 (SharedMleCheckProver::new(store, evaluators, point.to_vec()), claims)
87 });
88 let summands = iter::once((Either::Left(bitand), vec![bitand_claim]))
89 .chain(operand_provers.map(|(prover, claims)| (Either::Right(prover), claims)))
90 .collect::<Vec<_>>();
91
92 let n_vars = summands
93 .iter()
94 .map(|(prover, _)| prover.n_vars())
95 .max()
96 .expect("the BitAnd summand is always present");
97 let provers = summands
98 .into_iter()
99 .map(|(prover, claims)| {
100 let n_extra_vars = n_vars - prover.n_vars();
101 let prover = MleToSumCheckDecorator::new(prover);
102 PaddedSumcheckDecorator::new(prover, n_extra_vars, claims, SUMCHECK_DEGREE)
103 })
104 .collect();
105
106 let output = batch_prove_and_write_evals(provers, channel);
107
108 let mut evals = output.multilinear_evals.into_iter();
109 let bitand_evals = evals
110 .next()
111 .and_then(|evals| evals.try_into().ok())
112 .expect("the BitAnd summand reduces to its three columns");
113 let operand_evals = evals.flatten().collect();
114 let mut eval_point = output.challenges;
115 eval_point.reverse();
116 RerandOutput {
117 eval_point,
118 bitand_evals,
119 operand_evals,
120 }
121}
122
123#[cfg(test)]
124mod tests {
125 use std::{array, iter::repeat_with};
126
127 use binius_compute::GlobalAllocator;
128 use binius_field::{Field, Rijndael8b as B8, arch::OptimalPackedB128};
129 use binius_ip::channel::Error as ChannelError;
130 use binius_math::{
131 BinarySubspace, FieldBuffer,
132 multilinear::{eq::eq_ind_partial_eval_scalars, evaluate::evaluate},
133 test_utils::random_scalars,
134 univariate::EvaluationDomain,
135 };
136 use binius_transcript::{ProverTranscript, VerifierTranscript};
137 use binius_verifier::{
138 Error as VerifierError,
139 config::{B128, StdChallenger},
140 protocols::bitand::AndCheckOutput,
141 verify_bitand_reduction,
142 };
143 use rand::prelude::*;
144
145 use super::*;
146 use crate::protocols::bitand;
147
148 fn random_words(rng: &mut StdRng, n: usize) -> Vec<Word> {
149 repeat_with(|| Word(rng.random())).take(n).collect()
150 }
151
152 struct Reduction {
155 columns: Vec<Vec<Word>>,
156 point: Vec<B128>,
157 per_bit_claims: Vec<[B128; Word::BITS]>,
158 }
159
160 impl Reduction {
161 fn random(rng: &mut StdRng, n_vars: usize, n_columns: usize) -> Self {
162 let point = random_scalars::<B128>(&mut *rng, n_vars);
163 let eq = eq_ind_partial_eval_scalars(&point);
164 let columns = repeat_with(|| random_words(rng, 1 << n_vars))
165 .take(n_columns)
166 .collect::<Vec<_>>();
167 let per_bit_claims = columns
169 .iter()
170 .map(|words| {
171 array::from_fn(|bit| {
172 iter::zip(words, &eq)
173 .filter(|(word, _)| (word.0 >> bit) & 1 == 1)
174 .map(|(_, &eq)| eq)
175 .sum()
176 })
177 })
178 .collect();
179 Self {
180 columns,
181 point,
182 per_bit_claims,
183 }
184 }
185
186 fn claims(&self) -> OperandClaims<'_, B128> {
187 OperandClaims {
188 point: &self.point,
189 columns: self.per_bit_claims.iter().collect(),
190 }
191 }
192
193 fn witness(&self) -> OperandWitness<'_, B128> {
194 OperandWitness {
195 words: self.columns.iter().map(Vec::as_slice).collect(),
196 claims: self.claims(),
197 }
198 }
199 }
200
201 fn andcheck_domain() -> BinarySubspace<B128> {
202 BinarySubspace::<B8>::with_dim(Word::LOG_BITS + 1).isomorphic()
203 }
204
205 fn prove_random(
207 rng: &mut StdRng,
208 log_and: usize,
209 reductions: &[Reduction],
210 ) -> (AndCheckOutput<B128>, Vec<u8>) {
211 let columns = [
212 random_words(rng, 1 << log_and),
213 random_words(rng, 1 << log_and),
214 ];
215 let operands = reductions
216 .iter()
217 .map(Reduction::witness)
218 .collect::<Vec<_>>();
219 let mut transcript = ProverTranscript::new(StdChallenger::default());
220 let output = bitand::prove::<_, B128, OptimalPackedB128, _, _>(
221 columns,
222 &operands,
223 &mut transcript,
224 &GlobalAllocator,
225 );
226 (output, transcript.finalize())
227 }
228
229 fn verify_proof(
230 log_and: usize,
231 reductions: &[Reduction],
232 proof: Vec<u8>,
233 ) -> Result<AndCheckOutput<B128>, VerifierError> {
234 let claims = reductions.iter().map(Reduction::claims).collect::<Vec<_>>();
235 let mut transcript = VerifierTranscript::new(StdChallenger::default(), proof);
236 let output =
237 verify_bitand_reduction(log_and, &andcheck_domain(), &claims, &mut transcript)?;
238 transcript.finalize().expect("no trailing proof data");
239 Ok(output)
240 }
241
242 #[test]
243 fn prover_and_verifier_agree() {
244 let mut rng = StdRng::seed_from_u64(0);
245 let grid: [(usize, &[usize]); 6] = [
248 (5, &[3]),
249 (4, &[4]),
250 (3, &[5]),
251 (4, &[2, 6]),
252 (4, &[]),
253 (0, &[2]),
254 ];
255 for (log_and, log_zs) in grid {
256 let reductions = iter::zip(log_zs, [4, 6])
258 .map(|(&n_vars, n_columns)| Reduction::random(&mut rng, n_vars, n_columns))
259 .collect::<Vec<_>>();
260 let (output, proof) = prove_random(&mut rng, log_and, &reductions);
261 let verified = verify_proof(log_and, &reductions, proof).unwrap();
262 assert_eq!(output, verified, "at log_and = {log_and}, log_zs = {log_zs:?}");
263
264 let lagrange = andcheck_domain()
266 .reduce_dim(Word::LOG_BITS)
267 .lagrange_evals(&output.z_challenge);
268 let folder = &BitAxisFolder::new(&lagrange);
269 let rerand = output.rerand;
270 let expected = reductions
271 .iter()
272 .flat_map(|reduction| {
273 let point = &rerand.eval_point[..reduction.point.len()];
274 reduction.columns.iter().map(move |words| {
275 let folded: FieldBuffer<B128> = folder.fold(&GlobalAllocator, words);
276 evaluate(&folded, point)
277 })
278 })
279 .collect::<Vec<_>>();
280 assert_eq!(rerand.operand_evals, expected);
281 }
282 }
283
284 #[test]
285 fn mutated_operand_eval_is_rejected() {
286 let mut rng = StdRng::seed_from_u64(1);
287 let reductions = [Reduction::random(&mut rng, 5, 4)];
288 let (_, mut proof) = prove_random(&mut rng, 4, &reductions);
289
290 *proof.last_mut().unwrap() ^= 1;
292
293 let err = verify_proof(4, &reductions, proof).unwrap_err();
294 assert!(matches!(err, VerifierError::Channel(ChannelError::InvalidAssert)));
295 }
296
297 #[test]
298 fn mutated_per_bit_claim_is_rejected() {
299 let mut rng = StdRng::seed_from_u64(2);
300 let mut reductions = [Reduction::random(&mut rng, 3, 6)];
301 let (_, proof) = prove_random(&mut rng, 4, &reductions);
302
303 reductions[0].per_bit_claims[2][17] += B128::ONE;
304
305 let err = verify_proof(4, &reductions, proof).unwrap_err();
306 assert!(matches!(err, VerifierError::Channel(ChannelError::InvalidAssert)));
307 }
308}