1use std::marker::PhantomData;
9
10use binius_compute::Allocator;
11use binius_field::{BinaryField, PackedField};
12use getset::CopyGetters;
13
14use super::{FieldBuffer, FieldSlice, binary_subspace::BinarySubspace, ntt::AdditiveNTT};
15use crate::{
16 bit_reverse::bit_reverse_packed,
17 ntt::{DomainContext, domain_context::GaoMateerOnTheFly},
18};
19
20#[derive(Debug, Clone, CopyGetters)]
30pub struct ReedSolomonCode<F> {
31 log_dimension: usize,
32 #[get_copy = "pub"]
33 log_inv_rate: usize,
34 _marker: PhantomData<F>,
35}
36
37impl<F: BinaryField> ReedSolomonCode<F> {
38 pub const fn new(log_dimension: usize, log_inv_rate: usize) -> Self {
48 Self {
49 log_dimension,
50 log_inv_rate,
51 _marker: PhantomData,
52 }
53 }
54
55 pub fn subspace(&self) -> BinarySubspace<F> {
60 GaoMateerOnTheFly::<F>::generate(self.log_len()).subspace(self.log_len())
61 }
62
63 pub const fn dim(&self) -> usize {
65 1 << self.dim_bits()
66 }
67
68 pub const fn log_dim(&self) -> usize {
69 self.log_dimension
70 }
71
72 pub const fn log_len(&self) -> usize {
73 self.log_dimension + self.log_inv_rate
74 }
75
76 #[allow(clippy::len_without_is_empty)]
78 pub const fn len(&self) -> usize {
79 1 << (self.log_dimension + self.log_inv_rate)
80 }
81
82 const fn dim_bits(&self) -> usize {
84 self.log_dimension
85 }
86
87 pub const fn inv_rate(&self) -> usize {
89 1 << self.log_inv_rate
90 }
91
92 pub fn encode_batch<P, NTT, A>(
106 &self,
107 ntt: &NTT,
108 data: FieldSlice<'_, P>,
109 log_batch_size: usize,
110 alloc: &A,
111 ) -> FieldBuffer<P, A::Vec<P>>
112 where
113 P: PackedField<Scalar = F>,
114 NTT: AdditiveNTT<Field = F> + Sync,
115 A: Allocator,
116 {
117 assert_eq!(
118 ntt.subspace(self.log_len()),
119 self.subspace(),
120 "precondition: NTT subspace must match code subspace"
121 );
122 assert_eq!(
123 data.log_len(),
124 self.log_dim() + log_batch_size,
125 "precondition: data.log_len() must equal log_dim() + log_batch_size"
126 );
127
128 let _scope = tracing::trace_span!(
129 "Reed-Solomon encode",
130 log_len = self.log_len(),
131 log_batch_size = log_batch_size,
132 symbol_bits = F::N_BITS,
133 )
134 .entered();
135
136 let log_output_len = self.log_dim() + log_batch_size + self.log_inv_rate;
143 let mut output = FieldBuffer::from_view_with_capacity_in(alloc, data, log_output_len);
144
145 bit_reverse_packed(output.as_mut_view());
147 output.repeat_extend(log_output_len);
148
149 ntt.forward_transform(output.as_mut_view(), self.log_inv_rate, log_batch_size);
150 output
151 }
152}
153
154#[cfg(test)]
155mod tests {
156 use binius_compute::GlobalAllocator;
157 use binius_field::{BinaryField, PackedField, PackedGhash1x128b, PackedGhash4x128b};
158 use rand::{SeedableRng, rngs::StdRng};
159
160 use super::*;
161 use crate::{
162 FieldBuffer,
163 bit_reverse::reverse_bits,
164 ntt::{NeighborsLastReference, domain_context::GaoMateerPreExpanded},
165 test_utils::random_field_buffer,
166 };
167
168 fn test_encode_batch_helper<P: PackedField>(
169 log_dim: usize,
170 log_inv_rate: usize,
171 log_batch_size: usize,
172 ) where
173 P::Scalar: BinaryField,
174 {
175 let mut rng = StdRng::seed_from_u64(0);
176
177 let rs_code = ReedSolomonCode::<P::Scalar>::new(log_dim, log_inv_rate);
178
179 let domain_context = GaoMateerPreExpanded::<P::Scalar>::generate(rs_code.log_len());
181 let ntt = NeighborsLastReference {
182 domain_context: &domain_context,
183 };
184
185 let message = random_field_buffer::<P>(&mut rng, log_dim + log_batch_size);
187
188 let encoded_buffer =
190 rs_code.encode_batch(&ntt, message.as_view(), log_batch_size, &GlobalAllocator);
191
192 let mut reference_buffer = FieldBuffer::zeros(rs_code.log_len() + log_batch_size);
195 for (i, val) in message.iter_scalars().enumerate() {
196 let bits = (rs_code.log_dim() + log_batch_size) as u32;
197 reference_buffer.set(reverse_bits(i, bits), val);
198 }
199
200 ntt.forward_transform(reference_buffer.as_mut_view(), 0, log_batch_size);
202
203 assert_eq!(
205 encoded_buffer.as_ref(),
206 reference_buffer.as_ref(),
207 "encode_batch_inplace result differs from reference NTT implementation"
208 );
209 }
210
211 #[test]
212 fn test_encode_batch_above_packing_width() {
213 test_encode_batch_helper::<PackedGhash1x128b>(4, 2, 0);
215 test_encode_batch_helper::<PackedGhash1x128b>(6, 2, 1);
216 test_encode_batch_helper::<PackedGhash1x128b>(8, 3, 2);
217
218 test_encode_batch_helper::<PackedGhash4x128b>(4, 2, 0);
220 test_encode_batch_helper::<PackedGhash4x128b>(6, 2, 1);
221 test_encode_batch_helper::<PackedGhash4x128b>(8, 3, 2);
222 }
223
224 #[test]
225 fn test_encode_batch_below_packing_width() {
226 test_encode_batch_helper::<PackedGhash4x128b>(1, 2, 0);
228 }
229
230 fn test_lift_duplicate_identity_helper<P: PackedField>(
239 log_dim_small: usize,
240 log_dim_large: usize,
241 log_inv_rate: usize,
242 ) where
243 P::Scalar: BinaryField,
244 {
245 assert!(log_dim_small <= log_dim_large);
246 let eta = log_dim_large - log_dim_small;
247
248 let mut rng = StdRng::seed_from_u64(0);
249
250 let domain_context =
254 GaoMateerPreExpanded::<P::Scalar>::generate(log_dim_large + log_inv_rate);
255 let ntt = NeighborsLastReference {
256 domain_context: &domain_context,
257 };
258
259 let rs_small = ReedSolomonCode::new(log_dim_small, log_inv_rate);
260 let rs_large = ReedSolomonCode::new(log_dim_large, log_inv_rate);
261
262 let msg_small = random_field_buffer::<P>(&mut rng, log_dim_small);
264
265 let mut msg_large = FieldBuffer::<P>::zeros(log_dim_large);
268 for (i, val) in msg_small.iter_scalars().enumerate() {
269 msg_large.set(i, val);
270 }
271
272 let enc_small = rs_small.encode_batch(&ntt, msg_small.as_view(), 0, &GlobalAllocator);
273 let enc_large = rs_large.encode_batch(&ntt, msg_large.as_view(), 0, &GlobalAllocator);
274
275 let small_scalars = enc_small.iter_scalars().collect::<Vec<_>>();
276 let large_scalars = enc_large.iter_scalars().collect::<Vec<_>>();
277 assert_eq!(small_scalars.len(), 1 << (log_dim_small + log_inv_rate));
278 assert_eq!(large_scalars.len(), 1 << (log_dim_large + log_inv_rate));
279
280 for (j, &large) in large_scalars.iter().enumerate() {
281 assert_eq!(
282 large,
283 small_scalars[j >> eta],
284 "lift identity failed at index {j} (eta = {eta})"
285 );
286 }
287 }
288
289 #[test]
290 fn test_lift_duplicate_identity() {
291 test_lift_duplicate_identity_helper::<PackedGhash1x128b>(6, 6, 2);
293 test_lift_duplicate_identity_helper::<PackedGhash1x128b>(4, 6, 2);
295 test_lift_duplicate_identity_helper::<PackedGhash1x128b>(2, 8, 1);
296 test_lift_duplicate_identity_helper::<PackedGhash1x128b>(0, 4, 3);
297 test_lift_duplicate_identity_helper::<PackedGhash4x128b>(4, 8, 2);
299 }
300}