#![allow(missing_docs)]
use alloc::vec::Vec;
use msrtc_rans_core::sink::VecSink;
use msrtc_rans_core::source::SliceSource;
use msrtc_rans_core::source::Source;
use msrtc_rans_core::variant::{Rans64, RansByte, RansParams};
use msrtc_rans_core::{
Freq, Rans64DecSymbol, Rans64EncSymbol, Rans64Encoder, RansByteDecSymbol, RansByteEncSymbol,
RansByteEncoder, RawRansError,
};
const FREQ_BITS: u32 = (core::mem::size_of::<Freq>() * 8) as u32;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EntropyError {
InvalidPmf,
InvalidParams,
InvalidState,
InvalidStream,
RawRansError(RawRansError),
}
impl core::fmt::Display for EntropyError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
EntropyError::InvalidPmf => write!(f, "invalid PMF data"),
EntropyError::InvalidParams => write!(f, "invalid parameter value"),
EntropyError::InvalidState => write!(f, "invalid state (not initialized)"),
EntropyError::InvalidStream => write!(f, "invalid stream"),
EntropyError::RawRansError(e) => write!(f, "raw rANS error: {}", e),
}
}
}
#[cfg(feature = "std")]
impl std::error::Error for EntropyError {}
#[derive(Debug, Clone)]
struct DistributionDesc {
value_offset: i32,
bypass_sentinel: i32,
symbol_offset: usize,
}
fn initialize_distribution_desc(
distribution_descs: &mut Vec<DistributionDesc>,
pmf_lengths: &[i32],
pmf_offsets: &[i32],
pmf_table_size: usize,
) -> Result<(), EntropyError> {
let distribution_count = pmf_lengths.len();
if pmf_offsets.len() != distribution_count {
return Err(EntropyError::InvalidPmf);
}
distribution_descs.reserve(distribution_count);
let mut symbol_cursor: usize = 0;
for i in 0..distribution_count {
let length = pmf_lengths[i];
if length <= 1 || pmf_table_size - symbol_cursor < length as usize {
return Err(EntropyError::InvalidPmf);
}
distribution_descs.push(DistributionDesc {
value_offset: pmf_offsets[i],
bypass_sentinel: length - 1,
symbol_offset: symbol_cursor,
});
symbol_cursor += length as usize;
}
if symbol_cursor != pmf_table_size {
return Err(EntropyError::InvalidPmf);
}
Ok(())
}
#[inline]
fn check_bits(prob_bits: u32, max_scale_bits: u32) -> Result<(), EntropyError> {
if prob_bits < 2 || prob_bits > max_scale_bits {
return Err(EntropyError::InvalidParams);
}
Ok(())
}
#[inline]
fn bytes_to_u32_units(data: &[u8]) -> Vec<u32> {
data.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
pub(crate) trait RawEncoder {
type Unit: Copy + Default;
type Symbol;
fn put_raw(&mut self, start: Freq, freq: Freq, scale_bits: Freq);
fn put_symbol(&mut self, symbol: &Self::Symbol);
fn flush(&mut self);
fn into_units(self) -> Vec<Self::Unit>;
}
impl RawEncoder for RansByteEncoder<VecSink<u8>> {
type Unit = u8;
type Symbol = RansByteEncSymbol;
fn put_raw(&mut self, start: Freq, freq: Freq, scale_bits: Freq) {
self.put_raw(start, freq, scale_bits);
}
fn put_symbol(&mut self, symbol: &Self::Symbol) {
self.put(symbol);
}
fn flush(&mut self) {
self.flush();
}
fn into_units(self) -> Vec<u8> {
self.into_sink().encoded().to_vec()
}
}
impl RawEncoder for Rans64Encoder<VecSink<u32>> {
type Unit = u32;
type Symbol = Rans64EncSymbol;
fn put_raw(&mut self, start: Freq, freq: Freq, scale_bits: Freq) {
self.put_raw(start, freq, scale_bits);
}
fn put_symbol(&mut self, symbol: &Self::Symbol) {
self.put(symbol);
}
fn flush(&mut self) {
self.flush();
}
fn into_units(self) -> Vec<u32> {
self.into_sink().encoded().to_vec()
}
}
struct EncoderState<S: EncSymbol> {
symbol_bits: Freq,
distribution_descs: Vec<DistributionDesc>,
symbols: Vec<S>,
bypass_bits: Freq,
bypass_max_value: Freq,
}
pub(crate) trait EncSymbol: Sized {
fn try_new(start: Freq, freq: Freq, scale_bits: Freq) -> Result<Self, RawRansError>;
}
impl EncSymbol for RansByteEncSymbol {
fn try_new(start: Freq, freq: Freq, scale_bits: Freq) -> Result<Self, RawRansError> {
Self::try_new(start, freq, scale_bits)
}
}
impl EncSymbol for Rans64EncSymbol {
fn try_new(start: Freq, freq: Freq, scale_bits: Freq) -> Result<Self, RawRansError> {
Self::try_new(start, freq, scale_bits)
}
}
impl<S: EncSymbol> EncoderState<S> {
fn uninitialized() -> Self {
Self {
symbol_bits: 0,
distribution_descs: Vec::new(),
symbols: Vec::new(),
bypass_bits: 0,
bypass_max_value: 0,
}
}
fn initialize(
&mut self,
pmf_lengths: &[i32],
pmf_offsets: &[i32],
pmf_table: &[i32],
symbol_bits: i32,
bypass_bits: i32,
max_scale_bits: u32,
) -> Result<(), EntropyError> {
let sb = symbol_bits as Freq;
let bb = bypass_bits as Freq;
check_bits(sb, max_scale_bits)?;
check_bits(bb, max_scale_bits)?;
let is_byte_variant = max_scale_bits < 32;
let max_safe_bits = if is_byte_variant { 30u32 } else { 31u32 };
if sb > max_safe_bits || bb > max_safe_bits {
return Err(EntropyError::InvalidParams);
}
let mut distribution_descs = Vec::new();
initialize_distribution_desc(
&mut distribution_descs,
pmf_lengths,
pmf_offsets,
pmf_table.len(),
)?;
let max_freq = 1u64 << symbol_bits;
let mut symbols: Vec<S> = Vec::with_capacity(pmf_table.len());
let mut pmf_cursor: usize = 0;
for desc in &distribution_descs {
let mut start: u64 = 0;
for _i in 0..=desc.bypass_sentinel {
let freq = pmf_table[pmf_cursor] as u64;
pmf_cursor += 1;
if !(freq > 0 && freq <= max_freq - start) {
return Err(EntropyError::InvalidPmf);
}
let sym = S::try_new(start as Freq, freq as Freq, sb).map_err(|e| match e {
RawRansError::InvalidScaleBits { .. } => EntropyError::InvalidParams,
RawRansError::InvalidParameters => EntropyError::InvalidPmf,
})?;
symbols.push(sym);
start += freq;
}
}
self.distribution_descs = distribution_descs;
self.symbols = symbols;
self.symbol_bits = sb;
self.bypass_bits = bb;
self.bypass_max_value = ((1u64 << bb) - 1) as Freq;
Ok(())
}
fn encode_to_vec<E: RawEncoder<Symbol = S>>(
&self,
indices: &[i32],
values: &[i32],
make_encoder: impl FnOnce() -> E,
) -> Result<Vec<E::Unit>, EntropyError> {
if self.symbol_bits == 0 {
return Err(EntropyError::InvalidState);
}
if indices.len() != values.len() {
return Err(EntropyError::InvalidParams);
}
let mut encoder = make_encoder();
let data_size = indices.len();
let mut idx = data_size as isize - 1;
while idx >= 0 {
let index = indices[idx as usize];
let value = values[idx as usize];
if index < 0 {
idx -= 1;
continue;
}
let dist_len = self.distribution_descs.len();
let ui = if (index as usize) < dist_len {
index as usize
} else {
dist_len - 1
};
let desc = &self.distribution_descs[ui];
let adjusted = value
.checked_add(desc.value_offset)
.ok_or(EntropyError::InvalidParams)?;
let symbol_index: i32;
if adjusted < 0 || adjusted >= desc.bypass_sentinel {
let bypass_value: Freq = if adjusted < 0 {
let neg = adjusted.checked_neg().ok_or(EntropyError::InvalidParams)?;
2u64.wrapping_mul(neg as u64).wrapping_sub(1) as Freq
} else {
2u64.wrapping_mul((adjusted - desc.bypass_sentinel) as u64) as Freq
};
self.encode_bypass_value(&mut encoder, bypass_value);
symbol_index = desc.bypass_sentinel;
} else {
symbol_index = adjusted;
}
let sym_idx = desc.symbol_offset + symbol_index as usize;
encoder.put_symbol(&self.symbols[sym_idx]);
idx -= 1;
}
encoder.flush();
Ok(encoder.into_units())
}
#[inline]
fn encode_bypass_value<E: RawEncoder>(&self, encoder: &mut E, bypass_value: Freq) {
let max_parts = (FREQ_BITS as usize / self.bypass_bits as usize).max(2);
let mut bypass_buffer = Vec::with_capacity(max_parts);
let mut bv = bypass_value;
while bv != 0 {
bypass_buffer.push(bv & self.bypass_max_value);
bv >>= self.bypass_bits;
}
let mut bypass_count = bypass_buffer.len() as Freq;
for &digit in bypass_buffer.iter().rev() {
encoder.put_raw(digit, 1, self.bypass_bits);
}
let mut bypass_prefix_count: Freq = 0;
while bypass_count >= self.bypass_max_value {
bypass_count -= self.bypass_max_value;
bypass_prefix_count += 1;
}
encoder.put_raw(bypass_count, 1, self.bypass_bits);
for _ in 0..bypass_prefix_count {
encoder.put_raw(self.bypass_max_value, 1, self.bypass_bits);
}
}
}
struct DecoderState {
symbol_bits: Freq,
distribution_descs: Vec<DistributionDesc>,
cdf_table: Vec<Freq>,
bypass_bits: Freq,
bypass_max_value: Freq,
}
impl DecoderState {
fn uninitialized() -> Self {
Self {
symbol_bits: 0,
distribution_descs: Vec::new(),
cdf_table: Vec::new(),
bypass_bits: 0,
bypass_max_value: 0,
}
}
fn initialize(
&mut self,
pmf_lengths: &[i32],
pmf_offsets: &[i32],
pmf_table: &[i32],
symbol_bits: i32,
bypass_bits: i32,
max_scale_bits: u32,
) -> Result<(), EntropyError> {
let sb = symbol_bits as Freq;
let bb = bypass_bits as Freq;
check_bits(sb, max_scale_bits)?;
check_bits(bb, max_scale_bits)?;
let is_byte_variant = max_scale_bits < 32;
let max_safe_bits = if is_byte_variant { 30u32 } else { 31u32 };
if sb > max_safe_bits || bb > max_safe_bits {
return Err(EntropyError::InvalidParams);
}
let mut distribution_descs = Vec::new();
initialize_distribution_desc(
&mut distribution_descs,
pmf_lengths,
pmf_offsets,
pmf_table.len(),
)?;
let num_dist = distribution_descs.len();
let mut cdf_table = vec![0u32; pmf_table.len() + num_dist];
let max_freq = 1u64 << symbol_bits;
let mut cursor: usize = 0;
for dist_idx in 0..num_dist {
distribution_descs[dist_idx].symbol_offset = cursor + dist_idx;
let mut start: u64 = 0;
for _i in 0..=distribution_descs[dist_idx].bypass_sentinel {
let freq = pmf_table[cursor] as u64;
if !(freq > 0 && freq <= max_freq - start) {
return Err(EntropyError::InvalidPmf);
}
cdf_table[cursor + dist_idx] = start as Freq;
start += freq;
cursor += 1;
}
cdf_table[cursor + dist_idx] = start as Freq; }
self.distribution_descs = distribution_descs;
self.cdf_table = cdf_table;
self.symbol_bits = sb;
self.bypass_bits = bb;
self.bypass_max_value = ((1u64 << bb) - 1) as Freq;
Ok(())
}
fn decode_from_slice(
&self,
values: &mut [i32],
indices: &[i32],
data: &[u8],
is_byte_variant: bool,
) -> Result<(), EntropyError> {
if self.symbol_bits == 0 {
return Err(EntropyError::InvalidState);
}
if values.len() != indices.len() {
return Err(EntropyError::InvalidParams);
}
if is_byte_variant {
let units = data.to_vec();
let source = SliceSource::new(&units);
let mut decoder = msrtc_rans_core::RansByteDecoder::new(source);
if !decoder.init() {
return Err(EntropyError::InvalidStream);
}
self.decode_inner_byte(&mut decoder, values, indices)?;
if !decoder.source().is_exhausted() || !decoder.check_eof() {
return Err(EntropyError::InvalidStream);
}
} else {
if data.len() % 4 != 0 {
return Err(EntropyError::InvalidStream);
}
let units = bytes_to_u32_units(data);
let source = SliceSource::new(&units);
let mut decoder = msrtc_rans_core::Rans64Decoder::new(source);
if !decoder.init() {
return Err(EntropyError::InvalidStream);
}
self.decode_inner_64(&mut decoder, values, indices)?;
if !decoder.source().is_exhausted() || !decoder.check_eof() {
return Err(EntropyError::InvalidStream);
}
}
Ok(())
}
#[inline]
fn decode_bypass_count_byte(
&self,
decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
) -> Result<Freq, EntropyError> {
let mut total: Freq = 0;
loop {
let value = decoder.get(self.bypass_bits);
if !decoder.advance(value, 1, self.bypass_bits) {
return Err(EntropyError::InvalidStream);
}
total += value;
if value != self.bypass_max_value {
break;
}
if total > FREQ_BITS {
return Err(EntropyError::InvalidStream);
}
}
Ok(total)
}
#[inline]
fn decode_bypass_count_64(
&self,
decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
) -> Result<Freq, EntropyError> {
let mut total: Freq = 0;
loop {
let value = decoder.get(self.bypass_bits);
if !decoder.advance(value, 1, self.bypass_bits) {
return Err(EntropyError::InvalidStream);
}
total += value;
if value != self.bypass_max_value {
break;
}
if total > FREQ_BITS {
return Err(EntropyError::InvalidStream);
}
}
Ok(total)
}
#[inline]
fn decode_bypass_value_payload_byte(
&self,
decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
bypass_count: Freq,
) -> Result<Freq, EntropyError> {
let mut encoded_value: u64 = 0;
let total_bits = bypass_count as u64 * self.bypass_bits as u64;
let mut shift: u64 = 0;
while shift < total_bits {
let v = decoder.get(self.bypass_bits);
if !decoder.advance(v, 1, self.bypass_bits) {
return Err(EntropyError::InvalidStream);
}
encoded_value |= (v as u64) << shift;
shift += self.bypass_bits as u64;
}
Ok(encoded_value as Freq)
}
#[inline]
fn decode_bypass_value_payload_64(
&self,
decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
bypass_count: Freq,
) -> Result<Freq, EntropyError> {
let mut encoded_value: u64 = 0;
let total_bits = bypass_count as u64 * self.bypass_bits as u64;
let mut shift: u64 = 0;
while shift < total_bits {
let v = decoder.get(self.bypass_bits);
if !decoder.advance(v, 1, self.bypass_bits) {
return Err(EntropyError::InvalidStream);
}
encoded_value |= (v as u64) << shift;
shift += self.bypass_bits as u64;
}
Ok(encoded_value as Freq)
}
pub(crate) fn decode_inner_byte(
&self,
decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
values: &mut [i32],
indices: &[i32],
) -> Result<(), EntropyError> {
if self.symbol_bits == 0 {
return Err(EntropyError::InvalidState);
}
if values.len() != indices.len() {
return Err(EntropyError::InvalidParams);
}
for (i, &index) in indices.iter().enumerate() {
if index < 0 {
values[i] = 0;
continue;
}
let dist_len = self.distribution_descs.len();
let ui = if (index as usize) < dist_len {
index as usize
} else {
dist_len - 1
};
let desc = &self.distribution_descs[ui];
let cum_freq = decoder.get(self.symbol_bits);
debug_assert!(cum_freq < (1u32 << self.symbol_bits));
let base_offset = desc.symbol_offset;
let lo = base_offset + 1;
let hi = base_offset + desc.bypass_sentinel as usize + 1;
let upper_idx = {
let mut low = lo;
let mut high = hi;
while low < high {
let mid = low + (high - low) / 2;
if cum_freq < self.cdf_table[mid] {
high = mid;
} else {
low = mid + 1;
}
}
low
};
let start_idx = upper_idx - 1;
let s0 = self.cdf_table[start_idx];
let s1 = self.cdf_table[start_idx + 1];
let freq = s1 - s0;
if !decoder.advance_symbol(&RansByteDecSymbol::new(s0, freq), self.symbol_bits) {
return Err(EntropyError::InvalidStream);
}
let mut symbol = (start_idx - base_offset) as i32;
if symbol == desc.bypass_sentinel {
let bypass_count = self.decode_bypass_count_byte(decoder)?;
let bypass_value = self.decode_bypass_value_payload_byte(decoder, bypass_count)?;
let half = (bypass_value >> 1) as i64;
if bypass_value & 1 != 0 {
symbol = (-half)
.checked_sub(1)
.ok_or(EntropyError::InvalidStream)?
.try_into()
.map_err(|_| EntropyError::InvalidStream)?;
} else {
symbol = half
.checked_add(desc.bypass_sentinel as i64)
.ok_or(EntropyError::InvalidStream)?
.try_into()
.map_err(|_| EntropyError::InvalidStream)?;
}
}
values[i] = (symbol as i64)
.checked_sub(desc.value_offset as i64)
.ok_or(EntropyError::InvalidStream)?
.try_into()
.map_err(|_| EntropyError::InvalidStream)?;
}
Ok(())
}
pub(crate) fn decode_inner_64(
&self,
decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
values: &mut [i32],
indices: &[i32],
) -> Result<(), EntropyError> {
if self.symbol_bits == 0 {
return Err(EntropyError::InvalidState);
}
if values.len() != indices.len() {
return Err(EntropyError::InvalidParams);
}
for (i, &index) in indices.iter().enumerate() {
if index < 0 {
values[i] = 0;
continue;
}
let dist_len = self.distribution_descs.len();
let ui = if (index as usize) < dist_len {
index as usize
} else {
dist_len - 1
};
let desc = &self.distribution_descs[ui];
let cum_freq = decoder.get(self.symbol_bits);
debug_assert!(cum_freq < (1u32 << self.symbol_bits));
let base_offset = desc.symbol_offset;
let lo = base_offset + 1;
let hi = base_offset + desc.bypass_sentinel as usize + 1;
let upper_idx = {
let mut low = lo;
let mut high = hi;
while low < high {
let mid = low + (high - low) / 2;
if cum_freq < self.cdf_table[mid] {
high = mid;
} else {
low = mid + 1;
}
}
low
};
let start_idx = upper_idx - 1;
let s0 = self.cdf_table[start_idx];
let s1 = self.cdf_table[start_idx + 1];
let freq = s1 - s0;
if !decoder.advance_symbol(&Rans64DecSymbol::new(s0, freq), self.symbol_bits) {
return Err(EntropyError::InvalidStream);
}
let mut symbol = (start_idx - base_offset) as i32;
if symbol == desc.bypass_sentinel {
let bypass_count = self.decode_bypass_count_64(decoder)?;
let bypass_value = self.decode_bypass_value_payload_64(decoder, bypass_count)?;
let half = (bypass_value >> 1) as i64;
if bypass_value & 1 != 0 {
symbol = (-half)
.checked_sub(1)
.ok_or(EntropyError::InvalidStream)?
.try_into()
.map_err(|_| EntropyError::InvalidStream)?;
} else {
symbol = half
.checked_add(desc.bypass_sentinel as i64)
.ok_or(EntropyError::InvalidStream)?
.try_into()
.map_err(|_| EntropyError::InvalidStream)?;
}
}
values[i] = (symbol as i64)
.checked_sub(desc.value_offset as i64)
.ok_or(EntropyError::InvalidStream)?
.try_into()
.map_err(|_| EntropyError::InvalidStream)?;
}
Ok(())
}
}
pub trait EncoderVariantForS: RansParams {
type EncSymbol: EncSymbol;
type RawEnc: RawEncoder<Symbol = Self::EncSymbol>;
const MAX_SCALE_BITS: u32;
fn units_to_bytes(units: Vec<<Self::RawEnc as RawEncoder>::Unit>) -> Vec<u8>;
fn make_encoder() -> Self::RawEnc;
}
impl EncoderVariantForS for RansByte {
type EncSymbol = RansByteEncSymbol;
type RawEnc = RansByteEncoder<VecSink<u8>>;
const MAX_SCALE_BITS: u32 = 30;
fn units_to_bytes(units: Vec<u8>) -> Vec<u8> {
units
}
fn make_encoder() -> Self::RawEnc {
RansByteEncoder::new(VecSink::new(4096))
}
}
impl EncoderVariantForS for Rans64 {
type EncSymbol = Rans64EncSymbol;
type RawEnc = Rans64Encoder<VecSink<u32>>;
const MAX_SCALE_BITS: u32 = 32;
fn units_to_bytes(units: Vec<u32>) -> Vec<u8> {
let mut bytes = Vec::with_capacity(units.len() * 4);
for &u in &units {
bytes.extend_from_slice(&u.to_le_bytes());
}
bytes
}
fn make_encoder() -> Self::RawEnc {
Rans64Encoder::new(VecSink::new(4096))
}
}
pub struct EntropyEncoder<S: EncoderVariantForS> {
state: EncoderState<<S as EncoderVariantForS>::EncSymbol>,
}
impl<S: EncoderVariantForS> EntropyEncoder<S> {
pub fn new() -> Self {
Self {
state: EncoderState::uninitialized(),
}
}
pub fn initialize(
&mut self,
pmf_lengths: &[i32],
pmf_offsets: &[i32],
pmf_table: &[i32],
symbol_bits: u32,
bypass_bits: u32,
) -> Result<(), EntropyError> {
self.state.initialize(
pmf_lengths,
pmf_offsets,
pmf_table,
symbol_bits as i32,
bypass_bits as i32,
<S as EncoderVariantForS>::MAX_SCALE_BITS,
)
}
pub fn encode(
&self,
indices: &[i32],
values: &[i32],
buffer: &mut Vec<u8>,
) -> Result<(), EntropyError> {
let units = self.state.encode_to_vec(indices, values, S::make_encoder)?;
let bytes = S::units_to_bytes(units);
buffer.extend_from_slice(&bytes);
Ok(())
}
}
impl<S: EncoderVariantForS> Default for EntropyEncoder<S> {
fn default() -> Self {
Self::new()
}
}
fn _assert_encoder_bounds() {
fn _is_encoder<S: EncoderVariantForS>() {}
_is_encoder::<RansByte>();
_is_encoder::<Rans64>();
}
pub struct EntropyDecoder<S: RansParams> {
state: DecoderState,
_phantom: core::marker::PhantomData<S>,
}
impl<S: RansParams> EntropyDecoder<S> {
pub fn new() -> Self {
Self {
state: DecoderState::uninitialized(),
_phantom: core::marker::PhantomData,
}
}
pub fn initialize(
&mut self,
pmf_lengths: &[i32],
pmf_offsets: &[i32],
pmf_table: &[i32],
symbol_bits: u32,
bypass_bits: u32,
) -> Result<(), EntropyError> {
let max_scale_bits = match S::NAME {
"RansByte" => 30u32,
"Rans64" => 32u32,
_ => return Err(EntropyError::InvalidParams),
};
self.state.initialize(
pmf_lengths,
pmf_offsets,
pmf_table,
symbol_bits as i32,
bypass_bits as i32,
max_scale_bits,
)
}
pub fn decode(
&self,
values: &mut [i32],
indices: &[i32],
data: &[u8],
) -> Result<(), EntropyError> {
let is_byte = match S::NAME {
"RansByte" => true,
"Rans64" => false,
_ => return Err(EntropyError::InvalidParams),
};
self.state.decode_from_slice(values, indices, data, is_byte)
}
pub fn decode_partial(
&self,
values: &mut [i32],
indices: &[i32],
data: &[u8],
) -> Result<usize, EntropyError> {
if self.state.symbol_bits == 0 {
return Err(EntropyError::InvalidState);
}
if values.len() != indices.len() {
return Err(EntropyError::InvalidParams);
}
let consumed = match S::NAME {
"RansByte" => {
let units = data.to_vec();
let source = SliceSource::new(&units);
let mut decoder = msrtc_rans_core::RansByteDecoder::new(source);
if !decoder.init() {
return Err(EntropyError::InvalidStream);
}
self.state
.decode_inner_byte(&mut decoder, values, indices)?;
if !decoder.check_eof() {
return Err(EntropyError::InvalidStream);
}
decoder.source().position()
}
"Rans64" => {
if data.len() % 4 != 0 {
return Err(EntropyError::InvalidStream);
}
let units = bytes_to_u32_units(data);
let source = SliceSource::new(&units);
let mut decoder = msrtc_rans_core::Rans64Decoder::new(source);
if !decoder.init() {
return Err(EntropyError::InvalidStream);
}
self.state.decode_inner_64(&mut decoder, values, indices)?;
if !decoder.check_eof() {
return Err(EntropyError::InvalidStream);
}
decoder.source().position() * 4
}
_ => return Err(EntropyError::InvalidParams),
};
Ok(consumed)
}
}
impl<S: RansParams> Default for EntropyDecoder<S> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
const PMF_LENGTHS: [i32; 2] = [4, 6];
const PMF_OFFSETS: [i32; 2] = [1, 2];
const PMF_TABLE: [i32; 10] = [1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
const INDICES: [i32; 4] = [0, 1, 0, 1];
const VALUES: [i32; 4] = [-2, 1, 0, 1];
const SYMBOL_BITS: u32 = 16;
const BYPASS_BITS: u32 = 4;
const REF_HEX_BYTE: &str = "0500bd040001a10003000b00";
const REF_HEX_64: &str = "0500a1bd04000000110a002f03000300";
fn hex_decode(hex: &str) -> Vec<u8> {
(0..hex.len())
.step_by(2)
.map(|i| u8::from_str_radix(&hex[i..i + 2], 16).unwrap())
.collect()
}
#[test]
fn test_encoder_byte_initialize() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
assert!(
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS
)
.is_ok()
);
}
#[test]
fn test_encoder_64_initialize() {
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
assert!(
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS
)
.is_ok()
);
}
#[test]
fn test_encoder_rejects_invalid_pmf() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
assert_eq!(
enc.initialize(&[4], &[1, 2], &PMF_TABLE, SYMBOL_BITS, BYPASS_BITS),
Err(EntropyError::InvalidPmf)
);
}
#[test]
fn test_encoder_rejects_invalid_params() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
assert_eq!(
enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 1, BYPASS_BITS),
Err(EntropyError::InvalidParams)
);
}
#[test]
fn test_encoder_byte_rejects_length_leq_one() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
assert_eq!(
enc.initialize(&[1, 6], &[1, 2], &PMF_TABLE, SYMBOL_BITS, BYPASS_BITS),
Err(EntropyError::InvalidPmf)
);
}
#[test]
fn test_encode_byte_matches_reference() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut buffer = Vec::new();
enc.encode(&INDICES, &VALUES, &mut buffer).unwrap();
let expected = hex_decode(REF_HEX_BYTE);
assert_eq!(
buffer, expected,
"RansByte encode output does not match reference hex"
);
}
#[test]
fn test_encode_64_matches_reference() {
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut buffer = Vec::new();
enc.encode(&INDICES, &VALUES, &mut buffer).unwrap();
let expected = hex_decode(REF_HEX_64);
assert_eq!(
buffer, expected,
"Rans64 encode output does not match reference hex"
);
}
#[test]
fn test_encode_in_range_values_no_bypass() {
let in_range_values = [1i32, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut buffer = Vec::new();
let result = enc.encode(&INDICES, &in_range_values, &mut buffer);
assert!(result.is_ok(), "encode should succeed: {:?}", result);
assert!(!buffer.is_empty(), "encoded buffer should not be empty");
}
#[test]
fn test_decode_byte_roundtrip_in_range() {
let values = [1i32, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; values.len()];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(
decoded, values,
"roundtrip decode should match original values"
);
}
#[test]
fn test_decode_64_roundtrip_in_range() {
let values = [1i32, 1, 0, 1];
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; values.len()];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(
decoded, values,
"Rans64 roundtrip decode should match original values"
);
}
#[test]
fn test_decode_byte_roundtrip_bypass() {
let values = [-2i32, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; values.len()];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(
decoded, values,
"bypass roundtrip decode should match original values"
);
}
#[test]
fn test_decode_64_roundtrip_bypass() {
let values = [-2i32, 1, 0, 1];
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; values.len()];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(
decoded, values,
"Rans64 bypass roundtrip decode should match original values"
);
}
#[test]
fn test_encoder_64_symbol_bits_31_accepted() {
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
assert!(
enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 31, BYPASS_BITS)
.is_ok()
);
}
#[test]
fn test_encoder_64_symbol_bits_32_rejected() {
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
assert_eq!(
enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 32, BYPASS_BITS),
Err(EntropyError::InvalidParams)
);
}
#[test]
fn test_encoder_64_bypass_bits_32_rejected() {
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
assert_eq!(
enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 32),
Err(EntropyError::InvalidParams)
);
}
#[test]
fn test_decode_64_rejects_misaligned_1_extra_byte() {
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
.unwrap();
let mut misaligned = encoded.clone();
misaligned.push(0xAB);
let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
let result = dec.decode(&mut decoded, &INDICES, &misaligned);
assert_eq!(result, Err(EntropyError::InvalidStream));
}
#[test]
fn test_decode_64_rejects_misaligned_2_extra_bytes() {
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
.unwrap();
let mut misaligned = encoded.clone();
misaligned.extend_from_slice(&[0xAB, 0xCD]);
let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
let result = dec.decode(&mut decoded, &INDICES, &misaligned);
assert_eq!(result, Err(EntropyError::InvalidStream));
}
#[test]
fn test_decode_64_rejects_misaligned_3_extra_bytes() {
let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
.unwrap();
let mut misaligned = encoded.clone();
misaligned.extend_from_slice(&[0xAB, 0xCD, 0xEF]);
let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
let result = dec.decode(&mut decoded, &INDICES, &misaligned);
assert_eq!(result, Err(EntropyError::InvalidStream));
}
#[test]
fn test_decode_byte_accepts_extra_bytes() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
.unwrap();
let mut extended = encoded.clone();
extended.extend_from_slice(&[0xAB, 0xCD]);
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
let _ = dec.decode(&mut decoded, &INDICES, &extended);
}
#[test]
fn test_encode_bypass_positive_outlier() {
let values = [10i32, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_multi_digit_value() {
let values = [200i32, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_bits_2() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 2)
.unwrap();
let values = [10i32, 1, 0, 1]; let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 2)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_bits_3() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 3)
.unwrap();
let values = [10i32, 1, 0, 1];
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 3)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_bits_8() {
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 8)
.unwrap();
let values = [10i32, 1, 0, 1];
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 8)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_multiple_bypasses() {
let values = [-2i32, 10, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_mixed_in_range_and_bypass() {
let values = [0i32, 5, 1, -3];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_negative_outlier_at_boundary() {
let values = [-10i32, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_large_positive_outlier() {
let values = [10000i32, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
enc.encode(&INDICES, &values, &mut encoded).unwrap();
let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
dec.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut decoded = vec![0i32; 4];
dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
assert_eq!(decoded, values);
}
#[test]
fn test_encode_bypass_extreme_negative_i32_min_plus_one() {
let values = [i32::MIN + 1, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
let result = enc.encode(&INDICES, &values, &mut encoded);
assert!(result.is_ok() || result == Err(EntropyError::InvalidParams));
}
#[test]
fn test_encode_bypass_extreme_positive_i32_max() {
let values = [i32::MAX, 1, 0, 1];
let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
enc.initialize(
&PMF_LENGTHS,
&PMF_OFFSETS,
&PMF_TABLE,
SYMBOL_BITS,
BYPASS_BITS,
)
.unwrap();
let mut encoded = Vec::new();
let result = enc.encode(&INDICES, &values, &mut encoded);
assert_eq!(result, Err(EntropyError::InvalidParams));
}
}