1#[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#[cfg(feature = "alloc")]
20pub struct Decoder {
21 generation_size: usize,
22 symbol_size: usize,
23 rows: Vec<AlignedBuffer>,
25 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 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 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 pub fn generation_size(&self) -> usize {
57 self.generation_size
58 }
59 pub fn symbol_size(&self) -> usize {
61 self.symbol_size
62 }
63 pub fn rank(&self) -> usize {
65 self.rank
66 }
67 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 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 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 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 self.free_rows.push(row);
129 return Ok(false);
130 };
131
132 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 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 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 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 self.permute_rows_to_identity_pivots();
181 self.decoded = true;
182
183 Ok(Some(self.extract_symbols()))
184 }
185
186 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 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]); 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 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}