#![cfg(all(feature = "entropy", feature = "alloc"))]
use proptest::prelude::*;
use rustbinary::{EntropyConfig, Model, RansDecoder, RansEncoder};
fn config() -> EntropyConfig {
rustbinary::options().with_entropy_encoding()
}
fn encode_bytes(model: &Model, bytes: &[u8]) -> rustbinary::Result<Vec<u8>> {
config().compress(bytes, model)
}
fn decode_bytes(model: &Model, frame: &[u8]) -> rustbinary::Result<Vec<u8>> {
config().decompress(frame, model)
}
proptest! {
#[test]
fn byte_roundtrip_lossless(
bytes in proptest::collection::vec(any::<u8>(), 0..4096),
seed in any::<u32>(),
) {
let mut weights = vec![1u32; 256];
for (index, weight) in weights.iter_mut().enumerate() {
*weight = 1 + ((seed.wrapping_mul(index as u32 + 1)) % 1000);
}
let model = Model::from_weights(&weights).unwrap();
let frame = encode_bytes(&model, &bytes).unwrap();
let decoded = decode_bytes(&model, &frame).unwrap();
prop_assert_eq!(&decoded, &bytes);
prop_assert!(frame.len() <= 24 + bytes.len());
}
#[test]
fn uniform_roundtrip_lossless(
symbols in 1..2048u32,
length in 0..1024usize,
seed in any::<u64>(),
) {
let model = Model::from_uniform(symbols).unwrap();
let mut encoder = RansEncoder::new();
let mut data = Vec::with_capacity(length);
let mut state = seed;
for _ in 0..length {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
let symbol = (state % symbols as u64) as u32;
data.push(symbol);
encoder.put_symbol(&model, symbol).unwrap();
}
let (final_state, payload) = encoder.finish();
let mut decoder = RansDecoder::new(final_state, &payload);
let mut decoded = Vec::with_capacity(length);
for _ in 0..length {
decoded.push(decoder.get_symbol(&model).unwrap());
}
decoder.finish().unwrap();
decoded.reverse();
prop_assert_eq!(decoded, data);
}
#[test]
fn single_byte_corruption_is_detected(
length in 64..2048usize,
seed in any::<u64>(),
) {
let mut weights = vec![1u32; 256];
weights[(seed % 256) as usize] = 10000;
let model = Model::from_weights(&weights).unwrap();
let mut data = Vec::with_capacity(length);
let mut state = seed;
for _ in 0..length {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
let symbol = if state % 16 == 0 {
(state % 256) as u8
} else {
(seed % 256) as u8
};
data.push(symbol);
}
let frame = encode_bytes(&model, &data).unwrap();
if frame.len() >= 24 + data.len() {
return Ok(());
}
for offset in 24..frame.len() {
let mut corrupted = frame.clone();
corrupted[offset] ^= 0x80;
let decoded = decode_bytes(&model, &corrupted);
match decoded {
Err(_) => {}
Ok(bytes) => {
prop_assert!(
bytes != data,
"corruption at offset {} returned the original payload",
offset
);
}
}
}
}
}
#[test]
fn truncation_is_detected() {
let mut weights = vec![1u32; 256];
weights[b'x' as usize] = 10000;
let model = Model::from_weights(&weights).unwrap();
let data = vec![b'x'; 512];
let frame = encode_bytes(&model, &data).unwrap();
assert!(frame.len() < 24 + data.len());
for end in 24..frame.len() {
let truncated = &frame[..end];
assert!(
decode_bytes(&model, truncated).is_err(),
"truncation at {end} was accepted"
);
}
}