use alloc::vec::Vec;
use crate::entropy::{
EncoderVariantForS, EntropyDecoder, EntropyEncoder, EntropyError, RawEncoder,
};
use crate::source::SliceSource;
use crate::{Rans64Decoder, RansByteDecoder};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RansVariant {
RansByte,
Rans64,
}
impl RansVariant {
pub const fn as_int(&self) -> i32 {
match self {
RansVariant::RansByte => 1,
RansVariant::Rans64 => 0,
}
}
pub const fn from_int(v: i32) -> Option<Self> {
match v {
1 => Some(RansVariant::RansByte),
0 => Some(RansVariant::Rans64),
_ => None,
}
}
}
#[derive(Debug)]
pub struct RansEncoderStream<S: EncoderVariantForS> {
encoder: Option<S::RawEnc>,
_s: core::marker::PhantomData<S>,
}
impl<S: EncoderVariantForS> Default for RansEncoderStream<S> {
fn default() -> Self {
Self::new()
}
}
impl<S: EncoderVariantForS> RansEncoderStream<S> {
pub fn new() -> Self {
Self {
encoder: None,
_s: core::marker::PhantomData,
}
}
pub fn is_initialized(&self) -> bool {
self.encoder.is_some()
}
pub fn push(
&mut self,
encoder: &EntropyEncoder<S>,
indices: &[i32],
values: &[i32],
) -> Result<(), EntropyError> {
let raw = self.encoder_mut();
encoder.encode_batch(indices, values, raw)
}
pub fn flush(&mut self) -> Result<Vec<u8>, EntropyError> {
let mut raw = self.encoder.take().ok_or(EntropyError::InvalidState)?;
raw.flush();
let units = raw.into_units();
Ok(S::units_to_bytes(units))
}
pub fn reset(&mut self) {
self.encoder = None;
}
fn encoder_mut(&mut self) -> &mut S::RawEnc {
if self.encoder.is_none() {
self.encoder = Some(S::make_encoder());
}
self.encoder.as_mut().expect("just set")
}
}
#[derive(Debug)]
pub struct RansDecoderStream<S: EncoderVariantForS> {
data: Vec<u8>,
cursor: Option<(usize, u64)>,
_s: core::marker::PhantomData<S>,
}
impl<S: EncoderVariantForS> Default for RansDecoderStream<S> {
fn default() -> Self {
Self::new()
}
}
impl<S: EncoderVariantForS> RansDecoderStream<S> {
pub fn new() -> Self {
Self {
data: Vec::new(),
cursor: None,
_s: core::marker::PhantomData,
}
}
pub fn open_on(data: &[u8]) -> Self {
Self {
data: data.to_vec(),
cursor: None,
_s: core::marker::PhantomData,
}
}
pub fn is_open(&self) -> bool {
!self.data.is_empty()
}
pub fn check_eof(&self) -> bool {
let Some((pos, state)) = self.cursor else {
return !self.is_open();
};
let unit_len = match S::NAME {
"RansByte" => self.data.len(),
"Rans64" => self.data.len() / 4,
_ => 0,
};
let lower = match S::NAME {
"RansByte" => 1u64 << 23,
"Rans64" => 1u64 << 31,
_ => 0,
};
pos == unit_len && state == lower
}
pub fn open(&mut self, data: &[u8]) {
self.data = data.to_vec();
self.cursor = None;
}
pub fn close(&mut self) {
self.data.clear();
self.cursor = None;
}
pub fn decode_eof(&mut self) -> Result<(), EntropyError> {
if !self.check_eof() {
return Err(EntropyError::InvalidStream);
}
self.close();
Ok(())
}
pub fn decode(
&mut self,
decoder: &EntropyDecoder<S>,
values: &mut [i32],
indices: &[i32],
) -> Result<(), EntropyError> {
match S::NAME {
"RansByte" => {
let units = self.data.clone();
let mut source = SliceSource::new(&units);
let mut raw = match self.cursor {
Some((pos, state)) => {
source.seek(pos);
RansByteDecoder::from_state(source, state as u32)
}
None => {
let mut d = RansByteDecoder::new(source);
if !d.init() {
return Err(EntropyError::InvalidStream);
}
d
}
};
decoder.decode_byte_continue(&mut raw, values, indices)?;
self.cursor = Some((raw.source().position(), raw.state() as u64));
Ok(())
}
"Rans64" => {
if self.data.len() % 4 != 0 {
return Err(EntropyError::InvalidStream);
}
let units: Vec<u32> = self
.data
.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
let mut source = SliceSource::new(&units);
let mut raw = match self.cursor {
Some((pos, state)) => {
source.seek(pos);
Rans64Decoder::from_state(source, state)
}
None => {
let mut d = Rans64Decoder::new(source);
if !d.init() {
return Err(EntropyError::InvalidStream);
}
d
}
};
decoder.decode_64_continue(&mut raw, values, indices)?;
self.cursor = Some((raw.source().position(), raw.state()));
Ok(())
}
_ => Err(EntropyError::InvalidParams),
}
}
pub fn bytes_consumed(&self) -> usize {
match self.cursor {
Some((pos, _)) => match S::NAME {
"RansByte" => pos,
"Rans64" => pos * 4,
_ => 0,
},
None => 0,
}
}
pub fn data(&self) -> &[u8] {
&self.data
}
}
pub fn units_to_le_bytes(units: &[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
}
#[cfg(test)]
mod tests {
use super::*;
use crate::entropy::EntropyDecoder;
use crate::variant::{Rans64, RansByte};
#[test]
fn test_variant_values() {
assert_eq!(RansVariant::RansByte.as_int(), 1);
assert_eq!(RansVariant::Rans64.as_int(), 0);
assert_eq!(RansVariant::from_int(1), Some(RansVariant::RansByte));
assert_eq!(RansVariant::from_int(0), Some(RansVariant::Rans64));
assert_eq!(RansVariant::from_int(2), None);
}
#[test]
fn test_encoder_stream_multipart() {
let pmf_lengths1 = vec![4, 6];
let pmf_offsets1 = vec![1, 2];
let pmf_table1 = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
let values1 = vec![-2, 1, 0, 1];
let indices1 = vec![0, 1, 0, 1];
let pmf_lengths2 = vec![5];
let pmf_offsets2 = vec![1];
let pmf_table2 = vec![1, 3, 3, 1, 1];
let values2 = vec![-2, 1, 2];
let indices2 = vec![0, 0, 0];
let mut encoder1 = EntropyEncoder::<RansByte>::new();
encoder1
.initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
.expect("init1");
let mut encoder2 = EntropyEncoder::<RansByte>::new();
encoder2
.initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
.expect("init2");
let mut stream = RansEncoderStream::<RansByte>::new();
stream.push(&encoder2, &indices2, &values2).expect("push2");
stream.push(&encoder1, &indices1, &values1).expect("push1");
let data = stream.flush().expect("flush");
assert!(!data.is_empty());
let mut decoder1 = EntropyDecoder::<RansByte>::new();
decoder1
.initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
.expect("dec1 init");
let mut decoder2 = EntropyDecoder::<RansByte>::new();
decoder2
.initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
.expect("dec2 init");
let mut dstream = RansDecoderStream::<RansByte>::open_on(&data);
let mut decoded1 = vec![0i32; values1.len()];
dstream
.decode(&decoder1, &mut decoded1, &indices1)
.expect("decode1");
assert_eq!(decoded1, values1);
let mut decoded2 = vec![0i32; values2.len()];
dstream
.decode(&decoder2, &mut decoded2, &indices2)
.expect("decode2");
assert_eq!(decoded2, values2);
dstream.decode_eof().expect("eof");
}
#[test]
fn test_encoder_stream_multipart_64() {
let pmf_lengths1 = vec![4, 6];
let pmf_offsets1 = vec![1, 2];
let pmf_table1 = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
let values1 = vec![-2, 1, 0, 1];
let indices1 = vec![0, 1, 0, 1];
let pmf_lengths2 = vec![5];
let pmf_offsets2 = vec![1];
let pmf_table2 = vec![1, 3, 3, 1, 1];
let values2 = vec![-2, 1, 2];
let indices2 = vec![0, 0, 0];
let mut encoder1 = EntropyEncoder::<Rans64>::new();
encoder1
.initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
.expect("init1");
let mut encoder2 = EntropyEncoder::<Rans64>::new();
encoder2
.initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
.expect("init2");
let mut stream = RansEncoderStream::<Rans64>::new();
stream.push(&encoder2, &indices2, &values2).expect("push2");
stream.push(&encoder1, &indices1, &values1).expect("push1");
let data = stream.flush().expect("flush");
assert!(!data.is_empty());
assert_eq!(data.len() % 4, 0, "Rans64 stream must be 4-byte aligned");
let mut decoder1 = EntropyDecoder::<Rans64>::new();
decoder1
.initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
.expect("dec1 init");
let mut decoder2 = EntropyDecoder::<Rans64>::new();
decoder2
.initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
.expect("dec2 init");
let mut dstream = RansDecoderStream::<Rans64>::open_on(&data);
let mut decoded1 = vec![0i32; values1.len()];
dstream
.decode(&decoder1, &mut decoded1, &indices1)
.expect("decode1");
assert_eq!(decoded1, values1);
let mut decoded2 = vec![0i32; values2.len()];
dstream
.decode(&decoder2, &mut decoded2, &indices2)
.expect("decode2");
assert_eq!(decoded2, values2);
dstream.decode_eof().expect("eof");
}
#[test]
fn test_encoder_stream_reuse_after_flush() {
let pmf_lengths = vec![4, 6];
let pmf_offsets = vec![1, 2];
let pmf_table = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
let values = vec![-2, 1, 0, 1];
let indices = vec![0, 1, 0, 1];
let mut encoder = EntropyEncoder::<RansByte>::new();
encoder
.initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
.expect("init");
let mut stream = RansEncoderStream::<RansByte>::new();
stream.push(&encoder, &indices, &values).expect("push1");
let data1 = stream.flush().expect("flush1");
stream.push(&encoder, &indices, &values).expect("push2");
let data2 = stream.flush().expect("flush2");
let mut decoder = EntropyDecoder::<RansByte>::new();
decoder
.initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
.expect("dec init");
for data in [&data1, &data2] {
let mut decoded = vec![0i32; values.len()];
decoder
.decode(&mut decoded, &indices, data)
.expect("decode");
assert_eq!(decoded, values);
}
}
#[test]
fn test_encoder_stream_reset_aborts() {
let pmf_lengths = vec![4, 6];
let pmf_offsets = vec![1, 2];
let pmf_table = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
let values = vec![-2, 1, 0, 1];
let indices = vec![0, 1, 0, 1];
let mut encoder = EntropyEncoder::<RansByte>::new();
encoder
.initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
.expect("init");
let mut stream = RansEncoderStream::<RansByte>::new();
stream.push(&encoder, &indices, &values).expect("push");
stream.reset();
assert!(!stream.is_initialized(), "reset must clear state");
assert!(stream.flush().is_err());
}
#[test]
fn test_decoder_stream_lifecycle() {
let mut stream = RansDecoderStream::<RansByte>::new();
assert!(!stream.is_open());
stream.open(&[1, 2, 3, 4, 5]);
assert!(stream.is_open());
assert!(!stream.check_eof());
stream.close();
assert!(!stream.is_open());
assert!(stream.check_eof());
}
#[test]
fn test_units_to_le_bytes() {
let units = [0x01020304u32, 0x05060708];
let bytes = <Rans64 as EncoderVariantForS>::units_to_bytes(units.to_vec());
assert_eq!(bytes, vec![4, 3, 2, 1, 8, 7, 6, 5]);
}
}