Skip to main content

msrtc_rans/
stream.rs

1// Licensed under the MIT license.
2// Author: Riaan de Beer - github.com/infinityabundance - rdebeer.infinityabundance@gmail.com
3
4//! # rANS stream types
5//!
6//! Standalone `RansEncoderStream` and `RansDecoderStream` matching
7//! Microsoft's `RansEncoderStream` / `RansDecoderStream` from
8//! `EntropyCoder.h`.
9//!
10//! - Encoder stream keeps a **persistent raw encoder state** across
11//!   `push()` calls (matching Microsoft's `RawRansEncoderStream`),
12//!   flushing once at the end.
13//! - Decoder stream owns the encoded data and advances a persistent
14//!   decoder across sequential `decode` calls.
15//!
16//! Both types are generic over the rANS variant (`RansByte` or `Rans64`),
17//! making variant mismatches a compile-time error.
18
19use alloc::vec::Vec;
20
21use crate::entropy::{
22    EncoderVariantForS, EntropyDecoder, EntropyEncoder, EntropyError, RawEncoder,
23};
24use crate::source::SliceSource;
25use crate::{Rans64Decoder, RansByteDecoder};
26
27/// rANS variants for runtime dispatch (Python-facing).
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum RansVariant {
30    /// RansByte: u32 state, u8 unit (Microsoft value 1)
31    RansByte,
32    /// Rans64: u64 state, u32 unit (Microsoft value 0)
33    Rans64,
34}
35
36impl RansVariant {
37    /// Microsoft's integer value: Rans64=0, RansByte=1.
38    pub const fn as_int(&self) -> i32 {
39        match self {
40            RansVariant::RansByte => 1,
41            RansVariant::Rans64 => 0,
42        }
43    }
44
45    /// From Microsoft's integer value.
46    pub const fn from_int(v: i32) -> Option<Self> {
47        match v {
48            1 => Some(RansVariant::RansByte),
49            0 => Some(RansVariant::Rans64),
50            _ => None,
51        }
52    }
53}
54
55/// rANS encoder stream for a specific variant `S`.
56///
57/// Keeps a persistent raw encoder state across `push()` calls. `flush()`
58/// writes the final state and yields the encoded message; the stream can
59/// then be reused. `reset()` aborts the current session.
60#[derive(Debug)]
61pub struct RansEncoderStream<S: EncoderVariantForS> {
62    encoder: Option<S::RawEnc>,
63    _s: core::marker::PhantomData<S>,
64}
65
66impl<S: EncoderVariantForS> Default for RansEncoderStream<S> {
67    fn default() -> Self {
68        Self::new()
69    }
70}
71
72impl<S: EncoderVariantForS> RansEncoderStream<S> {
73    /// Create a new encoder stream.
74    pub fn new() -> Self {
75        Self {
76            encoder: None,
77            _s: core::marker::PhantomData,
78        }
79    }
80
81    /// Whether the stream has an active encoder session.
82    pub fn is_initialized(&self) -> bool {
83        self.encoder.is_some()
84    }
85
86    /// Push an entropy-encoded batch onto the stream.
87    ///
88    /// Encodes `values` using the given `encoder`'s PMF, continuing the
89    /// persistent raw encoder state. Matches `EntropyEncoder::push` in
90    /// Python and `IEntropyEncoderImpl::Encode(stream, indices, values)`.
91    pub fn push(
92        &mut self,
93        encoder: &EntropyEncoder<S>,
94        indices: &[i32],
95        values: &[i32],
96    ) -> Result<(), EntropyError> {
97        let raw = self.encoder_mut();
98        encoder.encode_batch(indices, values, raw)
99    }
100
101    /// Flush the current session and return the encoded message.
102    ///
103    /// After flush, the stream is reset for reuse. Returns an error if no
104    /// data has been pushed (matching the Python `flush` which raises on
105    /// empty output).
106    pub fn flush(&mut self) -> Result<Vec<u8>, EntropyError> {
107        let mut raw = self.encoder.take().ok_or(EntropyError::InvalidState)?;
108        raw.flush();
109        let units = raw.into_units();
110        Ok(S::units_to_bytes(units))
111    }
112
113    /// Abort the current session, discarding all pushed data.
114    pub fn reset(&mut self) {
115        self.encoder = None;
116    }
117
118    fn encoder_mut(&mut self) -> &mut S::RawEnc {
119        if self.encoder.is_none() {
120            self.encoder = Some(S::make_encoder());
121        }
122        self.encoder.as_mut().expect("just set")
123    }
124}
125
126/// rANS decoder stream for a specific variant `S`.
127///
128/// Owns the encoded data and keeps a **persistent decode cursor**
129/// (unit position + rANS state) across sequential `decode` calls.
130/// This matches Microsoft's `RansDecoderStream`, where the raw decoder
131/// is initialized once in `Open()` and reused by every `Decode` call.
132#[derive(Debug)]
133pub struct RansDecoderStream<S: EncoderVariantForS> {
134    data: Vec<u8>,
135    /// (unit position, decoder state). `None` before the first decode.
136    cursor: Option<(usize, u64)>,
137    _s: core::marker::PhantomData<S>,
138}
139
140impl<S: EncoderVariantForS> Default for RansDecoderStream<S> {
141    fn default() -> Self {
142        Self::new()
143    }
144}
145
146impl<S: EncoderVariantForS> RansDecoderStream<S> {
147    /// Create a decoder stream (closed).
148    pub fn new() -> Self {
149        Self {
150            data: Vec::new(),
151            cursor: None,
152            _s: core::marker::PhantomData,
153        }
154    }
155
156    /// Create a decoder stream opened on the given data.
157    pub fn open_on(data: &[u8]) -> Self {
158        Self {
159            data: data.to_vec(),
160            cursor: None,
161            _s: core::marker::PhantomData,
162        }
163    }
164
165    /// Whether the stream is open (has data).
166    pub fn is_open(&self) -> bool {
167        !self.data.is_empty()
168    }
169
170    /// Whether the current state can be at EOF.
171    ///
172    /// EOF requires the source to be exhausted and the decoder state to be
173    /// at `LowerBound` (or the stream to be unopened).
174    pub fn check_eof(&self) -> bool {
175        let Some((pos, state)) = self.cursor else {
176            return !self.is_open();
177        };
178        let unit_len = match S::NAME {
179            "RansByte" => self.data.len(),
180            "Rans64" => self.data.len() / 4,
181            _ => 0,
182        };
183        let lower = match S::NAME {
184            "RansByte" => 1u64 << 23,
185            "Rans64" => 1u64 << 31,
186            _ => 0,
187        };
188        pos == unit_len && state == lower
189    }
190
191    /// Open the stream on new data, resetting the cursor.
192    pub fn open(&mut self, data: &[u8]) {
193        self.data = data.to_vec();
194        self.cursor = None;
195    }
196
197    /// Close the stream, releasing data.
198    pub fn close(&mut self) {
199        self.data.clear();
200        self.cursor = None;
201    }
202
203    /// Check that decoding reached the end of the message and close on success.
204    pub fn decode_eof(&mut self) -> Result<(), EntropyError> {
205        if !self.check_eof() {
206            return Err(EntropyError::InvalidStream);
207        }
208        self.close();
209        Ok(())
210    }
211
212    /// Decode a batch of symbols from the stream, advancing the persistent cursor.
213    ///
214    /// The first call initializes the raw decoder from the stream's initial
215    /// state; subsequent calls continue from the saved cursor.
216    pub fn decode(
217        &mut self,
218        decoder: &EntropyDecoder<S>,
219        values: &mut [i32],
220        indices: &[i32],
221    ) -> Result<(), EntropyError> {
222        match S::NAME {
223            "RansByte" => {
224                let units = self.data.clone();
225                let mut source = SliceSource::new(&units);
226                let mut raw = match self.cursor {
227                    Some((pos, state)) => {
228                        source.seek(pos);
229                        RansByteDecoder::from_state(source, state as u32)
230                    }
231                    None => {
232                        let mut d = RansByteDecoder::new(source);
233                        if !d.init() {
234                            return Err(EntropyError::InvalidStream);
235                        }
236                        d
237                    }
238                };
239                decoder.decode_byte_continue(&mut raw, values, indices)?;
240                self.cursor = Some((raw.source().position(), raw.state() as u64));
241                Ok(())
242            }
243            "Rans64" => {
244                if self.data.len() % 4 != 0 {
245                    return Err(EntropyError::InvalidStream);
246                }
247                let units: Vec<u32> = self
248                    .data
249                    .chunks_exact(4)
250                    .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
251                    .collect();
252                let mut source = SliceSource::new(&units);
253                let mut raw = match self.cursor {
254                    Some((pos, state)) => {
255                        source.seek(pos);
256                        Rans64Decoder::from_state(source, state)
257                    }
258                    None => {
259                        let mut d = Rans64Decoder::new(source);
260                        if !d.init() {
261                            return Err(EntropyError::InvalidStream);
262                        }
263                        d
264                    }
265                };
266                decoder.decode_64_continue(&mut raw, values, indices)?;
267                self.cursor = Some((raw.source().position(), raw.state()));
268                Ok(())
269            }
270            _ => Err(EntropyError::InvalidParams),
271        }
272    }
273
274    /// Number of bytes consumed so far (unit position × unit size).
275    pub fn bytes_consumed(&self) -> usize {
276        match self.cursor {
277            Some((pos, _)) => match S::NAME {
278                "RansByte" => pos,
279                "Rans64" => pos * 4,
280                _ => 0,
281            },
282            None => 0,
283        }
284    }
285
286    /// The full stream data (borrowed).
287    pub fn data(&self) -> &[u8] {
288        &self.data
289    }
290}
291
292/// Convert u32 units to little-endian bytes.
293pub fn units_to_le_bytes(units: &[u32]) -> Vec<u8> {
294    let mut bytes = Vec::with_capacity(units.len() * 4);
295    for &u in units {
296        bytes.extend_from_slice(&u.to_le_bytes());
297    }
298    bytes
299}
300
301#[cfg(test)]
302mod tests {
303    use super::*;
304    use crate::entropy::EntropyDecoder;
305    use crate::variant::{Rans64, RansByte};
306
307    #[test]
308    fn test_variant_values() {
309        assert_eq!(RansVariant::RansByte.as_int(), 1);
310        assert_eq!(RansVariant::Rans64.as_int(), 0);
311        assert_eq!(RansVariant::from_int(1), Some(RansVariant::RansByte));
312        assert_eq!(RansVariant::from_int(0), Some(RansVariant::Rans64));
313        assert_eq!(RansVariant::from_int(2), None);
314    }
315
316    #[test]
317    fn test_encoder_stream_multipart() {
318        // Matches test_encode_decode_multi_part_0 from upstream
319        let pmf_lengths1 = vec![4, 6];
320        let pmf_offsets1 = vec![1, 2];
321        let pmf_table1 = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
322        let values1 = vec![-2, 1, 0, 1];
323        let indices1 = vec![0, 1, 0, 1];
324
325        let pmf_lengths2 = vec![5];
326        let pmf_offsets2 = vec![1];
327        let pmf_table2 = vec![1, 3, 3, 1, 1];
328        let values2 = vec![-2, 1, 2];
329        let indices2 = vec![0, 0, 0];
330
331        let mut encoder1 = EntropyEncoder::<RansByte>::new();
332        encoder1
333            .initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
334            .expect("init1");
335        let mut encoder2 = EntropyEncoder::<RansByte>::new();
336        encoder2
337            .initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
338            .expect("init2");
339
340        let mut stream = RansEncoderStream::<RansByte>::new();
341        stream.push(&encoder2, &indices2, &values2).expect("push2");
342        stream.push(&encoder1, &indices1, &values1).expect("push1");
343        let data = stream.flush().expect("flush");
344        assert!(!data.is_empty());
345
346        // Decode
347        let mut decoder1 = EntropyDecoder::<RansByte>::new();
348        decoder1
349            .initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
350            .expect("dec1 init");
351        let mut decoder2 = EntropyDecoder::<RansByte>::new();
352        decoder2
353            .initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
354            .expect("dec2 init");
355
356        let mut dstream = RansDecoderStream::<RansByte>::open_on(&data);
357
358        let mut decoded1 = vec![0i32; values1.len()];
359        dstream
360            .decode(&decoder1, &mut decoded1, &indices1)
361            .expect("decode1");
362        assert_eq!(decoded1, values1);
363
364        let mut decoded2 = vec![0i32; values2.len()];
365        dstream
366            .decode(&decoder2, &mut decoded2, &indices2)
367            .expect("decode2");
368        assert_eq!(decoded2, values2);
369
370        dstream.decode_eof().expect("eof");
371    }
372
373    #[test]
374    fn test_encoder_stream_multipart_64() {
375        // Same multipart test with Rans64
376        let pmf_lengths1 = vec![4, 6];
377        let pmf_offsets1 = vec![1, 2];
378        let pmf_table1 = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
379        let values1 = vec![-2, 1, 0, 1];
380        let indices1 = vec![0, 1, 0, 1];
381
382        let pmf_lengths2 = vec![5];
383        let pmf_offsets2 = vec![1];
384        let pmf_table2 = vec![1, 3, 3, 1, 1];
385        let values2 = vec![-2, 1, 2];
386        let indices2 = vec![0, 0, 0];
387
388        let mut encoder1 = EntropyEncoder::<Rans64>::new();
389        encoder1
390            .initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
391            .expect("init1");
392        let mut encoder2 = EntropyEncoder::<Rans64>::new();
393        encoder2
394            .initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
395            .expect("init2");
396
397        let mut stream = RansEncoderStream::<Rans64>::new();
398        stream.push(&encoder2, &indices2, &values2).expect("push2");
399        stream.push(&encoder1, &indices1, &values1).expect("push1");
400        let data = stream.flush().expect("flush");
401        assert!(!data.is_empty());
402        assert_eq!(data.len() % 4, 0, "Rans64 stream must be 4-byte aligned");
403
404        let mut decoder1 = EntropyDecoder::<Rans64>::new();
405        decoder1
406            .initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
407            .expect("dec1 init");
408        let mut decoder2 = EntropyDecoder::<Rans64>::new();
409        decoder2
410            .initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
411            .expect("dec2 init");
412
413        let mut dstream = RansDecoderStream::<Rans64>::open_on(&data);
414
415        let mut decoded1 = vec![0i32; values1.len()];
416        dstream
417            .decode(&decoder1, &mut decoded1, &indices1)
418            .expect("decode1");
419        assert_eq!(decoded1, values1);
420
421        let mut decoded2 = vec![0i32; values2.len()];
422        dstream
423            .decode(&decoder2, &mut decoded2, &indices2)
424            .expect("decode2");
425        assert_eq!(decoded2, values2);
426
427        dstream.decode_eof().expect("eof");
428    }
429
430    #[test]
431    fn test_encoder_stream_reuse_after_flush() {
432        let pmf_lengths = vec![4, 6];
433        let pmf_offsets = vec![1, 2];
434        let pmf_table = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
435        let values = vec![-2, 1, 0, 1];
436        let indices = vec![0, 1, 0, 1];
437
438        let mut encoder = EntropyEncoder::<RansByte>::new();
439        encoder
440            .initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
441            .expect("init");
442
443        let mut stream = RansEncoderStream::<RansByte>::new();
444        stream.push(&encoder, &indices, &values).expect("push1");
445        let data1 = stream.flush().expect("flush1");
446
447        // Reuse the stream for a second message
448        stream.push(&encoder, &indices, &values).expect("push2");
449        let data2 = stream.flush().expect("flush2");
450
451        let mut decoder = EntropyDecoder::<RansByte>::new();
452        decoder
453            .initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
454            .expect("dec init");
455
456        for data in [&data1, &data2] {
457            let mut decoded = vec![0i32; values.len()];
458            decoder
459                .decode(&mut decoded, &indices, data)
460                .expect("decode");
461            assert_eq!(decoded, values);
462        }
463    }
464
465    #[test]
466    fn test_encoder_stream_reset_aborts() {
467        let pmf_lengths = vec![4, 6];
468        let pmf_offsets = vec![1, 2];
469        let pmf_table = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
470        let values = vec![-2, 1, 0, 1];
471        let indices = vec![0, 1, 0, 1];
472
473        let mut encoder = EntropyEncoder::<RansByte>::new();
474        encoder
475            .initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
476            .expect("init");
477
478        let mut stream = RansEncoderStream::<RansByte>::new();
479        stream.push(&encoder, &indices, &values).expect("push");
480        stream.reset();
481        assert!(!stream.is_initialized(), "reset must clear state");
482        // flush after reset → no data → error
483        assert!(stream.flush().is_err());
484    }
485
486    #[test]
487    fn test_decoder_stream_lifecycle() {
488        let mut stream = RansDecoderStream::<RansByte>::new();
489        assert!(!stream.is_open());
490        stream.open(&[1, 2, 3, 4, 5]);
491        assert!(stream.is_open());
492        assert!(!stream.check_eof());
493        stream.close();
494        assert!(!stream.is_open());
495        assert!(stream.check_eof());
496    }
497
498    #[test]
499    fn test_units_to_le_bytes() {
500        let units = [0x01020304u32, 0x05060708];
501        let bytes = <Rans64 as EncoderVariantForS>::units_to_bytes(units.to_vec());
502        assert_eq!(bytes, vec![4, 3, 2, 1, 8, 7, 6, 5]);
503    }
504}