use alloc::vec;
use alloc::vec::Vec;
use crate::error::{CodeError, MatrixError, RecoverError};
use crate::kernel::Kernels;
use crate::matrix::Matrix;
#[derive(Debug, Clone)]
pub struct DecodePlan {
gftbls: Vec<u8>,
survivors: Vec<usize>,
rebuild: Vec<usize>,
n: usize,
}
#[derive(Debug, Clone)]
pub struct Coder {
matrix: Matrix,
gftbls: Vec<u8>,
kernels: Kernels,
}
impl Coder {
pub fn new(matrix: Matrix) -> Result<Self, MatrixError> {
Self::with_kernels(matrix, Kernels::scalar())
}
pub fn with_kernels(matrix: Matrix, kernels: Kernels) -> Result<Self, MatrixError> {
if matrix.rows() <= matrix.cols() {
return Err(MatrixError::Dimensions {
k: matrix.cols(),
p: matrix.rows().saturating_sub(matrix.cols()),
});
}
let gftbls = (kernels.init)(matrix.parity_bytes());
Ok(Self {
matrix,
gftbls,
kernels,
})
}
pub fn kernels(&self) -> &Kernels {
&self.kernels
}
pub fn k(&self) -> usize {
self.matrix.cols()
}
pub fn p(&self) -> usize {
self.matrix.rows() - self.matrix.cols()
}
pub fn matrix(&self) -> &Matrix {
&self.matrix
}
pub fn gftbls(&self) -> &[u8] {
&self.gftbls
}
fn check_data(&self, data: &[&[u8]], len: usize) -> Result<(), CodeError> {
if data.len() != self.k() {
return Err(CodeError::ShardCount {
expected: self.k(),
got: data.len(),
});
}
for (index, d) in data.iter().enumerate() {
if d.len() != len {
return Err(CodeError::ShardLength {
index,
expected: len,
got: d.len(),
});
}
}
Ok(())
}
pub fn encode(&self, data: &[&[u8]], parity: &mut [&mut [u8]]) -> Result<(), CodeError> {
if parity.len() != self.p() {
return Err(CodeError::ShardCount {
expected: self.p(),
got: parity.len(),
});
}
let len = parity.first().map_or(0, |b| b.len());
for (index, b) in parity.iter().enumerate() {
if b.len() != len {
return Err(CodeError::ShardLength {
index: self.k() + index,
expected: len,
got: b.len(),
});
}
}
self.check_data(data, len)?;
(self.kernels.encode)(&self.gftbls, data, parity);
Ok(())
}
pub fn update(
&self,
shard_index: usize,
data: &[u8],
parity: &mut [&mut [u8]],
) -> Result<(), CodeError> {
let k = self.k();
if shard_index >= k {
return Err(CodeError::ShardIndex {
index: shard_index,
k,
});
}
if parity.len() != self.p() {
return Err(CodeError::ShardCount {
expected: self.p(),
got: parity.len(),
});
}
for (index, b) in parity.iter().enumerate() {
if b.len() != data.len() {
return Err(CodeError::ShardLength {
index: k + index,
expected: data.len(),
got: b.len(),
});
}
}
(self.kernels.update)(&self.gftbls, k, shard_index, data, parity);
Ok(())
}
pub fn verify(&self, data: &[&[u8]], parity: &[&[u8]]) -> Result<bool, CodeError> {
if parity.len() != self.p() {
return Err(CodeError::ShardCount {
expected: self.p(),
got: parity.len(),
});
}
let len = parity.first().map_or(0, |b| b.len());
for (index, b) in parity.iter().enumerate() {
if b.len() != len {
return Err(CodeError::ShardLength {
index: self.k() + index,
expected: len,
got: b.len(),
});
}
}
self.check_data(data, len)?;
let p = self.p();
if len == 0 {
return Ok(true);
}
let mut scratch = vec![0u8; p * len];
{
let mut rows: Vec<&mut [u8]> = scratch.chunks_mut(len).collect();
(self.kernels.encode)(&self.gftbls, data, &mut rows);
}
for (l, expect) in parity.iter().enumerate() {
if &scratch[l * len..(l + 1) * len] != *expect {
return Ok(false);
}
}
Ok(true)
}
pub fn decode_plan(
&self,
present: &[bool],
rebuild: &[usize],
) -> Result<DecodePlan, RecoverError> {
let k = self.k();
let n = self.matrix.rows();
if present.len() != n {
return Err(CodeError::ShardCount {
expected: n,
got: present.len(),
}
.into());
}
for &x in rebuild {
if x >= n {
return Err(CodeError::ShardIndex { index: x, k: n }.into());
}
}
let mut survivors: Vec<usize> = Vec::with_capacity(k);
let mut have = 0usize;
for (i, &ok) in present.iter().enumerate() {
if ok {
have += 1;
if survivors.len() < k {
survivors.push(i);
}
}
}
if survivors.len() < k {
return Err(RecoverError::TooManyMissing {
missing: n - have,
p: self.p(),
});
}
let b = self.matrix.select_rows(&survivors)?;
let d = b.invert()?;
let mut coeffs = vec![0u8; rebuild.len() * k];
for (r, &x) in rebuild.iter().enumerate() {
let row = &mut coeffs[r * k..(r + 1) * k];
if x < k {
for (t, c) in row.iter_mut().enumerate() {
*c = d.get(x, t).expect("in range");
}
} else {
for (t, c) in row.iter_mut().enumerate() {
let mut s = 0u8;
for j in 0..k {
s ^= crate::gf::mul(
self.matrix.get(x, j).expect("in range"),
d.get(j, t).expect("in range"),
);
}
*c = s;
}
}
}
Ok(DecodePlan {
gftbls: (self.kernels.init)(&coeffs),
survivors,
rebuild: rebuild.to_vec(),
n,
})
}
pub fn recover_with(
&self,
plan: &DecodePlan,
shards: &[Option<&[u8]>],
out: &mut [&mut [u8]],
) -> Result<(), RecoverError> {
if shards.len() != plan.n {
return Err(CodeError::ShardCount {
expected: plan.n,
got: shards.len(),
}
.into());
}
if out.len() != plan.rebuild.len() {
return Err(CodeError::ShardCount {
expected: plan.rebuild.len(),
got: out.len(),
}
.into());
}
let mut src: Vec<&[u8]> = Vec::with_capacity(plan.survivors.len());
let len = out.first().map_or(0, |b| b.len());
for &i in &plan.survivors {
let s = shards[i].ok_or(RecoverError::TooManyMissing {
missing: 1,
p: self.p(),
})?;
if s.len() != len {
return Err(CodeError::ShardLength {
index: i,
expected: len,
got: s.len(),
}
.into());
}
src.push(s);
}
for b in out.iter() {
if b.len() != len {
return Err(CodeError::ShardLength {
index: 0,
expected: len,
got: b.len(),
}
.into());
}
}
(self.kernels.encode)(&plan.gftbls, &src, out);
Ok(())
}
pub fn recover(
&self,
shards: &[Option<&[u8]>],
rebuild: &[usize],
out: &mut [&mut [u8]],
) -> Result<(), RecoverError> {
let k = self.k();
let n = self.matrix.rows();
if shards.len() != n {
return Err(CodeError::ShardCount {
expected: n,
got: shards.len(),
}
.into());
}
if rebuild.len() != out.len() {
return Err(CodeError::ShardCount {
expected: rebuild.len(),
got: out.len(),
}
.into());
}
for &x in rebuild {
if x >= n {
return Err(CodeError::ShardIndex { index: x, k: n }.into());
}
}
let mut survivors: Vec<usize> = Vec::with_capacity(k);
let mut present = 0usize;
for (i, s) in shards.iter().enumerate() {
if s.is_some() {
present += 1;
if survivors.len() < k {
survivors.push(i);
}
}
}
if survivors.len() < k {
return Err(RecoverError::TooManyMissing {
missing: n - present,
p: self.p(),
});
}
let len = shards[survivors[0]].expect("survivor is present").len();
for (index, s) in shards.iter().enumerate() {
if let Some(s) = s
&& s.len() != len
{
return Err(CodeError::ShardLength {
index,
expected: len,
got: s.len(),
}
.into());
}
}
for (i, b) in out.iter().enumerate() {
if b.len() != len {
return Err(CodeError::ShardLength {
index: rebuild[i],
expected: len,
got: b.len(),
}
.into());
}
}
let b = self.matrix.select_rows(&survivors)?;
let d = b.invert()?;
let src: Vec<&[u8]> = survivors
.iter()
.map(|&i| shards[i].expect("survivor is present"))
.collect();
let mut coeffs = vec![0u8; rebuild.len() * k];
for (r, &x) in rebuild.iter().enumerate() {
let row = &mut coeffs[r * k..(r + 1) * k];
if x < k {
for (t, c) in row.iter_mut().enumerate() {
*c = d.get(x, t).expect("in range");
}
} else {
for (t, c) in row.iter_mut().enumerate() {
let mut s = 0u8;
for j in 0..k {
s ^= crate::gf::mul(
self.matrix.get(x, j).expect("in range"),
d.get(j, t).expect("in range"),
);
}
*c = s;
}
}
}
let gftbls = (self.kernels.init)(&coeffs);
(self.kernels.encode)(&gftbls, &src, out);
Ok(())
}
}