Skip to main content

rlnc_simdx/
decoder.rs

1//! RLNC Decoder — in-place pivot-based Gaussian elimination over GF(2⁸).
2//!
3//! Coefficients and payloads are stored separately. Incoming coefficients are
4//! reduced first so dependent packets never touch their large payload, while
5//! innovative packet buffers move directly into column-ordered pivot storage.
6
7#[cfg(feature = "alloc")]
8extern crate alloc;
9#[cfg(feature = "alloc")]
10use alloc::vec::Vec;
11
12use crate::aligned::AlignedBuffer;
13use crate::encoder::CodedPacket;
14use crate::error::RlncError;
15use crate::field::tables::{EXP, LOG};
16use crate::kernel;
17
18/// RLNC Decoder.
19#[cfg(feature = "alloc")]
20pub struct Decoder {
21    generation_size: usize,
22    symbol_size: usize,
23    coefficient_rows: Vec<AlignedBuffer>,
24    payload_rows: Vec<AlignedBuffer>,
25    pivot_col: Vec<Option<usize>>,
26    elimination_factors: Vec<u8>,
27    rank: usize,
28    decoded: bool,
29}
30
31#[cfg(feature = "alloc")]
32impl Decoder {
33    /// Create a new decoder for a generation.
34    pub fn new(generation_size: usize, symbol_size: usize) -> Result<Self, RlncError> {
35        if generation_size == 0 || symbol_size == 0 {
36            return Err(RlncError::InvalidParameters);
37        }
38        let k = generation_size;
39        let mut coefficient_rows = Vec::new();
40        let mut payload_rows = Vec::new();
41        let mut pivot_col = Vec::new();
42        let mut elimination_factors = Vec::new();
43        coefficient_rows
44            .try_reserve_exact(k)
45            .map_err(|_| RlncError::InvalidParameters)?;
46        payload_rows
47            .try_reserve_exact(k)
48            .map_err(|_| RlncError::InvalidParameters)?;
49        pivot_col
50            .try_reserve_exact(k)
51            .map_err(|_| RlncError::InvalidParameters)?;
52        elimination_factors
53            .try_reserve_exact(k)
54            .map_err(|_| RlncError::InvalidParameters)?;
55        pivot_col.resize(k, None);
56        elimination_factors.resize(k, 0);
57        Ok(Decoder {
58            generation_size,
59            symbol_size,
60            coefficient_rows,
61            payload_rows,
62            pivot_col,
63            elimination_factors,
64            rank: 0,
65            decoded: false,
66        })
67    }
68
69    /// Generation size (`k`).
70    pub fn generation_size(&self) -> usize {
71        self.generation_size
72    }
73    /// Symbol size in bytes.
74    pub fn symbol_size(&self) -> usize {
75        self.symbol_size
76    }
77    /// Current rank (innovative packets received).
78    pub fn rank(&self) -> usize {
79        self.rank
80    }
81    /// True when rank == `generation_size`.
82    pub fn is_complete(&self) -> bool {
83        self.rank == self.generation_size
84    }
85
86    /// Receive a coded packet.
87    ///
88    /// Returns `true` if the packet was innovative (increased rank).
89    ///
90    /// Innovative packet allocations are moved directly into pivot storage.
91    /// Dependent packets are rejected after coefficient-only reduction, without
92    /// payload-sized arithmetic or additional allocation.
93    pub fn receive(&mut self, pkt: CodedPacket) -> Result<bool, RlncError> {
94        let k = self.generation_size;
95        let n = self.symbol_size;
96
97        if pkt.coefficients.len() != k || pkt.payload.len() != n {
98            return Err(RlncError::PacketSizeMismatch {
99                expected_coeffs: k,
100                got_coeffs: pkt.coefficients.len(),
101                expected_payload: n,
102                got_payload: pkt.payload.len(),
103            });
104        }
105
106        if self.is_complete() {
107            return Ok(false);
108        }
109
110        let CodedPacket {
111            mut coefficients,
112            mut payload,
113        } = pkt;
114
115        // Reduce the small coefficient vector first and record the operations.
116        // A dependent packet can then be rejected without touching its payload.
117        for r in 0..self.rank {
118            let Some(col) = self.pivot_col[r] else {
119                continue;
120            };
121            let coeff = coefficients.as_slice()[col];
122            self.elimination_factors[r] = coeff;
123            if coeff == 0 {
124                continue;
125            }
126            // SAFETY: stored and incoming packet buffers are distinct owned
127            // allocations and the suffix lengths match.
128            unsafe {
129                kernel::axpy_unchecked(
130                    coeff,
131                    &self.coefficient_rows[r].as_slice()[col..],
132                    &mut coefficients.as_mut_slice()[col..],
133                );
134            }
135        }
136
137        let new_pivot = coefficients.as_slice().iter().position(|&b| b != 0);
138        let Some(pivot_col) = new_pivot else {
139            return Ok(false);
140        };
141
142        for r in 0..self.rank {
143            let coeff = self.elimination_factors[r];
144            if coeff != 0 {
145                // SAFETY: stored and incoming payloads are distinct and have
146                // the decoder's validated symbol length.
147                unsafe {
148                    kernel::axpy_unchecked(
149                        coeff,
150                        self.payload_rows[r].as_slice(),
151                        payload.as_mut_slice(),
152                    );
153                }
154            }
155        }
156
157        let pivot_val = coefficients.as_slice()[pivot_col];
158        if pivot_val != 1 {
159            let inv = EXP[255 - LOG[pivot_val as usize] as usize];
160            kernel::scale_inplace(inv, &mut coefficients.as_mut_slice()[pivot_col..]);
161            kernel::scale_inplace(inv, payload.as_mut_slice());
162        }
163
164        // Keep pivot rows ordered. Moving AlignedBuffer handles is cheap and
165        // makes the full-rank pivot order the identity, avoiding a decode-time
166        // O(k^2) selection sort and preserving the echelon invariant.
167        let insert_at = self.pivot_col[..self.rank]
168            .iter()
169            .position(|&col| col.is_some_and(|col| col > pivot_col))
170            .unwrap_or(self.rank);
171
172        self.coefficient_rows.push(coefficients);
173        self.payload_rows.push(payload);
174        self.pivot_col[self.rank] = Some(pivot_col);
175        for i in (insert_at..self.rank).rev() {
176            self.coefficient_rows.swap(i, i + 1);
177            self.payload_rows.swap(i, i + 1);
178            self.pivot_col.swap(i, i + 1);
179        }
180        self.rank += 1;
181        self.decoded = false;
182
183        debug_assert!(self.pivot_col[..self.rank]
184            .windows(2)
185            .all(|pair| pair[0] < pair[1]));
186
187        Ok(true)
188    }
189
190    /// Attempt to decode. Returns `Some(symbols)` when rank == `generation_size`.
191    pub fn decode(&mut self) -> Result<Option<Vec<Vec<u8>>>, RlncError> {
192        if !self.is_complete() {
193            return Ok(None);
194        }
195        if self.decoded {
196            return Ok(Some(self.extract_symbols()));
197        }
198
199        let k = self.generation_size;
200
201        // There are k distinct, ordered pivots drawn from k columns, so their
202        // only possible full-rank order is 0..k.
203        debug_assert!(self
204            .pivot_col
205            .iter()
206            .enumerate()
207            .all(|(col, &pivot)| pivot == Some(col)));
208
209        // Back-substitution — split_at_mut, no allocation
210        for r in (0..k).rev() {
211            let Some(col) = self.pivot_col[r] else {
212                continue;
213            };
214            let (coefficients_above, pivot_coefficients) = self.coefficient_rows.split_at_mut(r);
215            let coefficient_suffix = &pivot_coefficients[0].as_slice()[col..];
216            let (payloads_above, pivot_payloads) = self.payload_rows.split_at_mut(r);
217            let pivot_payload = pivot_payloads[0].as_slice();
218            for r2 in 0..r {
219                let coeff = coefficients_above[r2].as_slice()[col];
220                if coeff == 0 {
221                    continue;
222                }
223                // SAFETY: split_at_mut separates the pivot and destination
224                // rows; corresponding coefficient suffixes and payloads match.
225                unsafe {
226                    kernel::axpy_unchecked(
227                        coeff,
228                        coefficient_suffix,
229                        &mut coefficients_above[r2].as_mut_slice()[col..],
230                    );
231                    kernel::axpy_unchecked(coeff, pivot_payload, payloads_above[r2].as_mut_slice());
232                }
233            }
234        }
235
236        self.decoded = true;
237
238        Ok(Some(self.extract_symbols()))
239    }
240
241    fn extract_symbols(&self) -> Vec<Vec<u8>> {
242        self.payload_rows
243            .iter()
244            .map(AlignedBuffer::to_vec)
245            .collect()
246    }
247}
248
249#[cfg(test)]
250#[cfg(feature = "alloc")]
251mod tests {
252    use super::*;
253    use crate::encoder::{Encoder, SimpleRng};
254
255    fn make_source(k: usize, n: usize) -> Vec<Vec<u8>> {
256        (0..k)
257            .map(|i| (0..n).map(|j| (i * 7 + j * 3) as u8).collect())
258            .collect()
259    }
260
261    #[test]
262    fn encode_decode_round_trip() {
263        let k = 4usize;
264        let n = 64usize;
265        let source = make_source(k, n);
266        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
267
268        let enc = Encoder::new(k, n).unwrap();
269        let mut dec = Decoder::new(k, n).unwrap();
270        let mut rng = SimpleRng::new(0xDEAD_BEEF);
271
272        let mut innovative = 0;
273        for _ in 0..k + 2 {
274            let pkt = enc.encode_random(&refs, &mut rng).unwrap();
275            if dec.receive(pkt).unwrap() {
276                innovative += 1;
277            }
278        }
279        assert_eq!(innovative, k);
280        assert!(dec.is_complete());
281
282        let decoded = dec.decode().unwrap().unwrap();
283        assert_eq!(decoded.len(), k);
284        for i in 0..k {
285            assert_eq!(decoded[i], source[i], "symbol {i} mismatch");
286        }
287    }
288
289    #[test]
290    fn systematic_decode() {
291        let k = 3usize;
292        let n = 32usize;
293        let source = make_source(k, n);
294        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
295
296        let enc = Encoder::new(k, n).unwrap();
297        let mut dec = Decoder::new(k, n).unwrap();
298        for i in 0..k {
299            let pkt = enc.encode_systematic(&refs, i).unwrap();
300            assert!(dec.receive(pkt).unwrap());
301        }
302        assert!(dec.is_complete());
303        let decoded = dec.decode().unwrap().unwrap();
304        for i in 0..k {
305            assert_eq!(decoded[i], source[i]);
306        }
307    }
308
309    #[test]
310    fn redundant_packet_ignored() {
311        let k = 2usize;
312        let n = 8usize;
313        let source = make_source(k, n);
314        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
315
316        let enc = Encoder::new(k, n).unwrap();
317        let mut dec = Decoder::new(k, n).unwrap();
318        let pkt0 = enc.encode_systematic(&refs, 0).unwrap();
319        let pkt0_dup = enc.encode_systematic(&refs, 0).unwrap();
320        assert!(dec.receive(pkt0).unwrap());
321        assert!(!dec.receive(pkt0_dup).unwrap());
322        assert_eq!(dec.rank(), 1);
323    }
324
325    #[test]
326    fn decoder_rows_are_aligned() {
327        use crate::aligned::ALIGN;
328        let k = 4usize;
329        let n = 128usize;
330        let source = make_source(k, n);
331        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
332        let encoder = Encoder::new(k, n).unwrap();
333        let mut dec = Decoder::new(k, n).unwrap();
334        assert!(dec
335            .receive(encoder.encode_systematic(&refs, 0).unwrap())
336            .unwrap());
337        for (i, row) in dec.coefficient_rows.iter().enumerate() {
338            assert_eq!(
339                row.as_ptr() as usize % ALIGN,
340                0,
341                "decoder coefficient row {i} not {ALIGN}-byte aligned"
342            );
343        }
344        for (i, row) in dec.payload_rows.iter().enumerate() {
345            assert_eq!(
346                row.as_ptr() as usize % ALIGN,
347                0,
348                "decoder payload row {i} not {ALIGN}-byte aligned"
349            );
350        }
351    }
352
353    #[test]
354    fn redundant_packet_does_not_add_storage() {
355        let k = 2usize;
356        let n = 16usize;
357        let source = make_source(k, n);
358        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
359        let enc = Encoder::new(k, n).unwrap();
360        let mut dec = Decoder::new(k, n).unwrap();
361        let p = enc.encode_systematic(&refs, 0).unwrap();
362        assert!(dec.receive(p).unwrap());
363        let rows_before = dec.payload_rows.len();
364        let p2 = enc.encode_systematic(&refs, 0).unwrap();
365        assert!(!dec.receive(p2).unwrap());
366        assert_eq!(dec.payload_rows.len(), rows_before);
367    }
368
369    #[test]
370    fn new_rejects_zero_params() {
371        assert!(Decoder::new(0, 8).is_err());
372        assert!(Decoder::new(4, 0).is_err());
373    }
374
375    #[test]
376    fn receive_rejects_packet_size_mismatch() {
377        let mut dec = Decoder::new(2, 4).unwrap();
378        let bad = CodedPacket::from_slices(&[1], &[1, 2, 3, 4]); // wrong coeff len
379        let err = dec.receive(bad).unwrap_err();
380        match err {
381            crate::error::RlncError::PacketSizeMismatch {
382                expected_coeffs: 2,
383                got_coeffs: 1,
384                expected_payload: 4,
385                got_payload: 4,
386            } => {}
387            other => panic!("unexpected {other:?}"),
388        }
389    }
390
391    #[test]
392    fn decode_none_when_incomplete() {
393        let k = 3usize;
394        let n = 8usize;
395        let source = make_source(k, n);
396        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
397        let enc = Encoder::new(k, n).unwrap();
398        let mut dec = Decoder::new(k, n).unwrap();
399        // Only one systematic packet
400        let pkt = enc.encode_systematic(&refs, 0).unwrap();
401        assert!(dec.receive(pkt).unwrap());
402        assert!(!dec.is_complete());
403        assert_eq!(dec.rank(), 1);
404        let out = dec.decode().unwrap();
405        assert!(out.is_none(), "decode must be None before full rank");
406    }
407
408    #[test]
409    fn receive_after_complete_returns_false() {
410        let k = 2usize;
411        let n = 8usize;
412        let source = make_source(k, n);
413        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
414        let enc = Encoder::new(k, n).unwrap();
415        let mut dec = Decoder::new(k, n).unwrap();
416        for i in 0..k {
417            assert!(dec
418                .receive(enc.encode_systematic(&refs, i).unwrap())
419                .unwrap());
420        }
421        assert!(dec.is_complete());
422        let extra = enc.encode_systematic(&refs, 0).unwrap();
423        assert!(!dec.receive(extra).unwrap());
424    }
425}