use crate::sink::VecSink;
use crate::source::SliceSource;
use crate::{
Freq, Rans64Decoder, Rans64EncSymbol, Rans64Encoder, RansByteDecoder, RansByteEncSymbol,
RansByteEncoder,
};
use alloc::vec;
use alloc::vec::Vec;
struct Lcg(u64);
impl Lcg {
fn new(seed: u64) -> Self {
Lcg(seed.max(1))
}
fn next(&mut self) -> u64 {
let mut x = self.0;
x ^= x >> 12;
x ^= x << 25;
x ^= x >> 27;
self.0 = x;
x.wrapping_mul(0x2545F4914F6CDD1D)
}
fn below(&mut self, bound: u64) -> u64 {
if bound == 0 {
return 0;
}
self.next() % bound
}
}
fn gen_symbols(rng: &mut Lcg, scale: Freq, count: usize) -> Vec<(Freq, Freq)> {
let mut out = Vec::with_capacity(count);
for _ in 0..count {
let start = rng.below(scale as u64) as Freq;
let mut max_freq = scale - start;
if max_freq == scale {
max_freq -= 1;
}
let max_freq = max_freq.max(1);
let freq = match rng.below(4) {
0 => 1,
1 => max_freq,
2 => (max_freq / 2).max(1),
_ => 1 + rng.below(max_freq as u64) as Freq,
};
let freq = freq.clamp(1, max_freq);
out.push((start, freq));
}
out
}
fn roundtrip_byte(symbols: &[(Freq, Freq)], scale_bits: Freq) {
let sink = VecSink::<u8>::new(64);
let mut encoder = RansByteEncoder::new(sink);
for &(start, freq) in symbols {
encoder.put_raw(start, freq, scale_bits);
}
encoder.flush();
let encoded = encoder.into_sink().encoded().to_vec();
assert!(
!encoded.is_empty(),
"scale_bits={} symbols={}",
scale_bits,
symbols.len()
);
let source = SliceSource::new(&encoded[..]);
let mut decoder = RansByteDecoder::new(source);
assert!(decoder.init(), "init failed for scale_bits={}", scale_bits);
for &(start, freq) in symbols.iter().rev() {
let got = decoder.get(scale_bits);
assert!(
got >= start && got < start + freq,
"value out of symbol range at scale_bits={}: start={} freq={} got={}",
scale_bits,
start,
freq,
got
);
assert!(
decoder.advance(start, freq, scale_bits),
"advance failed at scale_bits={}",
scale_bits
);
}
assert!(
decoder.check_eof(),
"not at EOF for scale_bits={}",
scale_bits
);
}
fn roundtrip_64(symbols: &[(Freq, Freq)], scale_bits: Freq) {
let sink = VecSink::<u32>::new(64);
let mut encoder = Rans64Encoder::new(sink);
for &(start, freq) in symbols {
encoder.put_raw(start, freq, scale_bits);
}
encoder.flush();
let units = encoder.into_sink().encoded().to_vec();
assert!(!units.is_empty());
let source = SliceSource::new(&units[..]);
let mut decoder = Rans64Decoder::new(source);
assert!(decoder.init());
for &(start, freq) in symbols.iter().rev() {
let got = decoder.get(scale_bits);
assert!(
got >= start && got < start + freq,
"value out of symbol range at scale_bits={}: start={} freq={} got={}",
scale_bits,
start,
freq,
got
);
assert!(decoder.advance(start, freq, scale_bits));
}
assert!(decoder.check_eof());
}
fn prepared_matches_raw_byte(symbols: &[(Freq, Freq)], scale_bits: Freq) {
let sink = VecSink::<u8>::new(64);
let mut raw = RansByteEncoder::new(sink);
for &(start, freq) in symbols {
raw.put_raw(start, freq, scale_bits);
}
raw.flush();
let raw_out = raw.into_sink().encoded().to_vec();
let sink2 = VecSink::<u8>::new(64);
let mut prepared = RansByteEncoder::new(sink2);
for &(start, freq) in symbols {
let sym = RansByteEncSymbol::new(start, freq, scale_bits);
prepared.put(&sym);
}
prepared.flush();
let prep_out = prepared.into_sink().encoded().to_vec();
assert_eq!(
raw_out, prep_out,
"prepared != raw at scale_bits={}",
scale_bits
);
}
fn prepared_matches_raw_64(symbols: &[(Freq, Freq)], scale_bits: Freq) {
let sink = VecSink::<u32>::new(64);
let mut raw = Rans64Encoder::new(sink);
for &(start, freq) in symbols {
raw.put_raw(start, freq, scale_bits);
}
raw.flush();
let raw_out = raw.into_sink().encoded().to_vec();
let sink2 = VecSink::<u32>::new(64);
let mut prepared = Rans64Encoder::new(sink2);
for &(start, freq) in symbols {
let sym = Rans64EncSymbol::new(start, freq, scale_bits);
prepared.put(&sym);
}
prepared.flush();
let prep_out = prepared.into_sink().encoded().to_vec();
assert_eq!(
raw_out, prep_out,
"prepared != raw at scale_bits={}",
scale_bits
);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_ransbyte_scale_bits_sweep_roundtrip() {
let mut rng = Lcg::new(0x5EED_2026);
for scale_bits in 2..=23u32 {
let scale = 1u32 << scale_bits;
let count = (8 + (scale_bits as usize) * 4).min(512);
let symbols = gen_symbols(&mut rng, scale, count);
roundtrip_byte(&symbols, scale_bits);
}
}
#[test]
fn test_rans64_scale_bits_sweep_roundtrip() {
let mut rng = Lcg::new(0x64_2026);
for scale_bits in 2..=31u32 {
let scale = 1u32 << scale_bits;
let count = (8 + (scale_bits as usize) * 4).min(512);
let symbols = gen_symbols(&mut rng, scale, count);
roundtrip_64(&symbols, scale_bits);
}
}
#[test]
fn test_boundary_freq_patterns() {
let mut rng = Lcg::new(0xB0B);
for scale_bits in [2u32, 3, 4, 8, 16, 23, 24, 31] {
let scale = 1u32 << scale_bits;
let byte_ok = scale_bits <= 23;
let patterns: Vec<(Freq, Freq)> = vec![
(0, 1), (0, scale - 1), (scale - 1, 1), (scale / 2, scale / 2), (scale / 2, (scale / 2) - 1), (1, scale - 1), (scale - 2, 2), (0, scale - 2), (0, scale / 3), (scale - scale / 3, scale / 3),
];
let mut symbols = patterns.clone();
for _ in 0..32 {
symbols.extend(gen_symbols(&mut rng, scale, 1));
}
if byte_ok {
roundtrip_byte(&symbols, scale_bits);
}
roundtrip_64(&symbols, scale_bits);
}
}
#[test]
fn test_prepared_matches_raw_sweep() {
let mut rng = Lcg::new(0x7A7A);
for scale_bits in [2u32, 3, 4, 8, 12, 16, 20, 23, 24, 31] {
let scale = 1u32 << scale_bits;
let symbols = gen_symbols(&mut rng, scale, 64);
if scale_bits <= 23 {
prepared_matches_raw_byte(&symbols, scale_bits);
}
prepared_matches_raw_64(&symbols, scale_bits);
}
}
#[test]
fn test_truncated_stream_transactional() {
let mut rng = Lcg::new(0xDEAD);
let scale_bits = 16u32;
let scale = 1u32 << scale_bits;
let symbols = gen_symbols(&mut rng, scale, 64);
let sink = VecSink::<u8>::new(64);
let mut encoder = RansByteEncoder::new(sink);
for &(start, freq) in &symbols {
encoder.put_raw(start, freq, scale_bits);
}
encoder.flush();
let full = encoder.into_sink().encoded().to_vec();
for cut in 0..full.len() {
let data = &full[..cut];
let source = SliceSource::new(data);
let mut decoder = RansByteDecoder::new(source);
if !decoder.init() {
continue;
}
let mut ok = true;
for &(start, freq) in symbols.iter().rev() {
let _ = decoder.get(scale_bits);
let before = decoder.state();
if !decoder.advance(start, freq, scale_bits) {
assert_eq!(
decoder.state(),
before,
"advance failed but mutated state at cut={}",
cut
);
ok = false;
break;
}
}
if ok {
assert_eq!(cut, full.len(), "short stream decoded fully at cut={}", cut);
}
}
}
#[test]
fn test_vecsink_growth_sweep() {
for n in 1..=2000usize {
let sink = VecSink::<u8>::new(8);
let mut encoder = RansByteEncoder::new(sink);
for i in 0..n {
let start = (i % 200) as Freq;
let freq = 1u32; encoder.put_raw(start, freq, 8);
}
encoder.flush();
let encoded = encoder.into_sink().encoded().to_vec();
let source = SliceSource::new(&encoded[..]);
let mut decoder = RansByteDecoder::new(source);
assert!(decoder.init(), "init failed for n={}", n);
for i in (0..n).rev() {
let start = (i % 200) as Freq;
let got = decoder.get(8);
assert_eq!(got, start, "freq=1 value must equal start: n={} i={}", n, i);
assert!(
decoder.advance(start, 1, 8),
"advance failed n={} i={}",
n,
i
);
}
assert!(decoder.check_eof(), "eof failed n={}", n);
}
}
#[test]
fn test_raw_decoder_no_panic_on_corruption() {
let mut rng = Lcg::new(0xC0FFEE);
let scale_bits = 16u32;
let scale = 1u32 << scale_bits;
let symbols = gen_symbols(&mut rng, scale, 32);
let sink = VecSink::<u8>::new(64);
let mut encoder = RansByteEncoder::new(sink);
for &(start, freq) in &symbols {
encoder.put_raw(start, freq, scale_bits);
}
encoder.flush();
let full = encoder.into_sink().encoded().to_vec();
for flip in 0..full.len() {
let mut data = full.clone();
data[flip] ^= 0x5A;
let source = SliceSource::new(&data[..]);
let mut decoder = RansByteDecoder::new(source);
if !decoder.init() {
continue;
}
for &(start, freq) in symbols.iter().rev() {
let _ = decoder.get(scale_bits);
let before = decoder.state();
if !decoder.advance(start, freq, scale_bits) {
assert_eq!(decoder.state(), before, "state mutated on failure");
break;
}
}
}
}
#[test]
fn test_try_advance_invalid_parameters() {
use crate::error::RawRansError;
let sink = VecSink::<u8>::new(64);
let mut encoder = RansByteEncoder::new(sink);
encoder.put_raw(0, 128, 8);
encoder.flush();
let encoded = encoder.into_sink().encoded().to_vec();
let source = SliceSource::new(&encoded[..]);
let mut decoder = RansByteDecoder::new(source);
assert!(decoder.init());
assert!(matches!(
decoder.try_advance(256, 1, 8),
Err(RawRansError::InvalidParameters)
));
assert!(matches!(
decoder.try_advance(0, 0, 8),
Err(RawRansError::InvalidParameters)
));
assert!(matches!(
decoder.try_advance(255, 2, 8),
Err(RawRansError::InvalidParameters)
));
assert!(matches!(
decoder.try_advance(0, 1, 32),
Err(RawRansError::InvalidScaleBits { .. })
));
assert!(matches!(
decoder.try_advance(0, 1, 1),
Err(RawRansError::InvalidScaleBits { .. })
));
}
}