1use p3_dft::TwoAdicSubgroupDft;
4use p3_field::{Field, TwoAdicField};
5use p3_matrix::Matrix;
6use p3_matrix::dense::RowMajorMatrix;
7
8pub trait Encoder<F: Field> {
19 fn encode_batch(&self, message: RowMajorMatrix<F>, log_inv_rate: usize) -> RowMajorMatrix<F>;
28}
29
30impl<F: TwoAdicField, D: TwoAdicSubgroupDft<F>> Encoder<F> for D {
34 fn encode_batch(
35 &self,
36 mut message: RowMajorMatrix<F>,
37 log_inv_rate: usize,
38 ) -> RowMajorMatrix<F> {
39 if log_inv_rate > 0 {
40 let len = message.values.len();
42 let padded_len = u32::try_from(log_inv_rate)
43 .ok()
44 .and_then(|rate| len.checked_shl(rate))
45 .filter(|&padded| padded >> log_inv_rate == len)
49 .expect("codeword length overflows usize");
50 let mut values = F::zero_vec(padded_len);
51 values[..len].copy_from_slice(&message.values);
52 message.values = values;
53 }
54 self.dft_batch(message).to_row_major_matrix()
55 }
56}
57
58#[cfg(test)]
59mod tests {
60 use alloc::vec;
61
62 use p3_baby_bear::BabyBear;
63 use p3_dft::{Radix2DFTSmallBatch, Radix2DitParallel, TwoAdicSubgroupDft};
64 use p3_field::PrimeCharacteristicRing;
65 use p3_matrix::Matrix;
66 use p3_matrix::dense::RowMajorMatrix;
67 use rand::SeedableRng;
68 use rand::rngs::SmallRng;
69
70 use super::Encoder;
71
72 fn check_matches_padded_dft<D: TwoAdicSubgroupDft<BabyBear>>(dft: &D) {
74 let mut rng = SmallRng::seed_from_u64(1);
75 let message = RowMajorMatrix::<BabyBear>::rand(&mut rng, 8, 3);
76
77 let mut padded = message.clone();
78 padded
79 .values
80 .resize(message.values.len() * 4, BabyBear::ZERO);
81 let expected = dft.dft_batch(padded).to_row_major_matrix();
82
83 assert_eq!(dft.encode_batch(message, 2), expected);
84 }
85
86 #[test]
87 fn small_batch_encoder_matches_padded_dft() {
88 check_matches_padded_dft(&Radix2DFTSmallBatch::<BabyBear>::default());
89 }
90
91 #[test]
92 fn dit_parallel_encoder_matches_padded_dft() {
93 check_matches_padded_dft(&Radix2DitParallel::<BabyBear>::default());
94 }
95
96 #[test]
97 #[should_panic = "codeword length overflows usize"]
98 fn encode_batch_panics_when_the_codeword_length_overflows() {
99 let message = RowMajorMatrix::<BabyBear>::new(vec![BabyBear::ZERO; 2], 1);
100 let _ = Radix2DitParallel::<BabyBear>::default()
101 .encode_batch(message, usize::BITS as usize - 1);
102 }
103}