Skip to main content

rlnc_simdx/
decoder.rs

1//! RLNC Decoder — in-place pivot-based Gaussian elimination over GF(2⁸).
2//!
3//! Working rows are recycled from a free-list (no heap alloc in the steady-state
4//! receive path). Pivot reorder after back-substitution is done in-place via
5//! swaps (no second matrix allocation).
6
7#[cfg(feature = "alloc")]
8extern crate alloc;
9#[cfg(feature = "alloc")]
10use alloc::{vec, 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    /// Augmented rows: `[coefficients (k bytes) | payload (symbol_size bytes)]`.
24    rows: Vec<AlignedBuffer>,
25    /// Recycled working rows for `receive` (same length as matrix rows).
26    free_rows: Vec<AlignedBuffer>,
27    pivot_col: Vec<Option<usize>>,
28    rank: usize,
29    decoded: bool,
30}
31
32#[cfg(feature = "alloc")]
33impl Decoder {
34    /// Create a new decoder for a generation.
35    pub fn new(generation_size: usize, symbol_size: usize) -> Result<Self, RlncError> {
36        if generation_size == 0 || symbol_size == 0 {
37            return Err(RlncError::InvalidParameters);
38        }
39        let k = generation_size;
40        let row_len = k + symbol_size;
41        let rows = (0..k).map(|_| AlignedBuffer::zeroed(row_len)).collect();
42        // Pre-seed one free row so first `receive` need not allocate.
43        let free_rows = vec![AlignedBuffer::zeroed(row_len)];
44        Ok(Decoder {
45            generation_size,
46            symbol_size,
47            rows,
48            free_rows,
49            pivot_col: vec![None; k],
50            rank: 0,
51            decoded: false,
52        })
53    }
54
55    /// Generation size (`k`).
56    pub fn generation_size(&self) -> usize {
57        self.generation_size
58    }
59    /// Symbol size in bytes.
60    pub fn symbol_size(&self) -> usize {
61        self.symbol_size
62    }
63    /// Current rank (innovative packets received).
64    pub fn rank(&self) -> usize {
65        self.rank
66    }
67    /// True when rank == `generation_size`.
68    pub fn is_complete(&self) -> bool {
69        self.rank == self.generation_size
70    }
71
72    fn row_len(&self) -> usize {
73        self.generation_size + self.symbol_size
74    }
75
76    /// Take a working row from the free-list (or allocate if empty).
77    fn take_work_row(&mut self) -> AlignedBuffer {
78        let mut row = self
79            .free_rows
80            .pop()
81            .unwrap_or_else(|| AlignedBuffer::zeroed(self.row_len()));
82        row.as_mut_slice().fill(0);
83        row
84    }
85
86    /// Receive a coded packet.
87    ///
88    /// Returns `true` if the packet was innovative (increased rank).
89    ///
90    /// Steady-state hot path: **no heap allocation** (free-list row + in-place
91    /// [`kernel::scale_inplace`] with SIMD).
92    pub fn receive(&mut self, pkt: CodedPacket) -> Result<bool, RlncError> {
93        let k = self.generation_size;
94        let n = self.symbol_size;
95
96        if pkt.coefficients.len() != k || pkt.payload.len() != n {
97            return Err(RlncError::PacketSizeMismatch {
98                expected_coeffs: k,
99                got_coeffs: pkt.coefficients.len(),
100                expected_payload: n,
101                got_payload: pkt.payload.len(),
102            });
103        }
104
105        if self.is_complete() {
106            return Ok(false);
107        }
108
109        let mut row = self.take_work_row();
110        row.as_mut_slice()[..k].copy_from_slice(pkt.coefficients.as_slice());
111        row.as_mut_slice()[k..].copy_from_slice(pkt.payload.as_slice());
112
113        // Forward elimination — pivot rows and `row` are different buffers.
114        for r in 0..self.rank {
115            let Some(col) = self.pivot_col[r] else {
116                continue;
117            };
118            let coeff = row.as_slice()[col];
119            if coeff == 0 {
120                continue;
121            }
122            kernel::axpy(coeff, self.rows[r].as_slice(), row.as_mut_slice());
123        }
124
125        let new_pivot = row.as_slice()[..k].iter().position(|&b| b != 0);
126        let Some(pivot_col) = new_pivot else {
127            // Linearly dependent — recycle row
128            self.free_rows.push(row);
129            return Ok(false);
130        };
131
132        // Normalise pivot to 1 — SIMD in-place scale.
133        let pivot_val = row.as_slice()[pivot_col];
134        if pivot_val != 1 {
135            let inv = EXP[255 - LOG[pivot_val as usize] as usize];
136            if inv != 1 {
137                kernel::scale_inplace(inv, row.as_mut_slice());
138            }
139        }
140
141        // Install into matrix; previous placeholder row can be recycled.
142        let old = core::mem::replace(&mut self.rows[self.rank], row);
143        self.free_rows.push(old);
144        self.pivot_col[self.rank] = Some(pivot_col);
145        self.rank += 1;
146        self.decoded = false;
147
148        Ok(true)
149    }
150
151    /// Attempt to decode. Returns `Some(symbols)` when rank == `generation_size`.
152    pub fn decode(&mut self) -> Result<Option<Vec<Vec<u8>>>, RlncError> {
153        if !self.is_complete() {
154            return Ok(None);
155        }
156        if self.decoded {
157            return Ok(Some(self.extract_symbols()));
158        }
159
160        let k = self.generation_size;
161
162        // Back-substitution — split_at_mut, no allocation
163        for r in (0..k).rev() {
164            let Some(col) = self.pivot_col[r] else {
165                continue;
166            };
167            for r2 in 0..r {
168                let coeff = self.rows[r2].as_slice()[col];
169                if coeff == 0 {
170                    continue;
171                }
172                let (lo, hi) = self.rows.split_at_mut(r);
173                let pivot_slice: &[u8] = hi[0].as_slice();
174                kernel::axpy(coeff, pivot_slice, lo[r2].as_mut_slice());
175            }
176        }
177
178        // In-place permutation so that row i has pivot column i (cycle following).
179        // pivot_col[r] was the pivot of the r-th innovative row before reorder.
180        self.permute_rows_to_identity_pivots();
181        self.decoded = true;
182
183        Ok(Some(self.extract_symbols()))
184    }
185
186    /// In-place reorder: selection-sort rows by pivot column (O(k²), k small).
187    /// After full rank, row `i` has pivot column `i`. No second matrix allocation.
188    fn permute_rows_to_identity_pivots(&mut self) {
189        let k = self.generation_size;
190        for i in 0..k {
191            let mut best = i;
192            let mut best_col = self.pivot_col[i].unwrap_or(usize::MAX);
193            for j in (i + 1)..k {
194                let c = self.pivot_col[j].unwrap_or(usize::MAX);
195                if c < best_col {
196                    best = j;
197                    best_col = c;
198                }
199            }
200            if best != i {
201                self.rows.swap(i, best);
202                self.pivot_col.swap(i, best);
203            }
204        }
205        for i in 0..k {
206            self.pivot_col[i] = Some(i);
207        }
208    }
209
210    fn extract_symbols(&self) -> Vec<Vec<u8>> {
211        let k = self.generation_size;
212        let n = self.symbol_size;
213        self.rows
214            .iter()
215            .map(|row| row.as_slice()[k..k + n].to_vec())
216            .collect()
217    }
218}
219
220#[cfg(test)]
221#[cfg(feature = "alloc")]
222mod tests {
223    use super::*;
224    use crate::encoder::{Encoder, SimpleRng};
225
226    fn make_source(k: usize, n: usize) -> Vec<Vec<u8>> {
227        (0..k)
228            .map(|i| (0..n).map(|j| (i * 7 + j * 3) as u8).collect())
229            .collect()
230    }
231
232    #[test]
233    fn encode_decode_round_trip() {
234        let k = 4usize;
235        let n = 64usize;
236        let source = make_source(k, n);
237        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
238
239        let enc = Encoder::new(k, n).unwrap();
240        let mut dec = Decoder::new(k, n).unwrap();
241        let mut rng = SimpleRng::new(0xDEAD_BEEF);
242
243        let mut innovative = 0;
244        for _ in 0..k + 2 {
245            let pkt = enc.encode_random(&refs, &mut rng).unwrap();
246            if dec.receive(pkt).unwrap() {
247                innovative += 1;
248            }
249        }
250        assert_eq!(innovative, k);
251        assert!(dec.is_complete());
252
253        let decoded = dec.decode().unwrap().unwrap();
254        assert_eq!(decoded.len(), k);
255        for i in 0..k {
256            assert_eq!(decoded[i], source[i], "symbol {i} mismatch");
257        }
258    }
259
260    #[test]
261    fn systematic_decode() {
262        let k = 3usize;
263        let n = 32usize;
264        let source = make_source(k, n);
265        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
266
267        let enc = Encoder::new(k, n).unwrap();
268        let mut dec = Decoder::new(k, n).unwrap();
269        for i in 0..k {
270            let pkt = enc.encode_systematic(&refs, i).unwrap();
271            assert!(dec.receive(pkt).unwrap());
272        }
273        assert!(dec.is_complete());
274        let decoded = dec.decode().unwrap().unwrap();
275        for i in 0..k {
276            assert_eq!(decoded[i], source[i]);
277        }
278    }
279
280    #[test]
281    fn redundant_packet_ignored() {
282        let k = 2usize;
283        let n = 8usize;
284        let source = make_source(k, n);
285        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
286
287        let enc = Encoder::new(k, n).unwrap();
288        let mut dec = Decoder::new(k, n).unwrap();
289        let pkt0 = enc.encode_systematic(&refs, 0).unwrap();
290        let pkt0_dup = enc.encode_systematic(&refs, 0).unwrap();
291        assert!(dec.receive(pkt0).unwrap());
292        assert!(!dec.receive(pkt0_dup).unwrap());
293        assert_eq!(dec.rank(), 1);
294    }
295
296    #[test]
297    fn decoder_rows_are_aligned() {
298        use crate::aligned::ALIGN;
299        let k = 4usize;
300        let n = 128usize;
301        let dec = Decoder::new(k, n).unwrap();
302        for (i, row) in dec.rows.iter().enumerate() {
303            assert_eq!(
304                row.as_ptr() as usize % ALIGN,
305                0,
306                "decoder row {i} not {ALIGN}-byte aligned"
307            );
308        }
309    }
310
311    #[test]
312    fn free_list_recycles_on_redundant() {
313        let k = 2usize;
314        let n = 16usize;
315        let source = make_source(k, n);
316        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
317        let enc = Encoder::new(k, n).unwrap();
318        let mut dec = Decoder::new(k, n).unwrap();
319        let free_before = dec.free_rows.len();
320        let p = enc.encode_systematic(&refs, 0).unwrap();
321        assert!(dec.receive(p).unwrap());
322        let p2 = enc.encode_systematic(&refs, 0).unwrap();
323        assert!(!dec.receive(p2).unwrap());
324        // Redundant receive returns row to free list
325        assert!(dec.free_rows.len() >= free_before);
326    }
327
328    #[test]
329    fn new_rejects_zero_params() {
330        assert!(Decoder::new(0, 8).is_err());
331        assert!(Decoder::new(4, 0).is_err());
332    }
333
334    #[test]
335    fn receive_rejects_packet_size_mismatch() {
336        let mut dec = Decoder::new(2, 4).unwrap();
337        let bad = CodedPacket::from_slices(&[1], &[1, 2, 3, 4]); // wrong coeff len
338        let err = dec.receive(bad).unwrap_err();
339        match err {
340            crate::error::RlncError::PacketSizeMismatch {
341                expected_coeffs: 2,
342                got_coeffs: 1,
343                expected_payload: 4,
344                got_payload: 4,
345            } => {}
346            other => panic!("unexpected {other:?}"),
347        }
348    }
349
350    #[test]
351    fn decode_none_when_incomplete() {
352        let k = 3usize;
353        let n = 8usize;
354        let source = make_source(k, n);
355        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
356        let enc = Encoder::new(k, n).unwrap();
357        let mut dec = Decoder::new(k, n).unwrap();
358        // Only one systematic packet
359        let pkt = enc.encode_systematic(&refs, 0).unwrap();
360        assert!(dec.receive(pkt).unwrap());
361        assert!(!dec.is_complete());
362        assert_eq!(dec.rank(), 1);
363        let out = dec.decode().unwrap();
364        assert!(out.is_none(), "decode must be None before full rank");
365    }
366
367    #[test]
368    fn receive_after_complete_returns_false() {
369        let k = 2usize;
370        let n = 8usize;
371        let source = make_source(k, n);
372        let refs: Vec<&[u8]> = source.iter().map(Vec::as_slice).collect();
373        let enc = Encoder::new(k, n).unwrap();
374        let mut dec = Decoder::new(k, n).unwrap();
375        for i in 0..k {
376            assert!(dec
377                .receive(enc.encode_systematic(&refs, i).unwrap())
378                .unwrap());
379        }
380        assert!(dec.is_complete());
381        let extra = enc.encode_systematic(&refs, 0).unwrap();
382        assert!(!dec.receive(extra).unwrap());
383    }
384}