#![no_std]
#![forbid(unsafe_code)]
use core::fmt;
pub mod malformed;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EncodeError {
OutputTooSmall,
}
impl fmt::Display for EncodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
EncodeError::OutputTooSmall => write!(f, "output buffer too small"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DecodeError {
InputTooShort,
}
impl fmt::Display for DecodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
DecodeError::InputTooShort => write!(f, "truncated input stream"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelError {
EmptyInput,
ZeroTotal,
InvalidScaleBits,
ZeroFrequency,
FrequencyOutOfRange,
StartOutOfRange,
TotalMismatch,
WorkspaceTooSmall,
}
impl fmt::Display for ModelError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ModelError::EmptyInput => write!(f, "empty input sequence"),
ModelError::ZeroTotal => write!(f, "total frequency is zero"),
ModelError::InvalidScaleBits => write!(f, "scale_bits is out of valid range"),
ModelError::ZeroFrequency => write!(f, "symbol frequency is zero"),
ModelError::FrequencyOutOfRange => write!(f, "frequency exceeds allowed range"),
ModelError::StartOutOfRange => write!(f, "start value exceeds allowed range"),
ModelError::TotalMismatch => write!(f, "total frequency mismatch"),
ModelError::WorkspaceTooSmall => write!(f, "workspace buffer too small"),
}
}
}
#[cfg(feature = "std")]
extern crate std;
#[cfg(feature = "alloc")]
extern crate alloc;
#[cfg(feature = "std")]
impl std::error::Error for EncodeError {}
#[cfg(feature = "std")]
impl std::error::Error for DecodeError {}
#[cfg(feature = "std")]
impl std::error::Error for ModelError {}
pub const RANS_BYTE_L: u32 = 1u32 << 23;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct RansByteState(pub u32);
impl RansByteState {
#[inline]
pub const fn new() -> Self {
Self(RANS_BYTE_L)
}
#[inline]
pub const fn get(&self) -> u32 {
self.0
}
}
impl Default for RansByteState {
#[inline]
fn default() -> Self {
Self::new()
}
}
pub struct BackwardByteWriter<'a> {
buf: &'a mut [u8],
pos: usize, }
impl<'a> BackwardByteWriter<'a> {
#[inline]
pub fn new(buf: &'a mut [u8]) -> Self {
let len = buf.len();
Self { buf, pos: len }
}
#[inline]
pub fn write_byte(&mut self, b: u8) -> Result<(), ()> {
if self.pos == 0 {
return Err(());
}
self.pos -= 1;
self.buf[self.pos] = b;
Ok(())
}
#[inline]
pub fn write_u32_le(&mut self, v: u32) -> Result<(), ()> {
if self.pos < 4 {
return Err(());
}
self.pos -= 4;
self.buf[self.pos..self.pos + 4].copy_from_slice(&v.to_le_bytes());
Ok(())
}
#[inline]
pub fn position(&self) -> usize {
self.pos
}
#[inline]
pub fn bytes_written(&self) -> usize {
self.buf.len() - self.pos
}
#[inline]
pub fn encoded(&self) -> &[u8] {
&self.buf[self.pos..]
}
#[inline]
pub fn remaining(&self) -> usize {
self.pos
}
}
impl fmt::Debug for BackwardByteWriter<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BackwardByteWriter")
.field("pos", &self.pos)
.field("len", &self.buf.len())
.finish()
}
}
#[derive(Clone)]
pub struct ByteReader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> ByteReader<'a> {
#[inline]
pub fn new(buf: &'a [u8]) -> Self {
Self { buf, pos: 0 }
}
#[inline]
pub fn read_byte(&mut self) -> Option<u8> {
if self.pos >= self.buf.len() {
return None;
}
let b = self.buf[self.pos];
self.pos += 1;
Some(b)
}
#[inline]
pub fn read_u32_le(&mut self) -> Option<u32> {
if self.pos + 4 > self.buf.len() {
return None;
}
let v = u32::from_le_bytes([
self.buf[self.pos],
self.buf[self.pos + 1],
self.buf[self.pos + 2],
self.buf[self.pos + 3],
]);
self.pos += 4;
Some(v)
}
#[inline]
pub fn position(&self) -> usize {
self.pos
}
#[inline]
pub fn bytes_consumed(&self) -> usize {
self.pos
}
#[inline]
pub fn remaining(&self) -> usize {
self.buf.len().saturating_sub(self.pos)
}
}
impl fmt::Debug for ByteReader<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ByteReader")
.field("pos", &self.pos)
.field("len", &self.buf.len())
.finish()
}
}
#[derive(Debug, Clone, Copy)]
pub struct RansByteEncSymbol {
pub x_max: u32,
pub rcp_freq: u32,
pub bias: u32,
pub cmpl_freq: u16,
pub rcp_shift: u16,
}
impl RansByteEncSymbol {
#[inline]
pub fn new(start: u32, freq: u32, scale_bits: u32) -> Result<Self, ModelError> {
if !(1..=16).contains(&scale_bits) {
return Err(ModelError::InvalidScaleBits);
}
let max_start = 1u64 << scale_bits;
if (start as u64) > max_start {
return Err(ModelError::StartOutOfRange);
}
if freq == 0 {
return Err(ModelError::ZeroFrequency);
}
if (freq as u64) > max_start - (start as u64) {
return Err(ModelError::FrequencyOutOfRange);
}
Ok(Self::new_unchecked(start, freq, scale_bits))
}
#[inline]
pub(crate) fn new_unchecked(start: u32, freq: u32, scale_bits: u32) -> Self {
debug_assert!(scale_bits <= 16, "scale_bits must be <= 16");
debug_assert!(start <= (1u32 << scale_bits), "start out of range");
debug_assert!(freq <= (1u32 << scale_bits) - start, "freq out of range");
let x_max = ((RANS_BYTE_L >> scale_bits) << 8) * freq;
let cmpl_freq = ((1u32 << scale_bits) - freq) as u16;
if freq < 2 {
Self {
x_max,
rcp_freq: !0u32,
rcp_shift: 0,
bias: start + (1u32 << scale_bits) - 1,
cmpl_freq,
}
} else {
let mut shift = 0u32;
while freq > (1u32 << shift) {
shift += 1;
}
let rcp_freq = (((1u64 << (shift + 31)) + freq as u64 - 1) / freq as u64) as u32;
let rcp_shift = shift - 1;
Self {
x_max,
rcp_freq,
rcp_shift: rcp_shift as u16,
bias: start,
cmpl_freq,
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RansByteDecSymbol {
pub start: u16,
pub freq: u16,
}
impl RansByteDecSymbol {
#[inline]
pub fn new(start: u32, freq: u32) -> Result<Self, ModelError> {
if freq == 0 {
return Err(ModelError::ZeroFrequency);
}
if (start as u64) > (1u64 << 16) {
return Err(ModelError::StartOutOfRange);
}
if (freq as u64) > (1u64 << 16) - (start as u64) {
return Err(ModelError::FrequencyOutOfRange);
}
Ok(Self::new_unchecked(start, freq))
}
#[inline]
pub(crate) fn new_unchecked(start: u32, freq: u32) -> Self {
debug_assert!(start <= (1u32 << 16), "start out of range");
debug_assert!(freq <= (1u32 << 16) - start, "freq out of range");
Self {
start: start as u16,
freq: freq as u16,
}
}
}
#[inline]
pub fn rans_byte_enc_renorm<W: BackwardWriter>(
x: u32,
writer: &mut W,
freq: u32,
scale_bits: u32,
) -> Result<u32, EncodeError> {
let x_max = ((RANS_BYTE_L >> scale_bits) << 8) * freq;
let mut x = x;
if x >= x_max {
while x >= x_max {
writer
.write_byte((x & 0xff) as u8)
.map_err(|_| EncodeError::OutputTooSmall)?;
x >>= 8;
}
}
Ok(x)
}
#[inline]
pub fn rans_byte_enc_put<W: BackwardWriter>(
state: &mut RansByteState,
writer: &mut W,
start: u32,
freq: u32,
scale_bits: u32,
) -> Result<(), EncodeError> {
let x = rans_byte_enc_renorm(state.0, writer, freq, scale_bits)?;
state.0 = ((x / freq) << scale_bits) + (x % freq) + start;
Ok(())
}
#[inline]
pub fn rans_byte_enc_flush<W: BackwardWriter>(
state: &RansByteState,
writer: &mut W,
) -> Result<(), EncodeError> {
let x = state.0;
writer
.write_u32_le(x)
.map_err(|_| EncodeError::OutputTooSmall)
}
#[inline]
pub fn rans_byte_dec_init<R: ForwardReader>(reader: &mut R) -> Result<RansByteState, DecodeError> {
let x = reader.read_u32_le().ok_or(DecodeError::InputTooShort)?;
Ok(RansByteState(x))
}
#[inline]
pub fn rans_byte_dec_get(state: &RansByteState, scale_bits: u32) -> u32 {
state.0 & ((1u32 << scale_bits) - 1)
}
#[inline]
pub fn rans_byte_dec_advance<R: ForwardReader>(
state: &mut RansByteState,
reader: &mut R,
start: u32,
freq: u32,
scale_bits: u32,
) -> Result<(), DecodeError> {
let mask = (1u32 << scale_bits) - 1;
let x = state.0;
let mut x = freq * (x >> scale_bits) + (x & mask) - start;
if x < RANS_BYTE_L {
loop {
let b = reader.read_byte().ok_or(DecodeError::InputTooShort)?;
x = (x << 8) | (b as u32);
if x >= RANS_BYTE_L {
break;
}
}
}
state.0 = x;
Ok(())
}
#[inline]
pub fn rans_byte_enc_put_symbol<W: BackwardWriter>(
state: &mut RansByteState,
writer: &mut W,
sym: &RansByteEncSymbol,
) -> Result<(), EncodeError> {
debug_assert!(sym.x_max != 0, "cannot encode symbol with freq=0");
let mut x = state.0;
if x >= sym.x_max {
while x >= sym.x_max {
writer
.write_byte((x & 0xff) as u8)
.map_err(|_| EncodeError::OutputTooSmall)?;
x >>= 8;
}
}
let q = (((x as u64) * (sym.rcp_freq as u64)) >> 32) >> sym.rcp_shift;
state.0 = x + sym.bias + (q as u32) * (sym.cmpl_freq as u32);
Ok(())
}
#[inline]
pub fn rans_byte_dec_advance_symbol<R: ForwardReader>(
state: &mut RansByteState,
reader: &mut R,
sym: &RansByteDecSymbol,
scale_bits: u32,
) -> Result<(), DecodeError> {
rans_byte_dec_advance(state, reader, sym.start as u32, sym.freq as u32, scale_bits)
}
#[inline]
pub fn rans_byte_dec_advance_step(
state: &mut RansByteState,
start: u32,
freq: u32,
scale_bits: u32,
) {
let mask = (1u32 << scale_bits) - 1;
let x = state.0;
state.0 = freq * (x >> scale_bits) + (x & mask) - start;
}
#[inline]
pub fn rans_byte_dec_advance_symbol_step(
state: &mut RansByteState,
sym: &RansByteDecSymbol,
scale_bits: u32,
) {
rans_byte_dec_advance_step(state, sym.start as u32, sym.freq as u32, scale_bits);
}
#[inline]
pub fn rans_byte_dec_renorm<R: ForwardReader>(
state: &mut RansByteState,
reader: &mut R,
) -> Result<(), DecodeError> {
let mut x = state.0;
if x < RANS_BYTE_L {
loop {
let b = reader.read_byte().ok_or(DecodeError::InputTooShort)?;
x = (x << 8) | (b as u32);
if x >= RANS_BYTE_L {
break;
}
}
state.0 = x;
}
Ok(())
}
pub trait BackwardWriter {
fn write_byte(&mut self, b: u8) -> Result<(), ()>;
fn write_u32_le(&mut self, v: u32) -> Result<(), ()>;
}
pub trait ForwardReader {
fn read_byte(&mut self) -> Option<u8>;
fn read_u32_le(&mut self) -> Option<u32>;
}
impl<'a> BackwardWriter for BackwardByteWriter<'a> {
#[inline]
fn write_byte(&mut self, b: u8) -> Result<(), ()> {
self.write_byte(b)
}
#[inline]
fn write_u32_le(&mut self, v: u32) -> Result<(), ()> {
self.write_u32_le(v)
}
}
impl<'a> ForwardReader for ByteReader<'a> {
#[inline]
fn read_byte(&mut self) -> Option<u8> {
self.read_byte()
}
#[inline]
fn read_u32_le(&mut self) -> Option<u32> {
self.read_u32_le()
}
}
pub struct SliceBackwardWriter<'a>(pub &'a mut [u8]);
impl BackwardWriter for SliceBackwardWriter<'_> {
#[inline]
fn write_byte(&mut self, b: u8) -> Result<(), ()> {
let buf = core::mem::take(&mut self.0);
if buf.is_empty() {
self.0 = buf;
return Err(());
}
let len = buf.len();
buf[len - 1] = b;
self.0 = &mut buf[..len - 1];
Ok(())
}
#[inline]
fn write_u32_le(&mut self, v: u32) -> Result<(), ()> {
let buf = core::mem::take(&mut self.0);
let len = buf.len();
if len < 4 {
self.0 = buf;
return Err(());
}
let bytes = v.to_le_bytes();
buf[len - 4..len].copy_from_slice(&bytes);
self.0 = &mut buf[..len - 4];
Ok(())
}
}
impl<'a> ForwardReader for &'a [u8] {
#[inline]
fn read_byte(&mut self) -> Option<u8> {
if self.is_empty() {
return None;
}
let b = self[0];
*self = &self[1..];
Some(b)
}
#[inline]
fn read_u32_le(&mut self) -> Option<u32> {
if self.len() < 4 {
return None;
}
let v = u32::from_le_bytes([self[0], self[1], self[2], self[3]]);
*self = &self[4..];
Some(v)
}
}
pub struct ByteInterleavedEncoder<'a, W: BackwardWriter> {
state0: RansByteState,
state1: RansByteState,
writer: &'a mut W,
_scale_bits: u32,
_num_symbols: usize,
_odd_symbol: Option<u8>,
}
impl<'a, W: BackwardWriter> ByteInterleavedEncoder<'a, W> {
pub fn new(writer: &'a mut W, scale_bits: u32) -> Self {
Self {
state0: RansByteState::new(),
state1: RansByteState::new(),
writer,
_scale_bits: scale_bits,
_num_symbols: 0,
_odd_symbol: None,
}
}
pub fn encode_reverse(
&mut self,
symbols: &[u8],
esyms: &[RansByteEncSymbol],
) -> Result<(), EncodeError> {
let n = symbols.len();
self._num_symbols = n;
if n == 0 {
return Ok(());
}
if n & 1 != 0 {
let s = symbols[n - 1];
rans_byte_enc_put_symbol(&mut self.state0, &mut *self.writer, &esyms[s as usize])?;
}
let mut i = n & !1;
while i > 0 {
let s1 = symbols[i - 1] as usize;
let s0 = symbols[i - 2] as usize;
rans_byte_enc_put_symbol(&mut self.state1, &mut *self.writer, &esyms[s1])?;
rans_byte_enc_put_symbol(&mut self.state0, &mut *self.writer, &esyms[s0])?;
i = i.wrapping_sub(2);
}
Ok(())
}
pub fn flush(&mut self) -> Result<(), EncodeError> {
rans_byte_enc_flush(&self.state1, &mut *self.writer)?;
rans_byte_enc_flush(&self.state0, &mut *self.writer)?;
Ok(())
}
pub fn finalize(
mut self,
symbols: &[u8],
esyms: &[RansByteEncSymbol],
) -> Result<(), EncodeError> {
self.encode_reverse(symbols, esyms)?;
self.flush()
}
}
pub struct ByteInterleavedDecoder<'a, R: ForwardReader> {
state0: RansByteState,
state1: RansByteState,
reader: &'a mut R,
scale_bits: u32,
}
impl<'a, R: ForwardReader> ByteInterleavedDecoder<'a, R> {
pub fn new(reader: &'a mut R, scale_bits: u32) -> Result<Self, DecodeError> {
let state0 = rans_byte_dec_init(&mut *reader)?;
let state1 = rans_byte_dec_init(&mut *reader)?;
Ok(Self {
state0,
state1,
reader,
scale_bits,
})
}
pub fn decode(
&mut self,
output: &mut [u8],
cum2sym: &[u8],
dsyms: &[RansByteDecSymbol],
) -> Result<usize, DecodeError> {
let n = output.len();
if n == 0 {
return Ok(0);
}
let even_n = n & !1;
let mut i = 0usize;
while i < even_n {
let cf0 = rans_byte_dec_get(&self.state0, self.scale_bits);
let s0 = cum2sym[cf0 as usize] as usize;
let cf1 = rans_byte_dec_get(&self.state1, self.scale_bits);
let s1 = cum2sym[cf1 as usize] as usize;
output[i] = s0 as u8;
output[i + 1] = s1 as u8;
rans_byte_dec_advance_symbol_step(&mut self.state0, &dsyms[s0], self.scale_bits);
rans_byte_dec_advance_symbol_step(&mut self.state1, &dsyms[s1], self.scale_bits);
rans_byte_dec_renorm(&mut self.state0, &mut *self.reader)?;
rans_byte_dec_renorm(&mut self.state1, &mut *self.reader)?;
i += 2;
}
if n & 1 != 0 {
let cf0 = rans_byte_dec_get(&self.state0, self.scale_bits);
let s0 = cum2sym[cf0 as usize] as usize;
output[n - 1] = s0 as u8;
rans_byte_dec_advance_symbol(
&mut self.state0,
&mut *self.reader,
&dsyms[s0],
self.scale_bits,
)?;
}
Ok(n)
}
pub fn states(&self) -> (RansByteState, RansByteState) {
(self.state0, self.state1)
}
}
pub const RANS64_L: u64 = 1u64 << 31;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Rans64State(pub u64);
impl Rans64State {
#[inline]
pub const fn new() -> Self {
Self(RANS64_L)
}
#[inline]
pub const fn get(&self) -> u64 {
self.0
}
}
impl Default for Rans64State {
#[inline]
fn default() -> Self {
Self::new()
}
}
pub struct BackwardWord32Writer<'a> {
buf: &'a mut [u8],
pos: usize, }
impl<'a> BackwardWord32Writer<'a> {
#[inline]
pub fn new(buf: &'a mut [u8]) -> Self {
let len = buf.len();
debug_assert!(
len % 4 == 0,
"BackwardWord32Writer buffer len must be multiple of 4"
);
Self { buf, pos: len }
}
#[inline]
pub fn write_word32(&mut self, v: u32) -> Result<(), ()> {
if self.pos < 4 {
return Err(());
}
self.pos -= 4;
self.buf[self.pos..self.pos + 4].copy_from_slice(&v.to_le_bytes());
Ok(())
}
#[inline]
pub fn position(&self) -> usize {
self.pos
}
#[inline]
pub fn bytes_written(&self) -> usize {
self.buf.len() - self.pos
}
#[inline]
pub fn words_written(&self) -> usize {
(self.buf.len() - self.pos) / 4
}
#[inline]
pub fn encoded(&self) -> &[u8] {
&self.buf[self.pos..]
}
#[inline]
pub fn remaining(&self) -> usize {
self.pos
}
}
impl fmt::Debug for BackwardWord32Writer<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BackwardWord32Writer")
.field("pos", &self.pos)
.field("len", &self.buf.len())
.finish()
}
}
#[derive(Clone)]
pub struct Word32Reader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Word32Reader<'a> {
#[inline]
pub fn new(buf: &'a [u8]) -> Self {
Self { buf, pos: 0 }
}
#[inline]
pub fn read_word32(&mut self) -> Option<u32> {
if self.pos + 4 > self.buf.len() {
return None;
}
let v = u32::from_le_bytes([
self.buf[self.pos],
self.buf[self.pos + 1],
self.buf[self.pos + 2],
self.buf[self.pos + 3],
]);
self.pos += 4;
Some(v)
}
#[inline]
pub fn position(&self) -> usize {
self.pos
}
#[inline]
pub fn bytes_consumed(&self) -> usize {
self.pos
}
#[inline]
pub fn words_consumed(&self) -> usize {
self.pos / 4
}
#[inline]
pub fn remaining(&self) -> usize {
self.buf.len().saturating_sub(self.pos)
}
}
impl fmt::Debug for Word32Reader<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Word32Reader")
.field("pos", &self.pos)
.field("len", &self.buf.len())
.finish()
}
}
#[derive(Debug, Clone, Copy)]
pub struct Rans64EncSymbol {
pub x_max: u64,
pub rcp_freq: u64,
pub bias: u64,
pub cmpl_freq: u32,
pub rcp_shift: u32,
}
impl Rans64EncSymbol {
#[inline]
pub fn new(start: u32, freq: u32, scale_bits: u32) -> Result<Self, ModelError> {
if !(1..=31).contains(&scale_bits) {
return Err(ModelError::InvalidScaleBits);
}
let max_start = 1u64 << scale_bits;
if (start as u64) > max_start {
return Err(ModelError::StartOutOfRange);
}
if freq == 0 {
return Err(ModelError::ZeroFrequency);
}
if (freq as u64) > max_start - (start as u64) {
return Err(ModelError::FrequencyOutOfRange);
}
Ok(Self::new_unchecked(start, freq, scale_bits))
}
#[inline]
pub(crate) fn new_unchecked(start: u32, freq: u32, scale_bits: u32) -> Self {
debug_assert!(scale_bits <= 31, "scale_bits must be <= 31");
debug_assert!((start as u64) <= (1u64 << scale_bits), "start out of range");
debug_assert!(
(freq as u64) <= (1u64 << scale_bits) - (start as u64),
"freq out of range"
);
let x_max = ((RANS64_L >> scale_bits) << 32) * (freq as u64);
let cmpl_freq = ((1u64 << scale_bits) - freq as u64) as u32;
if freq < 2 {
Self {
x_max,
rcp_freq: !0u64,
rcp_shift: 0,
bias: (start as u64) + (1u64 << scale_bits) - 1,
cmpl_freq: (1u64 << scale_bits) as u32 - freq,
}
} else {
let mut shift = 0u32;
while freq > (1u32 << shift) {
shift += 1;
}
let rcp_freq = (((1u128 << (shift + 63)) + (freq as u128) - 1) / (freq as u128)) as u64;
let rcp_shift = shift - 1;
Self {
x_max,
rcp_freq,
rcp_shift,
bias: start as u64,
cmpl_freq,
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Rans64DecSymbol {
pub start: u32,
pub freq: u32,
}
impl Rans64DecSymbol {
#[inline]
pub fn new(start: u32, freq: u32) -> Result<Self, ModelError> {
if freq == 0 {
return Err(ModelError::ZeroFrequency);
}
if (start as u64) > (1u64 << 31) {
return Err(ModelError::StartOutOfRange);
}
if (freq as u64) > (1u64 << 31) - (start as u64) {
return Err(ModelError::FrequencyOutOfRange);
}
Ok(Self::new_unchecked(start, freq))
}
#[inline]
pub(crate) fn new_unchecked(start: u32, freq: u32) -> Self {
debug_assert!((start as u64) <= (1u64 << 31), "start out of range");
debug_assert!(
(freq as u64) <= (1u64 << 31) - (start as u64),
"freq out of range"
);
Self { start, freq }
}
}
#[inline]
pub fn rans64_mul_hi(a: u64, b: u64) -> u64 {
((a as u128) * (b as u128) >> 64) as u64
}
#[inline]
pub fn rans64_enc_renorm(
x: u64,
writer: &mut BackwardWord32Writer,
freq: u32,
scale_bits: u32,
) -> Result<u64, EncodeError> {
let x_max = ((RANS64_L >> scale_bits) << 32) * (freq as u64);
let mut x = x;
if x >= x_max {
while x >= x_max {
writer
.write_word32((x & 0xffffffff) as u32)
.map_err(|_| EncodeError::OutputTooSmall)?;
x >>= 32;
}
}
Ok(x)
}
#[inline]
pub fn rans64_enc_put(
state: &mut Rans64State,
writer: &mut BackwardWord32Writer,
start: u32,
freq: u32,
scale_bits: u32,
) -> Result<(), EncodeError> {
let x = rans64_enc_renorm(state.0, writer, freq, scale_bits)?;
state.0 = ((x / (freq as u64)) << scale_bits) + (x % (freq as u64)) + (start as u64);
Ok(())
}
#[inline]
pub fn rans64_enc_flush(
state: &Rans64State,
writer: &mut BackwardWord32Writer,
) -> Result<(), EncodeError> {
let x = state.0;
writer
.write_word32((x >> 32) as u32)
.map_err(|_| EncodeError::OutputTooSmall)?;
writer
.write_word32((x & 0xffffffff) as u32)
.map_err(|_| EncodeError::OutputTooSmall)?;
Ok(())
}
#[inline]
pub fn rans64_dec_init(reader: &mut Word32Reader) -> Result<Rans64State, DecodeError> {
let lo = reader.read_word32().ok_or(DecodeError::InputTooShort)?;
let hi = reader.read_word32().ok_or(DecodeError::InputTooShort)?;
Ok(Rans64State((lo as u64) | ((hi as u64) << 32)))
}
#[inline]
pub fn rans64_dec_get(state: &Rans64State, scale_bits: u32) -> u32 {
(state.0 & ((1u64 << scale_bits) - 1)) as u32
}
#[inline]
pub fn rans64_dec_advance(
state: &mut Rans64State,
reader: &mut Word32Reader,
start: u32,
freq: u32,
scale_bits: u32,
) -> Result<(), DecodeError> {
let mask = (1u64 << scale_bits) - 1;
let x = state.0;
let mut x = (freq as u64) * (x >> scale_bits) + (x & mask) - (start as u64);
if x < RANS64_L {
loop {
let word = reader.read_word32().ok_or(DecodeError::InputTooShort)?;
x = (x << 32) | (word as u64);
if x >= RANS64_L {
break;
}
}
}
state.0 = x;
Ok(())
}
#[inline]
pub fn rans64_enc_put_symbol(
state: &mut Rans64State,
writer: &mut BackwardWord32Writer,
sym: &Rans64EncSymbol,
) -> Result<(), EncodeError> {
debug_assert!(sym.x_max != 0, "cannot encode symbol with freq=0");
let mut x = state.0;
if x >= sym.x_max {
while x >= sym.x_max {
writer
.write_word32((x & 0xffffffff) as u32)
.map_err(|_| EncodeError::OutputTooSmall)?;
x >>= 32;
}
}
let q = rans64_mul_hi(x, sym.rcp_freq) >> sym.rcp_shift;
state.0 = x + sym.bias + q * (sym.cmpl_freq as u64);
Ok(())
}
#[inline]
pub fn rans64_dec_advance_symbol(
state: &mut Rans64State,
reader: &mut Word32Reader,
sym: &Rans64DecSymbol,
scale_bits: u32,
) -> Result<(), DecodeError> {
rans64_dec_advance(state, reader, sym.start, sym.freq, scale_bits)
}
#[inline]
pub fn rans64_dec_advance_step(state: &mut Rans64State, start: u32, freq: u32, scale_bits: u32) {
let mask = (1u64 << scale_bits) - 1;
let x = state.0;
state.0 = (freq as u64) * (x >> scale_bits) + (x & mask) - (start as u64);
}
#[inline]
pub fn rans64_dec_advance_symbol_step(
state: &mut Rans64State,
sym: &Rans64DecSymbol,
scale_bits: u32,
) {
rans64_dec_advance_step(state, sym.start, sym.freq, scale_bits);
}
#[inline]
pub fn rans64_dec_renorm(
state: &mut Rans64State,
reader: &mut Word32Reader,
) -> Result<(), DecodeError> {
let mut x = state.0;
if x < RANS64_L {
loop {
let word = reader.read_word32().ok_or(DecodeError::InputTooShort)?;
x = (x << 32) | (word as u64);
if x >= RANS64_L {
break;
}
}
state.0 = x;
}
Ok(())
}
pub const RANS_WORD_SCALE_BITS: u32 = 12;
pub const RANS_WORD_L: u32 = 1u32 << 16;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct RansWordState(pub u32);
impl RansWordState {
#[inline]
pub const fn new() -> Self {
Self(RANS_WORD_L)
}
#[inline]
pub const fn get(&self) -> u32 {
self.0
}
}
#[derive(Debug)]
pub struct BackwardWord16Writer<'a> {
buf: &'a mut [u8],
pos: usize,
}
impl<'a> BackwardWord16Writer<'a> {
#[inline]
pub fn new(buf: &'a mut [u8]) -> Self {
let len = buf.len();
Self { buf, pos: len }
}
#[inline]
pub fn write_word16(&mut self, v: u16) -> Result<(), ()> {
if self.pos < 2 {
return Err(());
}
self.pos -= 2;
self.buf[self.pos..self.pos + 2].copy_from_slice(&v.to_le_bytes());
Ok(())
}
#[inline]
pub fn position(&self) -> usize {
self.pos
}
#[inline]
pub fn encoded(&self) -> &[u8] {
&self.buf[self.pos..]
}
#[inline]
pub fn bytes_written(&self) -> usize {
self.buf.len() - self.pos
}
}
#[derive(Debug, Clone)]
pub struct Word16Reader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Word16Reader<'a> {
#[inline]
pub fn new(buf: &'a [u8]) -> Self {
Self { buf, pos: 0 }
}
#[inline]
pub fn read_word16(&mut self) -> Option<u16> {
if self.pos + 2 > self.buf.len() {
return None;
}
let v = u16::from_le_bytes([self.buf[self.pos], self.buf[self.pos + 1]]);
self.pos += 2;
Some(v)
}
#[inline]
pub fn position(&self) -> usize {
self.pos
}
#[inline]
pub fn bytes_consumed(&self) -> usize {
self.pos
}
}
#[derive(Debug, Clone, Copy)]
pub struct RansWordSlot {
pub freq: u16,
pub bias: u16,
}
pub struct RansWordTables<'a> {
pub slots: &'a [RansWordSlot],
pub slot2sym: &'a [u8],
}
#[inline]
pub fn rans_word_check_scale_bits(scale_bits: u32) -> Result<(), ModelError> {
if scale_bits != RANS_WORD_SCALE_BITS {
return Err(ModelError::InvalidScaleBits);
}
Ok(())
}
#[inline]
pub fn rans_word_enc_init() -> RansWordState {
RansWordState::new()
}
#[inline]
pub fn rans_word_enc_renorm(
state: &mut RansWordState,
writer: &mut BackwardWord16Writer,
freq: u32,
scale_bits: u32,
) -> Result<(), EncodeError> {
let x = state.0;
let threshold = ((RANS_WORD_L >> scale_bits) << 16) * freq;
if x >= threshold {
writer
.write_word16((x & 0xffff) as u16)
.map_err(|_| EncodeError::OutputTooSmall)?;
state.0 = x >> 16;
}
Ok(())
}
#[inline]
pub fn rans_word_enc_put(
state: &mut RansWordState,
writer: &mut BackwardWord16Writer,
start: u32,
freq: u32,
scale_bits: u32,
) -> Result<(), EncodeError> {
rans_word_check_scale_bits(scale_bits).map_err(|_| EncodeError::OutputTooSmall)?;
rans_word_enc_renorm(state, writer, freq, scale_bits)?;
let x = state.0;
state.0 = ((x / freq) << scale_bits) + (x % freq) + start;
Ok(())
}
#[inline]
pub fn rans_word_enc_flush(
state: &RansWordState,
writer: &mut BackwardWord16Writer,
) -> Result<(), EncodeError> {
let x = state.0;
writer
.write_word16((x >> 16) as u16)
.map_err(|_| EncodeError::OutputTooSmall)?;
writer
.write_word16((x & 0xffff) as u16)
.map_err(|_| EncodeError::OutputTooSmall)?;
Ok(())
}
#[inline]
pub fn rans_word_dec_init(reader: &mut Word16Reader) -> Result<RansWordState, DecodeError> {
let lo = reader.read_word16().ok_or(DecodeError::InputTooShort)? as u32;
let hi = reader.read_word16().ok_or(DecodeError::InputTooShort)? as u32;
Ok(RansWordState(lo | (hi << 16)))
}
#[inline]
pub fn rans_word_dec_sym(
state: &mut RansWordState,
tables: &RansWordTables<'_>,
scale_bits: u32,
) -> u8 {
debug_assert_eq!(
scale_bits, RANS_WORD_SCALE_BITS,
"word rANS requires scale_bits={}",
RANS_WORD_SCALE_BITS
);
let x = state.0;
let mask = (1u32 << scale_bits) - 1;
let slot = (x & mask) as usize;
state.0 =
(tables.slots[slot].freq as u32) * (x >> scale_bits) + (tables.slots[slot].bias as u32);
tables.slot2sym[slot]
}
#[inline]
pub fn rans_word_dec_renorm(
state: &mut RansWordState,
reader: &mut Word16Reader,
) -> Result<(), DecodeError> {
let x = state.0;
if x < RANS_WORD_L {
let w = reader.read_word16().ok_or(DecodeError::InputTooShort)? as u32;
state.0 = (x << 16) | w;
}
Ok(())
}
#[cfg(any(feature = "alloc", test))]
extern crate alloc as alloc_crate;
pub const ALIAS_LOG2_NSYMS: u32 = 8;
pub const ALIAS_NSYMS: usize = 1 << 8;
#[derive(Debug, Clone)]
#[cfg(any(feature = "alloc", test))]
pub struct AliasTable {
pub divider: [u32; ALIAS_NSYMS],
pub slot_freqs: [u32; ALIAS_NSYMS * 2],
pub slot_adjust: [u32; ALIAS_NSYMS * 2],
pub slot_symbols: [u8; ALIAS_NSYMS * 2],
pub alias_remap: alloc_crate::vec::Vec<u32>,
pub scale_bits: u32,
pub freqs: [u32; ALIAS_NSYMS],
pub cum_freqs: [u32; ALIAS_NSYMS + 1],
}
#[inline]
pub fn rans_byte_alias_normalize_freqs(
freqs: &[u32],
num_symbols: usize,
target_total: u32,
) -> Result<([u32; ALIAS_NSYMS], [u32; ALIAS_NSYMS + 1]), ModelError> {
if num_symbols == 0 || num_symbols > ALIAS_NSYMS {
return Err(ModelError::EmptyInput);
}
if target_total == 0 || (target_total & (target_total - 1)) != 0 {
return Err(ModelError::InvalidScaleBits);
}
let mut out_freqs = [0u32; ALIAS_NSYMS];
let mut cum = [0u32; ALIAS_NSYMS + 1];
cum[0] = 0;
for i in 0..num_symbols {
out_freqs[i] = freqs[i];
cum[i + 1] = cum[i]
.checked_add(freqs[i])
.ok_or(ModelError::TotalMismatch)?;
}
let cur_total = cum[num_symbols];
if cur_total == 0 {
return Err(ModelError::ZeroTotal);
}
let target = target_total as u64;
let cur_total_u64 = cur_total as u64;
for i in 1..=num_symbols {
cum[i] = ((target * cum[i] as u64) / cur_total_u64) as u32;
}
for i in 0..num_symbols {
if out_freqs[i] != 0 && cum[i + 1] == cum[i] {
let mut best_freq = u32::MAX;
let mut best_steal = None;
for j in 0..num_symbols {
let freq = cum[j + 1] - cum[j];
if freq > 1 && freq < best_freq {
best_freq = freq;
best_steal = Some(j);
}
}
let steal = best_steal.ok_or(ModelError::WorkspaceTooSmall)?;
if steal < i {
for j in (steal + 1)..=i {
cum[j] = cum[j].wrapping_sub(1);
}
} else {
for j in (i + 1)..=steal {
cum[j] = cum[j].wrapping_add(1);
}
}
}
}
for i in 0..num_symbols {
if out_freqs[i] == 0 {
debug_assert!(cum[i + 1] == cum[i]);
} else {
debug_assert!(cum[i + 1] > cum[i]);
}
out_freqs[i] = cum[i + 1] - cum[i];
}
Ok((out_freqs, cum))
}
#[inline]
#[cfg(any(feature = "alloc", test))]
pub fn rans_byte_alias_build_table(
freqs: &[u32; ALIAS_NSYMS],
cum_freqs: &[u32; ALIAS_NSYMS + 1],
scale_bits: u32,
) -> AliasTable {
let total = 1u64 << scale_bits;
let tgt_sum = (total / ALIAS_NSYMS as u64) as u32;
let mut divider = [0u32; ALIAS_NSYMS];
let mut slot_freqs = [0u32; ALIAS_NSYMS * 2];
let mut slot_adjust = [0u32; ALIAS_NSYMS * 2];
let mut slot_symbols = [0u8; ALIAS_NSYMS * 2];
let mut remaining = [0u32; ALIAS_NSYMS];
for i in 0..ALIAS_NSYMS {
remaining[i] = freqs[i];
divider[i] = tgt_sum;
slot_symbols[i * 2] = i as u8;
slot_symbols[i * 2 + 1] = i as u8;
}
let mut cur_large = 0;
while cur_large < ALIAS_NSYMS && remaining[cur_large] < tgt_sum {
cur_large += 1;
}
let mut cur_small = 0;
while cur_small < ALIAS_NSYMS && remaining[cur_small] >= tgt_sum {
cur_small += 1;
}
let mut next_small = cur_small + 1;
while cur_large < ALIAS_NSYMS && cur_small < ALIAS_NSYMS {
slot_symbols[cur_small * 2] = cur_large as u8;
divider[cur_small] = remaining[cur_small];
remaining[cur_large] -= tgt_sum - divider[cur_small];
if remaining[cur_large] >= tgt_sum || next_small <= cur_large {
cur_small = next_small;
while cur_small < ALIAS_NSYMS && remaining[cur_small] >= tgt_sum {
cur_small += 1;
}
next_small = cur_small + 1;
} else {
cur_small = cur_large;
}
while cur_large < ALIAS_NSYMS && remaining[cur_large] < tgt_sum {
cur_large += 1;
}
}
let mut assigned = [0u32; ALIAS_NSYMS];
let total_slots = total as usize;
let mut alias_remap = alloc_crate::vec![0u32; total_slots];
for i in 0..ALIAS_NSYMS {
let j = slot_symbols[i * 2] as usize;
let sym0_height = divider[i]; let sym1_height = tgt_sum - divider[i]; let base0 = assigned[i];
let base1 = assigned[j];
let cbase0 = cum_freqs[i] + base0;
let cbase1 = cum_freqs[j] + base1;
divider[i] = i as u32 * tgt_sum + sym0_height;
slot_freqs[i * 2 + 1] = freqs[i];
slot_freqs[i * 2] = freqs[j];
slot_adjust[i * 2 + 1] = (i as u32 * tgt_sum).wrapping_sub(base0);
slot_adjust[i * 2] = (i as u32 * tgt_sum).wrapping_sub(base1.wrapping_sub(sym0_height));
for k in 0..sym0_height {
let idx = (cbase0 + k) as usize;
if idx < total_slots {
alias_remap[idx] = k + i as u32 * tgt_sum;
}
}
for k in 0..sym1_height {
let idx = (cbase1 + k) as usize;
if idx < total_slots {
alias_remap[idx] = (k + sym0_height) + i as u32 * tgt_sum;
}
}
assigned[i] += sym0_height;
assigned[j] += sym1_height;
}
AliasTable {
divider,
slot_freqs,
slot_adjust,
slot_symbols,
alias_remap,
scale_bits,
freqs: *freqs,
cum_freqs: *cum_freqs,
}
}
#[inline]
#[cfg(feature = "alloc")]
pub fn rans_byte_alias_enc_put<W: BackwardWriter>(
state: &mut RansByteState,
writer: &mut W,
table: &AliasTable,
s: u8,
scale_bits: u32,
) -> Result<(), EncodeError> {
let freq = table.freqs[s as usize];
let x = rans_byte_enc_renorm(state.0, writer, freq, scale_bits)?;
let slot = table.alias_remap[(x % freq + table.cum_freqs[s as usize]) as usize];
state.0 = ((x / freq) << scale_bits) + slot;
Ok(())
}
#[inline]
#[cfg(feature = "alloc")]
pub fn rans_byte_alias_dec_get(
state: RansByteState,
table: &AliasTable,
scale_bits: u32,
) -> (u8, RansByteState) {
let x = state.0;
let mask = (1u32 << scale_bits).wrapping_sub(1);
let xm = x & mask;
let bucket_id = (xm >> (scale_bits - ALIAS_LOG2_NSYMS)) as usize;
let mut bucket2 = bucket_id * 2;
if xm < table.divider[bucket_id] {
bucket2 += 1;
}
let s = table.slot_symbols[bucket2];
let new_x = table.slot_freqs[bucket2] * (x >> scale_bits) + xm - table.slot_adjust[bucket2];
(s, RansByteState(new_x))
}
#[inline]
#[cfg(feature = "alloc")]
pub fn rans_byte_alias_dec_renorm<R: ForwardReader>(
state: &mut RansByteState,
reader: &mut R,
) -> Result<(), DecodeError> {
rans_byte_dec_renorm(state, reader)
}
#[inline]
#[cfg(feature = "alloc")]
pub fn rans_byte_alias_dec_advance<R: ForwardReader>(
state: &mut RansByteState,
reader: &mut R,
table: &AliasTable,
scale_bits: u32,
) -> Result<u8, DecodeError> {
let (s, new_state) = rans_byte_alias_dec_get(*state, table, scale_bits);
*state = new_state;
rans_byte_alias_dec_renorm(state, reader)?;
Ok(s)
}
#[cfg(test)]
mod tests {
use super::*;
extern crate alloc;
#[test]
fn test_state_init() {
let s = RansByteState::new();
assert_eq!(s.get(), RANS_BYTE_L);
}
#[test]
fn test_backward_writer_basic() {
let mut buf = [0u8; 10];
let pos;
{
let mut w = BackwardByteWriter::new(&mut buf);
assert!(w.write_byte(0xAB).is_ok());
assert_eq!(w.position(), 9);
assert!(w.write_byte(0xCD).is_ok());
pos = w.position();
}
assert_eq!(pos, 8);
assert_eq!(buf[8], 0xCD);
assert_eq!(buf[9], 0xAB);
assert_eq!(buf[8..10], [0xCD, 0xAB]);
}
#[test]
fn test_backward_writer_full() {
let mut buf = [0u8; 2];
let mut w = BackwardByteWriter::new(&mut buf);
assert!(w.write_byte(1).is_ok());
assert!(w.write_byte(2).is_ok());
assert!(w.write_byte(3).is_err());
}
#[test]
fn test_backward_writer_u32_le() {
let mut buf = [0u8; 8];
let pos1;
{
let mut w = BackwardByteWriter::new(&mut buf);
assert!(w.write_u32_le(0x01020304).is_ok());
pos1 = w.position();
assert!(w.write_u32_le(0x05060708).is_ok());
}
assert_eq!(pos1, 4);
assert_eq!(buf[4..8], [0x04, 0x03, 0x02, 0x01]);
assert_eq!(buf[0..4], [0x08, 0x07, 0x06, 0x05]);
}
#[test]
fn test_forward_reader_basic() {
let buf = [0x10, 0x20, 0x30, 0x40];
let mut r = ByteReader::new(&buf);
assert_eq!(r.read_byte(), Some(0x10));
assert_eq!(r.read_byte(), Some(0x20));
assert_eq!(r.position(), 2);
}
#[test]
fn test_forward_reader_u32_le() {
let buf = [0x04, 0x03, 0x02, 0x01, 0x08, 0x07, 0x06, 0x05];
let mut r = ByteReader::new(&buf);
assert_eq!(r.read_u32_le(), Some(0x01020304));
assert_eq!(r.read_u32_le(), Some(0x05060708));
assert_eq!(r.read_u32_le(), None);
}
#[test]
fn test_enc_symbol_init() {
let sym = RansByteEncSymbol::new(100, 2, 14).unwrap();
assert!(sym.x_max > 0);
assert!(sym.rcp_freq > 0);
assert_eq!(sym.bias, 100);
assert_eq!(sym.cmpl_freq, ((1u32 << 14) - 2) as u16);
assert_eq!(sym.rcp_shift, 0);
}
#[test]
fn test_enc_symbol_init_freq_one() {
let sym = RansByteEncSymbol::new(100, 1, 14).unwrap();
assert!(sym.x_max > 0);
assert_eq!(sym.rcp_freq, !0u32);
assert_eq!(sym.rcp_shift, 0);
assert_eq!(sym.bias, 100 + (1u32 << 14) - 1);
}
#[test]
fn test_enc_symbol_init_max_freq() {
let scale_bits = 14;
let start = 0;
let freq = 1u32 << scale_bits;
let sym = RansByteEncSymbol::new(start, freq, scale_bits).unwrap();
assert!(sym.x_max > 0);
}
#[test]
fn test_slice_backward_writer() {
let mut buf = [0u8; 10];
let mut writer = SliceBackwardWriter(&mut buf[..]);
assert!(writer.write_byte(0xAB).is_ok());
assert_eq!(writer.0.len(), 9);
assert!(writer.write_byte(0xCD).is_ok());
assert_eq!(writer.0.len(), 8);
assert_eq!(buf[8], 0xCD);
assert_eq!(buf[9], 0xAB);
}
#[test]
fn test_slice_forward_reader() {
let buf = [0x10, 0x20, 0x30];
let mut r = &buf[..];
assert_eq!(r.read_byte(), Some(0x10));
assert_eq!(r.read_byte(), Some(0x20));
assert_eq!(r.read_byte(), Some(0x30));
assert_eq!(r.read_byte(), None);
}
#[test]
fn test_roundtrip_single_symbol() {
let scale_bits = 14;
let symbols = [42u8; 100];
let mut out = [0u8; 1024];
let mut writer = BackwardByteWriter::new(&mut out);
let esym = RansByteEncSymbol::new(0, 1u32 << scale_bits, scale_bits).unwrap();
let mut state = RansByteState::new();
for _i in (0..symbols.len()).rev() {
rans_byte_enc_put_symbol(&mut state, &mut writer, &esym).unwrap();
}
rans_byte_enc_flush(&state, &mut writer).unwrap();
let encoded = writer.encoded();
assert!(!encoded.is_empty(), "encoded output should not be empty");
let mut reader = ByteReader::new(encoded);
let dsym = RansByteDecSymbol::new(0, 1u32 << scale_bits).unwrap();
let mut dec_state = rans_byte_dec_init(&mut reader).unwrap();
let cum2sym = [42u8; 1 << 14];
let mut output = alloc::vec![0u8; symbols.len()];
for i in 0..symbols.len() {
let cf = rans_byte_dec_get(&dec_state, scale_bits);
let s = cum2sym[cf as usize];
output[i] = s;
rans_byte_dec_advance_symbol(&mut dec_state, &mut reader, &dsym, scale_bits).unwrap();
}
assert_eq!(
output,
&symbols[..],
"single-symbol round-trip should match"
);
assert_eq!(
output,
&symbols[..],
"single-symbol round-trip should match"
);
}
#[test]
fn test_roundtrip_two_symbols() {
let scale_bits = 2;
let _total = 1u32 << scale_bits; let freq0 = 1u32;
let freq1 = 3u32;
let symbols: alloc::vec::Vec<u8> = (0..10).map(|i| (i % 2) as u8).collect();
let mut out = [0u8; 1024];
let mut writer = BackwardByteWriter::new(&mut out);
let mut state = RansByteState::new();
for idx in (0..symbols.len()).rev() {
let s = symbols[idx];
let start = if s == 0 { 0 } else { freq0 };
let freq = if s == 0 { freq0 } else { freq1 };
rans_byte_enc_put(&mut state, &mut writer, start, freq, scale_bits).unwrap();
}
rans_byte_enc_flush(&state, &mut writer).unwrap();
let encoded = writer.encoded();
assert!(!encoded.is_empty());
let mut reader = ByteReader::new(encoded);
let dsym0 = RansByteDecSymbol::new(0, freq0).unwrap();
let dsym1 = RansByteDecSymbol::new(freq0, freq1).unwrap();
let mut dec_state = rans_byte_dec_init(&mut reader).unwrap();
let cum2sym = [0u8, 0u8, 1u8, 1u8];
let mut output = alloc::vec![0u8; symbols.len()];
for i in 0..symbols.len() {
let cf = rans_byte_dec_get(&dec_state, scale_bits);
let s = cum2sym[cf as usize] as usize;
output[i] = s as u8;
let dsym = if s == 0 { &dsym0 } else { &dsym1 };
rans_byte_dec_advance_symbol(&mut dec_state, &mut reader, dsym, scale_bits).unwrap();
}
assert_eq!(output, symbols, "two-symbol round-trip should match");
}
#[test]
fn test_slice_trait_roundtrip() {
let scale_bits = 14;
let symbols: alloc::vec::Vec<u8> = (0..50).map(|i| (i % 17) as u8).collect();
let mut out = [0u8; 1024];
let mut writer = SliceBackwardWriter(&mut out[..]);
let mut state = RansByteState::new();
let total = 1u32 << scale_bits;
let n_syms = 17u32;
let base_freq = total / n_syms;
for i in (0..symbols.len()).rev() {
let s = symbols[i] as u32;
let start = s * base_freq;
let freq = base_freq;
rans_byte_enc_put(&mut state, &mut writer, start, freq, scale_bits).unwrap();
}
rans_byte_enc_flush(&state, &mut writer).unwrap();
let used = writer.0.len();
let encoded = &out[used..];
let mut reader = ByteReader::new(encoded);
let mut dec_state = rans_byte_dec_init(&mut reader).unwrap();
let mut output = alloc::vec![0u8; symbols.len()];
for i in 0..symbols.len() {
let cf = rans_byte_dec_get(&dec_state, scale_bits);
let s = cf / base_freq;
output[i] = s as u8;
let start = s * base_freq;
rans_byte_dec_advance(&mut dec_state, &mut reader, start, base_freq, scale_bits)
.unwrap();
}
assert_eq!(output, symbols, "uniform-symbol round-trip should match");
}
#[test]
fn test_reciprocal_roundtrip() {
let scale_bits = 14;
let total = 1u32 << scale_bits;
let freq0 = total / 3;
let freq1 = total / 3;
let freq2 = total - freq0 - freq1;
let esym0 = RansByteEncSymbol::new(0, freq0, scale_bits).unwrap();
let esym1 = RansByteEncSymbol::new(freq0, freq1, scale_bits).unwrap();
let esym2 = RansByteEncSymbol::new(freq0 + freq1, freq2, scale_bits).unwrap();
let dsym0 = RansByteDecSymbol::new(0, freq0).unwrap();
let dsym1 = RansByteDecSymbol::new(freq0, freq1).unwrap();
let dsym2 = RansByteDecSymbol::new(freq0 + freq1, freq2).unwrap();
let symbols: alloc::vec::Vec<u8> = (0..50).map(|i| (i % 3) as u8).collect();
let mut out = [0u8; 1024];
let mut writer = BackwardByteWriter::new(&mut out);
let mut state = RansByteState::new();
for idx in (0..symbols.len()).rev() {
let s = symbols[idx] as usize;
let esym = match s {
0 => &esym0,
1 => &esym1,
_ => &esym2,
};
rans_byte_enc_put_symbol(&mut state, &mut writer, esym).unwrap();
}
rans_byte_enc_flush(&state, &mut writer).unwrap();
let encoded = writer.encoded();
let mut reader = ByteReader::new(encoded);
let mut dec_state = rans_byte_dec_init(&mut reader).unwrap();
let cum2sym: alloc::vec::Vec<u8> = (0..total as usize)
.map(|i| {
if i < freq0 as usize {
0
} else if i < (freq0 + freq1) as usize {
1
} else {
2
}
})
.collect();
let mut output = alloc::vec![0u8; symbols.len()];
for i in 0..symbols.len() {
let cf = rans_byte_dec_get(&dec_state, scale_bits);
let s = cum2sym[cf as usize] as usize;
output[i] = s as u8;
let dsym = match s {
0 => &dsym0,
1 => &dsym1,
_ => &dsym2,
};
rans_byte_dec_advance_symbol(&mut dec_state, &mut reader, dsym, scale_bits).unwrap();
}
assert_eq!(output, symbols, "reciprocal round-trip should match");
}
#[test]
fn test_interleaved_roundtrip() {
let scale_bits = 14;
let symbols: alloc::vec::Vec<u8> = (0..77).map(|i| (i % 7) as u8).collect();
let total = 1u32 << scale_bits;
let base_freq = total / 7;
let esyms: alloc::vec::Vec<RansByteEncSymbol> = (0..7)
.map(|i| RansByteEncSymbol::new(i * base_freq, base_freq, scale_bits).unwrap())
.collect();
let dsyms: alloc::vec::Vec<RansByteDecSymbol> = (0..7)
.map(|i| RansByteDecSymbol::new(i * base_freq, base_freq).unwrap())
.collect();
let mut out = [0u8; 2048];
let mut writer = BackwardByteWriter::new(&mut out);
let mut s0 = RansByteState::new();
let mut s1 = RansByteState::new();
let n = symbols.len();
if n & 1 != 0 {
let s = symbols[n - 1] as usize;
rans_byte_enc_put_symbol(&mut s0, &mut writer, &esyms[s]).unwrap();
}
let mut i = n & !1;
while i > 0 {
let s1_idx = symbols[i - 1] as usize;
let s0_idx = symbols[i - 2] as usize;
rans_byte_enc_put_symbol(&mut s1, &mut writer, &esyms[s1_idx]).unwrap();
rans_byte_enc_put_symbol(&mut s0, &mut writer, &esyms[s0_idx]).unwrap();
i = i.wrapping_sub(2);
}
rans_byte_enc_flush(&s1, &mut writer).unwrap();
rans_byte_enc_flush(&s0, &mut writer).unwrap();
let encoded = writer.encoded();
let mut reader = ByteReader::new(encoded);
let mut d0 = rans_byte_dec_init(&mut reader).unwrap();
let mut d1 = rans_byte_dec_init(&mut reader).unwrap();
let cum2sym: alloc::vec::Vec<u8> = (0..total as usize)
.map(|i| (i / base_freq as usize) as u8)
.collect();
let mut output = alloc::vec![0u8; n];
let even_n = n & !1;
let mut pos = 0;
while pos < even_n {
let cf0 = rans_byte_dec_get(&d0, scale_bits);
let s0 = cum2sym[cf0 as usize] as usize;
let cf1 = rans_byte_dec_get(&d1, scale_bits);
let s1 = cum2sym[cf1 as usize] as usize;
output[pos] = s0 as u8;
output[pos + 1] = s1 as u8;
rans_byte_dec_advance_symbol_step(&mut d0, &dsyms[s0], scale_bits);
rans_byte_dec_advance_symbol_step(&mut d1, &dsyms[s1], scale_bits);
rans_byte_dec_renorm(&mut d0, &mut reader).unwrap();
rans_byte_dec_renorm(&mut d1, &mut reader).unwrap();
pos += 2;
}
if n & 1 != 0 {
let cf0 = rans_byte_dec_get(&d0, scale_bits);
let s0 = cum2sym[cf0 as usize] as usize;
output[n - 1] = s0 as u8;
rans_byte_dec_advance_symbol(&mut d0, &mut reader, &dsyms[s0], scale_bits).unwrap();
}
assert_eq!(output, symbols, "interleaved round-trip should match");
}
#[test]
fn test_reciprocal_equals_division() {
let scale_bits = 14;
let total = 1u32 << scale_bits;
fn div_put(x: u32, start: u32, freq: u32, scale_bits: u32) -> u32 {
((x / freq) << scale_bits) + (x % freq) + start
}
let test_freqs = [1, 2, 3, 5, 7, 10, 100, 1000, total / 2, total - 1];
for &freq in &test_freqs {
let start = 0;
let esym = RansByteEncSymbol::new(start, freq, scale_bits).unwrap();
let test_states = [
RANS_BYTE_L,
RANS_BYTE_L + 1,
RANS_BYTE_L * 2,
RANS_BYTE_L * 4,
RANS_BYTE_L * 8,
(1u32 << 31) - 1,
];
for &test_state in &test_states {
if test_state >= esym.x_max {
continue;
}
let expected = div_put(test_state, start, freq, scale_bits);
let mut state_fast = RansByteState(test_state);
let mut temp = [0u8; 8];
let mut w = BackwardByteWriter::new(&mut temp);
rans_byte_enc_put_symbol(&mut state_fast, &mut w, &esym).unwrap();
assert_eq!(
state_fast.0, expected,
"reciprocal mismatch for freq={}, start={}, state={}",
freq, start, test_state
);
}
}
}
#[test]
fn test_oracle_reciprocal_parameters() {
let sym = RansByteEncSymbol::new(0, 10, 14).unwrap();
assert_eq!(sym.x_max, 1310720);
assert_eq!(sym.rcp_freq, 3435973837);
assert_eq!(sym.bias, 0);
assert_eq!(sym.cmpl_freq as u32, 16374);
assert_eq!(sym.rcp_shift as u32, 3);
let sym = RansByteEncSymbol::new(100, 1, 14).unwrap();
assert_eq!(sym.x_max, 131072);
assert_eq!(sym.rcp_freq, 4294967295);
assert_eq!(sym.bias, 16483);
assert_eq!(sym.cmpl_freq as u32, 16383);
assert_eq!(sym.rcp_shift as u32, 0);
let sym = RansByteEncSymbol::new(0, 16384, 14).unwrap();
assert_eq!(sym.x_max, 2147483648);
assert_eq!(sym.rcp_freq, 2147483648);
assert_eq!(sym.bias, 0);
assert_eq!(sym.cmpl_freq as u32, 0);
assert_eq!(sym.rcp_shift as u32, 13);
let sym = RansByteEncSymbol::new(0, 2, 14).unwrap();
assert_eq!(sym.cmpl_freq as u32, 16382);
assert_eq!(sym.rcp_shift as u32, 0);
assert!(sym.rcp_freq == 2147483648 || sym.rcp_freq > 0);
}
#[test]
fn test_reciprocal_freq_one() {
let scale_bits = 14;
let freq = 1u32;
let start = 100;
let esym = RansByteEncSymbol::new(start, freq, scale_bits).unwrap();
fn expected(x: u32, start: u32, scale_bits: u32) -> u32 {
x * (1u32 << scale_bits) + start
}
let test_states = [RANS_BYTE_L, RANS_BYTE_L + 10, RANS_BYTE_L * 3, (1u32 << 30)];
for &test_state in &test_states {
if test_state >= esym.x_max {
continue;
}
let mut state = RansByteState(test_state);
let mut tmp = [0u8; 8];
let mut w = BackwardByteWriter::new(&mut tmp);
rans_byte_enc_put_symbol(&mut state, &mut w, &esym).unwrap();
assert_eq!(
state.0,
expected(test_state, start, scale_bits),
"freq=1 mismatch for state={}",
test_state
);
}
}
#[test]
fn test_decoder_symbol_init() {
let dsym = RansByteDecSymbol::new(100, 50).unwrap();
assert_eq!(dsym.start, 100);
assert_eq!(dsym.freq, 50);
}
#[test]
fn test_reader_exhaustion() {
let buf = [1u8; 3];
let mut reader = ByteReader::new(&buf);
assert!(reader.read_byte().is_some());
assert!(reader.read_byte().is_some());
assert!(reader.read_byte().is_some());
assert!(reader.read_byte().is_none());
assert!(reader.read_u32_le().is_none());
}
#[test]
fn test_writer_exhaustion() {
let mut buf = [0u8; 2];
let mut writer = BackwardByteWriter::new(&mut buf);
assert!(writer.write_u32_le(0x12345678).is_err());
assert!(writer.write_byte(1).is_ok());
assert!(writer.write_byte(2).is_ok());
assert!(writer.write_byte(3).is_err());
}
#[test]
fn test_rans64_state_init() {
let s = Rans64State::new();
assert_eq!(s.get(), RANS64_L);
assert_eq!(s, Rans64State::default());
}
#[test]
fn test_rans64_word32_writer_basic() {
let mut buf = [0u8; 12];
let pos;
let words;
{
let mut w = BackwardWord32Writer::new(&mut buf);
assert!(w.write_word32(0xDEADBEEF).is_ok());
assert!(w.write_word32(0xCAFEBABE).is_ok());
assert!(w.write_word32(0x12345678).is_ok());
pos = w.position();
words = w.words_written();
}
assert_eq!(pos, 0);
assert_eq!(words, 3);
assert_eq!(buf[8..12], 0xDEADBEEFu32.to_le_bytes());
assert_eq!(buf[4..8], 0xCAFEBABEu32.to_le_bytes());
assert_eq!(buf[0..4], 0x12345678u32.to_le_bytes());
}
#[test]
fn test_rans64_word32_reader_basic() {
let mut buf = [0u8; 12];
let v0 = 0xDEADBEEFu32;
let v1 = 0xCAFEBABEu32;
let v2 = 0x12345678u32;
buf[0..4].copy_from_slice(&v0.to_le_bytes());
buf[4..8].copy_from_slice(&v1.to_le_bytes());
buf[8..12].copy_from_slice(&v2.to_le_bytes());
let mut r = Word32Reader::new(&buf);
assert_eq!(r.read_word32(), Some(v0));
assert_eq!(r.read_word32(), Some(v1));
assert_eq!(r.read_word32(), Some(v2));
assert_eq!(r.read_word32(), None);
assert_eq!(r.words_consumed(), 3);
}
#[test]
fn test_rans64_enc_symbol_init() {
let sym = Rans64EncSymbol::new(100, 2, 14).unwrap();
assert!(sym.x_max > 0);
assert!(sym.rcp_freq > 0);
assert_eq!(sym.bias, 100);
assert_eq!(sym.cmpl_freq, ((1u32 << 14) - 2) as u32);
assert_eq!(sym.rcp_shift, 0);
let expected_x_max = ((RANS64_L >> 14) << 32) * 2;
assert_eq!(sym.x_max, expected_x_max);
}
#[test]
fn test_rans64_enc_symbol_init_freq_one() {
let sym = Rans64EncSymbol::new(100, 1, 14).unwrap();
assert!(sym.x_max > 0);
assert_eq!(sym.rcp_freq, !0u64);
assert_eq!(sym.rcp_shift, 0);
assert_eq!(sym.bias, 100 + (1u64 << 14) - 1);
}
#[test]
fn test_rans64_enc_symbol_init_large_scale() {
let scale_bits = 30;
let start = 0;
let freq = (1u32 << 29) + 1; let sym = Rans64EncSymbol::new(start, freq, scale_bits).unwrap();
assert!(sym.rcp_freq > 0);
assert!(sym.x_max > 0);
assert_eq!(sym.bias, 0);
}
#[test]
fn test_rans64_roundtrip_single_symbol_division() {
let scale_bits = 14;
let n = 50;
let symbols = [99u8; 50];
let mut out = [0u8; 4096];
let mut writer = BackwardWord32Writer::new(&mut out);
let mut state = Rans64State::new();
for _i in (0..n).rev() {
rans64_enc_put(
&mut state,
&mut writer,
0, 1u32 << scale_bits, scale_bits,
)
.unwrap();
}
rans64_enc_flush(&state, &mut writer).unwrap();
let encoded = writer.encoded();
assert!(
encoded.len() >= 8,
"encoded should have at least 8 bytes (2 words)"
);
assert!(!encoded.is_empty());
let mut reader = Word32Reader::new(encoded);
let dsym = Rans64DecSymbol::new(0, 1u32 << scale_bits).unwrap();
let mut dec_state = rans64_dec_init(&mut reader).unwrap();
let cum2sym = [99u8; 1 << 14];
let mut output = alloc::vec![0u8; n];
for i in 0..n {
let cf = rans64_dec_get(&dec_state, scale_bits);
let s = cum2sym[cf as usize];
output[i] = s;
rans64_dec_advance_symbol(&mut dec_state, &mut reader, &dsym, scale_bits).unwrap();
}
assert_eq!(output, symbols, "64-bit single-symbol division round-trip");
}
#[test]
fn test_rans64_roundtrip_two_symbols_division() {
let scale_bits = 14;
let total = 1u32 << scale_bits;
let freq0 = total / 4;
let freq1 = total - freq0;
let symbols: alloc::vec::Vec<u8> = (0..30).map(|i| (i % 2) as u8).collect();
let mut out = [0u8; 4096];
let mut writer = BackwardWord32Writer::new(&mut out);
let mut state = Rans64State::new();
for idx in (0..symbols.len()).rev() {
let s = symbols[idx];
let start = if s == 0 { 0 } else { freq0 };
let freq = if s == 0 { freq0 } else { freq1 };
rans64_enc_put(&mut state, &mut writer, start, freq, scale_bits).unwrap();
}
rans64_enc_flush(&state, &mut writer).unwrap();
let encoded = writer.encoded();
assert!(encoded.len() >= 8, "encoded length = {}", encoded.len());
let mut reader = Word32Reader::new(encoded);
let dsym0 = Rans64DecSymbol::new(0, freq0).unwrap();
let dsym1 = Rans64DecSymbol::new(freq0, freq1).unwrap();
let mut dec_state = rans64_dec_init(&mut reader).unwrap();
let mut cum2sym = alloc::vec![1u8; total as usize];
for i in 0..freq0 as usize {
cum2sym[i] = 0;
}
let mut output = alloc::vec![0u8; symbols.len()];
for i in 0..symbols.len() {
let cf = rans64_dec_get(&dec_state, scale_bits);
let s = cum2sym[cf as usize] as usize;
output[i] = s as u8;
let dsym = if s == 0 { &dsym0 } else { &dsym1 };
rans64_dec_advance_symbol(&mut dec_state, &mut reader, dsym, scale_bits).unwrap();
}
assert_eq!(output, symbols, "64-bit two-symbol division round-trip");
}
#[test]
fn test_rans64_reciprocal_equals_division() {
let scale_bits = 14;
let total = 1u32 << scale_bits;
fn div_put(x: u64, start: u32, freq: u32, scale_bits: u32) -> u64 {
((x / (freq as u64)) << scale_bits) + (x % (freq as u64)) + (start as u64)
}
let test_freqs = [1, 2, 3, 5, 7, 10, 100, 1000, total / 2, total - 1];
for &freq in &test_freqs {
let start = 0u32;
let esym = Rans64EncSymbol::new(start, freq, scale_bits).unwrap();
let test_states = [
RANS64_L,
RANS64_L + 1,
RANS64_L * 2,
RANS64_L * 4,
RANS64_L * 8,
(1u64 << 62) - 1,
];
for &test_state in &test_states {
if test_state >= esym.x_max {
continue;
}
let expected = div_put(test_state, start, freq, scale_bits);
let mut state_fast = Rans64State(test_state);
let mut temp = [0u8; 16];
let mut w = BackwardWord32Writer::new(&mut temp);
rans64_enc_put_symbol(&mut state_fast, &mut w, &esym).unwrap();
assert_eq!(
state_fast.0, expected,
"64-bit reciprocal mismatch for freq={}, start={}, state={}",
freq, start, test_state
);
}
}
}
#[test]
fn test_rans64_roundtrip_reciprocal() {
let scale_bits = 14;
let total = 1u32 << scale_bits;
let freq0 = total / 3;
let freq1 = total / 3;
let freq2 = total - freq0 - freq1;
let esym0 = Rans64EncSymbol::new(0, freq0, scale_bits).unwrap();
let esym1 = Rans64EncSymbol::new(freq0, freq1, scale_bits).unwrap();
let esym2 = Rans64EncSymbol::new(freq0 + freq1, freq2, scale_bits).unwrap();
let dsym0 = Rans64DecSymbol::new(0, freq0).unwrap();
let dsym1 = Rans64DecSymbol::new(freq0, freq1).unwrap();
let dsym2 = Rans64DecSymbol::new(freq0 + freq1, freq2).unwrap();
let symbols: alloc::vec::Vec<u8> = (0..50).map(|i| (i % 3) as u8).collect();
let mut out = [0u8; 4096];
let mut writer = BackwardWord32Writer::new(&mut out);
let mut state = Rans64State::new();
for idx in (0..symbols.len()).rev() {
let s = symbols[idx] as usize;
let esym = match s {
0 => &esym0,
1 => &esym1,
_ => &esym2,
};
rans64_enc_put_symbol(&mut state, &mut writer, esym).unwrap();
}
rans64_enc_flush(&state, &mut writer).unwrap();
let encoded = writer.encoded();
let mut reader = Word32Reader::new(encoded);
let mut dec_state = rans64_dec_init(&mut reader).unwrap();
let cum2sym: alloc::vec::Vec<u8> = (0..total as usize)
.map(|i| {
if i < freq0 as usize {
0
} else if i < (freq0 + freq1) as usize {
1
} else {
2
}
})
.collect();
let mut output = alloc::vec![0u8; symbols.len()];
for i in 0..symbols.len() {
let cf = rans64_dec_get(&dec_state, scale_bits);
let s = cum2sym[cf as usize] as usize;
output[i] = s as u8;
let dsym = match s {
0 => &dsym0,
1 => &dsym1,
_ => &dsym2,
};
rans64_dec_advance_symbol(&mut dec_state, &mut reader, dsym, scale_bits).unwrap();
}
assert_eq!(output, symbols, "64-bit reciprocal round-trip");
}
#[test]
fn test_rans64_step_operations() {
let scale_bits = 14;
let total = 1u32 << scale_bits;
let freq = total / 2;
let start = 0;
let dsym = Rans64DecSymbol::new(start, freq).unwrap();
let state_val = RANS64_L * 4;
let mut state_advance = Rans64State(state_val);
let mut state_step = Rans64State(state_val);
let dummy_buf = [0u8; 16];
let mut reader = Word32Reader::new(&dummy_buf);
rans64_dec_advance(&mut state_advance, &mut reader, start, freq, scale_bits).unwrap();
rans64_dec_advance_step(&mut state_step, start, freq, scale_bits);
assert_eq!(
state_advance.0, state_step.0,
"step-only advance should match regular advance when no renorm needed"
);
let mut state_adv_sym = Rans64State(state_val);
let mut state_step_sym = Rans64State(state_val);
let mut reader2 = Word32Reader::new(&dummy_buf);
rans64_dec_advance_symbol(&mut state_adv_sym, &mut reader2, &dsym, scale_bits).unwrap();
rans64_dec_advance_symbol_step(&mut state_step_sym, &dsym, scale_bits);
assert_eq!(
state_adv_sym.0, state_step_sym.0,
"step-only symbol advance should match regular"
);
}
#[test]
fn test_rans64_state_transition_cycle() {
let scale_bits = 14;
let _total = 1u32 << scale_bits;
let freq = 100u32;
let start = 500u32;
let esym = Rans64EncSymbol::new(start, freq, scale_bits).unwrap();
let _dsym = Rans64DecSymbol::new(start, freq).unwrap();
let x = RANS64_L;
let mut enc_state = Rans64State(x);
let mut tmp = [0u8; 16];
let mut w = BackwardWord32Writer::new(&mut tmp);
rans64_enc_put_symbol(&mut enc_state, &mut w, &esym).unwrap();
let encoded_x = enc_state.0;
let dummy = [0u8; 16];
let mut r = Word32Reader::new(&dummy);
let mut dec_state = Rans64State(encoded_x);
rans64_dec_advance(&mut dec_state, &mut r, start, freq, scale_bits).unwrap();
assert_eq!(
dec_state.0, x,
"decoding should invert encoding: D(s, C(s, x)) = x"
);
}
#[test]
fn test_rans64_flush_init_roundtrip() {
let test_state = 0xDEADBEEF_CAFEBABEu64;
let state_in = Rans64State(test_state);
let mut buf = [0u8; 16];
let mut writer = BackwardWord32Writer::new(&mut buf);
rans64_enc_flush(&state_in, &mut writer).unwrap();
let encoded = writer.encoded();
assert_eq!(encoded.len(), 8, "flush should write exactly 8 bytes");
let lo_expected = (test_state & 0xffffffff) as u32;
let hi_expected = (test_state >> 32) as u32;
let lo_actual = u32::from_le_bytes([encoded[0], encoded[1], encoded[2], encoded[3]]);
let hi_actual = u32::from_le_bytes([encoded[4], encoded[5], encoded[6], encoded[7]]);
assert_eq!(lo_actual, lo_expected, "low word should match");
assert_eq!(hi_actual, hi_expected, "high word should match");
let mut reader = Word32Reader::new(encoded);
let state_out = rans64_dec_init(&mut reader).unwrap();
assert_eq!(state_out.0, test_state, "flush+init round-trip");
}
#[test]
fn test_rans64_mul_hi() {
let a = 0xABCDEF0123456789u64;
let b = 0x9876543210FEDCBAu64;
let expected = ((a as u128) * (b as u128) >> 64) as u64;
assert_eq!(rans64_mul_hi(a, b), expected);
assert_eq!(rans64_mul_hi(1, 1), 0);
assert_eq!(rans64_mul_hi(1u64 << 63, 2), 1);
assert_eq!(rans64_mul_hi(!0u64, !0u64), !0u64 - 1);
}
#[test]
fn test_rans64_decoder_symbol_init() {
let dsym = Rans64DecSymbol::new(100, 50).unwrap();
assert_eq!(dsym.start, 100);
assert_eq!(dsym.freq, 50);
}
#[test]
fn test_rans64_freq_one_special() {
let scale_bits = 14;
let freq = 1u32;
let start = 100;
let esym = Rans64EncSymbol::new(start, freq, scale_bits).unwrap();
fn expected(x: u64, start: u64, scale_bits: u32) -> u64 {
x * (1u64 << scale_bits) + start
}
let test_states = [RANS64_L, RANS64_L + 10, RANS64_L * 3, (1u64 << 60)];
for &test_state in &test_states {
if test_state >= esym.x_max {
continue;
}
let mut state = Rans64State(test_state);
let mut tmp = [0u8; 16];
let mut w = BackwardWord32Writer::new(&mut tmp);
rans64_enc_put_symbol(&mut state, &mut w, &esym).unwrap();
assert_eq!(
state.0,
expected(test_state, start as u64, scale_bits),
"64-bit freq=1 mismatch for state={}",
test_state
);
}
}
#[test]
fn test_rans64_word32_writer_exhaustion() {
let mut buf = [0u8; 4]; let mut writer = BackwardWord32Writer::new(&mut buf);
assert!(writer.write_word32(0x12345678).is_ok());
assert!(writer.write_word32(0x9ABCDEF0).is_err());
}
#[test]
fn test_rans64_word32_reader_exhaustion() {
let buf = [0x01, 0x02, 0x03]; let mut reader = Word32Reader::new(&buf);
assert!(reader.read_word32().is_none());
}
#[test]
fn test_rans64_renorm_roundtrip() {
let scale_bits = 14;
let total = 1u32 << scale_bits;
let freq = 7u32;
let start = 100;
let esym = Rans64EncSymbol::new(start, freq, scale_bits).unwrap();
let dsym = Rans64DecSymbol::new(start, freq).unwrap();
let mut out = [0u8; 4096];
let mut writer = BackwardWord32Writer::new(&mut out);
let mut state = Rans64State::new();
let n = 100;
for _i in 0..n {
rans64_enc_put_symbol(&mut state, &mut writer, &esym).unwrap();
}
rans64_enc_flush(&state, &mut writer).unwrap();
let encoded = writer.encoded();
let mut reader = Word32Reader::new(encoded);
let mut dec_state = rans64_dec_init(&mut reader).unwrap();
let cum2sym: alloc::vec::Vec<u8> = (0..total as usize)
.map(|i| {
if (i as u32) < start {
255 } else if (i as u32) < start + freq {
42
} else {
255 }
})
.collect();
let mut output = alloc::vec![0u8; n];
for i in 0..n {
let cf = rans64_dec_get(&dec_state, scale_bits);
let s = cum2sym[cf as usize];
output[i] = s;
rans64_dec_advance_symbol(&mut dec_state, &mut reader, &dsym, scale_bits).unwrap();
}
assert_eq!(output.len(), n);
for &val in &output {
assert_eq!(val, 42, "all decoded symbols should be 42");
}
}
#[test]
fn test_rans64_renorm_only() {
let mut buf = [0u8; 12];
let w0 = 0x00000001u32;
let w1 = 0x00000002u32;
buf[0..4].copy_from_slice(&w0.to_le_bytes());
buf[4..8].copy_from_slice(&w1.to_le_bytes());
let mut reader = Word32Reader::new(&buf);
let mut state = Rans64State(RANS64_L - 1);
rans64_dec_renorm(&mut state, &mut reader).unwrap();
assert!(
state.0 >= RANS64_L,
"after renorm, state {} should be >= RANS64_L",
state.0
);
assert_eq!(state.0, ((RANS64_L - 1) << 32) | 1);
assert_eq!(reader.words_consumed(), 1);
}
#[test]
fn test_rans64_large_scale_reciprocal() {
use super::*;
for scale_bits in 17u32..=31u32 {
let total = 1u64 << scale_bits;
let test_cases = [
(0u32, 100u32), (0u32, 50000u32), (100u32, 1u32), (0u32, (1u32 << scale_bits.min(20))), (total as u32 / 3, total as u32 / 3), ];
for &(start, freq) in &test_cases {
if freq == 0 {
continue;
}
if (start as u64) + (freq as u64) > total {
continue;
}
let sym = Rans64EncSymbol::new(start, freq, scale_bits).unwrap();
let expected_cmpl = ((1u64 << scale_bits) - freq as u64) as u32;
assert_eq!(
sym.cmpl_freq, expected_cmpl,
"cmpl_freq mismatch for scale_bits={}, start={}, freq={}: expected {}, got {}",
scale_bits, start, freq, expected_cmpl, sym.cmpl_freq
);
let expected_x_max = ((RANS64_L >> scale_bits) << 32) * (freq as u64);
assert_eq!(
sym.x_max, expected_x_max,
"x_max mismatch for scale_bits={}, start={}, freq={}",
scale_bits, start, freq
);
if freq >= 2 {
assert!(sym.rcp_freq > 0, "rcp_freq must be > 0 for freq={}", freq);
let mut expected_shift = 0u32;
while freq > (1u32 << expected_shift) {
expected_shift += 1;
}
assert_eq!(
sym.rcp_shift,
expected_shift - 1,
"rcp_shift mismatch for freq={}",
freq
);
}
let expected_bias = if freq < 2 {
(start as u64) + (1u64 << scale_bits) - 1
} else {
start as u64
};
assert_eq!(
sym.bias, expected_bias,
"bias mismatch for scale_bits={}, start={}, freq={}",
scale_bits, start, freq
);
}
}
}
#[test]
fn test_rans64_reciprocal_equals_division_large() {
use super::*;
let scale_bits = 30;
let freqs = [1u32, 2, 100, 10000, 500000000, 1000000000, (1u32 << 30) - 1];
let total = 1u64 << scale_bits;
for &freq in &freqs {
let start = 0u32;
if (start as u64) + (freq as u64) > total {
continue;
}
let sym = Rans64EncSymbol::new(start, freq, scale_bits).unwrap();
let states = [RANS64_L, RANS64_L + 1, RANS64_L * 2, (1u64 << 62) - 1];
for &state_val in &states {
if state_val >= sym.x_max {
continue; }
let div_state = ((state_val / freq as u64) << scale_bits)
+ (state_val % freq as u64)
+ start as u64;
let q = rans64_mul_hi(state_val, sym.rcp_freq) >> sym.rcp_shift;
let fast_state = state_val + sym.bias + q * (sym.cmpl_freq as u64);
assert_eq!(
fast_state, div_state,
"reciprocal mismatch for scale_bits={}, freq={}, state={}: div={}, fast={}",
scale_bits, freq, state_val, div_state, fast_state
);
}
}
}
}