1use std::{marker::PhantomData, sync::Arc};
10
11use binius_compute::BufferPool;
12use binius_core::constraint_system::{ConstraintSystem, InoutSegment, ValueVec};
13use binius_field::{Ghash128b as B128, PackedField};
14use binius_hash_prover::ParallelHashSuite;
15use binius_iop_prover::basefold::compiler::BaseFoldProverCompiler;
16use binius_ip::channel::WordIPVerifierChannel;
17use binius_math::ntt::{NeighborsLastMultiThread, domain_context::GaoMateerPreExpanded};
18use binius_spartan_frontend::constraint_system::WitnessLayout;
19use binius_spartan_prover::wrapper::{ReplayChannel, ZKWrappedProverChannel};
20use binius_transcript::{ProverTranscript, fiat_shamir::Challenger};
21use binius_utils::{DeserializeBytes, SerializeBytes, serialization::SerializationError};
22use binius_verifier::{IOPVerifier, zk_config::ZKVerifier};
23use bytes::{Buf, BufMut};
24use digest::Output;
25use rand::CryptoRng;
26
27use crate::{IOPProver, protocols::shift::KeyCollection};
28
29type ProverNTT<F> = NeighborsLastMultiThread<GaoMateerPreExpanded<F>>;
30
31pub struct ZKProver<P, H>
36where
37 P: PackedField<Scalar = B128>,
38 H: ParallelHashSuite,
39{
40 inner_iop_prover: IOPProver,
41 inner_iop_verifier: IOPVerifier,
42 outer_iop_prover: binius_spartan_prover::IOPProver<B128>,
43 outer_layout: Arc<WitnessLayout<B128>>,
47 basefold_compiler: BaseFoldProverCompiler<P, ProverNTT<B128>>,
48 pool: BufferPool,
52 _hash_marker: PhantomData<H>,
54}
55
56impl<P, H> ZKProver<P, H>
57where
58 P: PackedField<Scalar = B128>,
59 H: ParallelHashSuite,
60 Output<H::LeafHash>: SerializeBytes,
61{
62 pub fn setup(zk_verifier: &ZKVerifier<H>) -> Result<Self, Error> {
64 let key_collection = {
65 let _guard = tracing::debug_span!("Build key collection").entered();
66 KeyCollection::build(
67 zk_verifier.inner_iop_verifier().constraint_system(),
68 InoutSegment::Public,
69 )
70 };
71 Self::setup_with_key_collection(zk_verifier, key_collection)
72 }
73
74 fn setup_with_key_collection(
78 zk_verifier: &ZKVerifier<H>,
79 key_collection: KeyCollection,
80 ) -> Result<Self, Error> {
81 let inner_iop_verifier = zk_verifier.inner_iop_verifier().clone();
83 let inner_iop_prover = IOPProver::new(inner_iop_verifier.clone(), key_collection);
84
85 let outer_cs = zk_verifier.outer_iop_verifier().constraint_system().clone();
94 let outer_layout = zk_verifier.outer_layout_arc();
95 let outer_iop_prover = binius_spartan_prover::IOPProver::new(outer_cs);
96
97 let log_domain_size = zk_verifier.basefold_compiler().max_log_domain_size();
99 let domain_context = {
100 let _guard = tracing::debug_span!("Precompute NTT domain").entered();
101 GaoMateerPreExpanded::generate(log_domain_size)
102 };
103 let log_num_shares = binius_utils::rayon::current_num_threads().ilog2() as usize;
104 let ntt = NeighborsLastMultiThread::new(domain_context, log_num_shares);
105 let basefold_compiler =
106 BaseFoldProverCompiler::from_verifier_compiler(zk_verifier.basefold_compiler(), ntt);
107
108 Ok(Self {
109 inner_iop_prover,
110 inner_iop_verifier,
111 outer_iop_prover,
112 outer_layout,
113 basefold_compiler,
114 pool: BufferPool::new(),
115 _hash_marker: PhantomData,
116 })
117 }
118
119 pub const fn inner_iop_prover(&self) -> &IOPProver {
121 &self.inner_iop_prover
122 }
123
124 pub const fn key_collection(&self) -> &crate::protocols::shift::KeyCollection {
126 self.inner_iop_prover.key_collection()
127 }
128
129 pub fn prove<Challenger_: Challenger>(
131 &self,
132 witness: &ValueVec,
133 mut rng: impl CryptoRng,
134 transcript: &mut ProverTranscript<Challenger_>,
135 ) -> Result<(), Error> {
136 let inout_words = witness.inout();
138
139 let alloc = &self.pool;
143
144 let basefold_channel = self
146 .basefold_compiler
147 .create_channel_from_transcript::<H, Challenger_, _, _>(transcript, &mut rng, alloc);
148 let mut wrapped_channel = ZKWrappedProverChannel::new(
149 basefold_channel,
150 &self.outer_iop_prover,
151 Arc::clone(&self.outer_layout),
152 &alloc,
153 &mut rng,
154 {
155 let inner_iop_verifier = &self.inner_iop_verifier;
156 move |replay_channel: &mut ReplayChannel<B128>| {
157 let inout = replay_channel.observe_words(inout_words);
160 let _ = inner_iop_verifier
164 .verify(&inout, replay_channel)
165 .expect("replay verification should not fail");
166 }
167 },
168 );
169
170 {
172 let inner_cs = self.inner_iop_prover.constraint_system();
173 let _scope = tracing::debug_span!(
174 "Binius64",
175 n_hidden_words = inner_cs.n_hidden_words(InoutSegment::Public),
176 n_bitand = inner_cs.and_constraints.len(),
177 n_intmul = inner_cs.imul_constraints.len(),
178 )
179 .entered();
180
181 self.inner_iop_prover
182 .prove::<_, P, _>(witness, &mut wrapped_channel, &alloc)?;
183 }
184
185 {
187 let outer_cs = self.outer_iop_prover.constraint_system();
188 let _scope = tracing::debug_span!(
189 "ZK Wrapper",
190 n_witness = outer_cs.n_private(),
191 n_constraints = outer_cs.mul_constraints().len(),
192 )
193 .entered();
194
195 wrapped_channel.finish(rng)?;
196 }
197
198 Ok(())
199 }
200
201 pub fn prove_sig<Challenger_: Challenger>(
206 &self,
207 witness: &ValueVec,
208 message: &[u8],
209 rng: impl CryptoRng,
210 transcript: &mut ProverTranscript<Challenger_>,
211 ) -> Result<(), Error> {
212 binius_verifier::signature::observe_message::<H, _>(&mut transcript.observe(), message);
213 self.prove(witness, rng, transcript)
214 }
215}
216
217impl<P, H> SerializeBytes for ZKProver<P, H>
221where
222 P: PackedField<Scalar = B128>,
223 H: ParallelHashSuite,
224 Output<H::LeafHash>: SerializeBytes + DeserializeBytes,
225{
226 fn serialize(&self, mut write_buf: impl BufMut) -> Result<(), SerializationError> {
227 const VERSION: u32 = 1;
228 VERSION.serialize(&mut write_buf)?;
229 self.inner_iop_verifier
230 .constraint_system()
231 .serialize(&mut write_buf)?;
232 self.basefold_compiler
233 .fri_params()
234 .rs_code()
235 .log_inv_rate()
236 .serialize(&mut write_buf)?;
237 self.inner_iop_prover.key_collection().serialize(write_buf)
238 }
239}
240
241impl<P, H> DeserializeBytes for ZKProver<P, H>
242where
243 P: PackedField<Scalar = B128>,
244 H: ParallelHashSuite,
245 Output<H::LeafHash>: SerializeBytes + DeserializeBytes,
246{
247 fn deserialize(mut read_buf: impl Buf) -> Result<Self, SerializationError> {
248 const VERSION: u32 = 1;
249 let version = u32::deserialize(&mut read_buf)?;
250 if version != VERSION {
251 return Err(SerializationError::InvalidConstruction {
252 name: "ZKProver::version",
253 });
254 }
255 let constraint_system = ConstraintSystem::deserialize(&mut read_buf)?;
256 let log_inv_rate = usize::deserialize(&mut read_buf)?;
257 let key_collection = KeyCollection::deserialize(&mut read_buf)?;
258 let zk_verifier = ZKVerifier::setup(constraint_system, log_inv_rate)
259 .map_err(|_| SerializationError::InvalidConstruction { name: "ZKProver" })?;
260 Self::setup_with_key_collection(&zk_verifier, key_collection)
261 .map_err(|_| SerializationError::InvalidConstruction { name: "ZKProver" })
262 }
263}
264
265#[derive(Debug, thiserror::Error)]
267pub enum Error {
268 #[error("inner proving error: {0}")]
269 InnerProving(#[from] crate::error::Error),
270 #[error("outer proving error: {0}")]
271 OuterProving(#[from] binius_spartan_prover::Error),
272}