1use std::{iter::repeat_with, sync::Arc};
15
16use binius_compute::Allocator;
17use binius_field::{BinaryField, PackedField};
18use binius_iop::channel::OracleSpec;
19use binius_iop_prover::{
20 basefold::channel::{BaseFoldOracle, BaseFoldProverChannel},
21 channel::IOPProverChannel,
22 merkle_channel::MerkleIPProverChannel,
23};
24use binius_ip_prover::channel::{IPProverChannel, WordIPProverChannel};
25use binius_math::{FieldSlice, FieldVec, StructuredBuffer, ntt::AdditiveNTT};
26use binius_spartan_frontend::constraint_system::WitnessLayout;
27use binius_spartan_verifier::IOPVerifier;
28use rand::CryptoRng;
29
30use crate::{Error, IOPProver, pack_and_blind_witness, wrapper::ReplayChannel};
31
32pub struct ZKWrappedProverChannel<'a, P, NTT, Channel, ReplayFn, A>
44where
45 P: PackedField<Scalar: BinaryField>,
46 NTT: AdditiveNTT<Field = P::Scalar> + Sync,
47 Channel: MerkleIPProverChannel<P::Scalar>,
48 A: Allocator,
49{
50 inner_channel: BaseFoldProverChannel<'a, P::Scalar, P, NTT, Channel, A>,
51 outer_prover: &'a IOPProver<P::Scalar>,
52 alloc: &'a A,
55 outer_layout: Arc<WitnessLayout<P::Scalar>>,
56 replay_fn: ReplayFn,
57 keys: Vec<P::Scalar>,
58 next_key_idx: usize,
59 interaction: Vec<P::Scalar>,
60 precommit_oracle: BaseFoldOracle,
65 precommit_packed: FieldVec<P, A>,
66 n_outer_suffix_oracles: usize,
69}
70
71impl<'a, F, P, NTT, Channel, ReplayFn, A> ZKWrappedProverChannel<'a, P, NTT, Channel, ReplayFn, A>
72where
73 F: BinaryField,
74 P: PackedField<Scalar = F>,
75 NTT: AdditiveNTT<Field = F> + Sync,
76 Channel: MerkleIPProverChannel<F>,
77 A: Allocator,
78{
79 pub fn new(
100 mut inner_channel: BaseFoldProverChannel<'a, F, P, NTT, Channel, A>,
101 outer_prover: &'a IOPProver<F>,
102 outer_layout: Arc<WitnessLayout<F>>,
103 alloc: &'a A,
104 rng: impl CryptoRng,
105 replay_fn: ReplayFn,
106 ) -> Self {
107 let outer_oracle_specs =
108 IOPVerifier::new(outer_prover.constraint_system().clone()).oracle_specs();
109 let all_specs = inner_channel.remaining_oracle_specs();
110 let n_outer = outer_oracle_specs.len();
111 assert!(
112 n_outer >= 1 && all_specs.len() >= n_outer,
113 "outer oracle specs ({n_outer}) exceed channel oracle specs ({}) or are empty",
114 all_specs.len(),
115 );
116 assert_eq!(
117 all_specs[0], outer_oracle_specs[0],
118 "outer precommit oracle spec must be the first spec on the channel",
119 );
120 let suffix_len = n_outer - 1;
121 assert_eq!(
122 &all_specs[all_specs.len() - suffix_len..],
123 &outer_oracle_specs[1..],
124 "outer private/mask oracle specs must be the final suffix of channel specs",
125 );
126
127 let (keys, precommit_oracle, precommit_packed) = {
128 let _scope = tracing::debug_span!("Commit Transcript Mask").entered();
129 Self::commit_transcript_mask(&mut inner_channel, outer_prover, alloc, rng)
130 };
131
132 Self {
133 inner_channel,
134 outer_prover,
135 alloc,
136 outer_layout,
137 replay_fn,
138 keys,
139 next_key_idx: 0,
140 interaction: Vec::new(),
141 precommit_oracle,
142 precommit_packed,
143 n_outer_suffix_oracles: suffix_len,
144 }
145 }
146
147 fn commit_transcript_mask(
152 inner_channel: &mut BaseFoldProverChannel<'a, F, P, NTT, Channel, A>,
153 outer_prover: &IOPProver<F>,
154 alloc: &A,
155 mut rng: impl CryptoRng,
156 ) -> (Vec<F>, BaseFoldOracle, FieldVec<P, A>) {
157 let cs = outer_prover.constraint_system();
158 let keys = repeat_with(|| F::random(&mut rng))
159 .take(cs.n_precommit() as usize)
160 .collect::<Vec<F>>();
161 let precommit_blinding = *cs.blinding_info();
162 let precommit_packed = pack_and_blind_witness::<_, _, P>(
163 alloc,
164 cs.log_precommit() as usize,
165 &keys,
166 cs.n_precommit() as usize,
167 &precommit_blinding,
168 &mut rng,
169 );
170 let precommit_oracle = inner_channel.send_oracle(precommit_packed.as_view());
171 (keys, precommit_oracle, precommit_packed)
172 }
173
174 fn next_key(&mut self) -> F {
175 let key = self.keys[self.next_key_idx];
176 self.next_key_idx += 1;
177 key
178 }
179
180 pub fn finish(self, rng: impl CryptoRng) -> Result<(), Error>
188 where
189 ReplayFn: FnOnce(&mut ReplayChannel<F>),
190 {
191 let Self {
192 mut inner_channel,
193 outer_prover,
194 alloc,
195 outer_layout,
196 replay_fn,
197 keys,
198 interaction,
199 precommit_oracle,
200 precommit_packed,
201 ..
202 } = self;
203
204 let witness = {
206 let _scope = tracing::debug_span!("Generating ZK wrapper witness").entered();
207 let mut replay_channel = ReplayChannel::new(outer_layout, keys, interaction);
208 replay_fn(&mut replay_channel);
209 replay_channel
210 .finish()
211 .expect("outer witness generation should not fail")
212 };
213
214 outer_prover.prove::<P, _, _>(
216 &witness,
217 precommit_oracle,
218 precommit_packed,
219 rng,
220 &mut inner_channel,
221 alloc,
222 )?;
223 inner_channel.finish();
226 Ok(())
227 }
228}
229
230impl<F, P, NTT, Channel, ReplayFn, A> IPProverChannel<F>
231 for ZKWrappedProverChannel<'_, P, NTT, Channel, ReplayFn, A>
232where
233 F: BinaryField,
234 P: PackedField<Scalar = F>,
235 NTT: AdditiveNTT<Field = F> + Sync,
236 Channel: MerkleIPProverChannel<F>,
237 A: Allocator,
238{
239 fn send_one(&mut self, elem: F) {
240 let key = self.next_key();
241 let encrypted = elem + key;
245 self.inner_channel.send_one(encrypted);
246 self.interaction.push(encrypted);
247 }
248
249 fn send_public_claim(&mut self, elem: F) {
250 self.inner_channel.send_one(elem);
254 self.interaction.push(elem);
255 }
256
257 fn observe_one(&mut self, val: F) {
258 self.inner_channel.observe_one(val);
259 self.interaction.push(val);
260 }
261
262 fn sample(&mut self) -> F {
263 let val = self.inner_channel.sample();
264 self.interaction.push(val);
265 val
266 }
267}
268
269impl<F, P, NTT, Channel, ReplayFn, A> WordIPProverChannel<F>
270 for ZKWrappedProverChannel<'_, P, NTT, Channel, ReplayFn, A>
271where
272 F: BinaryField,
273 P: PackedField<Scalar = F>,
274 NTT: AdditiveNTT<Field = F> + Sync,
275 Channel: MerkleIPProverChannel<F>,
276 A: Allocator,
277{
278 type Word = Channel::Word;
279
280 fn observe_words(&mut self, words: &[Self::Word]) {
281 self.inner_channel.observe_words(words);
284 }
285
286 fn sample_bits(&mut self, bits: usize) -> Self::Word {
287 self.inner_channel.sample_bits(bits)
288 }
289}
290
291impl<F, P, NTT, Channel, ReplayFn, A> IOPProverChannel<P, A>
292 for ZKWrappedProverChannel<'_, P, NTT, Channel, ReplayFn, A>
293where
294 F: BinaryField,
295 P: PackedField<Scalar = F>,
296 NTT: AdditiveNTT<Field = F> + Sync,
297 Channel: MerkleIPProverChannel<F>,
298 A: Allocator,
299{
300 type Oracle = BaseFoldOracle;
301
302 fn remaining_oracle_specs(&self) -> &[OracleSpec] {
303 let remaining = self.inner_channel.remaining_oracle_specs();
304 let n_inner_remaining = remaining.len() - self.n_outer_suffix_oracles;
305 &remaining[..n_inner_remaining]
306 }
307
308 fn send_oracle(&mut self, buffer: FieldSlice<'_, P>) -> Self::Oracle {
309 assert!(
310 !self.remaining_oracle_specs().is_empty(),
311 "send_oracle called but no inner oracle specs remaining"
312 );
313 self.inner_channel.send_oracle(buffer)
314 }
315
316 fn prove_oracle_relation(
317 &mut self,
318 oracle: Self::Oracle,
319 transparent: StructuredBuffer<P, A::Vec<P>>,
320 claim: P::Scalar,
321 ) {
322 self.inner_channel.send_one(claim);
326 self.interaction.push(claim);
327
328 self.inner_channel
329 .prove_oracle_relation(oracle, transparent, claim);
330 }
331
332 fn finalize_oracle(&mut self, oracle: Self::Oracle, buffer: FieldVec<P, A>) {
333 self.inner_channel.finalize_oracle(oracle, buffer);
334 }
335}