use tape_core::encoding::{EncodingProfile, EncodingType};
use tape_core::erasure::GROUP_SIZE;
use tape_core::types::SpoolIndex;
use tape_slicer::{
ClayCoder, ErasureCoder, ReedSolomonCoder, Slicer, SliceMetadata,
};
use crate::error::DownloadError;
pub struct BlobDecoder {
profile: EncodingProfile,
basic: Option<ReedSolomonCoder>,
clay: Option<Slicer<ClayCoder>>,
}
impl Default for BlobDecoder {
fn default() -> Self {
Self::new()
}
}
impl BlobDecoder {
pub fn new() -> Self {
Self::with_profile(EncodingProfile::clay_default())
}
pub fn with_profile(profile: EncodingProfile) -> Self {
let encoding_type = profile.encoding_type().unwrap_or(EncodingType::Unknown);
let mut decoder = Self {
profile,
basic: None,
clay: None,
};
match encoding_type {
EncodingType::Basic => {
let params = profile.rs_params();
decoder.basic = Some(ReedSolomonCoder::new(params.k() as usize, params.m() as usize));
}
EncodingType::Clay | EncodingType::Unknown => {
decoder.clay = Some(Slicer::with_profile(
ClayCoder::from_params(profile.clay_params()),
true, profile,
));
}
}
decoder
}
pub fn with_encoding(encoding_type: EncodingType) -> Self {
let profile = match encoding_type {
EncodingType::Basic => EncodingProfile::basic_default(),
EncodingType::Clay | EncodingType::Unknown => EncodingProfile::clay_default(),
};
Self::with_profile(profile)
}
pub fn encoding_type(&self) -> EncodingType {
self.profile.encoding_type().unwrap_or(EncodingType::Unknown)
}
pub fn profile(&self) -> EncodingProfile {
self.profile
}
fn min_slices_from_metadata(&self, slices: &[(SpoolIndex, Vec<u8>)]) -> Result<usize, DownloadError> {
match self.encoding_type() {
EncodingType::Clay => {
slices.first()
.and_then(|(_, data)| SliceMetadata::from_slice(data).ok())
.map(|meta| meta.profile().clay_params().k() as usize)
.ok_or_else(|| DownloadError::Decoding(
"Cannot determine k: no valid slice metadata".to_string()
))
}
EncodingType::Basic => Ok(self.profile.rs_params().k() as usize),
EncodingType::Unknown => Err(DownloadError::Decoding(
"Cannot decode with Unknown encoding type".to_string()
)),
}
}
fn decode_internal(&mut self, chunks: &[(usize, &[u8])]) -> Result<Vec<u8>, DownloadError> {
match self.encoding_type() {
EncodingType::Basic => {
let result = self.basic.as_mut().unwrap()
.decode(chunks)
.map_err(|e| DownloadError::Decoding(e.to_string()))?;
Ok(result)
}
EncodingType::Clay | EncodingType::Unknown => {
self.clay.as_mut().unwrap()
.decode(chunks)
.map_err(|e| DownloadError::Decoding(e.to_string()))
}
}
}
pub fn decode(&mut self, slices: Vec<(SpoolIndex, Vec<u8>)>) -> Result<Vec<u8>, DownloadError> {
let min_slices = self.min_slices_from_metadata(&slices)?;
if slices.len() < min_slices {
return Err(DownloadError::InsufficientSlices {
got: slices.len(),
need: min_slices,
});
}
for &(idx, _) in &slices {
if idx.as_usize() >= GROUP_SIZE {
return Err(DownloadError::InvalidSliceIndex(idx));
}
}
let chunks: Vec<(usize, &[u8])> = slices
.iter()
.map(|(idx, data)| (idx.as_usize(), data.as_slice()))
.collect();
self.decode_internal(&chunks)
}
pub fn decode_to_blob(&mut self, slices: Vec<(SpoolIndex, Vec<u8>)>) -> Result<Vec<u8>, DownloadError> {
self.decode(slices)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::codec::encoder::BlobEncoder;
fn test_encoder() -> BlobEncoder {
BlobEncoder::with_encoding(EncodingType::Basic)
}
fn test_decoder() -> BlobDecoder {
BlobDecoder::with_encoding(EncodingType::Basic)
}
#[test]
fn test_roundtrip() {
let original = vec![0xAB; 20_000];
let mut encoder = test_encoder();
let slices = encoder.encode(original.clone()).unwrap();
let mut decoder = test_decoder();
let recovered = decoder.decode(slices).unwrap();
assert_eq!(&recovered[..original.len()], &original);
}
#[test]
fn test_decode_with_only_data_slices() {
let original = vec![0xCD; 15_000];
let mut encoder = test_encoder();
let k = encoder.profile().k() as usize;
let slices = encoder.encode(original.clone()).unwrap();
let data_only: Vec<_> = slices.into_iter().take(k).collect();
let mut decoder = test_decoder();
let recovered = decoder.decode(data_only).unwrap();
assert_eq!(&recovered[..original.len()], &original);
}
#[test]
fn test_decode_with_missing_parity() {
let original = vec![0xEF; 20_000];
let mut encoder = test_encoder();
let k = encoder.profile().k() as usize;
let slices = encoder.encode(original.clone()).unwrap();
let data_only: Vec<_> = slices.into_iter().take(k).collect();
let mut decoder = test_decoder();
let recovered = decoder.decode(data_only).unwrap();
assert_eq!(&recovered[..original.len()], &original);
}
#[test]
fn test_decode_with_scattered_slices() {
let original = vec![0x12; 10_000];
let mut encoder = test_encoder();
let k = encoder.profile().k() as usize;
let slices = encoder.encode(original.clone()).unwrap();
let scattered: Vec<_> = slices
.into_iter()
.enumerate()
.filter(|(i, _)| i % 2 == 0)
.map(|(_, s)| s)
.collect();
assert!(scattered.len() >= k);
let mut decoder = test_decoder();
let recovered = decoder.decode(scattered).unwrap();
assert_eq!(&recovered[..original.len()], &original);
}
#[test]
fn test_decode_not_enough_slices() {
let original = vec![0x34; 10_000];
let mut encoder = test_encoder();
let k = encoder.profile().k() as usize;
let slices = encoder.encode(original).unwrap();
let too_few: Vec<_> = slices.into_iter().take(k - 1).collect();
let mut decoder = test_decoder();
let result = decoder.decode(too_few);
assert!(matches!(
result,
Err(DownloadError::InsufficientSlices { .. })
));
}
#[test]
fn test_decode_invalid_slice_index() {
let original = vec![0x99; 10_000];
let mut encoder = test_encoder();
let mut slices: Vec<_> = encoder.encode(original).unwrap();
slices[0].0 = SpoolIndex::from(9999);
let mut decoder = test_decoder();
let result = decoder.decode(slices);
assert!(matches!(
result,
Err(DownloadError::InvalidSliceIndex(idx)) if idx == SpoolIndex::from(9999)
));
}
#[test]
fn test_decode_empty_blob() {
let original = vec![];
let mut encoder = test_encoder();
let slices = encoder.encode(original.clone()).unwrap();
let mut decoder = test_decoder();
let recovered = decoder.decode(slices).unwrap();
assert!(recovered.iter().all(|&b| b == 0));
}
#[test]
fn test_decode_to_blob() {
let original = vec![0x56; 15_000];
let mut encoder = test_encoder();
let slices = encoder.encode(original.clone()).unwrap();
let mut decoder = test_decoder();
let blob = decoder.decode_to_blob(slices).unwrap();
assert_eq!(&blob[..original.len()], &original);
}
#[test]
fn test_encoding_type_default() {
let decoder = BlobDecoder::new();
assert_eq!(decoder.encoding_type(), EncodingType::Clay);
}
#[test]
fn test_encoding_type_basic() {
let decoder = BlobDecoder::with_encoding(EncodingType::Basic);
assert_eq!(decoder.encoding_type(), EncodingType::Basic);
}
#[test]
fn test_encoding_type_clay() {
let decoder = BlobDecoder::with_encoding(EncodingType::Clay);
assert_eq!(decoder.encoding_type(), EncodingType::Clay);
}
#[test]
fn test_clay_roundtrip() {
let original = vec![0xAB; 10_000];
let mut encoder = BlobEncoder::with_encoding(EncodingType::Clay);
let mut decoder = BlobDecoder::with_encoding(EncodingType::Clay);
let slices = encoder.encode(original.clone()).unwrap();
let recovered = decoder.decode(slices).unwrap();
assert_eq!(original, recovered);
}
#[test]
fn test_clay_decode_with_missing_slices() {
let original = vec![0xCD; 50_000];
let mut encoder = BlobEncoder::with_encoding(EncodingType::Clay);
let mut decoder = BlobDecoder::with_encoding(EncodingType::Clay);
let slices = encoder.encode(original.clone()).unwrap();
let k = encoder.profile().k() as usize;
let partial: Vec<_> = slices.into_iter().take(k).collect();
let recovered = decoder.decode(partial).unwrap();
assert_eq!(original, recovered);
}
}