1#[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#[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 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 pub fn generation_size(&self) -> usize {
71 self.generation_size
72 }
73 pub fn symbol_size(&self) -> usize {
75 self.symbol_size
76 }
77 pub fn rank(&self) -> usize {
79 self.rank
80 }
81 pub fn is_complete(&self) -> bool {
83 self.rank == self.generation_size
84 }
85
86 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 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 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 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 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 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 debug_assert!(self
204 .pivot_col
205 .iter()
206 .enumerate()
207 .all(|(col, &pivot)| pivot == Some(col)));
208
209 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 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]); 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 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}