use std::collections::HashMap;
use tape_reed_solomon::ReedSolomon;
use crate::coords::LayerGeometry;
use crate::encode::EncodeParams;
use crate::error::ClayError;
use crate::transforms::{compute_c_into, compute_u_into, pft_into, prt_into};
pub type DecodeParams = EncodeParams;
pub type RsCodec = ReedSolomon;
const SET_SKEW: usize = 320;
const SKEW_MIN_ROW: usize = 4 * 1024;
pub struct NodeRows {
rows: Vec<Vec<u8>>,
row_len: usize,
skew: usize,
}
impl NodeRows {
pub fn new(nodes: usize, row_len: usize) -> Self {
let skew = if row_len >= SKEW_MIN_ROW { SET_SKEW } else { 0 };
let rows = (0..nodes)
.map(|node| vec![0u8; row_len + node * skew])
.collect();
Self { rows, row_len, skew }
}
#[inline]
fn base(&self, node: usize) -> usize {
node * self.skew
}
#[inline]
pub fn len(&self) -> usize {
self.rows.len()
}
#[inline]
pub fn row(&self, node: usize) -> &[u8] {
&self.rows[node][self.base(node)..][..self.row_len]
}
#[inline]
pub fn row_mut(&mut self, node: usize) -> &mut [u8] {
let (base, len) = (self.base(node), self.row_len);
&mut self.rows[node][base..][..len]
}
pub fn pair_mut(&mut self, first: usize, second: usize) -> (&mut [u8], &mut [u8]) {
debug_assert_ne!(first, second);
let (len, skew) = (self.row_len, self.skew);
let (a, b) = pair_mut(&mut self.rows, first, second);
(&mut a[first * skew..][..len], &mut b[second * skew..][..len])
}
pub fn iter_mut(&mut self) -> impl Iterator<Item = &mut [u8]> {
let (len, skew) = (self.row_len, self.skew);
self.rows
.iter_mut()
.enumerate()
.map(move |(node, row)| &mut row[node * skew..][..len])
}
}
struct DecodeScratch {
u_buf: NodeRows,
u_computed: Vec<bool>,
needs_mds: Vec<bool>,
}
pub fn decode(
params: &DecodeParams,
rs: &RsCodec,
available: &HashMap<usize, Vec<u8>>,
erasures: &[usize],
) -> Result<Vec<u8>, ClayError> {
if available.is_empty() && erasures.is_empty() {
return Ok(Vec::new());
}
if available.is_empty() {
return Err(ClayError::InvalidParameters(
"No available chunks provided but erasures are non-empty".into(),
));
}
if erasures.len() > params.m {
return Err(ClayError::TooManyErasures {
max: params.m,
actual: erasures.len(),
});
}
let mut iter = available.iter();
let (_, first_chunk) = iter.next().unwrap();
let chunk_size = first_chunk.len();
if chunk_size == 0 || chunk_size % params.sub_chunk_no != 0 {
return Err(ClayError::InvalidChunkSize {
expected: params.sub_chunk_no,
actual: chunk_size,
});
}
for (&idx, chunk) in iter {
if chunk.len() != chunk_size {
return Err(ClayError::InconsistentChunkSizes {
first_size: chunk_size,
mismatched_idx: idx,
mismatched_size: chunk.len(),
});
}
}
for &idx in available.keys() {
if idx >= params.n {
return Err(ClayError::InvalidParameters(format!(
"Chunk index {} out of range [0, {})",
idx, params.n
)));
}
}
for &e in erasures {
if e >= params.n {
return Err(ClayError::InvalidParameters(format!(
"Erasure index {} out of range [0, {})",
e, params.n
)));
}
}
for &e in erasures {
if available.contains_key(&e) {
return Err(ClayError::InvalidParameters(format!(
"Node {} is both in available chunks and marked as erased",
e
)));
}
}
let expected_available = params.n - erasures.len();
if available.len() != expected_available {
return Err(ClayError::InvalidParameters(format!(
"Expected {} available chunks (n={} - erasures={}), but got {}",
expected_available,
params.n,
erasures.len(),
available.len()
)));
}
for node in 0..params.n {
if !erasures.contains(&node) && !available.contains_key(&node) {
return Err(ClayError::InvalidParameters(format!(
"Node {} is neither erased nor provided in available chunks",
node
)));
}
}
let mut chunks: Vec<Option<&[u8]>> = vec![None; params.n];
for (&idx, data) in available.iter() {
chunks[idx] = Some(data.as_slice());
}
decode_rows(params, rs, &chunks)
}
pub fn decode_rows(
params: &DecodeParams,
rs: &RsCodec,
chunks: &[Option<&[u8]>],
) -> Result<Vec<u8>, ClayError> {
if chunks.len() != params.n {
return Err(ClayError::InvalidParameters(format!(
"Expected {} chunk slots (n), got {}",
params.n,
chunks.len()
)));
}
let Some(chunk_size) = chunks.iter().flatten().map(|chunk| chunk.len()).next() else {
return Err(ClayError::InvalidParameters(
"No available chunks provided".into(),
));
};
if chunk_size == 0 || chunk_size % params.sub_chunk_no != 0 {
return Err(ClayError::InvalidChunkSize {
expected: params.sub_chunk_no,
actual: chunk_size,
});
}
let mut erased_count = 0usize;
for (idx, slot) in chunks.iter().enumerate() {
match slot {
Some(chunk) if chunk.len() != chunk_size => {
return Err(ClayError::InconsistentChunkSizes {
first_size: chunk_size,
mismatched_idx: idx,
mismatched_size: chunk.len(),
})
}
Some(_) => {}
None => erased_count += 1,
}
}
if erased_count > params.m {
return Err(ClayError::TooManyErasures {
max: params.m,
actual: erased_count,
});
}
let sub_chunk_size = chunk_size / params.sub_chunk_no;
let total_nodes = params.q * params.t;
let zero_row = vec![0u8; chunk_size];
let mut available_rows: Vec<Option<&[u8]>> = vec![None; total_nodes];
let mut erased_rows: Vec<Vec<u8>> = Vec::with_capacity(total_nodes);
for internal_idx in 0..total_nodes {
if internal_idx >= params.k && internal_idx < params.k + params.nu {
available_rows[internal_idx] = Some(&zero_row);
erased_rows.push(Vec::new());
continue;
}
let external_idx = if internal_idx < params.k {
internal_idx
} else {
internal_idx - params.nu
};
match chunks[external_idx] {
Some(data) => {
available_rows[internal_idx] = Some(data);
erased_rows.push(Vec::new());
}
None => erased_rows.push(vec![0u8; chunk_size]),
}
}
decode_layered(params, rs, &available_rows, &mut erased_rows, sub_chunk_size)?;
let mut result = Vec::with_capacity(params.k * chunk_size);
for i in 0..params.k {
match available_rows[i] {
Some(row) => result.extend_from_slice(row),
None => result.extend_from_slice(&erased_rows[i]),
}
}
Ok(result)
}
pub fn decode_layered(
params: &DecodeParams,
rs: &RsCodec,
available_rows: &[Option<&[u8]>],
erased_rows: &mut [Vec<u8>],
sub_chunk_size: usize,
) -> Result<(), ClayError> {
let total_nodes = params.q * params.t;
let alpha = params.sub_chunk_no;
let geometry = LayerGeometry::new(params.q, params.t, alpha);
let mut is_erased: Vec<bool> = vec![false; total_nodes];
let mut erased_nodes: Vec<usize> = Vec::new();
for (node, row) in available_rows.iter().enumerate() {
if row.is_none() {
is_erased[node] = true;
erased_nodes.push(node);
}
}
let chunk_size = sub_chunk_size * alpha;
let mut scratch = DecodeScratch {
u_buf: NodeRows::new(total_nodes, chunk_size),
u_computed: vec![false; total_nodes * alpha],
needs_mds: vec![false; total_nodes],
};
let max_iscore = get_max_iscore(params, &erased_nodes);
let mut layers_by_iscore: Vec<Vec<usize>> = vec![Vec::new(); max_iscore + 1];
for z in 0..alpha {
let digits = geometry.plane_digits(z);
let mut iscore = 0;
for &node in &erased_nodes {
if node % params.q == digits[node / params.q] {
iscore += 1;
}
}
layers_by_iscore[iscore].push(z);
}
let mut layer_patterns: Vec<(usize, Vec<bool>, usize)> = Vec::with_capacity(alpha);
for layers in &layers_by_iscore {
if layers.is_empty() {
continue;
}
layer_patterns.clear();
for &z in layers {
let mds_count = compute_layer_u(
params,
&geometry,
&is_erased,
z,
available_rows,
&mut scratch,
sub_chunk_size,
)?;
layer_patterns.push((z, scratch.needs_mds.clone(), mds_count));
}
let is_uniform = layer_patterns.windows(2).all(|pair| pair[0].1 == pair[1].1);
debug_assert!(is_uniform, "layers in one pass disagreed on the erasure set");
if is_uniform {
let pattern = &layer_patterns[0].1;
let erased_count = layer_patterns[0].2;
let mut layers: Vec<usize> = layer_patterns.iter().map(|entry| entry.0).collect();
layers.sort_unstable();
let mut cursor = 0;
while cursor < layers.len() {
let start = layers[cursor];
let mut end = start;
while cursor + 1 < layers.len() && layers[cursor + 1] == end + 1 {
cursor += 1;
end = layers[cursor];
}
decode_uncoupled_layer(
params,
rs,
pattern,
erased_count,
start,
start * sub_chunk_size,
(end - start + 1) * sub_chunk_size,
&mut scratch.u_buf,
)?;
cursor += 1;
}
} else {
for (z, pattern, erased) in &layer_patterns {
decode_uncoupled_layer(
params,
rs,
pattern,
*erased,
*z,
z * sub_chunk_size,
sub_chunk_size,
&mut scratch.u_buf,
)?;
}
}
for (z, pattern, _) in &layer_patterns {
for (node, &needed) in pattern.iter().enumerate() {
if needed {
scratch.u_computed[node * alpha + z] = true;
}
}
}
for &z in layers {
let digits = geometry.plane_digits(z);
for &node_xy in &erased_nodes {
let x = node_xy % params.q;
let y = node_xy / params.q;
let z_y = digits[y];
let node_sw = y * params.q + z_y;
let z_sw = geometry.companion_layer(z, x, y, z_y);
let offset_z = z * sub_chunk_size;
let offset_zsw = z_sw * sub_chunk_size;
if z_y != x {
if let Some(c_sw_row) = available_rows[node_sw] {
compute_c_into(
&scratch.u_buf.row(node_xy)[offset_z..offset_z + sub_chunk_size],
&c_sw_row[offset_zsw..offset_zsw + sub_chunk_size],
&mut erased_rows[node_xy][offset_z..offset_z + sub_chunk_size],
);
} else if z_y < x {
let u_xy = &scratch.u_buf.row(node_xy)[offset_z..offset_z + sub_chunk_size];
let u_sw = &scratch.u_buf.row(node_sw)[offset_zsw..offset_zsw + sub_chunk_size];
let (c_xy_row, c_sw_row) = pair_mut(erased_rows, node_xy, node_sw);
pft_into(
u_xy,
u_sw,
&mut c_xy_row[offset_z..offset_z + sub_chunk_size],
&mut c_sw_row[offset_zsw..offset_zsw + sub_chunk_size],
);
}
} else {
erased_rows[node_xy][offset_z..offset_z + sub_chunk_size].copy_from_slice(
&scratch.u_buf.row(node_xy)[offset_z..offset_z + sub_chunk_size],
);
}
}
}
}
Ok(())
}
fn compute_layer_u(
params: &DecodeParams,
geometry: &LayerGeometry,
is_erased: &[bool],
z: usize,
available_rows: &[Option<&[u8]>],
scratch: &mut DecodeScratch,
sub_chunk_size: usize,
) -> Result<usize, ClayError> {
let digits = geometry.plane_digits(z);
let alpha = params.sub_chunk_no;
let u_buf = &mut scratch.u_buf;
let u_computed = &mut scratch.u_computed;
scratch.needs_mds.copy_from_slice(is_erased);
let mut mds_count = 0;
for &erased in is_erased {
if erased {
mds_count += 1;
}
}
for x in 0..params.q {
for y in 0..params.t {
let node_xy = params.q * y + x;
let c_row = match available_rows[node_xy] {
Some(c_row) => c_row,
None => continue,
};
let z_y = digits[y];
let node_sw = params.q * y + z_y;
let z_sw = geometry.companion_layer(z, x, y, z_y);
let offset_z = z * sub_chunk_size;
let offset_zsw = z_sw * sub_chunk_size;
if z_y == x {
u_buf.row_mut(node_xy)[offset_z..offset_z + sub_chunk_size]
.copy_from_slice(&c_row[offset_z..offset_z + sub_chunk_size]);
u_computed[node_xy * alpha + z] = true;
} else if let Some(c_sw_row) = available_rows[node_sw] {
if z_y < x {
let c_xy = &c_row[offset_z..offset_z + sub_chunk_size];
let c_sw = &c_sw_row[offset_zsw..offset_zsw + sub_chunk_size];
let (u_xy_buf, u_sw_buf) = u_buf.pair_mut(node_xy, node_sw);
prt_into(
c_xy,
c_sw,
&mut u_xy_buf[offset_z..offset_z + sub_chunk_size],
&mut u_sw_buf[offset_zsw..offset_zsw + sub_chunk_size],
);
u_computed[node_xy * alpha + z] = true;
u_computed[node_sw * alpha + z_sw] = true;
}
} else if u_computed[node_sw * alpha + z_sw] {
let (u_xy_buf, u_sw_buf) = u_buf.pair_mut(node_xy, node_sw);
compute_u_into(
&c_row[offset_z..offset_z + sub_chunk_size],
&u_sw_buf[offset_zsw..offset_zsw + sub_chunk_size],
&mut u_xy_buf[offset_z..offset_z + sub_chunk_size],
);
u_computed[node_xy * alpha + z] = true;
} else {
scratch.needs_mds[node_xy] = true;
mds_count += 1;
}
}
}
Ok(mds_count)
}
pub fn decode_uncoupled_layer(
params: &DecodeParams,
rs: &RsCodec,
is_erased: &[bool],
erased_count: usize,
z: usize,
offset: usize,
sub_chunk_size: usize,
u_buf: &mut NodeRows,
) -> Result<(), ClayError> {
let parity_start = params.original_count;
if erased_count > params.m {
return Err(ClayError::TooManyErasures {
max: params.m,
actual: erased_count,
});
}
if erased_count == 0 {
return Ok(());
}
let mut has_erased_originals = false;
for &erased in &is_erased[..parity_start] {
if erased {
has_erased_originals = true;
break;
}
}
let mut shards: Vec<(&mut [u8], bool)> = Vec::with_capacity(u_buf.len());
for (i, node_buf) in u_buf.iter_mut().enumerate() {
shards.push((&mut node_buf[offset..offset + sub_chunk_size], !is_erased[i]));
}
if has_erased_originals {
rs.reconstruct(&mut shards).map_err(|e| {
ClayError::ReconstructionFailed(format!("Layer {} RS reconstruct failed: {:?}", z, e))
})?;
} else {
let mut slices: Vec<&mut [u8]> = Vec::with_capacity(shards.len());
for (slice, _) in shards {
slices.push(slice);
}
rs.encode(&mut slices).map_err(|e| {
ClayError::ReconstructionFailed(format!("Layer {} RS encode failed: {:?}", z, e))
})?;
}
Ok(())
}
pub fn pair_mut(
buffers: &mut [Vec<u8>],
first: usize,
second: usize,
) -> (&mut Vec<u8>, &mut Vec<u8>) {
debug_assert_ne!(first, second);
if first < second {
let (left, right) = buffers.split_at_mut(second);
(&mut left[first], &mut right[0])
} else {
let (left, right) = buffers.split_at_mut(first);
(&mut right[0], &mut left[second])
}
}
fn get_max_iscore(params: &DecodeParams, erased_nodes: &[usize]) -> usize {
let mut y_seen = vec![false; params.t];
let mut iscore = 0;
for &node in erased_nodes {
let y = node / params.q;
if !y_seen[y] {
y_seen[y] = true;
iscore += 1;
}
}
iscore
}
#[cfg(test)]
mod tests {
use super::*;
fn test_params() -> DecodeParams {
DecodeParams {
k: 4,
m: 2,
n: 6,
q: 2,
t: 3,
nu: 0,
sub_chunk_no: 8,
original_count: 4,
}
}
fn test_rs(params: &DecodeParams) -> RsCodec {
RsCodec::new(params.original_count, params.m).expect("test params should build a codec")
}
#[test]
fn decode_empty() {
let params = test_params();
let rs = test_rs(¶ms);
let available: HashMap<usize, Vec<u8>> = HashMap::new();
let result = decode(¶ms, &rs, &available, &[]);
assert!(result.is_ok());
assert!(result.unwrap().is_empty());
}
#[test]
fn max_iscore() {
let params = test_params();
assert_eq!(get_max_iscore(¶ms, &[]), 0);
assert_eq!(get_max_iscore(¶ms, &[0]), 1);
assert_eq!(get_max_iscore(¶ms, &[0, 1]), 1);
assert_eq!(get_max_iscore(¶ms, &[0, 2]), 2);
}
}