use crate::reed_solomon::{self, RsError};
#[derive(Clone, Debug)]
pub struct DegradedSlab<'a> {
pub data_shards: Vec<Option<&'a [u8]>>,
pub parity_shards: Vec<Option<&'a [u8]>>,
pub k: usize,
pub m: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RepairResult {
pub data_shards: Vec<Vec<u8>>,
pub parity_shards: Vec<Vec<u8>>,
pub reconstructed: Vec<usize>,
}
#[derive(Debug)]
pub enum RepairError {
Rs(RsError),
WrongArity { expected: usize, actual: usize },
InconsistentShardLen,
}
impl std::fmt::Display for RepairError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Rs(e) => write!(f, "repair: Reed-Solomon failure: {e}"),
Self::WrongArity { expected, actual } => {
write!(f, "repair: expected {expected} shards, got {actual}")
}
Self::InconsistentShardLen => write!(f, "repair: surviving shards disagree on length"),
}
}
}
impl std::error::Error for RepairError {}
impl From<RsError> for RepairError {
fn from(e: RsError) -> Self {
Self::Rs(e)
}
}
pub fn repair_shards(degraded: &DegradedSlab<'_>) -> Result<RepairResult, RepairError> {
if degraded.data_shards.len() != degraded.k {
return Err(RepairError::WrongArity {
expected: degraded.k,
actual: degraded.data_shards.len(),
});
}
if degraded.parity_shards.len() != degraded.m {
return Err(RepairError::WrongArity {
expected: degraded.m,
actual: degraded.parity_shards.len(),
});
}
let mut slots: Vec<Option<&[u8]>> = Vec::with_capacity(degraded.k + degraded.m);
for s in °raded.data_shards {
slots.push(*s);
}
for s in °raded.parity_shards {
slots.push(*s);
}
let missing: Vec<usize> = slots
.iter()
.enumerate()
.filter_map(|(i, s)| if s.is_none() { Some(i) } else { None })
.collect();
let data = reed_solomon::decode(&slots, degraded.k, degraded.m)?;
let data_refs: Vec<&[u8]> = data.iter().map(Vec::as_slice).collect();
let full = reed_solomon::encode(&data_refs, degraded.k, degraded.m)?;
let mut parity_shards = Vec::with_capacity(degraded.m);
for p in full.iter().skip(degraded.k) {
parity_shards.push(p.clone());
}
Ok(RepairResult {
data_shards: data,
parity_shards,
reconstructed: missing,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_shards(k: usize, len: usize) -> Vec<Vec<u8>> {
(0..k)
.map(|i| {
let base = u8::try_from(i).unwrap_or(0);
(0..len)
.map(|j| base.wrapping_add(u8::try_from(j).unwrap_or(0)))
.collect()
})
.collect()
}
fn as_refs(data: &[Vec<u8>]) -> Vec<&[u8]> {
data.iter().map(Vec::as_slice).collect()
}
#[test]
fn repair_one_missing_data_shard() {
let k = 4;
let m = 2;
let original = sample_shards(k, 32);
let encoded = reed_solomon::encode(&as_refs(&original), k, m).unwrap();
let data_shards: Vec<Option<&[u8]>> = (0..k)
.map(|i| {
if i == 2 {
None
} else {
Some(encoded[i].as_slice())
}
})
.collect();
let parity_shards: Vec<Option<&[u8]>> =
(k..k + m).map(|i| Some(encoded[i].as_slice())).collect();
let degraded = DegradedSlab {
data_shards,
parity_shards,
k,
m,
};
let repair = repair_shards(°raded).unwrap();
assert_eq!(repair.data_shards, original);
for i in 0..m {
assert_eq!(repair.parity_shards[i], encoded[k + i]);
}
assert_eq!(repair.reconstructed, vec![2]);
}
#[test]
fn repair_two_missing_parity_shards() {
let k = 3;
let m = 3;
let original = sample_shards(k, 16);
let encoded = reed_solomon::encode(&as_refs(&original), k, m).unwrap();
let data_shards: Vec<Option<&[u8]>> = (0..k).map(|i| Some(encoded[i].as_slice())).collect();
let parity_shards: Vec<Option<&[u8]>> = (0..m)
.map(|i| {
if i == 0 || i == 2 {
None
} else {
Some(encoded[k + i].as_slice())
}
})
.collect();
let degraded = DegradedSlab {
data_shards,
parity_shards,
k,
m,
};
let repair = repair_shards(°raded).unwrap();
assert_eq!(repair.data_shards, original);
for i in 0..m {
assert_eq!(repair.parity_shards[i], encoded[k + i]);
}
assert_eq!(repair.reconstructed, vec![k, k + 2]);
}
#[test]
fn repair_m_missing_shards_image_still_readable() {
let k = 4;
let m = 2;
let original = sample_shards(k, 64);
let encoded = reed_solomon::encode(&as_refs(&original), k, m).unwrap();
let data_shards: Vec<Option<&[u8]>> = (0..k)
.map(|i| {
if i == 1 {
None
} else {
Some(encoded[i].as_slice())
}
})
.collect();
let parity_shards: Vec<Option<&[u8]>> = (k..k + m)
.map(|i| {
if i == k + 1 {
None
} else {
Some(encoded[i].as_slice())
}
})
.collect();
let degraded = DegradedSlab {
data_shards,
parity_shards,
k,
m,
};
let repair = repair_shards(°raded).unwrap();
assert_eq!(repair.data_shards, original);
assert_eq!(repair.reconstructed.len(), m);
}
#[test]
fn repair_too_many_erasures_fails() {
let k = 4;
let m = 2;
let original = sample_shards(k, 32);
let encoded = reed_solomon::encode(&as_refs(&original), k, m).unwrap();
let data_shards: Vec<Option<&[u8]>> = (0..k)
.map(|i| {
if i < 3 {
None
} else {
Some(encoded[i].as_slice())
}
})
.collect();
let parity_shards: Vec<Option<&[u8]>> =
(k..k + m).map(|i| Some(encoded[i].as_slice())).collect();
let degraded = DegradedSlab {
data_shards,
parity_shards,
k,
m,
};
let err = repair_shards(°raded).unwrap_err();
assert!(matches!(
err,
RepairError::Rs(RsError::InsufficientShards { .. })
));
}
#[test]
fn repair_no_erasures_is_identity() {
let k = 3;
let m = 2;
let original = sample_shards(k, 16);
let encoded = reed_solomon::encode(&as_refs(&original), k, m).unwrap();
let data_shards: Vec<Option<&[u8]>> = (0..k).map(|i| Some(encoded[i].as_slice())).collect();
let parity_shards: Vec<Option<&[u8]>> =
(k..k + m).map(|i| Some(encoded[i].as_slice())).collect();
let degraded = DegradedSlab {
data_shards,
parity_shards,
k,
m,
};
let repair = repair_shards(°raded).unwrap();
assert_eq!(repair.data_shards, original);
for i in 0..m {
assert_eq!(repair.parity_shards[i], encoded[k + i]);
}
assert!(
repair.reconstructed.is_empty(),
"no reconstructions when nothing was missing"
);
}
#[test]
fn repair_wrong_arity_rejected() {
let k = 4;
let m = 2;
let original = sample_shards(k, 16);
let encoded = reed_solomon::encode(&as_refs(&original), k, m).unwrap();
let data_shards: Vec<Option<&[u8]>> = (0..=k)
.map(|i| Some(encoded[i.min(k + m - 1)].as_slice()))
.collect();
let parity_shards: Vec<Option<&[u8]>> =
(k..k + m).map(|i| Some(encoded[i].as_slice())).collect();
let degraded = DegradedSlab {
data_shards,
parity_shards,
k,
m,
};
let err = repair_shards(°raded).unwrap_err();
assert!(matches!(err, RepairError::WrongArity { .. }));
}
#[test]
fn repair_preserves_drop_ids() {
let k = 5;
let m = 3;
let original = sample_shards(k, 32);
let encoded = reed_solomon::encode(&as_refs(&original), k, m).unwrap();
let data_shards: Vec<Option<&[u8]>> = (0..k)
.map(|i| {
if i == 0 || i == 2 {
None
} else {
Some(encoded[i].as_slice())
}
})
.collect();
let parity_shards: Vec<Option<&[u8]>> = (k..k + m)
.map(|i| {
if i == k + 1 {
None
} else {
Some(encoded[i].as_slice())
}
})
.collect();
let degraded = DegradedSlab {
data_shards,
parity_shards,
k,
m,
};
let repair = repair_shards(°raded).unwrap();
for (i, shard) in repair.data_shards.iter().enumerate() {
let orig_hash = blake3::hash(&original[i]);
let repaired_hash = blake3::hash(shard);
assert_eq!(orig_hash, repaired_hash, "data shard {i} DropId mismatch");
}
}
}