1use std::{marker::PhantomData, mem::MaybeUninit};
5
6use binius_compute::{Allocator, BufferPool, VecLike};
7use binius_core::{
8 constraint_system::{ConstraintSystem, InoutSegment, Operand, ValueVec},
9 word::Word,
10};
11use binius_field::{PackedField, Rijndael8b as B8};
12use binius_hash_prover::ParallelHashSuite;
13use binius_iop_prover::{basefold::compiler::BaseFoldProverCompiler, channel::IOPProverChannel};
14use binius_ip::sumcheck::SumcheckOutput;
15use binius_ip_prover::channel::WordIPProverChannel;
16use binius_math::{
17 BinarySubspace, FieldBuffer, FieldVec,
18 ntt::{NeighborsLastMultiThread, domain_context::GaoMateerPreExpanded},
19};
20use binius_transcript::{ProverTranscript, fiat_shamir::Challenger};
21use binius_utils::{
22 SerializeBytes,
23 rayon::{prelude::*, task_size::IndexedParallelIteratorExt},
24};
25use binius_verifier::{
26 IOPVerifier, Verifier,
27 config::{B128, LOG_WORDS_PER_ELEM},
28 protocols::bitand::AndCheckOutput,
29};
30use digest::Output;
31
32use super::error::Error;
33use crate::{
34 protocols::{
35 binmul, bitand, intmul,
36 rerand::OperandWitness,
37 shift::{self, KeyCollection, OperandClaims, ShiftOutput},
38 },
39 ring_switch,
40};
41
42type ProverNTT<F> = NeighborsLastMultiThread<GaoMateerPreExpanded<F>>;
44
45#[derive(Debug)]
51pub struct IOPProver {
52 constraint_system: ConstraintSystem,
53 log_witness_elems: usize,
54 key_collection: KeyCollection,
55}
56
57impl IOPProver {
58 pub fn new(iop_verifier: IOPVerifier, key_collection: KeyCollection) -> Self {
60 let log_witness_elems = iop_verifier.log_witness_elems();
61 let constraint_system = iop_verifier.into_constraint_system();
62 Self {
63 constraint_system,
64 log_witness_elems,
65 key_collection,
66 }
67 }
68
69 pub const fn constraint_system(&self) -> &ConstraintSystem {
71 &self.constraint_system
72 }
73
74 pub const fn key_collection(&self) -> &KeyCollection {
78 &self.key_collection
79 }
80
81 pub fn prove<A, P, Channel>(
86 &self,
87 witness: &ValueVec,
88 channel: &mut Channel,
89 alloc: &A,
90 ) -> Result<(), Error>
91 where
92 A: Allocator,
93 P: PackedField<Scalar = B128>,
94 Channel: IOPProverChannel<P, A> + WordIPProverChannel<B128, Word = Word>,
95 {
96 let cs = &self.constraint_system;
97
98 let setup_guard = tracing::debug_span!("Prepare witness").entered();
103 let witness_packed =
104 pack_witness::<P, _>(alloc, self.log_witness_elems, witness.non_public())?;
105 drop(setup_guard);
106
107 channel.observe_words(witness.inout());
110
111 let witness_commit_guard = tracing::info_span!("Commit witness").entered();
113
114 let trace_oracle = channel.send_oracle(witness_packed.as_view());
116
117 drop(witness_commit_guard);
118
119 let intmul = if cs.n_imul_constraints() > 0 {
127 let intmul_guard = tracing::info_span!(
128 "[phase] IntMul check",
129 n_constraints = cs.imul_constraints.len()
130 )
131 .entered();
132 let mul_columns = tracing::debug_span!("Assemble columns")
133 .in_scope(|| build_operation_columns(&cs.imul_constraints, witness, alloc));
134
135 let [a, b, lo, hi] = &mul_columns;
136 let intmul_output = intmul::prove::<_, _, P, _>([a, b, lo, hi], &mut *channel, alloc)?;
137 drop(intmul_guard);
138 Some((mul_columns, intmul_output))
139 } else {
140 None
141 };
142
143 let binmul = if cs.n_bmul_constraints() > 0 {
150 let binmul_guard = tracing::info_span!(
151 "[phase] BinMul check",
152 n_constraints = cs.bmul_constraints.len()
153 )
154 .entered();
155 let binmul_columns = tracing::debug_span!("Assemble columns")
156 .in_scope(|| build_operation_columns(&cs.bmul_constraints, witness, alloc));
157
158 let [a_lo, a_hi, b_lo, b_hi, c_lo, c_hi] = &binmul_columns;
159 let binmul_output = binmul::prove::<_, _, P, _>(
160 [a_lo, a_hi, b_lo, b_hi, c_lo, c_hi],
161 &mut *channel,
162 alloc,
163 );
164 drop(binmul_guard);
165 Some((binmul_columns, binmul_output))
166 } else {
167 None
168 };
169
170 let operands = [
174 intmul.as_ref().map(|(columns, output)| OperandWitness {
175 words: columns.iter().map(|column| &**column).collect(),
176 claims: output.operand_claims(),
177 }),
178 binmul.as_ref().map(|(columns, output)| OperandWitness {
179 words: columns.iter().map(|column| &**column).collect(),
180 claims: output.operand_claims(),
181 }),
182 ]
183 .into_iter()
184 .flatten()
185 .collect::<Vec<_>>();
186 let bitand_guard =
187 tracing::info_span!("[phase] BitAnd check", n_constraints = cs.and_constraints.len())
188 .entered();
189 let AndCheckOutput {
190 z_challenge,
191 rerand,
192 } = {
193 let bitand_columns = tracing::debug_span!("Assemble columns")
195 .in_scope(|| build_operation_columns(&cs.and_constraints, witness, alloc));
196 bitand::prove::<_, B128, P, _, _>(bitand_columns, &operands, &mut *channel, alloc)
197 };
198 drop(bitand_guard);
199
200 let claims = OperandClaims::from_rerand(cs, 0, z_challenge, &rerand, || channel.sample());
206
207 let subspace = BinarySubspace::<B8>::with_dim(Word::LOG_BITS).isomorphic();
210
211 let shift_guard = tracing::info_span!(
213 "[phase] Shift Reduction",
214 phase = "shift_reduction",
215 perfetto_category = "phase"
216 )
217 .entered();
218 let ShiftOutput {
219 sumcheck: SumcheckOutput {
220 challenges: eval_point,
221 eval: _,
222 },
223 wiring_eval,
224 } = shift::prove::<_, P, _, _>(
225 &self.key_collection,
226 witness.public(),
227 witness.non_public(),
228 claims,
229 &subspace,
230 &mut *channel,
231 alloc,
232 );
233 drop(shift_guard);
234
235 let witness_point = &eval_point[..eval_point.len() - 1];
239 let (r_j, r_y) = witness_point.split_at(Word::LOG_BITS);
240
241 ring_switch::prove_public_eval::<_, P, _>(alloc, witness.public(), r_j, r_y, &mut *channel);
244
245 channel.send_public_claim(wiring_eval);
248
249 let pcs_guard = tracing::info_span!(
251 "[phase] PCS Opening",
252 phase = "pcs_opening",
253 perfetto_category = "phase"
254 )
255 .entered();
256
257 let ring_switch::RingSwitchOutput {
260 rs_eq_ind,
261 sumcheck_claim,
262 } = ring_switch::prove(alloc, witness_packed.as_view(), witness_point, &mut *channel);
263
264 channel.prove_oracle_relation(trace_oracle.clone(), rs_eq_ind.into(), sumcheck_claim);
267 channel.finalize_oracle(trace_oracle, witness_packed);
268
269 drop(pcs_guard);
270
271 Ok(())
272 }
273}
274
275#[cfg(target_arch = "x86_64")]
282fn warn_on_software_field_arithmetic() {
283 use std::{arch::is_x86_feature_detected, sync::Once};
284
285 static ONCE: Once = Once::new();
286 ONCE.call_once(|| {
287 if !cfg!(target_feature = "pclmulqdq") && is_x86_feature_detected!("pclmulqdq") {
288 tracing::warn!(
289 "this CPU supports carryless multiply (PCLMULQDQ), but the build does not \
290 enable it, so field arithmetic will run in software; rebuild with \
291 `-C target-cpu=native` or `-C target-feature=+pclmulqdq`"
292 );
293 }
294 });
295}
296
297#[cfg(not(target_arch = "x86_64"))]
298const fn warn_on_software_field_arithmetic() {}
299
300pub struct Prover<P, H>
306where
307 P: PackedField<Scalar = B128>,
308 H: ParallelHashSuite,
309{
310 iop_prover: IOPProver,
312 iop_compiler: BaseFoldProverCompiler<P, ProverNTT<B128>>,
314 pool: BufferPool,
317 _hash_marker: PhantomData<H>,
319}
320
321impl<P, H> Prover<P, H>
322where
323 P: PackedField<Scalar = B128>,
324 H: ParallelHashSuite,
325 Output<H::LeafHash>: SerializeBytes,
326{
327 pub fn setup(verifier: Verifier<H>) -> Result<Self, Error> {
331 let key_collection =
332 KeyCollection::build(verifier.constraint_system(), InoutSegment::Public);
333 Self::setup_with_key_collection(verifier, key_collection)
334 }
335
336 pub fn setup_with_key_collection(
341 verifier: Verifier<H>,
342 key_collection: KeyCollection,
343 ) -> Result<Self, Error> {
344 warn_on_software_field_arithmetic();
345
346 let domain_context =
349 GaoMateerPreExpanded::generate(verifier.iop_compiler().max_log_domain_size());
350 let log_num_shares = binius_utils::rayon::current_num_threads().ilog2() as usize;
354 let ntt = NeighborsLastMultiThread::new(domain_context, log_num_shares);
355
356 let iop_compiler =
358 BaseFoldProverCompiler::from_verifier_compiler(verifier.iop_compiler(), ntt);
359
360 let iop_prover = IOPProver::new(verifier.into_iop_verifier(), key_collection);
361
362 Ok(Prover {
363 iop_prover,
364 iop_compiler,
365 pool: BufferPool::new(),
366 _hash_marker: PhantomData,
367 })
368 }
369
370 pub const fn iop_prover(&self) -> &IOPProver {
372 &self.iop_prover
373 }
374
375 pub const fn key_collection(&self) -> &KeyCollection {
379 self.iop_prover.key_collection()
380 }
381
382 pub fn prove<Challenger_: Challenger + Clone>(
383 &self,
384 witness: &ValueVec,
385 transcript: &mut ProverTranscript<Challenger_>,
386 ) -> Result<(), Error> {
387 let cs = self.iop_prover.constraint_system();
388
389 let _prove_guard = tracing::info_span!(
390 "Prove",
391 n_hidden_words = cs.n_hidden_words(InoutSegment::Public),
392 n_bitand = cs.and_constraints.len(),
393 n_intmul = cs.imul_constraints.len(),
394 )
395 .entered();
396
397 let alloc = &self.pool;
401
402 let mut channel = self
406 .iop_compiler
407 .create_channel_without_zk_from_transcript::<H, Challenger_, _, _>(transcript, alloc);
408 self.iop_prover
409 .prove::<_, P, _>(witness, &mut channel, &alloc)?;
410 channel.finish();
411 Ok(())
412 }
413}
414
415pub fn pack_witness<P: PackedField<Scalar = B128>, A: Allocator>(
434 alloc: &A,
435 log_witness_elems: usize,
436 witness: &[Word],
437) -> Result<FieldVec<P, A>, Error> {
438 let n_witness_elems = witness.len().div_ceil(1 << LOG_WORDS_PER_ELEM);
440 if n_witness_elems > 1 << log_witness_elems {
441 return Err(Error::ArgumentError {
442 arg: "witness".to_string(),
443 msg: "witness element count is incompatible with the constraint system".to_string(),
444 });
445 }
446
447 let len = 1 << log_witness_elems.saturating_sub(P::LOG_WIDTH);
448 let mut padded_witness_elems = alloc.alloc::<P>(len);
449
450 let (pairs, word_remaining) = witness.as_chunks::<2>();
453 let aligned_len = pairs.len() / P::WIDTH * P::WIDTH;
454 let (pairs_aligned, word_pair_remaining) = pairs.split_at(aligned_len);
455 let n_aligned_elems = aligned_len / P::WIDTH;
459 (
460 pairs_aligned.par_chunks(P::WIDTH),
461 padded_witness_elems.spare_capacity_mut()[..n_aligned_elems].par_iter_mut(),
462 )
463 .into_par_iter()
464 .with_min_task_bytes::<P>()
465 .for_each(|(word_pairs, out)| {
466 out.write(P::from_scalars(
467 word_pairs
468 .iter()
469 .map(|[w0, w1]| B128::new(((w1.0 as u128) << 64) | (w0.0 as u128))),
470 ));
471 });
472 unsafe { padded_witness_elems.set_len(n_aligned_elems) };
478
479 if !word_pair_remaining.is_empty() || !word_remaining.is_empty() {
484 let word_pairs = word_pair_remaining
485 .iter()
486 .copied()
487 .chain(word_remaining.iter().map(|&word| [word, Word::ZERO]));
488 padded_witness_elems.push(P::from_scalars(
489 word_pairs.map(|[w0, w1]| B128::new(((w1.0 as u128) << 64) | (w0.0 as u128))),
490 ));
491 }
492
493 padded_witness_elems.resize(len, P::default());
494
495 Ok(FieldBuffer::new(log_witness_elems, padded_witness_elems))
496}
497
498fn build_operation_columns<C, A, const ARITY: usize, const N_COLS: usize>(
518 constraints: &[C],
519 witness: &ValueVec,
520 alloc: &A,
521) -> [A::Vec<Word>; N_COLS]
522where
523 C: AsRef<[Operand; ARITY]> + Sync,
524 A: Allocator,
525{
526 const {
527 assert!(N_COLS <= ARITY, "N_COLS must not exceed the constraint arity");
528 }
529
530 let n_constraints = constraints.len();
531 let n_rows = n_constraints.max(1);
533 (0..N_COLS)
534 .into_par_iter()
535 .map(|op_idx| {
536 let mut column = alloc.alloc::<Word>(n_rows);
537 let rows = &mut column.spare_capacity_mut()[..n_rows];
540 if n_constraints == 0 {
542 rows.fill(MaybeUninit::new(Word::ZERO));
543 }
544 (constraints, &mut *rows)
545 .into_par_iter()
546 .for_each(|(constraint, out)| {
547 out.write(witness.eval_operand(&constraint.as_ref()[op_idx]));
548 });
549 unsafe { column.set_len(n_rows) };
553 column
554 })
555 .collect::<Vec<_>>()
556 .try_into()
557 .unwrap_or_else(|_| unreachable!("source iterator has N_COLS elements"))
558}
559
560#[cfg(test)]
561mod tests {
562 use binius_compute::GlobalAllocator;
563 use binius_core::constraint_system::{AndConstraint, ConstraintSystem};
564 use binius_field::{Field, PackedGhash2x128b};
565 use binius_frontend::CircuitBuilder;
566
567 use super::{B128, ValueVec, Word, build_operation_columns, pack_witness};
568
569 fn and_gate_system(n_gates: usize) -> (ConstraintSystem, ValueVec) {
574 let builder = CircuitBuilder::new();
575 let wires: Vec<_> = (0..n_gates)
576 .map(|_| {
577 let x = builder.add_witness();
578 let y = builder.add_witness();
579 builder.force_commit(builder.band(x, y));
580 (x, y)
581 })
582 .collect();
583 let circuit = builder.build();
584
585 let mut w = circuit.new_witness_filler();
586 for (i, &(x, y)) in wires.iter().enumerate() {
587 w[x] = Word(0x0123_4567_89AB_CDEF | (i as u64) << 32 | 1);
589 w[y] = Word(0xFEDC_BA98_7654_3210 | (i as u64) | 1);
590 }
591 circuit.populate_wire_witness(&mut w).unwrap();
592
593 let cs = circuit.constraint_system().clone();
594 cs.validate().unwrap();
595 (cs, w.into_value_vec())
596 }
597
598 fn and_gate_witness(n_gates: usize) -> (Vec<AndConstraint>, ValueVec) {
602 let (cs, witness) = and_gate_system(n_gates);
603 assert_eq!(cs.n_and_constraints(), n_gates);
604 (cs.and_constraints, witness)
605 }
606
607 #[test]
611 fn build_operation_columns_stops_at_the_last_constraint() {
612 let (constraints, witness) = and_gate_witness(3);
613 let columns = build_operation_columns::<AndConstraint, _, 3, 2>(
614 &constraints,
615 &witness,
616 &GlobalAllocator,
617 );
618
619 for column in &columns {
622 assert_eq!(column.len(), 3);
623 assert!(column.iter().all(|&word| word != Word::ZERO));
624 }
625 }
626
627 #[test]
631 fn build_operation_columns_gives_an_empty_set_one_zero_row() {
632 let (_, witness) = and_gate_witness(1);
633 let columns =
634 build_operation_columns::<AndConstraint, _, 3, 2>(&[], &witness, &GlobalAllocator);
635
636 for column in &columns {
637 assert_eq!(column.len(), 1);
638 assert_eq!(column[0], Word::ZERO);
639 }
640 }
641
642 fn expected_scalars(words: &[Word], n_elems: usize) -> Vec<B128> {
646 let mut scalars = vec![B128::ZERO; n_elems];
647 for (elem, pair) in scalars.iter_mut().zip(words.chunks(2)) {
648 let lo = pair[0].0 as u128;
649 let hi = pair.get(1).map_or(0, |w| w.0 as u128);
650 *elem = B128::new((hi << 64) | lo);
651 }
652 scalars
653 }
654
655 #[test]
660 fn test_pack_witness_unaligned_pair_count_with_remainder() {
661 type P = PackedGhash2x128b;
662 assert_eq!(P::WIDTH, 2, "this test is meaningful only when the packing width is 2");
663
664 let words: Vec<Word> = (1..=7u64).map(Word).collect();
665 let log_witness_elems = 3; let packed = pack_witness::<P, _>(&GlobalAllocator, log_witness_elems, &words).unwrap();
668 let got: Vec<B128> = packed.iter_scalars().collect();
669
670 assert_eq!(got, expected_scalars(&words, 1 << log_witness_elems));
671 }
672
673 #[test]
676 fn test_pack_witness_various_lengths() {
677 type P = PackedGhash2x128b;
678
679 for n_words in [1usize, 2, 3, 4, 5, 6, 7, 8, 9, 13, 17] {
680 let words: Vec<Word> = (0..n_words as u64).map(|i| Word(i + 100)).collect();
681 let n_elems = n_words.div_ceil(2);
682 let log_witness_elems = n_elems.max(P::WIDTH).next_power_of_two().ilog2() as usize;
684
685 let packed = pack_witness::<P, _>(&GlobalAllocator, log_witness_elems, &words).unwrap();
686 let got: Vec<B128> = packed.iter_scalars().collect();
687
688 assert_eq!(
689 got,
690 expected_scalars(&words, 1 << log_witness_elems),
691 "n_words = {n_words}"
692 );
693 }
694 }
695}