use super::*;
use reed_solomon_erasure::galois_8::ReedSolomon;
use rs_merkle::*;
use sha2::{Digest, Sha256};
use std::collections::HashMap;
pub fn encode_rs(
payload: Vec<u8>,
data_shards: usize,
parity_shards: usize,
) -> Result<Vec<Vec<u8>>, ShardError> {
if data_shards == 0 || parity_shards == 0 {
return Err(ShardError::Config("Shard counts must be > 0".to_string()));
}
let original_len = payload.len() as u64;
let mut auth_payload = original_len.to_le_bytes().to_vec();
auth_payload.extend_from_slice(&payload);
let shard_size = (auth_payload.len() + data_shards - 1) / data_shards;
let mut shards = Vec::with_capacity(data_shards + parity_shards);
for i in 0..data_shards {
let start = i * shard_size;
let shard = if start < auth_payload.len() {
let end = usize::min(start + shard_size, auth_payload.len());
let mut s = auth_payload[start..end].to_vec();
s.resize(shard_size, 0); s
} else {
vec![0u8; shard_size] };
shards.push(shard);
}
for _ in 0..parity_shards {
shards.push(vec![0u8; shard_size]);
}
let r = ReedSolomon::new(data_shards, parity_shards)
.map_err(|e| ShardError::Config(e.to_string()))?;
r.encode(&mut shards)
.map_err(|e| ShardError::Failed(e.to_string()))?;
Ok(shards)
}
pub fn decode_rs(
shards_map: HashMap<usize, Vec<u8>>,
data_shards: usize,
parity_shards: usize,
) -> Result<Vec<Vec<u8>>, ShardError> {
let total_shards = data_shards + parity_shards;
let r = ReedSolomon::new(data_shards, parity_shards)
.map_err(|e| ShardError::Config(e.to_string()))?;
let mut shards: Vec<Option<Vec<u8>>> = vec![None; total_shards];
let max_shard_size = (MAX_PAYLOAD_SIZE + 8 + data_shards - 1) / data_shards;
for (&idx, shard) in &shards_map {
if shard.len() > max_shard_size {
return Err(ShardError::Config(format!(
"Shard {} exceeds maximum allowed size ({} > {})",
idx,
shard.len(),
max_shard_size
)));
}
if idx < total_shards {
shards[idx] = Some(shard.clone());
} else {
return Err(ShardError::OutOfBounds(idx, total_shards - 1));
}
}
r.reconstruct(&mut shards)
.map_err(|e| ShardError::Failed(e.to_string()))?;
let full_shards: Vec<Vec<u8>> = shards
.into_iter()
.map(|opt| opt.ok_or(ShardError::Incomplete))
.collect::<Result<_, _>>()?;
if !r
.verify(&full_shards)
.map_err(|e| ShardError::Failed(e.to_string()))?
{
return Err(ShardError::Failed(
"Reed-Solomon verification failed: shards are not a valid codeword".into(),
));
}
Ok(full_shards)
}
pub fn reconstruct_payload(
decoded_shards: Vec<Vec<u8>>,
data_shards: usize,
) -> Result<Vec<u8>, ShardError> {
if decoded_shards.len() < data_shards {
return Err(ShardError::Incomplete);
}
let total_len: usize = decoded_shards
.iter()
.take(data_shards)
.map(|s| s.len())
.sum();
if total_len > MAX_PAYLOAD_SIZE {
return Err(ShardError::Config(
"Reconstructed payload exceeds maximum size".to_string(),
));
}
let mut payload = decoded_shards
.into_iter()
.take(data_shards)
.flatten()
.collect::<Vec<u8>>();
if payload.len() < 8 {
return Err(ShardError::Config(
"Payload too short for length prefix".to_string(),
));
}
let mut len_bytes = [0u8; 8];
len_bytes.copy_from_slice(&payload[0..8]);
let original_len = u64::from_le_bytes(len_bytes) as usize;
payload.drain(0..8);
if original_len > payload.len() {
return Err(ShardError::Config(
"Original length exceeds payload".to_string(),
));
}
payload.truncate(original_len);
Ok(payload)
}
#[derive(Clone)]
pub struct Sha256Algorithm {}
impl Hasher for Sha256Algorithm {
type Hash = [u8; 32];
fn hash(data: &[u8]) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(data);
hasher.finalize().into()
}
}
pub fn gen_merkletree(shards: &[Vec<u8>]) -> MerkleTree<Sha256Algorithm> {
let leaves: Vec<[u8; 32]> = shards.iter().map(|x| Sha256Algorithm::hash(&x)).collect();
MerkleTree::<Sha256Algorithm>::from_leaves(&leaves)
}
pub fn get_merkle_proof(proof: Vec<u8>) -> Result<MerkleProof<Sha256Algorithm>, ShardError> {
if proof.len() <= 32 {
return Err(ShardError::Merkle("Invalid fingerprint length".to_string()));
}
MerkleProof::<Sha256Algorithm>::try_from(&proof[32..])
.map_err(|e| ShardError::Merkle(e.to_string()))
}
pub fn verify_merkle(
id: usize,
n: usize,
proof: Vec<u8>,
shard: Vec<u8>,
) -> Result<bool, ShardError> {
if proof.len() < 32 {
return Err(ShardError::Merkle("Invalid fingerprint length".to_string()));
}
let root: [u8; 32] = proof[0..32]
.try_into()
.map_err(|_| ShardError::Merkle("Failed to extract Merkle root".to_string()))?;
let proof = get_merkle_proof(proof)?;
let leaf_hash = Sha256Algorithm::hash(&shard);
Ok(proof.verify(root, &vec![id], &[leaf_hash], n))
}
pub fn generate_merkle_proofs_map(
shards: &[Vec<u8>],
) -> Result<HashMap<usize, Vec<u8>>, ShardError> {
let n = shards.len();
let tree = gen_merkletree(shards);
if tree.root().is_none() {
return Err(ShardError::Merkle(
"Failed to extract Merkle root".to_string(),
));
}
let mut proofs_map = HashMap::with_capacity(n);
for i in 0..n {
let proof_bytes = tree.proof(&[i]).to_bytes();
proofs_map.insert(i, proof_bytes);
}
Ok(proofs_map)
}
pub fn set_value_round(bit: bool, number: u32) -> Vec<u8> {
let mut vec = Vec::with_capacity(5);
let first_byte = if bit { 1u8 } else { 0u8 };
vec.push(first_byte);
vec.extend_from_slice(&number.to_le_bytes());
vec
}
pub fn get_value_round(data: &[u8]) -> Option<(bool, u32)> {
if data.len() < 5 {
return None;
}
let bit = data[0] & 1 != 0;
let number_bytes = &data[1..5];
let number = u32::from_le_bytes(number_bytes.try_into().ok()?);
Some((bit, number))
}
pub fn get_value(data: &[u8]) -> Option<bool> {
data.get(0).map(|byte| byte & 1 != 0)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_reed_solomon_encode_decode_roundtrip() {
let payload = b"Hello, world!".to_vec();
let data_shards = 4;
let parity_shards = 2;
let shards = encode_rs(payload.clone(), data_shards, parity_shards)
.expect("Encoding should succeed");
let mut shards_map = HashMap::new();
for (i, shard) in shards.iter().enumerate() {
if i != 1 {
shards_map.insert(i, shard.clone());
}
}
let decoded_shards =
decode_rs(shards_map, data_shards, parity_shards).expect("Decoding should succeed");
let reconstructed = reconstruct_payload(decoded_shards, data_shards)
.expect("Reconstruction should succeed");
assert_eq!(payload, reconstructed);
}
#[test]
fn test_merkle_tree_and_proof_verification() {
let payload = b"Merkle tree test!".to_vec();
let data_shards = 4;
let parity_shards = 2;
let shards = encode_rs(payload.clone(), data_shards, parity_shards)
.expect("Encoding should succeed");
let tree = gen_merkletree(&shards);
let proofs_map =
generate_merkle_proofs_map(&shards).expect("Proof generation should succeed");
for (i, shard) in shards.iter().enumerate() {
let mut proof_with_root = vec![];
proof_with_root.extend_from_slice(&tree.root().unwrap()); proof_with_root.extend_from_slice(&proofs_map[&i]);
let verified = verify_merkle(i, shards.len(), proof_with_root, shard.clone())
.expect("Verification should succeed");
assert!(verified, "Merkle proof for shard {} failed", i);
}
}
#[test]
fn test_decode_failure_with_insufficient_data() {
let payload = b"Test failure!".to_vec();
let data_shards = 3;
let parity_shards = 2;
let shards = encode_rs(payload.clone(), data_shards, parity_shards)
.expect("Encoding should succeed");
let mut shards_map = HashMap::new();
shards_map.insert(0, shards[0].clone());
shards_map.insert(1, shards[1].clone());
let result = decode_rs(shards_map, data_shards, parity_shards);
assert!(
result.is_err(),
"Decoding should fail due to insufficient data shards"
);
}
}