Skip to main content

binius_iop/fri/
common.rs

1// Copyright 2024-2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::marker::PhantomData;
5
6use binius_field::{BinaryField, Field};
7use binius_math::reed_solomon::ReedSolomonCode;
8use binius_utils::checked_arithmetics::log2_ceil_usize;
9use getset::{CopyGetters, Getters};
10
11use crate::{channel::OracleSpec, merkle_tree::MerkleTreeScheme};
12
13/// Parameters for an FRI interleaved code proximity protocol.
14///
15/// ## Invariants
16///
17/// The dimension of the first-round (reduced) FRI oracle is
18/// `rs_code.log_dim() == log_terminal_dim + sum(fold_arities)`. For all oracle specs in
19/// `input_oracles`:
20/// - `log_batch_size <= log_msg_len`
21/// - `log_msg_len <= rs_code.log_dim() + log_batch_size` (equivalently, `log_msg_len <=
22///   log_terminal_dim + sum(fold_arities) + log_batch_size`)
23#[derive(Debug, Clone, Getters, CopyGetters)]
24pub struct FRIParams<F> {
25	/// The Reed-Solomon code the verifier is testing proximity to.
26	#[getset(get = "pub")]
27	rs_code: ReedSolomonCode<F>,
28	/// Guaranteed to be non-empty.
29	input_oracles: Vec<CodewordSpec>,
30	/// log2 the maximum message length of all input oracles, after lifting each to the reduced
31	/// dimension. Equals `rs_code.log_dim() + max_early + max_later`, where `max_early` (resp.
32	/// `max_later`) is the maximum `log_early_batch_size` (resp. `log_later_batch_size`) over the
33	/// input oracles (the within-oracle batch challenges the first fold must draw).
34	max_log_msg_len: usize,
35	/// log2 ceiling of the number of input oracles.
36	log_n_oracles: usize,
37	/// The reduction arities between each oracle sent to the verifier.
38	fold_arities: Vec<usize>,
39	/// log2 the dimension of the terminal codeword.
40	log_terminal_dim: usize,
41	/// The number oracle consistency queries required during the query phase.
42	#[getset(get_copy = "pub")]
43	n_test_queries: usize,
44}
45
46/// Specification of one committed codeword batched into the first-round FRI oracle.
47///
48/// Each input oracle commits an interleaved Reed–Solomon codeword whose own dimension is
49/// `rs_code.log_dim() - log_lift`; it is lifted to the shared reduced dimension by duplicating
50/// each entry `2^log_lift` times before the oracles are batched together. This is distinct from
51/// [`crate::channel::OracleSpec`], the higher-level description of an oracle to be committed (its
52/// message length and whether it is masked); a `CodewordSpec` is the resolved, FRI-level layout
53/// the [`FRIParams`] selection computes from a batch of those.
54#[derive(Debug, Clone)]
55pub struct CodewordSpec {
56	/// log2 the number of times each committed codeword entry is duplicated to lift it to the
57	/// shared first-round (reduced) dimension.
58	///
59	/// The committed codeword's Reed–Solomon dimension is `rs_code.log_dim() - log_lift`, so its
60	/// message length is `rs_code.log_dim() - log_lift + log_batch_size`. It is `0` when the
61	/// oracle already sits at the reduced dimension.
62	pub log_lift: usize,
63	/// log2 the number of *early* batch-fold challenges this oracle's interleaving folds with.
64	///
65	/// The first fold draws its within-oracle batch challenges in two groups: `max_early =
66	/// max(log_early_batch_size)` *early* challenges, sampled before the `log_n_oracles` outer
67	/// (oracle-combine) challenges, followed by `max_later = max(log_later_batch_size)` *later*
68	/// challenges, sampled after them. The full first-fold challenge slice is therefore
69	/// `[early (max_early)] ++ [outer (log_n_oracles)] ++ [later (max_later)]`.
70	///
71	/// Oracle `i` folds its interleaving with the concatenation `early_window ++ later_window`,
72	/// where `early_window` is the `log_early_batch_size`-length *suffix* of the early challenges
73	/// and `later_window` is the `log_later_batch_size`-length *suffix* of the later challenges.
74	/// The early group carries the shared masking challenge γ of ZK BaseFold oracles (each such
75	/// oracle is purely early, `log_later_batch_size == 0`); the later group carries the non-ZK
76	/// oracles' flexible batch folds (each such oracle is purely later, `log_early_batch_size ==
77	/// 0`). The total interleave batch size of an oracle is `log_early_batch_size +
78	/// log_later_batch_size` (see [`Self::log_batch_size`]).
79	pub log_early_batch_size: usize,
80	/// log2 the number of *later* batch-fold challenges this oracle's interleaving folds with,
81	/// sampled after the outer (oracle-combine) challenges. See [`Self::log_early_batch_size`].
82	pub log_later_batch_size: usize,
83}
84
85impl CodewordSpec {
86	/// log2 the interleaved batch size: the early plus the later batch challenges.
87	pub const fn log_batch_size(&self) -> usize {
88		self.log_early_batch_size + self.log_later_batch_size
89	}
90}
91
92impl<F> FRIParams<F>
93where
94	F: BinaryField,
95{
96	/// ## Preconditions
97	///
98	/// * `sum(fold_arities)` must be at most `rs_code.log_dim()`.
99	pub fn new(
100		rs_code: ReedSolomonCode<F>,
101		log_batch_size: usize,
102		fold_arities: Vec<usize>,
103		n_test_queries: usize,
104	) -> Self {
105		// A single oracle already sits at the reduced dimension (no lifting) and is non-ZK /
106		// homogeneous: its whole batch is "later" (with `log_n_oracles == 0` there is no outer
107		// group between early and later, so this is identical to the old prefix convention).
108		let oracle_spec = CodewordSpec {
109			log_lift: 0,
110			log_early_batch_size: 0,
111			log_later_batch_size: log_batch_size,
112		};
113		Self::new_batch(rs_code, vec![oracle_spec], fold_arities, n_test_queries)
114	}
115
116	/// Create parameters for a batch of committed codewords with an explicit per-codeword layout.
117	///
118	/// This is the low-level constructor: the caller supplies the reduced Reed–Solomon code, each
119	/// codeword's lift / early & later batch-fold routing, and the fold arities. The
120	/// proof-size-minimizing selection of those values from a batch of higher-level oracle
121	/// descriptions lives in [`Self::optimal_for_batch`].
122	///
123	/// ## Preconditions
124	///
125	/// * `oracles` is non-empty.
126	/// * `sum(fold_arities)` must be at most `rs_code.log_dim()`.
127	/// * For each oracle, `log_lift <= rs_code.log_dim()`.
128	pub fn new_batch(
129		rs_code: ReedSolomonCode<F>,
130		oracles: Vec<CodewordSpec>,
131		fold_arities: Vec<usize>,
132		n_test_queries: usize,
133	) -> Self {
134		assert!(!oracles.is_empty(), "precondition: oracles must be non-empty");
135
136		let fold_arities_sum = fold_arities.iter().sum();
137		let log_terminal_dim = rs_code
138			.log_dim()
139			.checked_sub(fold_arities_sum)
140			.expect("precondition: sum(fold_arities) must be at most rs_code.log_dim()");
141
142		let log_n_oracles = log2_ceil_usize(oracles.len());
143
144		// The first fold draws its within-oracle batch challenges as `max_early` early challenges
145		// followed (after the `log_n_oracles` outer oracle-combine folds) by `max_later` later
146		// challenges, so the full first-fold slice is `[early ++ outer ++ later]`. The lifted
147		// message length each oracle reaches is `rs_code.log_dim() + max_early + max_later`.
148		let max_early = oracles
149			.iter()
150			.map(|spec| spec.log_early_batch_size)
151			.max()
152			.expect("oracles is non-empty");
153		let max_later = oracles
154			.iter()
155			.map(|spec| spec.log_later_batch_size)
156			.max()
157			.expect("oracles is non-empty");
158		let max_log_msg_len = rs_code.log_dim() + max_early + max_later;
159
160		Self {
161			rs_code,
162			input_oracles: oracles,
163			max_log_msg_len,
164			log_n_oracles,
165			fold_arities,
166			log_terminal_dim,
167			n_test_queries,
168		}
169	}
170
171	/// Create parameters using the given arity selection strategy.
172	///
173	/// ## Arguments
174	///
175	/// * `merkle_scheme` - the Merkle tree scheme used for commitments.
176	/// * `log_msg_len` - the binary logarithm of the length of the message to commit.
177	/// * `log_batch_size` - if `Some`, fixes the batch size; if `None`, the batch size is chosen
178	///   optimally along with the fold arities.
179	/// * `log_inv_rate` - the binary logarithm of the inverse Reed–Solomon code rate.
180	/// * `n_test_queries` - the number of test queries for the FRI protocol.
181	/// * `strategy` - the strategy for selecting fold arities.
182	///
183	/// ## Preconditions
184	///
185	/// * If `log_batch_size` is `Some(b)`, then `b <= log_msg_len`.
186	pub fn with_strategy<MerkleScheme, Strategy>(
187		merkle_scheme: &MerkleScheme,
188		log_msg_len: usize,
189		log_batch_size: Option<usize>,
190		log_inv_rate: usize,
191		n_test_queries: usize,
192		strategy: &Strategy,
193	) -> Self
194	where
195		MerkleScheme: MerkleTreeScheme<F>,
196		Strategy: AritySelectionStrategy,
197	{
198		assert!(log_batch_size.is_none_or(|b| b <= log_msg_len)); // precondition
199
200		let mut fold_arities = strategy.choose_arities::<F, _>(
201			merkle_scheme,
202			log_msg_len - log_batch_size.unwrap_or(0),
203			log_inv_rate,
204			n_test_queries,
205		);
206		// Without a fixed batch size, the first chosen arity becomes the batch size.
207		let log_batch_size = log_batch_size.unwrap_or_else(|| {
208			// Edge case: no folds were chosen, so batch down to a log_dim = 0 code.
209			if fold_arities.is_empty() {
210				log_msg_len
211			} else {
212				fold_arities.remove(0)
213			}
214		});
215
216		let log_dim = log_msg_len - log_batch_size;
217		let rs_code = ReedSolomonCode::new(log_dim, log_inv_rate);
218		Self::new(rs_code, log_batch_size, fold_arities, n_test_queries)
219	}
220
221	/// Create parameters for a batch of input oracles, minimizing the estimated proof size.
222	///
223	/// The input oracles may have differing message lengths. Each oracle is reduced into a common
224	/// first-round FRI oracle, whose dimension is chosen to minimize the estimated proof size; the
225	/// per-oracle batch sizes and the subsequent fold arities are chosen along with it.
226	///
227	/// Returns the parameters together with the estimated proof size in bytes.
228	///
229	/// ## Arguments
230	///
231	/// * `merkle_scheme` - the Merkle tree scheme used for commitments.
232	/// * `oracles` - the oracles to batch. A ZK oracle commits its message interleaved with an
233	///   equal-length mask (fixed `log_batch_size = 1`); a non-ZK oracle commits the bare message
234	///   with a batch size chosen optimally.
235	/// * `log_inv_rate` - the binary logarithm of the inverse Reed–Solomon code rate.
236	/// * `n_test_queries` - the number of test queries for the FRI protocol.
237	///
238	/// ## Preconditions
239	///
240	/// * `oracles` is non-empty.
241	pub fn optimal_for_batch<MerkleScheme>(
242		merkle_scheme: &MerkleScheme,
243		oracles: &[OracleSpec],
244		log_inv_rate: usize,
245		n_test_queries: usize,
246	) -> (Self, usize)
247	where
248		MerkleScheme: MerkleTreeScheme<F>,
249	{
250		assert!(!oracles.is_empty()); // precondition
251
252		let ChooseCodewordSpecsOutput {
253			proof_size,
254			reduced_log_dim,
255			oracle_specs,
256			fold_arities,
257		} = choose_codeword_specs_for_oracles(merkle_scheme, oracles, log_inv_rate, n_test_queries);
258
259		let rs_code = ReedSolomonCode::new(reduced_log_dim, log_inv_rate);
260
261		let params = Self::new_batch(rs_code, oracle_specs, fold_arities, n_test_queries);
262		(params, proof_size)
263	}
264
265	/// Number of folding rounds in the FRI protocol.
266	///
267	/// This is the largest input message length, plus `log_n_oracles` extra rounds that fold the
268	/// distinct input oracles together into the batched codeword.
269	pub const fn n_fold_rounds(&self) -> usize {
270		self.max_log_msg_len + self.log_n_oracles
271	}
272
273	/// Number of oracles sent during the fold rounds.
274	pub const fn n_oracles(&self) -> usize {
275		// One for the batched codeword commitment, and one for each subsequent one.
276		1 + self.fold_arities.len()
277	}
278
279	/// Number of bits in the query indices sampled during the query phase.
280	pub const fn index_bits(&self) -> usize {
281		self.rs_code.log_len()
282	}
283
284	/// Number of folding challenges the verifier sends after receiving the last oracle.
285	pub const fn n_final_challenges(&self) -> usize {
286		self.log_terminal_dim
287	}
288
289	/// The reduction arities between each oracle sent to the verifier.
290	pub fn fold_arities(&self) -> &[usize] {
291		&self.fold_arities
292	}
293
294	/// The specifications of the input oracles batched into the first-round FRI oracle.
295	pub fn input_oracles(&self) -> &[CodewordSpec] {
296		&self.input_oracles
297	}
298
299	/// The arity of the reduction to the first round oracle.
300	pub fn log_batch_size(&self) -> usize {
301		self.log_msg_len() - self.rs_code().log_dim()
302	}
303
304	/// The binary logarithm of the length of the initial oracle.
305	pub fn log_len(&self) -> usize {
306		self.log_msg_len() + self.rs_code().log_inv_rate()
307	}
308
309	/// The binary logarithm of the length of the initial message.
310	///
311	/// This includes the `log_n_oracles` extra rounds used to fold the distinct input oracles
312	/// together, so it equals [`Self::n_fold_rounds`].
313	pub const fn log_msg_len(&self) -> usize {
314		self.max_log_msg_len + self.log_n_oracles
315	}
316}
317
318struct ChooseBatchSizeAndAritiesOutput {
319	proof_size: usize,
320	reduced_log_dim: usize,
321	fold_arities: Vec<usize>,
322}
323
324/// Choose the shared reduced dimension, fold arities, and resulting proof size for a batch of
325/// oracles, given each oracle's committed message length and (optionally fixed) batch size.
326///
327/// This is unaware of ZK: an oracle with a fixed `log_batch_size` (`Some`) has its batch size
328/// pinned by the caller, while an oracle with `None` takes a flexible batch size, decided here to
329/// reduce it to the shared dimension. For all input oracles we need their log_batch_size <=
330/// committed message length and committed message length <= reduced_log_dim + log_batch_size. We
331/// allow the committed message length to be less than reduced_log_dim + log_batch_size because we
332/// can lift Reed-Solomon encoded oracles.
333///
334/// `committed`: per oracle, its committed message length and its fixed log_batch_size (`Some`), or
335/// `None` for a flexible batch size.
336fn choose_batch_size_and_arities_multi<F, MerkleScheme>(
337	merkle_scheme: &MerkleScheme,
338	committed: &[(usize, Option<usize>)],
339	log_inv_rate: usize,
340	n_test_queries: usize,
341) -> ChooseBatchSizeAndAritiesOutput
342where
343	F: BinaryField,
344	MerkleScheme: MerkleTreeScheme<F>,
345{
346	// First, figure out lower and upper bounds on the reduced_log_dim. If there are any input
347	// oracles with a fixed log_batch_size, then their dimension lower bounds the reduced oracle
348	// dimension.
349	let min_reduced_log_dim = committed
350		.iter()
351		.filter_map(|&(committed_log_msg_len, log_batch_size)| {
352			Some(committed_log_msg_len - log_batch_size?)
353		})
354		.max()
355		.unwrap_or(0);
356	// The upper bound is then the committed message length of the largest flexible oracle, if
357	// larger.
358	let max_reduced_log_dim = committed
359		.iter()
360		.filter(|(_, log_batch_size)| log_batch_size.is_none())
361		.map(|&(committed_log_msg_len, _)| committed_log_msg_len)
362		.max()
363		.unwrap_or(0)
364		.max(min_reduced_log_dim);
365
366	let optimizer = ReductionOptimizer::<F, _>::new(merkle_scheme, n_test_queries);
367	let min_sizes = optimizer.compute_optimal_arities(max_reduced_log_dim, log_inv_rate);
368
369	// Compute the contribution of the oracles with a fixed batch size to the initial fold reduction
370	// size. It's not necessary to compute this for the purpose of parameter minimization, but it's
371	// nice to get the resulting proof size estimate from this function.
372	let fixed_reduction_size = committed
373		.iter()
374		.filter_map(|&(committed_log_msg_len, log_batch_size)| {
375			log_batch_size.map(|log_batch_size| {
376				optimizer.compute_layer_reduction_size(
377					committed_log_msg_len + log_inv_rate,
378					log_batch_size,
379				)
380			})
381		})
382		.sum::<usize>();
383
384	let (reduced_log_dim, proof_size) =
385		min_concave((min_reduced_log_dim..=max_reduced_log_dim).rev(), |reduced_log_dim| {
386			// Compute the reduction sizes of the oracles with a flexible batch size, assuming
387			// the first FRI round oracle has dimension reduced_log_dim.
388			let non_fixed_reduction_size = committed
389				.iter()
390				.filter(|(_, log_batch_size)| log_batch_size.is_none())
391				.map(|&(committed_log_msg_len, _)| {
392					optimizer.compute_layer_reduction_size(
393						committed_log_msg_len + log_inv_rate,
394						committed_log_msg_len.saturating_sub(reduced_log_dim),
395					)
396				})
397				.sum::<usize>();
398
399			let reduction_size = fixed_reduction_size + non_fixed_reduction_size;
400			let reduced_proof_size = min_sizes[reduced_log_dim].proof_size;
401			reduction_size + reduced_proof_size
402		})
403		.expect("range is non-empty because it's inclusive of an upper bound >= the lower bound");
404
405	let fold_arities =
406		optimizer.optimizer_entries_to_fold_arities(&min_sizes[..reduced_log_dim + 1]);
407
408	ChooseBatchSizeAndAritiesOutput {
409		proof_size,
410		reduced_log_dim,
411		fold_arities,
412	}
413}
414
415struct ChooseCodewordSpecsOutput {
416	proof_size: usize,
417	reduced_log_dim: usize,
418	oracle_specs: Vec<CodewordSpec>,
419	fold_arities: Vec<usize>,
420}
421
422fn choose_codeword_specs_for_oracles<F, MerkleScheme>(
423	merkle_scheme: &MerkleScheme,
424	oracles: &[OracleSpec],
425	log_inv_rate: usize,
426	n_test_queries: usize,
427) -> ChooseCodewordSpecsOutput
428where
429	F: BinaryField,
430	MerkleScheme: MerkleTreeScheme<F>,
431{
432	// We want to determine the dimension of the first folded FRI oracle, which we'll call the
433	// "reduced" oracle. This is reduced_log_dim. For each input oracle, we will determine a
434	// log_batch_size. A ZK oracle commits its message interleaved with an equal-length mask, so its
435	// committed message length is `log_msg_len + 1` and its batch size is fixed at 1. A non-ZK
436	// oracle commits the bare message and takes a flexible batch size.
437	//
438	// `committed`: per oracle, its committed message length and its fixed log_batch_size (`Some`
439	// for ZK oracles, `None` for flexible non-ZK ones).
440	let committed: Vec<(usize, Option<usize>)> = oracles
441		.iter()
442		.map(|oracle| {
443			let committed_log_msg_len = oracle.log_msg_len + usize::from(oracle.is_zk);
444			let log_batch_size = oracle.is_zk.then_some(1);
445			(committed_log_msg_len, log_batch_size)
446		})
447		.collect();
448
449	let ChooseBatchSizeAndAritiesOutput {
450		proof_size,
451		reduced_log_dim,
452		fold_arities,
453	} = choose_batch_size_and_arities_multi(merkle_scheme, &committed, log_inv_rate, n_test_queries);
454
455	// Resolve each oracle's concrete log_batch_size (fixed, or chosen to reduce to the shared
456	// dimension), paired with its committed message length and is_zk flag.
457	let resolved: Vec<(usize, usize, bool)> = oracles
458		.iter()
459		.zip(&committed)
460		.map(|(oracle, &(committed_log_msg_len, log_batch_size))| {
461			let log_batch_size = log_batch_size
462				.unwrap_or_else(|| committed_log_msg_len.saturating_sub(reduced_log_dim));
463			(committed_log_msg_len, log_batch_size, oracle.is_zk)
464		})
465		.collect();
466
467	// ZK-aware early/later batch-fold routing. A ZK oracle's batch is the shared masking challenge
468	// γ, sampled *before* the outer (oracle-combine) challenges, so it is entirely "early". A
469	// non-ZK oracle's flexible batch is sampled *after* the outer challenges, so it is entirely
470	// "later". The first-fold challenge slice is therefore `[early ++ outer ++ later]`, and each
471	// oracle folds its interleaving with a suffix of whichever group it belongs to.
472	let oracle_specs = resolved
473		.iter()
474		.map(|&(committed_log_msg_len, log_batch_size, is_zk)| {
475			let log_early_batch_size = if is_zk { log_batch_size } else { 0 };
476			let log_later_batch_size = if is_zk { 0 } else { log_batch_size };
477			// The committed codeword's own dimension is `committed_log_msg_len - log_batch_size`;
478			// it is lifted to the reduced dimension, so `log_lift` is the gap between the two.
479			let oracle_log_dim = committed_log_msg_len - log_batch_size;
480			let log_lift = reduced_log_dim - oracle_log_dim;
481			CodewordSpec {
482				log_lift,
483				log_early_batch_size,
484				log_later_batch_size,
485			}
486		})
487		.collect();
488
489	ChooseCodewordSpecsOutput {
490		proof_size,
491		reduced_log_dim,
492		oracle_specs,
493		fold_arities,
494	}
495}
496
497/// Calculates the number of test queries required to achieve a target soundness error.
498///
499/// This chooses a number of test queries so that the soundness error of the FRI query phase is
500/// at most $2^{-t}$, where $t$ is the threshold `security_bits`. This _does not_ account for the
501/// soundness error from the FRI folding phase or any other protocols, only the query phase. This
502/// sets the proximity parameter for FRI to the code's unique decoding radius. See [DP24],
503/// Section 5.2, for concrete soundness analysis.
504///
505/// [DP24]: <https://eprint.iacr.org/2024/504>
506pub fn calculate_n_test_queries(security_bits: usize, log_inv_rate: usize) -> usize {
507	let rate = 2.0f64.powi(-(log_inv_rate as i32));
508	let per_query_err = 0.5 * (1f64 + rate);
509	(security_bits as f64 / -per_query_err.log2()).ceil() as usize
510}
511
512/// Strategy for selecting fold arities in the FRI protocol.
513pub trait AritySelectionStrategy {
514	fn choose_arities<F, MerkleScheme>(
515		&self,
516		merkle_scheme: &MerkleScheme,
517		log_msg_len: usize,
518		log_inv_rate: usize,
519		n_test_queries: usize,
520	) -> Vec<usize>
521	where
522		F: Field,
523		MerkleScheme: MerkleTreeScheme<F>;
524}
525
526#[derive(Debug)]
527struct ReductionOptimizerEntry {
528	// The minimum proof size attainable for the indexed value of i.
529	proof_size: usize,
530	// The first reduction arity to achieve the minimum proof size. If the value is none,
531	// then the best reduction sequence is to skip all folding and send the full codeword.
532	arity: Option<usize>,
533}
534
535struct ReductionOptimizer<'a, F, MTScheme> {
536	merkle_scheme: &'a MTScheme,
537	n_test_queries: usize,
538	_marker: PhantomData<F>,
539}
540
541impl<'a, F, MTScheme> ReductionOptimizer<'a, F, MTScheme>
542where
543	F: Field,
544	MTScheme: MerkleTreeScheme<F>,
545{
546	const fn new(merkle_scheme: &'a MTScheme, n_test_queries: usize) -> Self {
547		Self {
548			merkle_scheme,
549			n_test_queries,
550			_marker: PhantomData,
551		}
552	}
553
554	/// The proof bytes one reduction of the given arity contributes.
555	///
556	/// Each test query sends one opened coset and its Merkle branch:
557	///
558	/// ```text
559	///     coset     2^arity field elements
560	///     branch    one hash per tree level
561	/// ```
562	///
563	/// The oracle commits one coset per leaf.
564	/// So its tree holds `2^(log_code_len - arity)` leaves, not `2^log_code_len`.
565	///
566	/// Sizing the tree by the codeword length would charge `arity` extra hashes per branch.
567	/// The arities chosen would then minimize a proof size no prover produces.
568	fn compute_layer_reduction_size(&self, log_code_len: usize, arity: usize) -> usize {
569		// Each queried coset contains 2^arity values.
570		let leaf_size = F::BYTE_SIZE << arity;
571		// One coset per test query.
572		let leaves_size = leaf_size * self.n_test_queries;
573
574		// One leaf per coset, so the tree is `arity` levels shorter than the codeword.
575		let log_n_cosets = log_code_len - arity;
576
577		// Size of the Merkle multi-proof.
578		let optimal_layer = self
579			.merkle_scheme
580			.optimal_verify_layer(self.n_test_queries, log_n_cosets);
581		let merkle_size =
582			self.merkle_scheme
583				.proof_size(1 << log_n_cosets, self.n_test_queries, optimal_layer);
584
585		leaves_size + merkle_size
586	}
587
588	fn compute_optimal_arities(
589		&self,
590		log_msg_len: usize,
591		log_inv_rate: usize,
592	) -> Vec<ReductionOptimizerEntry> {
593		type Entry = ReductionOptimizerEntry;
594
595		// This algorithm uses a dynamic programming approach to determine the sequence of arities
596		// that minimizes proof size. For each i in [0, log_msg_len], we determine the minimum
597		// proof size attainable when for a batched codeword with message size 2^i. This is
598		// determined by minimizing over the first reduction arity, using the values already
599		// determined for the smaller values of i.
600
601		// This vec maps log_msg_len values to the minimum proof size attainable for a batched FRI
602		// protocol committing a message with that length.
603		let mut min_sizes = Vec::<Entry>::with_capacity(log_msg_len + 1);
604
605		for i in 0..=log_msg_len {
606			// Length of the batched codeword.
607			let log_code_len = i + log_inv_rate;
608
609			let non_terminal_entry = min_concave(1..=i, |arity| {
610				// The additional proof bytes for the reduction by arity.
611				let reduction_proof_size = self.compute_layer_reduction_size(log_code_len, arity);
612				let reduced_proof_size = min_sizes[i - arity].proof_size;
613				reduction_proof_size + reduced_proof_size
614			})
615			.map(|(arity, proof_size)| Entry {
616				arity: Some(arity),
617				proof_size,
618			});
619
620			// Determine the proof size if this is the terminal codeword. In that case, the proof
621			// simply consists of the 2^(i + log_inv_rate) leaf values.
622			let terminal_proof_size = F::BYTE_SIZE << log_code_len;
623			let terminal_entry = Entry {
624				proof_size: terminal_proof_size,
625				arity: None,
626			};
627
628			let optimal_entry = if let Some(non_terminal_entry) = non_terminal_entry
629				&& non_terminal_entry.proof_size < terminal_entry.proof_size
630			{
631				non_terminal_entry
632			} else {
633				terminal_entry
634			};
635
636			min_sizes.push(optimal_entry);
637		}
638
639		min_sizes
640	}
641
642	fn optimizer_entries_to_fold_arities(
643		&self,
644		min_sizes: &[ReductionOptimizerEntry],
645	) -> Vec<usize> {
646		let mut fold_arities = Vec::new();
647
648		let mut i = min_sizes.len() - 1;
649		let mut entry = &min_sizes[i];
650		while let Some(arity) = entry.arity {
651			fold_arities.push(arity);
652			i -= arity;
653			entry = &min_sizes[i];
654		}
655		fold_arities
656	}
657}
658
659/// Minimizes `f` over the values yielded by `params`, returning the minimizing argument and value.
660///
661/// This assumes `f` is unimodal (quasi-convex) over the iteration order: non-increasing up to the
662/// minimum and non-decreasing afterwards. It scans in order and stops as soon as `f` strictly
663/// increases. On ties it keeps the later argument. Returns `None` if `params` is empty.
664fn min_concave<A: Copy, B: Ord>(
665	mut params: impl Iterator<Item = A>,
666	f: impl Fn(A) -> B,
667) -> Option<(A, B)> {
668	let mut min_a = params.next()?;
669	let mut min_b = f(min_a);
670	for a in params {
671		let b = f(a);
672		if b <= min_b {
673			min_a = a;
674			min_b = b;
675		} else {
676			// The function f is concave in the sequence of params, so break if it begins
677			// increasing.
678			break;
679		}
680	}
681	Some((min_a, min_b))
682}
683
684/// Strategy that minimizes proof size using dynamic programming.
685#[derive(Debug, Clone, Copy, Default)]
686pub struct MinProofSizeStrategy;
687
688impl AritySelectionStrategy for MinProofSizeStrategy {
689	fn choose_arities<F, MerkleScheme>(
690		&self,
691		merkle_scheme: &MerkleScheme,
692		log_msg_len: usize,
693		log_inv_rate: usize,
694		n_test_queries: usize,
695	) -> Vec<usize>
696	where
697		F: Field,
698		MerkleScheme: MerkleTreeScheme<F>,
699	{
700		let optimizer = ReductionOptimizer::<F, _>::new(merkle_scheme, n_test_queries);
701		let min_sizes = optimizer.compute_optimal_arities(log_msg_len, log_inv_rate);
702		optimizer.optimizer_entries_to_fold_arities(&min_sizes)
703	}
704}
705
706/// Strategy that uses a constant fold arity.
707#[derive(Debug, Clone, Copy)]
708pub struct ConstantArityStrategy {
709	/// The fold arity to use for each reduction step.
710	pub arity: usize,
711}
712
713impl ConstantArityStrategy {
714	/// Creates a new strategy with the given arity.
715	pub const fn new(arity: usize) -> Self {
716		Self { arity }
717	}
718
719	/// Creates a strategy with an estimated optimal arity.
720	///
721	/// Uses a heuristic to estimate the optimal FRI folding arity that minimizes proof size.
722	///
723	/// ## Arguments
724	///
725	/// * `_merkle_scheme` - the Merkle tree scheme (used to infer digest size)
726	/// * `approx_log_code_len` - approximate log2 of the codeword length
727	pub fn with_optimal_arity<F, MerkleScheme>(
728		_merkle_scheme: &MerkleScheme,
729		approx_log_code_len: usize,
730	) -> Self
731	where
732		F: Field,
733		MerkleScheme: MerkleTreeScheme<F>,
734	{
735		let digest_size = std::mem::size_of::<MerkleScheme::Digest>() * 8;
736		let field_size = std::mem::size_of::<F>() * 8;
737
738		// Estimate optimal arity using a heuristic based on the approximation of a single
739		// query_proof_size, where θ is the arity:
740		// ((n-θ) + (n-2θ) + ...) * digest_size + ((n-θ)/θ) * 2^θ * field_size
741		let arity = (1..=approx_log_code_len)
742			.map(|arity| {
743				(
744					arity,
745					((approx_log_code_len) / 2 * digest_size + (1 << arity) * field_size)
746						* (approx_log_code_len - arity)
747						/ arity,
748				)
749			})
750			// Scan and terminate when query_proof_size increases.
751			.scan(None, |old: &mut Option<(usize, usize)>, new| {
752				let should_continue = !matches!(*old, Some(ref old) if new.1 > old.1);
753				*old = Some(new);
754				should_continue.then_some(new)
755			})
756			.last()
757			.map(|(arity, _)| arity)
758			.unwrap_or(1);
759
760		Self { arity }
761	}
762}
763
764impl AritySelectionStrategy for ConstantArityStrategy {
765	fn choose_arities<F, MerkleScheme>(
766		&self,
767		merkle_scheme: &MerkleScheme,
768		log_msg_len: usize,
769		log_inv_rate: usize,
770		n_test_queries: usize,
771	) -> Vec<usize>
772	where
773		F: Field,
774		MerkleScheme: MerkleTreeScheme<F>,
775	{
776		let log_code_len = log_msg_len + log_inv_rate;
777		let cap_height = merkle_scheme.optimal_verify_layer(n_test_queries, log_code_len);
778		let log_terminal_len = cap_height.max(log_inv_rate);
779
780		let mut fold_arities = Vec::new();
781		let mut i = log_code_len;
782		while i > log_terminal_len {
783			if let Some(next_i) = i.checked_sub(self.arity) {
784				fold_arities.push(self.arity);
785				i = next_i;
786			} else {
787				break;
788			}
789		}
790		fold_arities
791	}
792}
793
794#[cfg(test)]
795mod tests {
796	use binius_field::Ghash128b as B128;
797	use binius_hash::StdHashSuite;
798
799	use super::*;
800	use crate::{fri::proof_size, merkle_tree::BinaryMerkleTreeScheme};
801
802	type TestMerkleScheme = BinaryMerkleTreeScheme<B128, StdHashSuite>;
803
804	fn test_merkle_scheme() -> TestMerkleScheme {
805		BinaryMerkleTreeScheme::new()
806	}
807
808	/// Security level the shipped verifier targets.
809	///
810	/// Restated rather than imported.
811	/// `binius_verifier::SECURITY_BITS` lives in a crate that depends on this one.
812	const SECURITY_BITS: usize = 96;
813
814	/// Candidate inverse rates, wide enough to bracket where the proof size turns around.
815	const LOG_INV_RATES: [usize; 6] = [1, 2, 3, 4, 5, 6];
816
817	/// Exact proof size per shape and rate, as `(shape, log_inv_rate, n_test_queries, bytes)`.
818	///
819	/// Every shape bottoms out inside the candidate range.
820	/// The smallest oracle bottoms out at rate 1/8, and the larger two at 1/16.
821	const PINNED_PROOF_SIZE_BY_RATE: [(&str, usize, usize, usize); 18] = [
822		("single/17", 1, 232, 205152),
823		("single/17", 2, 142, 160128),
824		("single/17", 3, 116, 154240),
825		("single/17", 4, 106, 161088),
826		("single/17", 5, 101, 172544),
827		("single/17", 6, 99, 186304),
828		("single/24", 1, 232, 472480),
829		("single/24", 2, 142, 340128),
830		("single/24", 3, 116, 307264),
831		("single/24", 4, 106, 307008),
832		("single/24", 5, 101, 318240),
833		("single/24", 6, 99, 336192),
834		("zk/24+21", 1, 232, 730208),
835		("zk/24+21", 2, 142, 513344),
836		("zk/24+21", 3, 116, 458432),
837		("zk/24+21", 4, 106, 452640),
838		("zk/24+21", 5, 101, 463856),
839		("zk/24+21", 6, 99, 485424),
840	];
841
842	/// The exact proof size at one rate, with the arities re-optimized for it.
843	///
844	/// The query count is not a parameter, but the one that rate needs for `SECURITY_BITS`.
845	/// A rate buys proof bytes precisely by changing it, so pairing the two is the point.
846	fn proof_size_at_rate(
847		merkle_scheme: &TestMerkleScheme,
848		oracles: &[OracleSpec],
849		log_inv_rate: usize,
850	) -> usize {
851		let n_test_queries = calculate_n_test_queries(SECURITY_BITS, log_inv_rate);
852		let (params, _) =
853			FRIParams::optimal_for_batch(merkle_scheme, oracles, log_inv_rate, n_test_queries);
854		proof_size(&params, merkle_scheme)
855	}
856
857	// Invariant: the size the arity search minimizes is the size a prover actually sends.
858	//
859	//     cost model    what `compute_layer_reduction_size` charges, and the search minimizes
860	//     proof_size    the exact byte count
861	//
862	// The cost model omits the commitment digests, which do not vary with the arity choice.
863	// A batch of N input oracles carries `N + 1 + fold_arities.len()` of them.
864	#[test]
865	fn optimizer_estimate_matches_exact_proof_size() {
866		let merkle_scheme = test_merkle_scheme();
867		let digest_size = size_of::<<TestMerkleScheme as MerkleTreeScheme<B128>>::Digest>();
868
869		// Single oracles across the size range, then shapes that stress the batch layout: lifting,
870		// ZK mixed with flexible, non-power-of-two counts.
871		//
872		// A ZK oracle pins its batch size at 1; a non-ZK oracle takes a flexible one, so the two
873		// exercise different branches of the selection.
874		let mut batches: Vec<Vec<OracleSpec>> = Vec::new();
875		for log_msg_len in [0, 1, 4, 8, 12, 16, 20] {
876			batches.push(vec![OracleSpec::new(log_msg_len)]);
877			batches.push(vec![OracleSpec::new_zk(log_msg_len)]);
878		}
879		batches.extend([
880			vec![OracleSpec::new(16), OracleSpec::new(16)],
881			vec![OracleSpec::new(16), OracleSpec::new(12)],
882			vec![OracleSpec::new_zk(11), OracleSpec::new(16)],
883			vec![
884				OracleSpec::new_zk(9),
885				OracleSpec::new_zk(11),
886				OracleSpec::new(16),
887			],
888			vec![
889				OracleSpec::new(20),
890				OracleSpec::new_zk(15),
891				OracleSpec::new(8),
892				OracleSpec::new_zk(4),
893			],
894		]);
895
896		for log_inv_rate in [1, 2, 3] {
897			for n_test_queries in [32, 128, 232] {
898				for oracles in &batches {
899					let (params, estimate) = FRIParams::optimal_for_batch(
900						&merkle_scheme,
901						oracles,
902						log_inv_rate,
903						n_test_queries,
904					);
905
906					let digests = (oracles.len() + 1 + params.fold_arities().len()) * digest_size;
907					assert_eq!(
908						estimate + digests,
909						proof_size(&params, &merkle_scheme),
910						"oracles={oracles:?} log_inv_rate={log_inv_rate} \
911						 n_test_queries={n_test_queries} arities={:?}",
912						params.fold_arities(),
913					);
914				}
915			}
916		}
917	}
918
919	// Invariant: lowering the rate trades encoding work for proof bytes, and overshoots.
920	//
921	//     queries     fall monotonically as the rate falls, then flatten out
922	//     bytes       fall with the queries, then climb again
923	//
924	// The climb is the terminal codeword and every opened coset growing with the rate.
925	// The candidates bracket the turning point on both sides, and it moves with the oracle size.
926	#[test]
927	fn pinned_proof_size_by_rate() {
928		let merkle_scheme = test_merkle_scheme();
929
930		// A small and a large single oracle, plus a ZK batch.
931		// The ZK oracle pins its batch size at 1, taking the fixed-batch-size branch of selection.
932		// The shorter non-ZK oracle beside it exercises lifting.
933		let batches: [(&str, Vec<OracleSpec>); 3] = [
934			("single/17", vec![OracleSpec::new(17)]),
935			("single/24", vec![OracleSpec::new(24)]),
936			("zk/24+21", vec![OracleSpec::new_zk(24), OracleSpec::new(21)]),
937		];
938
939		let observed = batches
940			.iter()
941			.flat_map(|(label, oracles)| {
942				LOG_INV_RATES.map(|log_inv_rate| {
943					(
944						*label,
945						log_inv_rate,
946						calculate_n_test_queries(SECURITY_BITS, log_inv_rate),
947						proof_size_at_rate(&merkle_scheme, oracles, log_inv_rate),
948					)
949				})
950			})
951			.collect::<Vec<_>>();
952
953		assert_eq!(observed, PINNED_PROOF_SIZE_BY_RATE);
954	}
955
956	// Reports the rate trade-off, rather than asserting it, for picking `log_inv_rate` by hand.
957	//
958	//     cargo test -p binius-iop --lib -- --ignored --nocapture report_rate_trade_off
959	//
960	// The encode column is derived, not measured, and is exact.
961	// `ReedSolomonCode::encode_batch` skips its first `log_inv_rate` NTT layers.
962	// It therefore runs `log_dim` layers over a codeword of `2^(log_dim + log_inv_rate)`:
963	//
964	//     butterflies = log_dim * 2^(log_dim + log_inv_rate - 1)
965	//
966	// The rate enters as a bare factor of `2^log_inv_rate`, with no log term on it.
967	// One step down the rate is exactly twice the encoding work, so the column is a doubling.
968	//
969	// `crates/math/benches/reed_solomon.rs` measures how far memory traffic bends that.
970	#[test]
971	#[ignore = "prints a table instead of asserting; needs --nocapture to be useful"]
972	fn report_rate_trade_off() {
973		let merkle_scheme = test_merkle_scheme();
974
975		for log_msg_len in [17, 20, 24] {
976			let oracles = [OracleSpec::new(log_msg_len)];
977			let priced = LOG_INV_RATES
978				.map(|log_inv_rate| proof_size_at_rate(&merkle_scheme, &oracles, log_inv_rate));
979			let smallest = priced
980				.iter()
981				.copied()
982				.min()
983				.expect("LOG_INV_RATES is non-empty");
984
985			println!();
986			println!(
987				"one non-ZK oracle, log_msg_len = {log_msg_len}, {SECURITY_BITS}-bit security"
988			);
989			println!("  rate  queries        proof  vs best encode");
990			for (log_inv_rate, proof_size) in std::iter::zip(LOG_INV_RATES, priced) {
991				let n_test_queries = calculate_n_test_queries(SECURITY_BITS, log_inv_rate);
992				let kib = proof_size as f64 / 1024.0;
993				let vs_best = proof_size as f64 / smallest as f64;
994				// Encoding work relative to the first candidate, which is the cheapest to encode.
995				let encode = 1usize << (log_inv_rate - LOG_INV_RATES[0]);
996				println!(
997					"  1/{:<3} {n_test_queries:>7}  {kib:>7.2} KiB  {vs_best:>6.2}x  {encode:>4}x",
998					1 << log_inv_rate
999				);
1000			}
1001		}
1002	}
1003
1004	// Invariant: a lower rate needs fewer queries, but the saving flattens out fast.
1005	//
1006	//     1/2 -> 232     1/16 -> 106
1007	//     1/4 -> 142     1/32 -> 101
1008	//     1/8 -> 116     1/64 ->  99
1009	//
1010	// Past 1/8, halving the rate again buys ten queries while doubling the codeword.
1011	// That flattening is why `pinned_proof_size_by_rate` turns around instead of falling forever.
1012	#[test]
1013	fn test_calculate_n_test_queries() {
1014		let observed =
1015			LOG_INV_RATES.map(|log_inv_rate| calculate_n_test_queries(SECURITY_BITS, log_inv_rate));
1016		assert_eq!(observed, [232, 142, 116, 106, 101, 99]);
1017	}
1018
1019	#[test]
1020	fn test_min_proof_size_strategy() {
1021		let merkle_scheme = test_merkle_scheme();
1022		let log_inv_rate = 2;
1023		let n_test_queries = 128;
1024		let strategy = MinProofSizeStrategy;
1025
1026		// log_msg_len = 0: no folding needed, terminal codeword is optimal
1027		let arities =
1028			strategy.choose_arities::<B128, _>(&merkle_scheme, 0, log_inv_rate, n_test_queries);
1029		assert_eq!(arities, vec![]);
1030
1031		// log_msg_len = 3: no folding needed, terminal codeword is optimal
1032		let arities =
1033			strategy.choose_arities::<B128, _>(&merkle_scheme, 3, log_inv_rate, n_test_queries);
1034		assert_eq!(arities, vec![]);
1035
1036		// log_msg_len = 24
1037		let arities =
1038			strategy.choose_arities::<B128, _>(&merkle_scheme, 24, log_inv_rate, n_test_queries);
1039		assert_eq!(arities, vec![4, 4, 4, 4]);
1040	}
1041
1042	#[test]
1043	fn test_with_strategy_min_proof_size() {
1044		let merkle_scheme = test_merkle_scheme();
1045		let log_inv_rate = 2;
1046		let n_test_queries = 128;
1047
1048		// log_msg_len = 0
1049		{
1050			let fri_params = FRIParams::with_strategy(
1051				&merkle_scheme,
1052				0,
1053				None,
1054				log_inv_rate,
1055				n_test_queries,
1056				&MinProofSizeStrategy,
1057			);
1058			assert_eq!(fri_params.fold_arities(), &[]);
1059			assert_eq!(fri_params.log_batch_size(), 0);
1060		}
1061
1062		// log_msg_len = 3
1063		{
1064			let fri_params = FRIParams::with_strategy(
1065				&merkle_scheme,
1066				3,
1067				None,
1068				log_inv_rate,
1069				n_test_queries,
1070				&MinProofSizeStrategy,
1071			);
1072			assert_eq!(fri_params.fold_arities(), &[]);
1073			assert_eq!(fri_params.log_batch_size(), 3);
1074		}
1075
1076		// log_msg_len = 24
1077		{
1078			let fri_params = FRIParams::with_strategy(
1079				&merkle_scheme,
1080				24,
1081				None,
1082				log_inv_rate,
1083				n_test_queries,
1084				&MinProofSizeStrategy,
1085			);
1086			assert_eq!(fri_params.fold_arities(), &[4, 4, 4]);
1087			assert_eq!(fri_params.log_batch_size(), 4);
1088		}
1089	}
1090
1091	#[test]
1092	fn test_optimal_for_batch_three_oracles() {
1093		let merkle_scheme = test_merkle_scheme();
1094		let log_inv_rate = 2;
1095		let n_test_queries = 128;
1096
1097		// Two masked ZK oracles (fixed batch size 1, committed lengths 10 and 12) and one non-ZK
1098		// oracle with a flexible batch size (committed length 16). The ZK oracles lower-bound the
1099		// reduced dimension; the flexible oracle folds down to it.
1100		let oracles = vec![
1101			OracleSpec::new_zk(9),
1102			OracleSpec::new_zk(11),
1103			OracleSpec::new(16),
1104		];
1105
1106		let (fri_params, proof_size) =
1107			FRIParams::optimal_for_batch(&merkle_scheme, &oracles, log_inv_rate, n_test_queries);
1108
1109		// The reduced oracle dimension is the dimension of the first FRI round oracle, equal to
1110		// log_terminal_dim + sum(fold_arities).
1111		let reduced_log_dim = fri_params.rs_code().log_dim();
1112		assert_eq!(
1113			reduced_log_dim,
1114			fri_params.log_terminal_dim + fri_params.fold_arities().iter().sum::<usize>()
1115		);
1116
1117		// Each input oracle satisfies the FRIParams invariants.
1118		assert_eq!(fri_params.input_oracles.len(), oracles.len());
1119		for (spec, oracle) in fri_params.input_oracles.iter().zip(&oracles) {
1120			// A ZK oracle interleaves the message with an equal-length mask, so its committed
1121			// message length is `log_msg_len + 1`; a non-ZK oracle commits the bare message.
1122			let committed_log_msg_len = oracle.log_msg_len + usize::from(oracle.is_zk);
1123			// The committed codeword (dimension `committed_log_msg_len - log_batch_size`) is lifted
1124			// to the reduced dimension, recovering the committed message length.
1125			let oracle_log_dim = reduced_log_dim - spec.log_lift;
1126			assert_eq!(oracle_log_dim + spec.log_batch_size(), committed_log_msg_len);
1127			if oracle.is_zk {
1128				// ZK oracles keep their fixed batch size of 1.
1129				assert_eq!(spec.log_batch_size(), 1);
1130			}
1131			// log_batch_size <= committed message length
1132			assert!(spec.log_batch_size() <= committed_log_msg_len);
1133		}
1134
1135		// The largest input oracle is the non-ZK flexible one, so its batch size folds it down
1136		// exactly to the reduced dimension (no lifting).
1137		assert_eq!(fri_params.input_oracles[2].log_lift, 0);
1138		assert_eq!(fri_params.input_oracles[2].log_batch_size(), 16 - reduced_log_dim);
1139
1140		// Pin the estimated proof size, to catch unintended changes in the optimizer.
1141		//
1142		// This sums one reduction per committed oracle, as the exact byte count does.
1143		// `optimizer_estimate_matches_exact_proof_size` ties the two together.
1144		assert_eq!(proof_size, 188416);
1145	}
1146
1147	#[test]
1148	fn test_with_strategy_fixed_batch_size() {
1149		let merkle_scheme = test_merkle_scheme();
1150		let log_inv_rate = 2;
1151		let n_test_queries = 128;
1152
1153		// log_msg_len = 3
1154		{
1155			let fri_params = FRIParams::with_strategy(
1156				&merkle_scheme,
1157				3,
1158				Some(1),
1159				log_inv_rate,
1160				n_test_queries,
1161				&MinProofSizeStrategy,
1162			);
1163			assert_eq!(fri_params.fold_arities(), &[]);
1164			assert_eq!(fri_params.log_batch_size(), 1);
1165		}
1166
1167		// log_msg_len = 24
1168		{
1169			let fri_params = FRIParams::with_strategy(
1170				&merkle_scheme,
1171				24,
1172				Some(1),
1173				log_inv_rate,
1174				n_test_queries,
1175				&MinProofSizeStrategy,
1176			);
1177			assert_eq!(fri_params.fold_arities(), &[4, 4, 4, 3]);
1178			assert_eq!(fri_params.log_batch_size(), 1);
1179		}
1180	}
1181}