Skip to main content

binius_prover/protocols/
rerand.rs

1// Copyright 2026 The Binius Developers
2
3//! Prover for the BitAnd sumcheck batched with the operand-column MLE-checks.
4//!
5//! See [`binius_verifier::protocols::rerand`] for the protocol.
6
7use 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
26/// One reduction's operand word columns with their per-bit claims.
27pub struct OperandWitness<'a, F> {
28	/// One column per operand, in the same order as `claims.columns`.
29	pub words: Vec<&'a [Word]>,
30	/// The per-bit claims on the columns, at the reduction's point.
31	pub claims: OperandClaims<'a, F>,
32}
33
34/// Proves the BitAnd sumcheck batched with the operand-column MLE-checks.
35///
36/// Every summand runs as a sumcheck in plain form, zero-padded on its high variables to the
37/// longest summand's variable count. The evaluations go out flat in summand order:
38/// `[a, b, c, sigma_1, ..., sigma_n]`.
39///
40/// # Arguments
41///
42/// * `bitand` - the BitAnd MLE-check prover, from the univariate skip's fold.
43/// * `bitand_claim` - its claim, the univariate-skip polynomial at `z`.
44/// * `lagrange` - the Lagrange weights at `z` on the 64-point domain.
45/// * `operands` - the operand columns of each multiplication reduction, with their claims.
46///
47/// The per-bit claims of `operands` must be in the transcript before `z` was drawn, matching
48/// [`binius_verifier::protocols::rerand::verify`].
49///
50/// # Preconditions
51///
52/// * each operand's column count equals its claim count
53/// * each column's length rounds up to `2^point.len()` for its reduction's point
54pub 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	// ponytail: ten columns in the sumcheck where two theta-combined tables would do (spec ยง4.7.1).
71	// Combine them only if this sumcheck shows in a profile.
72	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	/// A synthetic multiplication reduction: random word columns, and their honest per-bit claims
153	/// at a random point.
154	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			// Bit `i` of every word, as a multilinear, evaluated at the point.
168			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	/// Proves BitAnd over random columns of `2^log_and` rows, batched with `reductions`.
206	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		// `(log_and, [log_z per reduction])`: shorter, equal and longer than BitAnd, two reductions
246		// of different lengths, no reductions, and an empty BitAnd axis.
247		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			// IntMul's and BinMul's column counts, alternating.
257			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			// Each operand eval is its folded column at its point's prefix of the unified point.
265			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		// The last operand eval is the proof's final element.
291		*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}