binius_iop_prover/basefold/
compiler.rs1use std::{borrow::BorrowMut, marker::PhantomData};
6
7use binius_compute::Allocator;
8use binius_field::{BinaryField, PackedField};
9use binius_hash_prover::ParallelHashSuite;
10use binius_iop::{
11 basefold::compiler::BaseFoldVerifierCompiler, channel::OracleSpec, fri::FRIParams,
12 merkle_tree::BinaryMerkleTreeScheme,
13};
14use binius_math::ntt::AdditiveNTT;
15use binius_transcript::{ProverTranscript, fiat_shamir::Challenger};
16use binius_utils::SerializeBytes;
17use digest::Output;
18use rand::{CryptoRng, SeedableRng, rngs::StdRng};
19
20use crate::{
21 basefold::channel::BaseFoldProverChannel,
22 merkle_channel::{MerkleIPProverChannel, ProverMerkleTranscriptChannel},
23 merkle_tree::prover::BinaryMerkleTreeProver,
24};
25
26pub type TranscriptBaseFoldProverChannel<'a, F, P, NTT, T, Challenger_, H, A> =
31 BaseFoldProverChannel<'a, F, P, NTT, ProverMerkleTranscriptChannel<T, Challenger_, F, H, A>, A>;
32
33#[derive(Debug)]
38pub struct BaseFoldProverCompiler<P, NTT>
39where
40 P: PackedField<Scalar: BinaryField>,
41 NTT: AdditiveNTT<Field = P::Scalar> + Sync,
42{
43 ntt: NTT,
44 oracle_specs: Vec<OracleSpec>,
45 fri_params: FRIParams<P::Scalar>,
47 _marker: PhantomData<P>,
48}
49
50impl<F, P, NTT> BaseFoldProverCompiler<P, NTT>
51where
52 F: BinaryField,
53 P: PackedField<Scalar = F>,
54 NTT: AdditiveNTT<Field = F> + Sync,
55{
56 pub fn new<H>(
63 ntt: NTT,
64 merkle_scheme: &BinaryMerkleTreeScheme<F, H>,
65 oracle_specs: Vec<OracleSpec>,
66 log_inv_rate: usize,
67 n_test_queries: usize,
68 ) -> Self
69 where
70 H: ParallelHashSuite,
71 {
72 assert!(
73 !oracle_specs.is_empty(),
74 "BaseFoldProverCompiler requires at least one oracle spec"
75 );
76
77 let (fri_params, _) = FRIParams::optimal_for_batch(
81 merkle_scheme,
82 &oracle_specs,
83 log_inv_rate,
84 n_test_queries,
85 );
86
87 Self {
88 ntt,
89 oracle_specs,
90 fri_params,
91 _marker: PhantomData,
92 }
93 }
94
95 pub fn from_verifier_compiler(
99 verifier_compiler: &BaseFoldVerifierCompiler<F>,
100 ntt: NTT,
101 ) -> Self {
102 Self {
103 ntt,
104 oracle_specs: verifier_compiler.oracle_specs().to_vec(),
105 fri_params: verifier_compiler.fri_params().clone(),
106 _marker: PhantomData,
107 }
108 }
109
110 pub const fn ntt(&self) -> &NTT {
112 &self.ntt
113 }
114
115 pub fn oracle_specs(&self) -> &[OracleSpec] {
117 &self.oracle_specs
118 }
119
120 pub const fn fri_params(&self) -> &FRIParams<F> {
122 &self.fri_params
123 }
124
125 pub fn create_channel<Channel, A>(
135 &self,
136 channel: Channel,
137 rng: impl CryptoRng,
138 alloc: A,
139 ) -> BaseFoldProverChannel<'_, F, P, NTT, Channel, A>
140 where
141 Channel: MerkleIPProverChannel<F>,
142 A: Allocator,
143 {
144 BaseFoldProverChannel::new(
145 channel,
146 &self.ntt,
147 self.oracle_specs.clone(),
148 self.fri_params.clone(),
149 rng,
150 alloc,
151 )
152 }
153
154 pub fn create_channel_without_zk<Channel, A>(
166 &self,
167 channel: Channel,
168 alloc: A,
169 ) -> BaseFoldProverChannel<'_, F, P, NTT, Channel, A>
170 where
171 Channel: MerkleIPProverChannel<F>,
172 A: Allocator,
173 {
174 assert!(
176 self.oracle_specs().iter().all(|spec| !spec.is_zk),
177 "create_channel_without_zk requires every oracle to be non-ZK"
178 );
179
180 self.create_channel(channel, StdRng::seed_from_u64(0), alloc)
182 }
183
184 pub fn create_channel_from_transcript<H, Challenger_, T, A>(
192 &self,
193 transcript: T,
194 rng: impl CryptoRng,
195 alloc: A,
196 ) -> TranscriptBaseFoldProverChannel<'_, F, P, NTT, T, Challenger_, H, A>
197 where
198 H: ParallelHashSuite,
199 Challenger_: Challenger,
200 T: BorrowMut<ProverTranscript<Challenger_>>,
201 Output<H::LeafHash>: SerializeBytes,
202 A: Allocator,
203 {
204 self.create_channel(
205 ProverMerkleTranscriptChannel::with_merkle_prover(
206 transcript,
207 BinaryMerkleTreeProver::with_allocator(alloc),
208 ),
209 rng,
210 alloc,
211 )
212 }
213
214 pub fn create_channel_without_zk_from_transcript<H, Challenger_, T, A>(
219 &self,
220 transcript: T,
221 alloc: A,
222 ) -> TranscriptBaseFoldProverChannel<'_, F, P, NTT, T, Challenger_, H, A>
223 where
224 H: ParallelHashSuite,
225 Challenger_: Challenger,
226 T: BorrowMut<ProverTranscript<Challenger_>>,
227 Output<H::LeafHash>: SerializeBytes,
228 A: Allocator,
229 {
230 self.create_channel_without_zk(
231 ProverMerkleTranscriptChannel::with_merkle_prover(
232 transcript,
233 BinaryMerkleTreeProver::with_allocator(alloc),
234 ),
235 alloc,
236 )
237 }
238}