use std::num::{NonZeroU16, NonZeroUsize};
use std::sync::OnceLock;
use commonware_codec::{Decode, Encode};
use commonware_coding::{CodecConfig, Scheme as _};
use commonware_cryptography::Blake3;
use commonware_parallel::{Rayon, Sequential, Strategy};
use crate::hash::{self, HASH_LEN, Hash};
pub use commonware_coding::Config;
pub use commonware_parallel::{Rayon as ParallelStrategy, Sequential as SequentialStrategy};
type RsScheme = commonware_coding::ReedSolomon<Blake3>;
type Commitment = <RsScheme as commonware_coding::Scheme>::Commitment;
type RsChunk = <RsScheme as commonware_coding::Scheme>::Shard;
pub const PARALLEL_STRATEGY_THRESHOLD: usize = 4 * 1024 * 1024;
fn should_use_parallel_strategy(pack_len: usize) -> bool {
pack_len >= PARALLEL_STRATEGY_THRESHOLD
}
fn shared_parallel_strategy() -> Option<Rayon> {
static POOL: OnceLock<Option<Rayon>> = OnceLock::new();
POOL.get_or_init(|| {
let threads = std::thread::available_parallelism().map_or(1, NonZeroUsize::get);
NonZeroUsize::new(threads).and_then(|n| Rayon::new(n).ok())
})
.clone()
}
fn default_parallel_strategy_for_len(pack_len: usize) -> Option<Rayon> {
if should_use_parallel_strategy(pack_len) {
shared_parallel_strategy()
} else {
None
}
}
const MAX_SHARD_BYTES: usize = 4 * 1024 * 1024 * 1024;
pub const SHARD_SIZE_THRESHOLD: u64 = 1024 * 1024;
pub const MANIFEST_MAGIC: [u8; 4] = *b"MKSH";
pub const MANIFEST_VERSION: u8 = 0x02;
const MANIFEST_PROLOGUE_LEN: usize = 5;
pub const MANIFEST_MAX_BYTES: usize = 1024 * 1024;
#[must_use]
pub fn default_config() -> Config {
Config {
minimum_shards: NonZeroU16::new(16).expect("16 != 0"),
extra_shards: NonZeroU16::new(4).expect("4 != 0"),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Shard {
pub index: u16,
pub bytes: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ShardSet {
pub pack_hash: Hash,
pub config: Config,
pub shard_hashes: Vec<Hash>,
pub commitment: Hash,
}
#[derive(Debug, thiserror::Error)]
pub enum ShardError {
#[error("reed-solomon encode failed: {0}")]
EncodeFailed(String),
#[error("reed-solomon decode failed: {0}")]
DecodeFailed(String),
#[error("shard codec decode failed at index {index}: {source}")]
ShardCodecFailed {
index: u16,
#[source]
source: commonware_codec::Error,
},
#[error("shard {index} BLAKE3 mismatch (manifest tampered or shard corrupted)")]
ShardHashMismatch { index: u16 },
#[error("shard index {index} is out of range for config (total = {total})")]
IndexOutOfRange { index: u16, total: u32 },
#[error("duplicate shard index {index}")]
DuplicateIndex { index: u16 },
#[error(
"manifest has {actual} shard_hashes, expected {expected} \
(config.total_shards())"
)]
ManifestShardCountMismatch { actual: usize, expected: usize },
#[error("reconstructed pack hash does not match manifest.pack_hash")]
PackHashMismatch,
#[error("insufficient shards: {provided} < {minimum}")]
InsufficientShards { provided: usize, minimum: u16 },
#[error("invalid manifest prologue: {0}")]
InvalidManifestPrologue(&'static str),
#[error("unexpected eof while decoding manifest")]
ManifestUnexpectedEof,
#[error("trailing bytes after manifest body")]
ManifestTrailingBytes,
#[error("manifest declares zero shard count (min={minimum}, extra={extra})")]
ManifestZeroShardCount { minimum: u16, extra: u16 },
#[error("manifest is too large: {actual} > {max}")]
ManifestTooLarge { actual: usize, max: usize },
}
pub fn encode_pack_to_shards(
pack: &[u8],
config: Config,
) -> Result<(Vec<Shard>, ShardSet), ShardError> {
match default_parallel_strategy_for_len(pack.len()) {
Some(strategy) => encode_pack_to_shards_with_strategy(pack, config, &strategy),
None => encode_pack_to_shards_with_strategy(pack, config, &Sequential),
}
}
pub fn encode_pack_to_shards_with_strategy<S: Strategy>(
pack: &[u8],
config: Config,
strategy: &S,
) -> Result<(Vec<Shard>, ShardSet), ShardError> {
let (commitment, chunks) = RsScheme::encode(&config, pack, strategy)
.map_err(|e| ShardError::EncodeFailed(format!("{e:?}")))?;
let total = config.total_shards() as usize;
debug_assert_eq!(chunks.len(), total);
let results: Vec<(u16, Vec<u8>, Hash)> =
strategy.map_collect_vec(chunks.into_iter().enumerate(), |(i, chunk)| {
let index = u16::try_from(i).expect("commonware emits <= u16::MAX shards");
let bytes = chunk.encode().to_vec();
let h = hash::hash(&bytes);
(index, bytes, h)
});
let mut shards = Vec::with_capacity(total);
let mut shard_hashes = Vec::with_capacity(total);
for (index, bytes, h) in results {
shards.push(Shard { index, bytes });
shard_hashes.push(h);
}
let manifest = ShardSet {
pack_hash: hash::hash(pack),
config,
shard_hashes,
commitment: digest_to_bytes(&commitment),
};
Ok((shards, manifest))
}
pub fn decode_pack_from_shards(
shards: &[Shard],
manifest: &ShardSet,
) -> Result<Vec<u8>, ShardError> {
let size_hint: usize = shards.iter().map(|s| s.bytes.len()).sum();
match default_parallel_strategy_for_len(size_hint) {
Some(strategy) => decode_pack_from_shards_with_strategy(shards, manifest, &strategy),
None => decode_pack_from_shards_with_strategy(shards, manifest, &Sequential),
}
}
pub fn decode_pack_from_shards_with_strategy<S: Strategy>(
shards: &[Shard],
manifest: &ShardSet,
strategy: &S,
) -> Result<Vec<u8>, ShardError> {
let total = manifest.config.total_shards();
if manifest.shard_hashes.len() != total as usize {
return Err(ShardError::ManifestShardCountMismatch {
actual: manifest.shard_hashes.len(),
expected: total as usize,
});
}
let minimum = manifest.config.minimum_shards.get();
let commitment = bytes_to_digest(&manifest.commitment);
let codec_cfg = CodecConfig {
maximum_shard_size: MAX_SHARD_BYTES,
};
let mut seen = vec![false; total as usize];
let mut checked = Vec::with_capacity(shards.len());
for shard in shards {
if u32::from(shard.index) >= total {
return Err(ShardError::IndexOutOfRange {
index: shard.index,
total,
});
}
let slot = &mut seen[shard.index as usize];
if *slot {
return Err(ShardError::DuplicateIndex { index: shard.index });
}
*slot = true;
let expected = &manifest.shard_hashes[shard.index as usize];
if &hash::hash(&shard.bytes) != expected {
return Err(ShardError::ShardHashMismatch { index: shard.index });
}
let chunk = RsChunk::decode_cfg(shard.bytes.as_slice(), &codec_cfg).map_err(|e| {
ShardError::ShardCodecFailed {
index: shard.index,
source: e,
}
})?;
let checked_shard = RsScheme::check(&manifest.config, &commitment, shard.index, &chunk)
.map_err(|e| ShardError::DecodeFailed(format!("check({}): {e:?}", shard.index)))?;
checked.push(checked_shard);
}
if checked.len() < usize::from(minimum) {
return Err(ShardError::InsufficientShards {
provided: checked.len(),
minimum,
});
}
let pack = RsScheme::decode(&manifest.config, &commitment, checked.iter(), strategy)
.map_err(|e| ShardError::DecodeFailed(format!("{e:?}")))?;
if hash::hash(&pack) != manifest.pack_hash {
return Err(ShardError::PackHashMismatch);
}
Ok(pack)
}
fn digest_to_bytes(d: &Commitment) -> [u8; HASH_LEN] {
let slice: &[u8] = d.as_ref();
let mut out = [0u8; HASH_LEN];
out.copy_from_slice(slice);
out
}
fn bytes_to_digest(b: &[u8; HASH_LEN]) -> Commitment {
use commonware_codec::FixedSize;
debug_assert_eq!(<Commitment as FixedSize>::SIZE, HASH_LEN);
Commitment::from(*b)
}
pub fn encode_manifest(manifest: &ShardSet) -> Result<Vec<u8>, ShardError> {
let total = manifest.config.total_shards() as usize;
if manifest.shard_hashes.len() != total {
return Err(ShardError::ManifestShardCountMismatch {
actual: manifest.shard_hashes.len(),
expected: total,
});
}
let body_len = MANIFEST_PROLOGUE_LEN + HASH_LEN + 2 + 2 + HASH_LEN + 4 + total * HASH_LEN;
let mut out = Vec::with_capacity(body_len);
out.extend_from_slice(&MANIFEST_MAGIC);
out.push(MANIFEST_VERSION);
out.extend_from_slice(&manifest.pack_hash);
out.extend_from_slice(&manifest.config.minimum_shards.get().to_le_bytes());
out.extend_from_slice(&manifest.config.extra_shards.get().to_le_bytes());
out.extend_from_slice(&manifest.commitment);
out.extend_from_slice(
&u32::try_from(total)
.expect("total_shards fits in u32")
.to_le_bytes(),
);
for h in &manifest.shard_hashes {
out.extend_from_slice(h);
}
debug_assert_eq!(out.len(), body_len);
Ok(out)
}
pub fn decode_manifest(bytes: &[u8]) -> Result<ShardSet, ShardError> {
if bytes.len() > MANIFEST_MAX_BYTES {
return Err(ShardError::ManifestTooLarge {
actual: bytes.len(),
max: MANIFEST_MAX_BYTES,
});
}
if bytes.len() < MANIFEST_PROLOGUE_LEN {
return Err(ShardError::InvalidManifestPrologue(
"input shorter than prologue",
));
}
if bytes[..4] != MANIFEST_MAGIC {
return Err(ShardError::InvalidManifestPrologue("bad magic"));
}
if bytes[4] != MANIFEST_VERSION {
if bytes[4] == 0x01 {
return Err(ShardError::InvalidManifestPrologue(
"manifest version 0x01 (Sha256-era) — re-shard with a current mkit",
));
}
return Err(ShardError::InvalidManifestPrologue("unsupported version"));
}
let mut pos = MANIFEST_PROLOGUE_LEN;
if bytes.len() - pos < HASH_LEN {
return Err(ShardError::ManifestUnexpectedEof);
}
let mut pack_hash = [0u8; HASH_LEN];
pack_hash.copy_from_slice(&bytes[pos..pos + HASH_LEN]);
pos += HASH_LEN;
if bytes.len() - pos < 4 {
return Err(ShardError::ManifestUnexpectedEof);
}
let minimum = u16::from_le_bytes([bytes[pos], bytes[pos + 1]]);
let extra = u16::from_le_bytes([bytes[pos + 2], bytes[pos + 3]]);
pos += 4;
let minimum_nz =
NonZeroU16::new(minimum).ok_or(ShardError::ManifestZeroShardCount { minimum, extra })?;
let extra_nz =
NonZeroU16::new(extra).ok_or(ShardError::ManifestZeroShardCount { minimum, extra })?;
let config = Config {
minimum_shards: minimum_nz,
extra_shards: extra_nz,
};
let total = config.total_shards();
if bytes.len() - pos < HASH_LEN {
return Err(ShardError::ManifestUnexpectedEof);
}
let mut commitment = [0u8; HASH_LEN];
commitment.copy_from_slice(&bytes[pos..pos + HASH_LEN]);
pos += HASH_LEN;
if bytes.len() - pos < 4 {
return Err(ShardError::ManifestUnexpectedEof);
}
let declared_len =
u32::from_le_bytes([bytes[pos], bytes[pos + 1], bytes[pos + 2], bytes[pos + 3]]);
pos += 4;
if declared_len != total {
return Err(ShardError::ManifestShardCountMismatch {
actual: declared_len as usize,
expected: total as usize,
});
}
if (declared_len as usize).saturating_mul(HASH_LEN) > bytes.len() - pos {
return Err(ShardError::ManifestUnexpectedEof);
}
let mut shard_hashes = Vec::with_capacity(declared_len as usize);
for _ in 0..declared_len {
let mut h = [0u8; HASH_LEN];
h.copy_from_slice(&bytes[pos..pos + HASH_LEN]);
pos += HASH_LEN;
shard_hashes.push(h);
}
if pos != bytes.len() {
return Err(ShardError::ManifestTrailingBytes);
}
Ok(ShardSet {
pack_hash,
config,
shard_hashes,
commitment,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn synthetic_pack(bytes: usize) -> Vec<u8> {
let mut x: u64 = 0x9E37_79B9_7F4A_7C15;
let mut out = Vec::with_capacity(bytes);
while out.len() < bytes {
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
out.extend_from_slice(&x.to_le_bytes());
}
out.truncate(bytes);
out
}
#[derive(Clone, Debug)]
struct CountingStrategy {
calls: std::sync::Arc<std::sync::atomic::AtomicUsize>,
}
impl CountingStrategy {
fn new() -> Self {
Self {
calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
}
}
fn calls(&self) -> usize {
self.calls.load(std::sync::atomic::Ordering::SeqCst)
}
}
impl Strategy for CountingStrategy {
fn manual(&self) -> commonware_parallel::Manual<Self> {
commonware_parallel::Manual::new(self.clone(), std::num::NonZeroUsize::new(1).unwrap())
}
fn spawn<F, T>(&self, f: F) -> impl core::future::Future<Output = T> + Send + 'static
where
F: FnOnce(Self) -> T + Send + 'static,
T: Send + 'static,
{
let result = f(self.clone());
async move { result }
}
fn run<R, SEQ, PAR>(&self, _len: usize, serial: SEQ, _parallel: PAR) -> R
where
R: Send,
SEQ: FnOnce() -> R + Send,
PAR: FnOnce() -> R + Send,
{
serial()
}
fn try_run<R, E, SEQ, PAR>(&self, _len: usize, serial: SEQ, _parallel: PAR) -> Result<R, E>
where
R: Send,
E: Send,
SEQ: FnOnce() -> Result<R, E> + Send,
PAR: FnOnce() -> Result<R, E> + Send,
{
serial()
}
fn fold_init<I, INIT, T, R, ID, F, RD>(
&self,
iter: I,
init: INIT,
identity: ID,
fold_op: F,
reduce_op: RD,
) -> R
where
I: IntoIterator<IntoIter: Send, Item: Send> + Send,
INIT: Fn() -> T + Send + Sync,
T: Send,
R: Send,
ID: Fn() -> R + Send + Sync,
F: Fn(R, &mut T, I::Item) -> R + Send + Sync,
RD: Fn(R, R) -> R + Send + Sync,
{
self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Sequential.fold_init(iter, init, identity, fold_op, reduce_op)
}
fn try_fold<I, R, E, ID, F, RD>(
&self,
iter: I,
identity: ID,
fold_op: F,
reduce_op: RD,
) -> Result<R, E>
where
I: IntoIterator<IntoIter: Send, Item: Send> + Send,
R: Send,
E: Send,
ID: Fn() -> R + Send + Sync,
F: Fn(R, I::Item) -> Result<R, E> + Send + Sync,
RD: Fn(R, R) -> R + Send + Sync,
{
Sequential.try_fold(iter, identity, fold_op, reduce_op)
}
fn join<A, B, RA, RB>(&self, a: A, b: B) -> (RA, RB)
where
A: FnOnce() -> RA + Send,
B: FnOnce() -> RB + Send,
RA: Send,
RB: Send,
{
Sequential.join(a, b)
}
fn sort_by<T, C>(&self, items: &mut [T], compare: C)
where
T: Send,
C: Fn(&T, &T) -> std::cmp::Ordering + Send + Sync,
{
Sequential.sort_by(items, compare);
}
}
#[test]
fn explicit_strategy_is_actually_exercised_by_encode_and_decode() {
let pack = synthetic_pack(64 * 1024);
let config = default_config();
let spy = CountingStrategy::new();
let (shards, manifest) = encode_pack_to_shards_with_strategy(&pack, config, &spy).unwrap();
let calls_after_encode = spy.calls();
assert!(
calls_after_encode > 0,
"encode_pack_to_shards_with_strategy never invoked the supplied strategy"
);
let subset: Vec<Shard> = shards.into_iter().take(16).collect();
let recovered = decode_pack_from_shards_with_strategy(&subset, &manifest, &spy).unwrap();
assert_eq!(recovered, pack);
assert!(
spy.calls() > calls_after_encode,
"decode_pack_from_shards_with_strategy never invoked the supplied strategy"
);
}
#[test]
fn round_trip_with_explicit_parallel_strategy() {
let pack = synthetic_pack(256 * 1024);
let config = default_config();
let strategy = Rayon::new(NonZeroUsize::new(2).unwrap()).expect("build rayon pool");
let (shards, manifest) =
encode_pack_to_shards_with_strategy(&pack, config, &strategy).unwrap();
let subset: Vec<Shard> = shards.into_iter().take(16).collect();
let recovered =
decode_pack_from_shards_with_strategy(&subset, &manifest, &strategy).unwrap();
assert_eq!(recovered, pack);
}
#[test]
fn default_strategy_selection_is_a_runtime_threshold_not_a_const() {
assert!(!should_use_parallel_strategy(0));
assert!(!should_use_parallel_strategy(
PARALLEL_STRATEGY_THRESHOLD - 1
));
assert!(should_use_parallel_strategy(PARALLEL_STRATEGY_THRESHOLD));
assert!(should_use_parallel_strategy(
PARALLEL_STRATEGY_THRESHOLD + 1
));
}
#[test]
fn default_encode_decode_round_trip_at_parallel_threshold() {
let pack = synthetic_pack(PARALLEL_STRATEGY_THRESHOLD);
let config = default_config();
let (shards, manifest) = encode_pack_to_shards(&pack, config).unwrap();
let subset: Vec<Shard> = shards.into_iter().take(16).collect();
let recovered = decode_pack_from_shards(&subset, &manifest).unwrap();
assert_eq!(recovered, pack);
}
#[test]
fn round_trip_default_config_1_mib_first_n_shards() {
let pack = synthetic_pack(1024 * 1024);
let config = default_config();
let (shards, manifest) = encode_pack_to_shards(&pack, config).unwrap();
assert_eq!(shards.len(), 20);
assert_eq!(manifest.shard_hashes.len(), 20);
assert_eq!(manifest.pack_hash, hash::hash(&pack));
let subset: Vec<Shard> = shards.into_iter().take(16).collect();
let recovered = decode_pack_from_shards(&subset, &manifest).unwrap();
assert_eq!(recovered, pack);
}
#[test]
fn lossy_round_trip_drops_shards_0_5_10_17() {
let pack = synthetic_pack(1024 * 1024);
let config = default_config();
let (shards, manifest) = encode_pack_to_shards(&pack, config).unwrap();
let dropped = [0u16, 5, 10, 17];
let subset: Vec<Shard> = shards
.into_iter()
.filter(|s| !dropped.contains(&s.index))
.collect();
assert_eq!(subset.len(), 16);
let recovered = decode_pack_from_shards(&subset, &manifest).unwrap();
assert_eq!(recovered, pack);
}
#[test]
fn tampered_shard_is_rejected_before_decode() {
let pack = synthetic_pack(256 * 1024);
let config = default_config();
let (mut shards, manifest) = encode_pack_to_shards(&pack, config).unwrap();
let last = shards[0].bytes.len() - 1;
shards[0].bytes[last] ^= 0x01;
let subset: Vec<Shard> = shards.into_iter().take(16).collect();
let err = decode_pack_from_shards(&subset, &manifest).unwrap_err();
assert!(
matches!(err, ShardError::ShardHashMismatch { index: 0 }),
"expected ShardHashMismatch{{index: 0}}, got {err:?}"
);
}
#[test]
fn index_out_of_range_is_rejected() {
let pack = synthetic_pack(64 * 1024);
let (_, manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
let total = manifest.config.total_shards();
let bogus = Shard {
index: u16::try_from(total).unwrap(),
bytes: vec![0u8; 32],
};
let err = decode_pack_from_shards(&[bogus], &manifest).unwrap_err();
assert!(
matches!(
err,
ShardError::IndexOutOfRange { index, total: t } if index == u16::try_from(total).unwrap() && t == total
),
"expected IndexOutOfRange, got {err:?}"
);
}
#[test]
fn duplicate_index_is_rejected() {
let pack = synthetic_pack(64 * 1024);
let (shards, manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
let real_shard_0 = shards[0].clone();
let impostor = Shard {
index: 0,
bytes: vec![0xFFu8; real_shard_0.bytes.len()],
};
let err = decode_pack_from_shards(&[real_shard_0, impostor], &manifest).unwrap_err();
assert!(
matches!(err, ShardError::DuplicateIndex { index: 0 }),
"expected DuplicateIndex{{index: 0}}, got {err:?}"
);
}
#[test]
fn pack_hash_mismatch_on_forged_but_consistent_shard_set() {
let pack = synthetic_pack(256 * 1024);
let (shards, mut manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
manifest.pack_hash = hash::hash(b"not the real pack");
let subset: Vec<Shard> = shards.into_iter().take(16).collect();
let err = decode_pack_from_shards(&subset, &manifest).unwrap_err();
assert!(
matches!(err, ShardError::PackHashMismatch),
"expected PackHashMismatch, got {err:?}"
);
}
#[test]
fn manifest_wire_format_round_trip_default_config() {
let pack = synthetic_pack(64 * 1024);
let (_, manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
let bytes = encode_manifest(&manifest).unwrap();
assert_eq!(bytes.len(), 717);
assert_eq!(&bytes[..4], &MANIFEST_MAGIC);
assert_eq!(bytes[4], MANIFEST_VERSION);
let decoded = decode_manifest(&bytes).unwrap();
assert_eq!(decoded, manifest);
}
#[test]
fn manifest_decode_rejects_bad_magic() {
let pack = synthetic_pack(32 * 1024);
let (_, manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
let mut bytes = encode_manifest(&manifest).unwrap();
bytes[0] = b'X';
let err = decode_manifest(&bytes).unwrap_err();
assert!(
matches!(err, ShardError::InvalidManifestPrologue("bad magic")),
"expected InvalidManifestPrologue(bad magic), got {err:?}"
);
}
#[test]
fn manifest_decode_rejects_unsupported_version() {
let pack = synthetic_pack(32 * 1024);
let (_, manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
let mut bytes = encode_manifest(&manifest).unwrap();
bytes[4] = 0xFF;
let err = decode_manifest(&bytes).unwrap_err();
assert!(
matches!(
err,
ShardError::InvalidManifestPrologue("unsupported version")
),
"expected InvalidManifestPrologue(unsupported version), got {err:?}"
);
}
#[test]
fn manifest_decode_rejects_trailing_bytes() {
let pack = synthetic_pack(32 * 1024);
let (_, manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
let mut bytes = encode_manifest(&manifest).unwrap();
bytes.push(0xAB);
let err = decode_manifest(&bytes).unwrap_err();
assert!(
matches!(err, ShardError::ManifestTrailingBytes),
"expected ManifestTrailingBytes, got {err:?}"
);
}
#[test]
fn manifest_decode_rejects_truncated_body() {
let pack = synthetic_pack(32 * 1024);
let (_, manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
let mut bytes = encode_manifest(&manifest).unwrap();
bytes.truncate(bytes.len() - 1);
let err = decode_manifest(&bytes).unwrap_err();
assert!(
matches!(err, ShardError::ManifestUnexpectedEof),
"expected ManifestUnexpectedEof, got {err:?}"
);
}
#[test]
fn manifest_decode_rejects_oversize_input() {
let bytes = vec![0u8; MANIFEST_MAX_BYTES + 1];
let err = decode_manifest(&bytes).unwrap_err();
assert!(
matches!(err, ShardError::ManifestTooLarge { .. }),
"expected ManifestTooLarge, got {err:?}"
);
}
#[test]
fn manifest_decode_rejects_zero_config() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&MANIFEST_MAGIC);
bytes.push(MANIFEST_VERSION);
bytes.extend_from_slice(&[0u8; HASH_LEN]); bytes.extend_from_slice(&0u16.to_le_bytes()); bytes.extend_from_slice(&4u16.to_le_bytes()); bytes.extend_from_slice(&[0u8; HASH_LEN]); bytes.extend_from_slice(&0u32.to_le_bytes()); let err = decode_manifest(&bytes).unwrap_err();
assert!(
matches!(err, ShardError::ManifestZeroShardCount { .. }),
"expected ManifestZeroShardCount, got {err:?}"
);
}
#[test]
fn insufficient_shards_returns_error() {
let pack = synthetic_pack(64 * 1024);
let config = default_config();
let (shards, manifest) = encode_pack_to_shards(&pack, config).unwrap();
let subset: Vec<Shard> = shards.into_iter().take(15).collect();
let err = decode_pack_from_shards(&subset, &manifest).unwrap_err();
assert!(
matches!(
err,
ShardError::InsufficientShards {
provided: 15,
minimum: 16,
}
),
"expected InsufficientShards{{15, 16}}, got {err:?}"
);
}
#[test]
fn manifest_version_is_0x02_and_v01_is_rejected() {
assert_eq!(
MANIFEST_VERSION, 0x02,
"MANIFEST_VERSION must be bumped to 0x02 for the Blake3 cutover"
);
let pack = synthetic_pack(32 * 1024);
let (_, manifest) = encode_pack_to_shards(&pack, default_config()).unwrap();
let mut bytes = encode_manifest(&manifest).unwrap();
assert_eq!(
bytes[4], 0x02,
"encode_manifest must emit the current MANIFEST_VERSION"
);
bytes[4] = 0x01;
let err = decode_manifest(&bytes).unwrap_err();
match err {
ShardError::InvalidManifestPrologue(msg) => {
assert!(
msg.contains("0x01"),
"expected the version-specific message to name 0x01, got {msg:?}"
);
assert!(
msg.to_ascii_lowercase().contains("sha256"),
"expected the version-specific message to call out the \
retired Sha256-era scheme, got {msg:?}"
);
}
other => panic!("expected InvalidManifestPrologue, got {other:?}"),
}
}
#[test]
fn blake3_scheme_roundtrips() {
let pack = synthetic_pack(1024 * 1024);
let config = default_config();
let (shards, manifest) = encode_pack_to_shards(&pack, config).unwrap();
assert_eq!(shards.len(), 20);
let dropped = [1u16, 6, 11, 18];
let subset: Vec<Shard> = shards
.into_iter()
.filter(|s| !dropped.contains(&s.index))
.collect();
assert_eq!(subset.len(), 16);
let recovered = decode_pack_from_shards(&subset, &manifest).unwrap();
assert_eq!(recovered, pack);
assert_eq!(hash::hash(&recovered), manifest.pack_hash);
}
#[test]
fn commitment_from_a_different_scheme_fails_the_merkle_check_not_silently() {
use commonware_cryptography::Sha256;
type OldRsScheme = commonware_coding::ReedSolomon<Sha256>;
let pack = synthetic_pack(256 * 1024);
let config = default_config();
let (shards, mut manifest) = encode_pack_to_shards(&pack, config).unwrap();
let (old_commitment, _old_chunks) =
OldRsScheme::encode(&config, pack.as_slice(), &Sequential)
.expect("old-scheme (Sha256) encode must still succeed");
let old_bytes: &[u8] = old_commitment.as_ref();
let mut forged = [0u8; HASH_LEN];
forged.copy_from_slice(old_bytes);
manifest.commitment = forged;
let subset: Vec<Shard> = shards.into_iter().take(16).collect();
let err = decode_pack_from_shards(&subset, &manifest).unwrap_err();
assert!(
matches!(err, ShardError::DecodeFailed(_)),
"expected a typed DecodeFailed error at the Merkle-proof check \
step, got {err:?}"
);
}
}