1use 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
13pub 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 = ¶ms.input_oracles()[oracle_index];
50 let log_batch_size = oracle_spec.log_batch_size();
51 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 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#[derive(Debug)]
84pub struct MaskedCodeword<P: PackedField, Data: Deref<Target = [P]> = Vec<P>> {
85 pub codeword: FieldBuffer<P, Data>,
87 pub mask: FieldBuffer<P, Data>,
89}
90
91pub 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 = ¶ms.input_oracles()[oracle_index];
122 assert_eq!(oracle_spec.log_batch_size(), 1, "encode_masked requires log_batch_size == 1");
123 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 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 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(¶ms, 0, &ntt, message.as_view(), &mut rng, &GlobalAllocator);
212
213 assert_eq!(output.mask.log_len(), log_dim);
215
216 assert_eq!(output.codeword.log_len(), log_dim + log_batch_size + log_inv_rate);
218 }
219}