use alloc::borrow::Cow;
use alloc::format;
use alloc::sync::Arc;
use alloc::vec::Vec;
use core::fmt;
use super::util::{random_u64, zero_buffer};
use super::{FragmentStrategy, Fragments};
use crate::Result;
use crate::error::Error;
use crate::fetcher::RawKey;
use crate::memory::LockedBytes;
#[derive(Clone)]
pub struct LayeredFragmenter {
sub_strategies: Vec<Arc<dyn FragmentStrategy>>,
}
impl LayeredFragmenter {
pub fn new(sub_strategies: Vec<Arc<dyn FragmentStrategy>>) -> Result<Self> {
if sub_strategies.is_empty() {
return Err(Error::InvalidConfig(alloc::string::ToString::to_string(
"LayeredFragmenter requires at least one sub-strategy",
)));
}
if sub_strategies.len() > u32::MAX as usize {
return Err(Error::InvalidConfig(alloc::string::ToString::to_string(
"LayeredFragmenter sub-strategy count exceeds u32",
)));
}
Ok(Self { sub_strategies })
}
#[must_use]
pub fn sub_strategy_count(&self) -> usize {
self.sub_strategies.len()
}
}
impl fmt::Debug for LayeredFragmenter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let names: Vec<_> = self
.sub_strategies
.iter()
.map(|s| s.describe().into_owned())
.collect();
f.debug_struct("LayeredFragmenter")
.field("sub_strategies", &names)
.finish()
}
}
impl FragmentStrategy for LayeredFragmenter {
fn fragment(&self, key: &RawKey) -> Result<Fragments> {
#[allow(clippy::cast_possible_truncation)]
let n = self.sub_strategies.len() as u64;
#[allow(clippy::cast_possible_truncation)]
let pick = (random_u64()? % n) as usize;
let sub_fragments = self.sub_strategies[pick].fragment(key)?;
let (chunks, sub_layout, total_len) = sub_fragments.into_parts();
let sub_layout_bytes = sub_layout.as_bytes();
let mut new_layout_bytes: Vec<u8> = Vec::with_capacity(4 + sub_layout_bytes.len());
let pick_u32 = u32::try_from(pick)
.map_err(|_| Error::Internal("LayeredFragmenter sub-strategy index exceeded u32"))?;
new_layout_bytes.extend_from_slice(&pick_u32.to_le_bytes());
new_layout_bytes.extend_from_slice(sub_layout_bytes);
let new_layout = LockedBytes::from_slice(&new_layout_bytes);
zero_buffer(&mut new_layout_bytes);
drop(new_layout_bytes);
drop(sub_layout);
Ok(Fragments::from_parts(chunks, new_layout, total_len))
}
fn defragment(&self, fragments: &Fragments) -> Result<RawKey> {
let layout = fragments.layout().as_bytes();
if layout.len() < 4 {
return Err(Error::Defragment(alloc::string::ToString::to_string(
"layered layout shorter than 4-byte header",
)));
}
let pick_raw: [u8; 4] = layout[0..4]
.try_into()
.map_err(|_| Error::Defragment(alloc::string::ToString::to_string("layout slice")))?;
let pick = u32::from_le_bytes(pick_raw) as usize;
if pick >= self.sub_strategies.len() {
return Err(Error::Defragment(format!(
"layered layout strategy index {pick} out of range",
)));
}
let sub_layout = LockedBytes::from_slice(&layout[4..]);
let chunks_copy: Vec<LockedBytes> = fragments
.chunks()
.iter()
.map(|c| LockedBytes::from_slice(c.as_bytes()))
.collect();
let sub_fragments = Fragments::from_parts(chunks_copy, sub_layout, fragments.total_len());
self.sub_strategies[pick].defragment(&sub_fragments)
}
fn describe(&self) -> Cow<'_, str> {
Cow::Borrowed("layered")
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
mod tests {
use super::*;
use crate::fragment::{InterleavedFragmenter, RandomFragmenter, StandardFragmenter};
fn key(bytes: &[u8]) -> RawKey {
RawKey::new(bytes.to_vec())
}
#[test]
fn rejects_empty_sub_strategy_list() {
let err = LayeredFragmenter::new(Vec::new()).unwrap_err();
assert!(matches!(err, Error::InvalidConfig(_)));
}
#[test]
fn round_trip_with_three_sub_strategies() {
let frag = LayeredFragmenter::new(alloc::vec![
Arc::new(StandardFragmenter::new()) as Arc<dyn FragmentStrategy>,
Arc::new(InterleavedFragmenter::new()) as Arc<dyn FragmentStrategy>,
Arc::new(RandomFragmenter::new()) as Arc<dyn FragmentStrategy>,
])
.unwrap();
let bytes: Vec<u8> = (0u8..64).collect();
let original = key(&bytes);
for _ in 0..30 {
let fragments = frag.fragment(&original).unwrap();
let recovered = frag.defragment(&fragments).unwrap();
assert_eq!(recovered.as_bytes(), &bytes[..]);
}
}
#[test]
fn round_trip_with_single_sub_strategy() {
let frag = LayeredFragmenter::new(alloc::vec![
Arc::new(StandardFragmenter::new()) as Arc<dyn FragmentStrategy>,
])
.unwrap();
let bytes: Vec<u8> = (0u8..32).collect();
let original = key(&bytes);
let fragments = frag.fragment(&original).unwrap();
let recovered = frag.defragment(&fragments).unwrap();
assert_eq!(recovered.as_bytes(), &bytes[..]);
}
#[test]
fn describe_returns_layered() {
let frag = LayeredFragmenter::new(alloc::vec![
Arc::new(StandardFragmenter::new()) as Arc<dyn FragmentStrategy>,
])
.unwrap();
assert_eq!(frag.describe(), "layered");
}
#[test]
fn sub_strategy_count_is_correct() {
let frag = LayeredFragmenter::new(alloc::vec![
Arc::new(StandardFragmenter::new()) as Arc<dyn FragmentStrategy>,
Arc::new(RandomFragmenter::new()) as Arc<dyn FragmentStrategy>,
])
.unwrap();
assert_eq!(frag.sub_strategy_count(), 2);
}
}