Skip to main content

binius_iop_prover/fri/
encode.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::{iter, ops::Deref};
5
6use binius_compute::{Allocator, CollectIntoAllocVec, VecLike};
7use binius_field::{BinaryField, PackedField};
8use binius_iop::fri::FRIParams;
9use binius_math::{FieldBuffer, FieldSlice, ntt::AdditiveNTT, reed_solomon::ReedSolomonCode};
10use binius_utils::rand::par_rand;
11use rand::{CryptoRng, rngs::StdRng};
12
13/// Reed-Solomon encodes one input oracle's interleaved message.
14///
15/// `params` are the (possibly batched) FRI parameters and `oracle_index` selects which oracle of
16/// [`FRIParams::input_oracles`] is encoded. The oracle is Reed–Solomon encoded at its own
17/// dimension — which may be smaller than the batched code's reduced dimension — over the same
18/// subspace and rate; the lift to the reduced dimension is applied later during the combined fold.
19///
20/// The returned codeword is committed by sending it over a Merkle channel (via
21/// `MerkleIPProverChannel::send_merkle_commitment`) with one interleaved coset of
22/// `2^log_batch_size` scalars per leaf.
23///
24/// ## Arguments
25///
26/// * `params` - the (possibly batched) FRI protocol parameters.
27/// * `oracle_index` - the index into [`FRIParams::input_oracles`] of the oracle being encoded.
28/// * `ntt` - the additive NTT for Reed-Solomon encoding.
29/// * `message` - the interleaved message to encode.
30/// * `alloc` - the allocator the codeword is drawn from.
31///
32/// ## Preconditions
33///
34/// * `message.log_len()` must equal the oracle's committed message length,
35///   `params.rs_code().log_dim() - log_lift + log_batch_size`.
36pub fn encode_interleaved<F, P, NTT, A>(
37	params: &FRIParams<F>,
38	oracle_index: usize,
39	ntt: &NTT,
40	message: FieldSlice<'_, P>,
41	alloc: &A,
42) -> FieldBuffer<P, A::Vec<P>>
43where
44	F: BinaryField,
45	P: PackedField<Scalar = F>,
46	NTT: AdditiveNTT<Field = F> + Sync,
47	A: Allocator,
48{
49	let oracle_spec = &params.input_oracles()[oracle_index];
50	let log_batch_size = oracle_spec.log_batch_size();
51	// The oracle's own codeword dimension is the reduced dimension minus its lift; its committed
52	// (interleaved) message length is that plus the batch size.
53	let oracle_log_dim = params.rs_code().log_dim() - oracle_spec.log_lift;
54
55	assert_eq!(
56		message.log_len(),
57		oracle_log_dim + log_batch_size,
58		"precondition: interleaved message length must match the oracle's spec"
59	);
60
61	// Encode this oracle at its own dimension (≤ the batched code's reduced dimension), at the
62	// same rate as the batched code. Both codes evaluate over the Gao-Mateer basis, the shorter
63	// one over a prefix of the longer's, which is how the per-oracle codeword stays consistent
64	// with the combined fold's lift. `encode_batch` checks the NTT agrees with that domain.
65	let rs_code = ReedSolomonCode::new(oracle_log_dim, params.rs_code().log_inv_rate());
66
67	let _scope = tracing::debug_span!(
68		"Reed–Solomon Encode",
69		log_batch_size,
70		log_dim = rs_code.log_dim(),
71		log_inv_rate = rs_code.log_inv_rate(),
72		field_bits = F::N_BITS,
73	)
74	.entered();
75
76	rs_code.encode_batch(ntt, message.as_view(), log_batch_size, alloc)
77}
78
79/// Output of [`encode_masked`]: the interleaved (message ‖ mask) codeword and the generated mask.
80///
81/// Both buffers come from the allocator the encode was given, so a pooled prover holds neither on
82/// the global heap.
83#[derive(Debug)]
84pub struct MaskedCodeword<P: PackedField, Data: Deref<Target = [P]> = Vec<P>> {
85	/// The Reed-Solomon encoding of the interleaved (message ‖ mask) buffer.
86	pub codeword: FieldBuffer<P, Data>,
87	/// The generated random mask, of equal length to the message.
88	pub mask: FieldBuffer<P, Data>,
89}
90
91/// Generates a random mask, interleaves it with the message, and Reed-Solomon encodes.
92///
93/// This is used for zero-knowledge FRI commitments. The function generates a random mask of
94/// equal length to the input message, concatenates `message || mask` as the interleaved message
95/// (with `log_batch_size = 1`), and performs Reed-Solomon encoding. The returned codeword is
96/// committed by sending it over a Merkle channel, like [`encode_interleaved`]'s.
97///
98/// ## Arguments
99///
100/// * `params` - the (possibly batched) FRI parameters.
101/// * `oracle_index` - the index into [`FRIParams::input_oracles`] of the oracle being encoded; its
102///   spec must have `log_batch_size == 1` and `log_msg_len - 1 == message.log_len()`.
103/// * `ntt` - the additive NTT for Reed-Solomon encoding
104/// * `message` - the raw message to encode (not doubled)
105/// * `rng` - cryptographic RNG for mask generation
106/// * `alloc` - the allocator the mask, the codeword, and the concatenation temporary are drawn from
107pub fn encode_masked<F, P, NTT, A>(
108	params: &FRIParams<F>,
109	oracle_index: usize,
110	ntt: &NTT,
111	message: FieldSlice<'_, P>,
112	mut rng: impl CryptoRng,
113	alloc: &A,
114) -> MaskedCodeword<P, A::Vec<P>>
115where
116	F: BinaryField,
117	P: PackedField<Scalar = F>,
118	NTT: AdditiveNTT<Field = F> + Sync,
119	A: Allocator,
120{
121	let oracle_spec = &params.input_oracles()[oracle_index];
122	assert_eq!(oracle_spec.log_batch_size(), 1, "encode_masked requires log_batch_size == 1");
123	// With batch size 1, the oracle's own codeword dimension (reduced dimension minus its lift) is
124	// exactly the bare message length; the mask is interleaved internally.
125	let oracle_log_dim = params.rs_code().log_dim() - oracle_spec.log_lift;
126	assert_eq!(
127		oracle_log_dim,
128		message.log_len(),
129		"encode_masked requires the oracle's message dimension to match the message length"
130	);
131
132	// Generate random mask of equal length to message.
133	let log_len = message.log_len();
134	let packed_len = 1usize << log_len.saturating_sub(P::LOG_WIDTH);
135
136	let gen_mask_scope = tracing::debug_span!("Generate random mask").entered();
137	let mask_values =
138		par_rand::<StdRng, _, _>(packed_len, &mut rng, P::random).collect_into_alloc_vec(alloc);
139	let mask = FieldBuffer::new(log_len, mask_values);
140	drop(gen_mask_scope);
141
142	let combined_values = if log_len < P::LOG_WIDTH {
143		let combined_value =
144			P::from_scalars(iter::chain(message.iter_scalars(), mask.iter_scalars()));
145		let mut values = alloc.alloc::<P>(1);
146		values.push(combined_value);
147		values
148	} else {
149		let _scope = tracing::debug_span!("Concatenate message and mask").entered();
150		// TODO: Ideally, encoding should not allocate and copy the memory into a temp buffer.
151		// Until then the temporary at least comes from the pool rather than the global heap.
152		let mut combined_values = alloc.alloc::<P>(2 * packed_len);
153		combined_values.extend_from_slice(message.as_ref());
154		combined_values.extend_from_slice(mask.as_ref());
155		combined_values
156	};
157	let combined = FieldBuffer::new(log_len + 1, combined_values);
158
159	let codeword = encode_interleaved(params, oracle_index, ntt, combined.as_view(), alloc);
160
161	MaskedCodeword { codeword, mask }
162}
163
164#[cfg(test)]
165mod tests {
166	use binius_compute::GlobalAllocator;
167	use binius_field::{Ghash128b as B128, PackedGhash1x128b};
168	use binius_hash::StdHashSuite;
169	use binius_iop::fri::FRIParams;
170	use binius_math::{
171		ntt::{NeighborsLastSingleThread, domain_context::GaoMateerOnTheFly},
172		test_utils::random_field_buffer,
173	};
174	use rand::{SeedableRng, rngs::StdRng};
175
176	use super::*;
177	use crate::merkle_tree::prover::BinaryMerkleTreeProver;
178
179	#[test]
180	fn test_encode_masked() {
181		type F = B128;
182		type P = PackedGhash1x128b;
183
184		let mut rng = StdRng::seed_from_u64(42);
185
186		let log_dim = 6;
187		let log_inv_rate = 1;
188		let log_batch_size = 1;
189		let n_test_queries = 3;
190
191		let merkle_prover = BinaryMerkleTreeProver::<F, StdHashSuite>::new();
192
193		let domain_context = GaoMateerOnTheFly::generate(log_dim + log_inv_rate);
194		let ntt = NeighborsLastSingleThread::new(domain_context);
195
196		let params = FRIParams::with_strategy(
197			merkle_prover.scheme(),
198			log_dim + log_batch_size,
199			Some(log_batch_size),
200			log_inv_rate,
201			n_test_queries,
202			&binius_iop::fri::ConstantArityStrategy::new(2),
203		);
204
205		assert_eq!(params.log_batch_size(), 1);
206		assert_eq!(params.rs_code().log_dim(), log_dim);
207
208		let message = random_field_buffer::<P>(&mut rng, log_dim);
209
210		let output: MaskedCodeword<P> =
211			encode_masked(&params, 0, &ntt, message.as_view(), &mut rng, &GlobalAllocator);
212
213		// Verify mask has correct dimensions.
214		assert_eq!(output.mask.log_len(), log_dim);
215
216		// Verify the codeword has expected length (log_dim + log_batch_size + log_inv_rate).
217		assert_eq!(output.codeword.log_len(), log_dim + log_batch_size + log_inv_rate);
218	}
219}