1#![warn(rustdoc::missing_crate_level_docs)]
29
30mod error;
31mod wiring;
32pub mod wrapper;
33
34use std::{
35 iter::{repeat_n, repeat_with},
36 marker::PhantomData,
37 ops::Deref,
38};
39
40use binius_compute::{Allocator, BufferPool, VecLike};
41use binius_field::{BinaryField, Field, PackedField};
42use binius_hash_prover::ParallelHashSuite;
43use binius_iop_prover::{basefold::compiler::BaseFoldProverCompiler, channel::IOPProverChannel};
44use binius_ip_prover::{
45 channel::IPProverChannel,
46 sumcheck::{quadratic_mlecheck_prover, zk_mlecheck},
47};
48use binius_math::{
49 FieldBuffer, FieldSlice, FieldVec,
50 inner_product::inner_product_buffers,
51 multilinear::eq::eq_ind_partial_eval,
52 ntt::{NeighborsLastMultiThread, domain_context::GaoMateerPreExpanded},
53 univariate::evaluate_univariate,
54};
55use binius_spartan_frontend::constraint_system::{
56 MulConstraint, Witness, WitnessIndex, WitnessSegment,
57};
58use binius_spartan_verifier::{
59 Verifier,
60 constraint_system::{BlindingInfo, ConstraintSystemPadded},
61 wiring::evaluate_wiring_mle_public,
62};
63use binius_transcript::{ProverTranscript, fiat_shamir::Challenger};
64use binius_utils::{SerializeBytes, checked_arithmetics::checked_log_2, rayon::prelude::*};
65use digest::Output;
66pub use error::*;
67use itertools::chain;
68use rand::CryptoRng;
69
70use crate::wiring::{WiringTranspose, fold_constraints};
71
72type ProverNTT<F> = NeighborsLastMultiThread<GaoMateerPreExpanded<F>>;
73
74#[derive(Debug)]
80pub struct IOPProver<F: Field> {
81 constraint_system: ConstraintSystemPadded<F>,
82 precommit_wiring_transpose: WiringTranspose,
83 private_wiring_transpose: WiringTranspose,
84}
85
86pub struct Prover<P, H>
92where
93 P: PackedField<Scalar: BinaryField>,
94 H: ParallelHashSuite,
95{
96 iop_prover: IOPProver<P::Scalar>,
97 basefold_compiler: BaseFoldProverCompiler<P, ProverNTT<P::Scalar>>,
98 pool: BufferPool,
101 _hash_marker: PhantomData<H>,
103}
104
105impl<F: Field> IOPProver<F> {
106 pub fn new(constraint_system: ConstraintSystemPadded<F>) -> Self {
108 let precommit_wiring_transpose = WiringTranspose::transpose(
109 WitnessSegment::Precommit,
110 constraint_system.precommit_size(),
111 constraint_system.mul_constraints(),
112 );
113 let private_wiring_transpose = WiringTranspose::transpose(
114 WitnessSegment::Private,
115 constraint_system.private_size(),
116 constraint_system.mul_constraints(),
117 );
118 Self {
119 constraint_system,
120 precommit_wiring_transpose,
121 private_wiring_transpose,
122 }
123 }
124
125 pub const fn constraint_system(&self) -> &ConstraintSystemPadded<F> {
126 &self.constraint_system
127 }
128
129 pub fn commit_precommit<P, Channel, A>(
136 &self,
137 witness: &Witness<F>,
138 rng: &mut impl CryptoRng,
139 channel: &mut Channel,
140 alloc: &A,
141 ) -> (Channel::Oracle, FieldVec<P, A>)
142 where
143 F: BinaryField,
144 P: PackedField<Scalar = F>,
145 Channel: IOPProverChannel<P, A>,
146 A: Allocator,
147 {
148 let cs = &self.constraint_system;
149 let precommit_blinding = *cs.blinding_info();
150 let precommit_packed = pack_and_blind_witness::<_, _, P>(
151 alloc,
152 cs.log_precommit() as usize,
153 witness.precommit(),
154 cs.n_precommit() as usize,
155 &precommit_blinding,
156 rng,
157 );
158 let precommit_oracle = channel.send_oracle(precommit_packed.as_view());
159 (precommit_oracle, precommit_packed)
160 }
161
162 pub fn prove<P, Channel, A>(
177 &self,
178 witness: &Witness<F>,
179 precommit_oracle: Channel::Oracle,
180 precommit_packed: FieldVec<P, A>,
181 mut rng: impl CryptoRng,
182 channel: &mut Channel,
183 alloc: &A,
184 ) -> Result<(), Error>
185 where
186 F: BinaryField,
187 P: PackedField<Scalar = F>,
188 Channel: IOPProverChannel<P, A>,
189 A: Allocator,
190 {
191 let _prove_guard =
192 tracing::info_span!("Prove", operation = "prove", perfetto_category = "operation")
193 .entered();
194
195 let cs = &self.constraint_system;
196
197 let expected_public_size = 1 << cs.log_public() as usize;
199 let expected_precommit_size = cs.precommit_size();
200 let expected_private_size = cs.private_size();
201 if witness.public().len() != expected_public_size {
202 return Err(Error::ArgumentError {
203 arg: "witness".to_string(),
204 msg: format!(
205 "public segment has {} elements, expected {}",
206 witness.public().len(),
207 expected_public_size
208 ),
209 });
210 }
211 if witness.precommit().len() != expected_precommit_size {
212 return Err(Error::ArgumentError {
213 arg: "witness".to_string(),
214 msg: format!(
215 "precommit segment has {} elements, expected {}",
216 witness.precommit().len(),
217 expected_precommit_size
218 ),
219 });
220 }
221 if witness.private().len() != expected_private_size {
222 return Err(Error::ArgumentError {
223 arg: "witness".to_string(),
224 msg: format!(
225 "private segment has {} elements, expected {}",
226 witness.private().len(),
227 expected_private_size
228 ),
229 });
230 }
231
232 let log_mul_constraints = checked_log_2(cs.mul_constraints().len());
233
234 let (m_n, m_d) = cs.mask_dims();
236 let mask_degree = 2; let log_masks_buffer_size = m_n + m_d;
238
239 let masks_buffer = {
240 let packed_len = 1 << log_masks_buffer_size.saturating_sub(P::LOG_WIDTH);
244 let mut values = alloc.alloc::<P>(packed_len);
245 values.extend(repeat_with(|| P::random(&mut rng)).take(packed_len));
246 FieldBuffer::new(log_masks_buffer_size, values)
247 };
248
249 let mulcheck_mask =
250 zk_mlecheck::Mask::new(log_mul_constraints, mask_degree, masks_buffer.as_view());
251
252 let blinding_info = cs.blinding_info();
254 let private_packed = pack_and_blind_witness::<_, _, P>(
255 alloc,
256 cs.log_private() as usize,
257 witness.private(),
258 cs.n_private() as usize,
259 blinding_info,
260 &mut rng,
261 );
262
263 let private_oracle = channel.send_oracle(private_packed.as_view());
266 let mask_oracle = channel.send_oracle(masks_buffer.as_view());
267
268 let (mulcheck_evals, mask_eval, r_x) = prove_mulcheck::<F, P, _, _>(
270 cs.mul_constraints(),
271 witness.public(),
272 precommit_packed.as_view(),
273 private_packed.as_view(),
274 mulcheck_mask,
275 &mut *channel,
276 alloc,
277 );
278
279 let lambda = channel.sample();
281
282 let batched_sum = evaluate_univariate(&mulcheck_evals, &lambda);
284
285 let r_x_tensor = eq_ind_partial_eval::<F>(&r_x);
287
288 let public_eval = evaluate_wiring_mle_public(
290 cs.mul_constraints(),
291 witness.public(),
292 &lambda,
293 r_x_tensor.as_ref(),
294 );
295
296 let precommit_wiring_poly =
299 fold_constraints(alloc, &self.precommit_wiring_transpose, lambda, r_x_tensor.as_ref());
300 let precommit_claim = inner_product_buffers(&precommit_packed, &precommit_wiring_poly);
301 channel.send_one(precommit_claim);
302
303 let private_claim = batched_sum - public_eval - precommit_claim;
304
305 let private_wiring_poly =
307 fold_constraints(alloc, &self.private_wiring_transpose, lambda, r_x_tensor.as_ref());
308
309 let n_vars = r_x.len();
311 let libra_eval_tensor =
312 zk_mlecheck::expand_libra_eval::<A, P>(alloc, &r_x, n_vars, mask_degree, m_n, m_d);
313
314 channel.prove_oracle_relation(
317 precommit_oracle.clone(),
318 precommit_wiring_poly.into(),
319 precommit_claim,
320 );
321 channel.finalize_oracle(precommit_oracle, precommit_packed);
322 channel.prove_oracle_relation(
323 private_oracle.clone(),
324 private_wiring_poly.into(),
325 private_claim,
326 );
327 channel.finalize_oracle(private_oracle, private_packed);
328 channel.prove_oracle_relation(mask_oracle.clone(), libra_eval_tensor.into(), mask_eval);
329 channel.finalize_oracle(mask_oracle, masks_buffer);
330
331 Ok(())
332 }
333}
334
335impl<F, P, H> Prover<P, H>
336where
337 F: BinaryField,
338 P: PackedField<Scalar = F>,
339 H: ParallelHashSuite,
340 Output<H::LeafHash>: SerializeBytes,
341{
342 pub fn setup(verifier: &Verifier<F, H>) -> Result<Self, Error> {
346 let log_num_shares = binius_utils::rayon::current_num_threads().ilog2() as usize;
347
348 let domain_context =
351 GaoMateerPreExpanded::generate(verifier.iop_compiler().max_log_domain_size());
352 let ntt = NeighborsLastMultiThread::new(domain_context, log_num_shares);
353
354 let basefold_compiler =
357 BaseFoldProverCompiler::from_verifier_compiler(verifier.iop_compiler(), ntt);
358
359 let iop_prover = IOPProver::new(verifier.constraint_system().clone());
360
361 Ok(Prover {
362 iop_prover,
363 basefold_compiler,
364 pool: BufferPool::new(),
365 _hash_marker: PhantomData,
366 })
367 }
368
369 pub const fn iop_prover(&self) -> &IOPProver<P::Scalar> {
371 &self.iop_prover
372 }
373
374 pub const fn iop_compiler(&self) -> &BaseFoldProverCompiler<P, ProverNTT<F>> {
376 &self.basefold_compiler
377 }
378
379 pub fn prove<Challenger_: Challenger>(
391 &self,
392 witness: &Witness<F>,
393 mut rng: impl CryptoRng,
394 transcript: &mut ProverTranscript<Challenger_>,
395 ) -> Result<(), Error> {
396 let public = witness.public();
398 transcript.observe().write_slice(public);
399
400 let alloc = &self.pool;
404 let mut channel = self
407 .basefold_compiler
408 .create_channel_from_transcript::<H, Challenger_, _, _>(transcript, &mut rng, alloc);
409 let (precommit_oracle, precommit_packed) =
410 self.iop_prover
411 .commit_precommit::<P, _, _>(witness, &mut rng, &mut channel, &alloc);
412 self.iop_prover.prove::<P, _, _>(
415 witness,
416 precommit_oracle,
417 precommit_packed,
418 rng,
419 &mut channel,
420 &alloc,
421 )?;
422 channel.finish();
423 Ok(())
424 }
425}
426
427fn prove_mulcheck<F, P, Channel, A>(
428 mul_constraints: &[MulConstraint<WitnessIndex>],
429 public: &[F],
430 precommit_packed: FieldSlice<'_, P>,
431 private_packed: FieldSlice<'_, P>,
432 mask: zk_mlecheck::Mask<P, impl Deref<Target = [P]>>,
433 channel: &mut Channel,
434 alloc: &A,
435) -> ([F; 3], F, Vec<F>)
436where
437 F: BinaryField,
438 P: PackedField<Scalar = F>,
439 Channel: IPProverChannel<F>,
440 A: Allocator,
441{
442 let mulcheck_witness = wiring::build_mulcheck_witness(
443 alloc,
444 mul_constraints,
445 public,
446 precommit_packed,
447 private_packed,
448 );
449
450 let r_mulcheck = channel.sample_many(mask.n_vars());
452
453 let mlecheck_prover = quadratic_mlecheck_prover(
455 alloc,
456 [mulcheck_witness.a, mulcheck_witness.b, mulcheck_witness.c],
457 |[a, b, c]| a * b - c, |[a, b, _c]| a * b, r_mulcheck,
460 F::ZERO, );
462
463 let mlecheck_output = zk_mlecheck::prove(mlecheck_prover, mask, channel);
465
466 let mut r_x = mlecheck_output.challenges;
468 r_x.reverse(); let [a_eval, b_eval, c_eval]: [F; 3] = mlecheck_output
471 .multilinear_evals
472 .try_into()
473 .expect("mlecheck returns 3 evaluations");
474
475 channel.send_many(&[a_eval, b_eval, c_eval]);
477
478 let mulcheck_evals = [a_eval, b_eval, c_eval];
479 let mask_eval = mlecheck_output.mask_eval;
480
481 (mulcheck_evals, mask_eval, r_x)
482}
483
484fn pack_and_blind_witness<A: Allocator, F: Field, P: PackedField<Scalar = F>>(
486 alloc: &A,
487 log_private: usize,
488 private: &[F],
489 n_private: usize,
490 blinding_info: &BlindingInfo,
491 mut rng: impl CryptoRng,
492) -> FieldVec<P, A> {
493 let packed_len = 1 << log_private.saturating_sub(P::LOG_WIDTH);
496 let mut packed = alloc.alloc::<P>(packed_len);
497 if log_private < P::LOG_WIDTH {
498 debug_assert_eq!(packed_len, 1);
500 let elems_iter = private.iter().copied();
501 let zeros_iter = repeat_n(F::ZERO, (1 << log_private) - private.len());
502
503 packed.push(P::from_scalars(chain!(elems_iter, zeros_iter)));
504 } else {
505 debug_assert!(private.len() <= 1 << log_private);
510 packed.resize(packed_len, P::zero());
511 private
512 .par_chunks(P::WIDTH)
513 .zip(packed.par_iter_mut())
514 .for_each(|(chunk, out)| *out = P::from_scalars(chunk.iter().copied()));
515 }
516
517 let mut buffer = FieldBuffer::new(log_private, packed);
518
519 for i in 0..blinding_info.n_dummy_wires {
522 buffer.set(n_private + i, F::random(&mut rng));
523 }
524
525 let constraint_wire_base = n_private + blinding_info.n_dummy_wires;
527 for i in 0..blinding_info.n_dummy_constraints {
528 let a = F::random(&mut rng);
529 let b = F::random(&mut rng);
530 let c = a * b;
531
532 buffer.set(constraint_wire_base + 3 * i, a);
533 buffer.set(constraint_wire_base + 3 * i + 1, b);
534 buffer.set(constraint_wire_base + 3 * i + 2, c);
535 }
536
537 buffer
538}