use alloc::borrow::Cow;
use alloc::sync::Arc;
use alloc::vec::Vec;
use core::fmt;
use super::util::{fisher_yates, sample_range, zero_buffer, zero_buffer_owned};
use super::{FragmentStrategy, Fragments};
use crate::Result;
use crate::decoy::DecoyStrategy;
use crate::error::Error;
use crate::fetcher::RawKey;
use crate::memory::LockedBytes;
const DECOY_OFFSET: u32 = u32::MAX;
const DEFAULT_MIN_CHUNK: usize = 1;
const DEFAULT_MAX_CHUNK: usize = 8;
#[derive(Clone)]
pub struct StandardFragmenter {
min_chunk: usize,
max_chunk: usize,
decoy: Option<Arc<dyn DecoyStrategy>>,
}
impl StandardFragmenter {
#[must_use]
pub fn new() -> Self {
Self {
min_chunk: DEFAULT_MIN_CHUNK,
max_chunk: DEFAULT_MAX_CHUNK,
decoy: None,
}
}
#[must_use]
pub fn with_chunk_range(min: usize, max: usize) -> Self {
let min = min.max(1);
let max = max.max(min);
Self {
min_chunk: min,
max_chunk: max,
decoy: None,
}
}
#[must_use]
pub fn with_decoy<D>(mut self, decoy: D) -> Self
where
D: DecoyStrategy + 'static,
{
self.decoy = Some(Arc::new(decoy));
self
}
}
impl Default for StandardFragmenter {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for StandardFragmenter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StandardFragmenter")
.field("min_chunk", &self.min_chunk)
.field("max_chunk", &self.max_chunk)
.field("decoy", &self.decoy.as_ref().map(|d| d.describe()))
.finish()
}
}
impl FragmentStrategy for StandardFragmenter {
fn fragment(&self, key: &RawKey) -> Result<Fragments> {
let bytes = key.as_bytes();
let total_len = bytes.len();
if total_len == 0 {
return Err(Error::Fragment(alloc::string::ToString::to_string(
"empty key cannot be fragmented",
)));
}
if total_len >= DECOY_OFFSET as usize {
return Err(Error::Fragment(alloc::string::ToString::to_string(
"key too large for fragmentation",
)));
}
let sizes = sample_chunk_sizes(total_len, self.min_chunk, self.max_chunk)?;
let n_real = sizes.len();
let mut real_pairs: Vec<(u32, &[u8])> = Vec::with_capacity(n_real);
{
let mut offset = 0usize;
for &size in &sizes {
let offset_u32 = u32::try_from(offset).map_err(|_| {
Error::Fragment(alloc::string::ToString::to_string(
"key too large for fragmentation",
))
})?;
real_pairs.push((offset_u32, &bytes[offset..offset + size]));
offset += size;
}
}
let n_decoy = if self.decoy.is_some() { n_real } else { 0 };
let mut decoy_chunks: Vec<LockedBytes> = Vec::with_capacity(n_decoy);
if let Some(ref decoy) = self.decoy {
for _ in 0..n_decoy {
let size = sample_range(self.min_chunk, self.max_chunk)?;
let bytes = decoy.generate(key, size)?;
decoy_chunks.push(LockedBytes::from_slice(&bytes));
zero_buffer_owned(bytes);
}
}
let total_chunks = n_real + n_decoy;
let mut order: Vec<ChunkKind> = Vec::with_capacity(total_chunks);
for i in 0..n_real {
order.push(ChunkKind::Real(i));
}
for i in 0..n_decoy {
order.push(ChunkKind::Decoy(i));
}
fisher_yates(&mut order)?;
let mut chunks: Vec<LockedBytes> = Vec::with_capacity(total_chunks);
let mut layout_bytes: Vec<u8> = Vec::with_capacity(total_chunks * 4);
let mut decoy_slots: Vec<Option<LockedBytes>> =
decoy_chunks.into_iter().map(Some).collect();
for kind in &order {
match *kind {
ChunkKind::Real(idx) => {
let (offset, slice) = real_pairs[idx];
chunks.push(LockedBytes::from_slice(slice));
layout_bytes.extend_from_slice(&offset.to_le_bytes());
}
ChunkKind::Decoy(idx) => {
let lb = decoy_slots[idx].take().ok_or(Error::Internal(
"decoy slot taken twice during fragmentation",
))?;
chunks.push(lb);
layout_bytes.extend_from_slice(&DECOY_OFFSET.to_le_bytes());
}
}
}
let layout = LockedBytes::from_slice(&layout_bytes);
zero_buffer(&mut layout_bytes);
drop(layout_bytes);
drop(real_pairs);
drop(decoy_slots);
drop(order);
Ok(Fragments::from_parts(chunks, layout, total_len))
}
fn defragment(&self, fragments: &Fragments) -> Result<RawKey> {
let mut out = alloc::vec![0u8; fragments.total_len()];
self.defragment_into(fragments, &mut out)?;
Ok(RawKey::new(out))
}
fn defragment_into(&self, fragments: &Fragments, out: &mut [u8]) -> Result<()> {
let n_chunks = fragments.chunk_count();
let layout = fragments.layout().as_bytes();
let total_len = fragments.total_len();
if layout.len() != n_chunks * 4 {
return Err(Error::Defragment(alloc::string::ToString::to_string(
"layout buffer length does not match chunk count",
)));
}
if out.len() != total_len {
return Err(Error::Defragment(alloc::string::ToString::to_string(
"scratch buffer size does not match fragments.total_len()",
)));
}
let mut written = 0usize;
for (i, chunk) in fragments.chunks().iter().enumerate() {
let raw: [u8; 4] = layout[i * 4..i * 4 + 4].try_into().map_err(|_| {
Error::Defragment(alloc::string::ToString::to_string(
"layout buffer slice did not size to u32",
))
})?;
let offset = u32::from_le_bytes(raw);
if offset == DECOY_OFFSET {
continue;
}
let chunk_bytes = chunk.as_bytes();
let start = offset as usize;
let end = start.checked_add(chunk_bytes.len()).ok_or_else(|| {
Error::Defragment(alloc::string::ToString::to_string(
"chunk offset overflowed when added to chunk length",
))
})?;
if end > total_len {
return Err(Error::Defragment(alloc::string::ToString::to_string(
"chunk would write past end of output buffer",
)));
}
out[start..end].copy_from_slice(chunk_bytes);
written = written.saturating_add(chunk_bytes.len());
}
if written != total_len {
return Err(Error::Defragment(alloc::string::ToString::to_string(
"reassembled length does not match recorded total",
)));
}
Ok(())
}
fn describe(&self) -> Cow<'_, str> {
Cow::Borrowed("standard")
}
}
#[derive(Clone, Copy)]
enum ChunkKind {
Real(usize),
Decoy(usize),
}
fn sample_chunk_sizes(total: usize, min: usize, max: usize) -> Result<Vec<usize>> {
if min == 0 || max < min {
return Err(Error::Fragment(alloc::string::ToString::to_string(
"invalid chunk-size range",
)));
}
let mut sizes: Vec<usize> = Vec::new();
let mut remaining = total;
while remaining > 0 {
if remaining <= max {
sizes.push(remaining);
remaining = 0;
} else {
let pick = sample_range(min, max)?;
let pick = pick.min(remaining.saturating_sub(min));
let pick = pick.max(min).min(max).min(remaining);
sizes.push(pick);
remaining -= pick;
}
}
Ok(sizes)
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
mod tests {
use super::*;
fn key(bytes: &[u8]) -> RawKey {
RawKey::new(bytes.to_vec())
}
#[test]
fn round_trip_short_key() {
let frag = StandardFragmenter::new();
let original = key(&[0u8, 1, 2, 3, 4, 5, 6, 7, 8, 9]);
let fragments = frag.fragment(&original).unwrap();
let recovered = frag.defragment(&fragments).unwrap();
assert_eq!(recovered.len(), 10);
assert_eq!(recovered.as_bytes(), original.as_bytes());
}
#[test]
fn round_trip_256_bit_key() {
let frag = StandardFragmenter::new();
let bytes: Vec<u8> = (0..32).map(|i| (i * 7) as u8).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 round_trip_for_many_sizes() {
let frag = StandardFragmenter::new();
for len in [1usize, 7, 16, 32, 64, 128, 255, 256, 500, 1024, 4096] {
let bytes: Vec<u8> = (0..len).map(|i| (i & 0xff) as u8).collect();
let original = key(&bytes);
let fragments = frag.fragment(&original).expect("fragment");
let recovered = frag.defragment(&fragments).expect("defragment");
assert_eq!(
recovered.as_bytes(),
&bytes[..],
"round-trip mismatch for len = {len}"
);
}
}
#[test]
fn two_calls_produce_different_layouts() {
let frag = StandardFragmenter::new();
let bytes: Vec<u8> = (0..32).map(|i| (i ^ 0x5a) as u8).collect();
let original = key(&bytes);
let a = frag.fragment(&original).unwrap();
let b = frag.fragment(&original).unwrap();
let same_count = a.chunk_count() == b.chunk_count();
let same_layout = same_count && a.layout().as_bytes() == b.layout().as_bytes();
assert!(
!(same_count && same_layout),
"two consecutive fragmentations produced the same layout"
);
assert_eq!(frag.defragment(&a).unwrap().as_bytes(), &bytes[..]);
assert_eq!(frag.defragment(&b).unwrap().as_bytes(), &bytes[..]);
}
#[test]
fn chunk_sizes_respect_configured_range() {
let frag = StandardFragmenter::with_chunk_range(2, 4);
let bytes: Vec<u8> = (0..32).collect();
let original = key(&bytes);
let fragments = frag.fragment(&original).unwrap();
let chunks = fragments.chunks();
let mut below_min = 0;
let mut total = 0usize;
for c in chunks {
assert!(
c.len() >= 1 && c.len() <= 4,
"chunk size {} not in [1,4]",
c.len()
);
if c.len() < 2 {
below_min += 1;
}
total += c.len();
}
assert!(
below_min <= 1,
"more than one chunk below min size: {below_min}"
);
assert_eq!(total, 32);
assert_eq!(frag.defragment(&fragments).unwrap().as_bytes(), &bytes[..]);
}
#[test]
fn empty_key_rejected() {
let frag = StandardFragmenter::new();
let empty = key(&[]);
let err = frag.fragment(&empty).unwrap_err();
assert!(matches!(err, Error::Fragment(_)));
}
#[test]
fn describe_returns_standard() {
let frag = StandardFragmenter::new();
assert_eq!(frag.describe(), "standard");
}
#[test]
fn stress_round_trip_thousand_iterations() {
let frag = StandardFragmenter::new();
let bytes: Vec<u8> = (0..32).map(|i| ((i * 13) ^ 0xa5) as u8).collect();
let original = key(&bytes);
for _ in 0..1000 {
let fragments = frag.fragment(&original).expect("fragment");
let recovered = frag.defragment(&fragments).expect("defragment");
assert_eq!(recovered.as_bytes(), &bytes[..]);
}
}
}