Skip to main content

msrtc_rans/
entropy.rs

1// Licensed under the MIT license.
2// Author: Riaan de Beer - github.com/infinityabundance - rdebeer.infinityabundance@gmail.com
3// Derived from Microsoft MLVC msrtc_rans (MIT)
4// See NOTICE file for attribution.
5
6//! # Entropy encoder/decoder (high-level PMF/bypass/CDF pipeline)
7//!
8//! Implements the high-level entropy coder matching Microsoft's
9//! `EntropyEncoder` / `EntropyDecoder` in `EntropyCoder.cpp`.
10//!
11//! Builds on the raw rANS primitives from `msrtc-rans-core`.
12
13#![allow(missing_docs)]
14
15use alloc::vec::Vec;
16
17use msrtc_rans_core::sink::VecSink;
18use msrtc_rans_core::source::SliceSource;
19use msrtc_rans_core::source::Source;
20use msrtc_rans_core::variant::{Rans64, RansByte, RansParams};
21use msrtc_rans_core::{
22    Freq, Rans64DecSymbol, Rans64EncSymbol, Rans64Encoder, RansByteDecSymbol, RansByteEncSymbol,
23    RansByteEncoder, RawRansError,
24};
25
26// ---------------------------------------------------------------------------
27// Constants
28// ---------------------------------------------------------------------------
29
30/// Size of Freq in bits (sizeof(u32) * 8 = 32).
31const FREQ_BITS: u32 = (core::mem::size_of::<Freq>() * 8) as u32;
32
33// ---------------------------------------------------------------------------
34// Error types
35// ---------------------------------------------------------------------------
36
37/// Errors that can occur during entropy encode/decode operations.
38#[derive(Debug, Clone, PartialEq, Eq)]
39pub enum EntropyError {
40    /// Invalid PMF data (lengths, offsets, or table).
41    InvalidPmf,
42    /// Invalid parameter value (symbolBits, bypassBits, etc.).
43    InvalidParams,
44    /// Encoder/decoder is not initialized.
45    InvalidState,
46    /// Stream data is truncated or corrupted.
47    InvalidStream,
48    /// Raw rANS primitive error.
49    RawRansError(RawRansError),
50}
51
52impl core::fmt::Display for EntropyError {
53    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
54        match self {
55            EntropyError::InvalidPmf => write!(f, "invalid PMF data"),
56            EntropyError::InvalidParams => write!(f, "invalid parameter value"),
57            EntropyError::InvalidState => write!(f, "invalid state (not initialized)"),
58            EntropyError::InvalidStream => write!(f, "invalid stream"),
59            EntropyError::RawRansError(e) => write!(f, "raw rANS error: {}", e),
60        }
61    }
62}
63
64#[cfg(feature = "std")]
65impl std::error::Error for EntropyError {}
66
67// ---------------------------------------------------------------------------
68// Distribution descriptor
69// ---------------------------------------------------------------------------
70
71/// Describes one probability distribution (one set of PMF symbols).
72#[derive(Debug, Clone)]
73struct DistributionDesc {
74    /// Offset applied to input values for this distribution.
75    value_offset: i32,
76    /// Sentinel index = length - 1 (last PMF element is tail mass for bypass).
77    bypass_sentinel: i32,
78    /// Starting offset of this distribution in the global symbol/CDF table.
79    symbol_offset: usize,
80}
81
82// ---------------------------------------------------------------------------
83// Helper: validate distribution descriptors from PMF arrays
84// ---------------------------------------------------------------------------
85
86fn initialize_distribution_desc(
87    distribution_descs: &mut Vec<DistributionDesc>,
88    pmf_lengths: &[i32],
89    pmf_offsets: &[i32],
90    pmf_table_size: usize,
91) -> Result<(), EntropyError> {
92    let distribution_count = pmf_lengths.len();
93    if pmf_offsets.len() != distribution_count {
94        return Err(EntropyError::InvalidPmf);
95    }
96    distribution_descs.reserve(distribution_count);
97
98    let mut symbol_cursor: usize = 0;
99    for i in 0..distribution_count {
100        let length = pmf_lengths[i];
101        // Each length must be > 1 (last element is bypass tail mass)
102        if length <= 1 || pmf_table_size - symbol_cursor < length as usize {
103            return Err(EntropyError::InvalidPmf);
104        }
105        distribution_descs.push(DistributionDesc {
106            value_offset: pmf_offsets[i],
107            bypass_sentinel: length - 1,
108            symbol_offset: symbol_cursor,
109        });
110        symbol_cursor += length as usize;
111    }
112
113    if symbol_cursor != pmf_table_size {
114        return Err(EntropyError::InvalidPmf);
115    }
116    Ok(())
117}
118
119// ---------------------------------------------------------------------------
120// Helper: check probability bits against max
121// ---------------------------------------------------------------------------
122
123#[inline]
124fn check_bits(prob_bits: u32, max_scale_bits: u32) -> Result<(), EntropyError> {
125    if prob_bits < 2 || prob_bits > max_scale_bits {
126        return Err(EntropyError::InvalidParams);
127    }
128    Ok(())
129}
130
131// ---------------------------------------------------------------------------
132// Helper: convert byte slice to u32 units (LE)
133// ---------------------------------------------------------------------------
134
135#[inline]
136fn bytes_to_u32_units(data: &[u8]) -> Vec<u32> {
137    data.chunks_exact(4)
138        .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
139        .collect()
140}
141
142// ---------------------------------------------------------------------------
143// Raw rANS encoder trait — abstracts RansByteEncoder and Rans64Encoder
144// ---------------------------------------------------------------------------
145
146/// Trait abstracting over raw rANS encoder variants for the entropy coder.
147pub(crate) trait RawEncoder {
148    type Unit: Copy + Default;
149    type Symbol;
150    fn put_raw(&mut self, start: Freq, freq: Freq, scale_bits: Freq);
151    fn put_symbol(&mut self, symbol: &Self::Symbol);
152    fn flush(&mut self);
153    fn into_units(self) -> Vec<Self::Unit>;
154}
155
156impl RawEncoder for RansByteEncoder<VecSink<u8>> {
157    type Unit = u8;
158    type Symbol = RansByteEncSymbol;
159
160    fn put_raw(&mut self, start: Freq, freq: Freq, scale_bits: Freq) {
161        self.put_raw(start, freq, scale_bits);
162    }
163
164    fn put_symbol(&mut self, symbol: &Self::Symbol) {
165        self.put(symbol);
166    }
167
168    fn flush(&mut self) {
169        self.flush();
170    }
171
172    fn into_units(self) -> Vec<u8> {
173        self.into_sink().encoded().to_vec()
174    }
175}
176
177impl RawEncoder for Rans64Encoder<VecSink<u32>> {
178    type Unit = u32;
179    type Symbol = Rans64EncSymbol;
180
181    fn put_raw(&mut self, start: Freq, freq: Freq, scale_bits: Freq) {
182        self.put_raw(start, freq, scale_bits);
183    }
184
185    fn put_symbol(&mut self, symbol: &Self::Symbol) {
186        self.put(symbol);
187    }
188
189    fn flush(&mut self) {
190        self.flush();
191    }
192
193    fn into_units(self) -> Vec<u32> {
194        self.into_sink().encoded().to_vec()
195    }
196}
197
198// ---------------------------------------------------------------------------
199// Internal encoder state — generic across variants
200// ---------------------------------------------------------------------------
201
202struct EncoderState<S: EncSymbol> {
203    symbol_bits: Freq,
204    distribution_descs: Vec<DistributionDesc>,
205    symbols: Vec<S>,
206    bypass_bits: Freq,
207    bypass_max_value: Freq,
208}
209
210/// Trait for encoder symbol types.
211pub(crate) trait EncSymbol: Sized {
212    fn try_new(start: Freq, freq: Freq, scale_bits: Freq) -> Result<Self, RawRansError>;
213}
214
215impl EncSymbol for RansByteEncSymbol {
216    fn try_new(start: Freq, freq: Freq, scale_bits: Freq) -> Result<Self, RawRansError> {
217        Self::try_new(start, freq, scale_bits)
218    }
219}
220
221impl EncSymbol for Rans64EncSymbol {
222    fn try_new(start: Freq, freq: Freq, scale_bits: Freq) -> Result<Self, RawRansError> {
223        Self::try_new(start, freq, scale_bits)
224    }
225}
226
227impl<S: EncSymbol> EncoderState<S> {
228    fn uninitialized() -> Self {
229        Self {
230            symbol_bits: 0,
231            distribution_descs: Vec::new(),
232            symbols: Vec::new(),
233            bypass_bits: 0,
234            bypass_max_value: 0,
235        }
236    }
237
238    fn initialize(
239        &mut self,
240        pmf_lengths: &[i32],
241        pmf_offsets: &[i32],
242        pmf_table: &[i32],
243        symbol_bits: i32,
244        bypass_bits: i32,
245        max_scale_bits: u32,
246    ) -> Result<(), EntropyError> {
247        let sb = symbol_bits as Freq;
248        let bb = bypass_bits as Freq;
249        check_bits(sb, max_scale_bits)?;
250        check_bits(bb, max_scale_bits)?;
251
252        // Issue 2a: safe maximum — RansByte allows 30, Rans64 allows 31 for bypass
253        let is_byte_variant = max_scale_bits < 32;
254        let max_safe_bits = if is_byte_variant { 30u32 } else { 31u32 };
255        if sb > max_safe_bits || bb > max_safe_bits {
256            return Err(EntropyError::InvalidParams);
257        }
258
259        let mut distribution_descs = Vec::new();
260        initialize_distribution_desc(
261            &mut distribution_descs,
262            pmf_lengths,
263            pmf_offsets,
264            pmf_table.len(),
265        )?;
266
267        // Issue 2b: use u64 for max_freq computation to avoid i32 overflow
268        let max_freq = 1u64 << symbol_bits;
269        let mut symbols: Vec<S> = Vec::with_capacity(pmf_table.len());
270        let mut pmf_cursor: usize = 0;
271
272        for desc in &distribution_descs {
273            // Issue 2b: track cumulative start in u64
274            let mut start: u64 = 0;
275            for _i in 0..=desc.bypass_sentinel {
276                let freq = pmf_table[pmf_cursor] as u64;
277                pmf_cursor += 1;
278                if !(freq > 0 && freq <= max_freq - start) {
279                    return Err(EntropyError::InvalidPmf);
280                }
281                let sym = S::try_new(start as Freq, freq as Freq, sb).map_err(|e| match e {
282                    RawRansError::InvalidScaleBits { .. } => EntropyError::InvalidParams,
283                    RawRansError::InvalidParameters => EntropyError::InvalidPmf,
284                })?;
285                symbols.push(sym);
286                start += freq;
287            }
288        }
289
290        self.distribution_descs = distribution_descs;
291        self.symbols = symbols;
292        self.symbol_bits = sb;
293        self.bypass_bits = bb;
294        // Issue 2c: compute from u64 to avoid overflow at bb=32
295        self.bypass_max_value = ((1u64 << bb) - 1) as Freq;
296        Ok(())
297    }
298
299    /// Encode symbols onto an existing encoder (batch/persistent mode).
300    ///
301    /// Unlike `encode_to_vec`, this does NOT flush the encoder or extract units,
302    /// enabling persistent streaming where multiple batches feed the same encoder.
303    fn encode_batch<E: RawEncoder<Symbol = S>>(
304        &self,
305        indices: &[i32],
306        values: &[i32],
307        encoder: &mut E,
308    ) -> Result<(), EntropyError> {
309        if self.symbol_bits == 0 {
310            return Err(EntropyError::InvalidState);
311        }
312        if indices.len() != values.len() {
313            return Err(EntropyError::InvalidParams);
314        }
315
316        // Encode in reverse order (matching C++: iterate from last to first)
317        let data_size = indices.len();
318        let dist_len = self.distribution_descs.len();
319        let mut idx = data_size as isize - 1;
320        while idx >= 0 {
321            let index = indices[idx as usize];
322            let value = values[idx as usize];
323
324            if index < 0 {
325                // Skip encoding, decoder returns 0 for skipped indices
326                idx -= 1;
327                continue;
328            }
329
330            // Clamp distribution index to a valid range
331            let ui = if (index as usize) < dist_len {
332                index as usize
333            } else {
334                dist_len - 1
335            };
336            let desc = &self.distribution_descs[ui];
337
338            // Issue 5: checked add to avoid i32 overflow
339            let adjusted = value
340                .checked_add(desc.value_offset)
341                .ok_or(EntropyError::InvalidParams)?;
342            let symbol_index: i32;
343            if adjusted < 0 || adjusted >= desc.bypass_sentinel {
344                // Out of PMF range — use bypass
345                let bypass_value: Freq = if adjusted < 0 {
346                    let neg = adjusted.checked_neg().ok_or(EntropyError::InvalidParams)?;
347                    2u64.wrapping_mul(neg as u64).wrapping_sub(1) as Freq
348                } else {
349                    2u64.wrapping_mul((adjusted - desc.bypass_sentinel) as u64) as Freq
350                };
351                self.encode_bypass_value(encoder, bypass_value);
352                symbol_index = desc.bypass_sentinel;
353            } else {
354                symbol_index = adjusted;
355            }
356
357            let sym_idx = desc.symbol_offset + symbol_index as usize;
358            encoder.put_symbol(&self.symbols[sym_idx]);
359            idx -= 1;
360        }
361
362        Ok(())
363    }
364
365    fn encode_to_vec<E: RawEncoder<Symbol = S>>(
366        &self,
367        indices: &[i32],
368        values: &[i32],
369        make_encoder: impl FnOnce() -> E,
370    ) -> Result<Vec<E::Unit>, EntropyError> {
371        let mut encoder = make_encoder();
372        self.encode_batch(indices, values, &mut encoder)?;
373        encoder.flush();
374        Ok(encoder.into_units())
375    }
376
377    #[inline]
378    fn encode_bypass_value<E: RawEncoder>(&self, encoder: &mut E, bypass_value: Freq) {
379        // Split bypassValue into bypassBits-sized digits (LSB first) into a
380        // FIXED stack buffer — no heap allocation per bypassed value.
381        //
382        // Max bypass_value needs 33 bits (i32 value range + offset), i.e. at
383        // most 17 digits at bypass_bits=2; 40 entries is a safe upper bound.
384        // (The C++ uses a 16-part std::array which overflows for wide bypass
385        // values — latent UB that Rust does not replicate.)
386        let mut bypass_buffer = [0u32; 40];
387        let mut parts = 0usize;
388
389        let mut bv = bypass_value;
390        while bv != 0 {
391            bypass_buffer[parts] = bv & self.bypass_max_value;
392            bv >>= self.bypass_bits;
393            parts += 1;
394        }
395
396        let mut bypass_count = parts as Freq;
397
398        // Put digits in reverse order (MSB first in the bitstream)
399        // since the rANS encoder writes from end to start
400        while parts > 0 {
401            parts -= 1;
402            encoder.put_raw(bypass_buffer[parts], 1, self.bypass_bits);
403        }
404
405        // Encode bypass count as remainder-coded prefix
406        // (each maxValue digit means "more to follow")
407        let mut bypass_prefix_count: Freq = 0;
408        while bypass_count >= self.bypass_max_value {
409            bypass_count -= self.bypass_max_value;
410            bypass_prefix_count += 1;
411        }
412        // Put bypassCount remainder (terminal digit)
413        encoder.put_raw(bypass_count, 1, self.bypass_bits);
414        // Put bypassCount prefix markers
415        for _ in 0..bypass_prefix_count {
416            encoder.put_raw(self.bypass_max_value, 1, self.bypass_bits);
417        }
418    }
419}
420
421// ---------------------------------------------------------------------------
422// Internal decoder state — generic across variants
423// ---------------------------------------------------------------------------
424
425struct DecoderState {
426    symbol_bits: Freq,
427    distribution_descs: Vec<DistributionDesc>,
428    cdf_table: Vec<Freq>,
429    bypass_bits: Freq,
430    bypass_max_value: Freq,
431}
432
433impl DecoderState {
434    fn uninitialized() -> Self {
435        Self {
436            symbol_bits: 0,
437            distribution_descs: Vec::new(),
438            cdf_table: Vec::new(),
439            bypass_bits: 0,
440            bypass_max_value: 0,
441        }
442    }
443
444    fn initialize(
445        &mut self,
446        pmf_lengths: &[i32],
447        pmf_offsets: &[i32],
448        pmf_table: &[i32],
449        symbol_bits: i32,
450        bypass_bits: i32,
451        max_scale_bits: u32,
452    ) -> Result<(), EntropyError> {
453        let sb = symbol_bits as Freq;
454        let bb = bypass_bits as Freq;
455        check_bits(sb, max_scale_bits)?;
456        check_bits(bb, max_scale_bits)?;
457
458        // Issue 2a: safe maximum
459        let is_byte_variant = max_scale_bits < 32;
460        let max_safe_bits = if is_byte_variant { 30u32 } else { 31u32 };
461        if sb > max_safe_bits || bb > max_safe_bits {
462            return Err(EntropyError::InvalidParams);
463        }
464
465        let mut distribution_descs = Vec::new();
466        initialize_distribution_desc(
467            &mut distribution_descs,
468            pmf_lengths,
469            pmf_offsets,
470            pmf_table.len(),
471        )?;
472
473        // Build CDF table: for each distribution, store cumulative starts
474        // plus one extra entry per distribution for the total sum
475        let num_dist = distribution_descs.len();
476        let mut cdf_table = vec![0u32; pmf_table.len() + num_dist];
477        // Issue 2b: use u64 for max_freq
478        let max_freq = 1u64 << symbol_bits;
479
480        let mut cursor: usize = 0;
481        for dist_idx in 0..num_dist {
482            // Update symbol_offset to point into the CDF table (not the PMF table)
483            distribution_descs[dist_idx].symbol_offset = cursor + dist_idx;
484
485            let mut start: u64 = 0;
486            for _i in 0..=distribution_descs[dist_idx].bypass_sentinel {
487                let freq = pmf_table[cursor] as u64;
488                if !(freq > 0 && freq <= max_freq - start) {
489                    return Err(EntropyError::InvalidPmf);
490                }
491                cdf_table[cursor + dist_idx] = start as Freq;
492                start += freq;
493                cursor += 1;
494            }
495            cdf_table[cursor + dist_idx] = start as Freq; // total sum
496        }
497
498        self.distribution_descs = distribution_descs;
499        self.cdf_table = cdf_table;
500        self.symbol_bits = sb;
501        self.bypass_bits = bb;
502        // Issue 2c: compute from u64
503        self.bypass_max_value = ((1u64 << bb) - 1) as Freq;
504        Ok(())
505    }
506
507    fn decode_from_slice(
508        &self,
509        values: &mut [i32],
510        indices: &[i32],
511        data: &[u8],
512        is_byte_variant: bool,
513    ) -> Result<(), EntropyError> {
514        if self.symbol_bits == 0 {
515            return Err(EntropyError::InvalidState);
516        }
517        if values.len() != indices.len() {
518            return Err(EntropyError::InvalidParams);
519        }
520
521        if is_byte_variant {
522            let units = data.to_vec();
523            let source = SliceSource::new(&units);
524            let mut decoder = msrtc_rans_core::RansByteDecoder::new(source);
525            if !decoder.init() {
526                return Err(EntropyError::InvalidStream);
527            }
528            self.decode_inner_byte(&mut decoder, values, indices)?;
529            if !decoder.source().is_exhausted() || !decoder.check_eof() {
530                return Err(EntropyError::InvalidStream);
531            }
532        } else {
533            // Issue 3: Reject misaligned Rans64 streams (must be 4-byte aligned)
534            if data.len() % 4 != 0 {
535                return Err(EntropyError::InvalidStream);
536            }
537            let units = bytes_to_u32_units(data);
538            let source = SliceSource::new(&units);
539            let mut decoder = msrtc_rans_core::Rans64Decoder::new(source);
540            if !decoder.init() {
541                return Err(EntropyError::InvalidStream);
542            }
543            self.decode_inner_64(&mut decoder, values, indices)?;
544            if !decoder.source().is_exhausted() || !decoder.check_eof() {
545                return Err(EntropyError::InvalidStream);
546            }
547        }
548        Ok(())
549    }
550
551    /// Helper: decode bypass count using remainder-coded prefix for RansByte decoder.
552    #[inline]
553    fn decode_bypass_count_byte(
554        &self,
555        decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
556    ) -> Result<Freq, EntropyError> {
557        let mut total: Freq = 0;
558        loop {
559            let value = decoder.get(self.bypass_bits);
560            if !decoder.advance(value, 1, self.bypass_bits) {
561                return Err(EntropyError::InvalidStream);
562            }
563            total += value;
564            if value != self.bypass_max_value {
565                break;
566            }
567            if total > FREQ_BITS {
568                return Err(EntropyError::InvalidStream);
569            }
570        }
571        Ok(total)
572    }
573
574    /// Helper: decode bypass count for Rans64 decoder.
575    #[inline]
576    fn decode_bypass_count_64(
577        &self,
578        decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
579    ) -> Result<Freq, EntropyError> {
580        let mut total: Freq = 0;
581        loop {
582            let value = decoder.get(self.bypass_bits);
583            if !decoder.advance(value, 1, self.bypass_bits) {
584                return Err(EntropyError::InvalidStream);
585            }
586            total += value;
587            if value != self.bypass_max_value {
588                break;
589            }
590            if total > FREQ_BITS {
591                return Err(EntropyError::InvalidStream);
592            }
593        }
594        Ok(total)
595    }
596
597    /// Helper: decode bypass value for RansByte decoder.
598    #[inline]
599    fn decode_bypass_value_payload_byte(
600        &self,
601        decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
602        bypass_count: Freq,
603    ) -> Result<Freq, EntropyError> {
604        // Issue 2c: use u64 for intermediate value to avoid overflow on shift
605        let mut encoded_value: u64 = 0;
606        let total_bits = bypass_count as u64 * self.bypass_bits as u64;
607        // Corrupt-stream hardening: a bounded bypass count can still imply
608        // shift >= 64 for wide bypass_bits; reject instead of panicking.
609        // (C++ computes `freq_t << shift` — undefined at shift >= 32.)
610        if total_bits >= 64 {
611            return Err(EntropyError::InvalidStream);
612        }
613        let mut shift: u64 = 0;
614        while shift < total_bits {
615            let v = decoder.get(self.bypass_bits);
616            if !decoder.advance(v, 1, self.bypass_bits) {
617                return Err(EntropyError::InvalidStream);
618            }
619            encoded_value |= (v as u64) << shift;
620            shift += self.bypass_bits as u64;
621        }
622        Ok(encoded_value as Freq)
623    }
624
625    /// Helper: decode bypass value for Rans64 decoder.
626    #[inline]
627    fn decode_bypass_value_payload_64(
628        &self,
629        decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
630        bypass_count: Freq,
631    ) -> Result<Freq, EntropyError> {
632        // Issue 2c: use u64 for intermediate value to avoid overflow on shift
633        let mut encoded_value: u64 = 0;
634        let total_bits = bypass_count as u64 * self.bypass_bits as u64;
635        // Corrupt-stream hardening (same as byte path).
636        if total_bits >= 64 {
637            return Err(EntropyError::InvalidStream);
638        }
639        let mut shift: u64 = 0;
640        while shift < total_bits {
641            let v = decoder.get(self.bypass_bits);
642            if !decoder.advance(v, 1, self.bypass_bits) {
643                return Err(EntropyError::InvalidStream);
644            }
645            encoded_value |= (v as u64) << shift;
646            shift += self.bypass_bits as u64;
647        }
648        Ok(encoded_value as Freq)
649    }
650
651    pub(crate) fn decode_inner_byte(
652        &self,
653        decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
654        values: &mut [i32],
655        indices: &[i32],
656    ) -> Result<(), EntropyError> {
657        if self.symbol_bits == 0 {
658            return Err(EntropyError::InvalidState);
659        }
660        if values.len() != indices.len() {
661            return Err(EntropyError::InvalidParams);
662        }
663
664        for (i, &index) in indices.iter().enumerate() {
665            if index < 0 {
666                values[i] = 0;
667                continue;
668            }
669
670            let dist_len = self.distribution_descs.len();
671            let ui = if (index as usize) < dist_len {
672                index as usize
673            } else {
674                dist_len - 1
675            };
676            let desc = &self.distribution_descs[ui];
677
678            // Get cumulative frequency from state (low symbolBits bits)
679            let cum_freq = decoder.get(self.symbol_bits);
680            debug_assert!(cum_freq < (1u32 << self.symbol_bits));
681
682            // Binary search in CDF table to find the symbol
683            let base_offset = desc.symbol_offset;
684            let lo = base_offset + 1;
685            let hi = base_offset + desc.bypass_sentinel as usize + 1;
686
687            // upper_bound: first element > cum_freq
688            let upper_idx = {
689                let mut low = lo;
690                let mut high = hi;
691                while low < high {
692                    let mid = low + (high - low) / 2;
693                    if cum_freq < self.cdf_table[mid] {
694                        high = mid;
695                    } else {
696                        low = mid + 1;
697                    }
698                }
699                low
700            };
701            // upper_bound - 1 gives the last element ≤ cum_freq
702            let start_idx = upper_idx - 1;
703
704            let s0 = self.cdf_table[start_idx];
705            let s1 = self.cdf_table[start_idx + 1];
706            let freq = s1 - s0;
707
708            if !decoder.advance_symbol(&RansByteDecSymbol::new(s0, freq), self.symbol_bits) {
709                return Err(EntropyError::InvalidStream);
710            }
711
712            let mut symbol = (start_idx - base_offset) as i32;
713            if symbol == desc.bypass_sentinel {
714                let bypass_count = self.decode_bypass_count_byte(decoder)?;
715                let bypass_value = self.decode_bypass_value_payload_byte(decoder, bypass_count)?;
716                // Issue 5: safe conversion with overflow checks using i64
717                let half = (bypass_value >> 1) as i64;
718                if bypass_value & 1 != 0 {
719                    // Negative: 2*(-value) - 1 -> value = -(bypassValue >> 1) - 1
720                    // = -(half as i64) - 1
721                    symbol = (-half)
722                        .checked_sub(1)
723                        .ok_or(EntropyError::InvalidStream)?
724                        .try_into()
725                        .map_err(|_| EntropyError::InvalidStream)?;
726                } else {
727                    // Positive: 2*(value - sentinel) -> value = (bypassValue >> 1) + sentinel
728                    symbol = half
729                        .checked_add(desc.bypass_sentinel as i64)
730                        .ok_or(EntropyError::InvalidStream)?
731                        .try_into()
732                        .map_err(|_| EntropyError::InvalidStream)?;
733                }
734            }
735
736            values[i] = (symbol as i64)
737                .checked_sub(desc.value_offset as i64)
738                .ok_or(EntropyError::InvalidStream)?
739                .try_into()
740                .map_err(|_| EntropyError::InvalidStream)?;
741        }
742        Ok(())
743    }
744
745    pub(crate) fn decode_inner_64(
746        &self,
747        decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
748        values: &mut [i32],
749        indices: &[i32],
750    ) -> Result<(), EntropyError> {
751        if self.symbol_bits == 0 {
752            return Err(EntropyError::InvalidState);
753        }
754        if values.len() != indices.len() {
755            return Err(EntropyError::InvalidParams);
756        }
757
758        for (i, &index) in indices.iter().enumerate() {
759            if index < 0 {
760                values[i] = 0;
761                continue;
762            }
763
764            let dist_len = self.distribution_descs.len();
765            let ui = if (index as usize) < dist_len {
766                index as usize
767            } else {
768                dist_len - 1
769            };
770            let desc = &self.distribution_descs[ui];
771
772            let cum_freq = decoder.get(self.symbol_bits);
773            debug_assert!(cum_freq < (1u32 << self.symbol_bits));
774
775            let base_offset = desc.symbol_offset;
776            let lo = base_offset + 1;
777            let hi = base_offset + desc.bypass_sentinel as usize + 1;
778
779            let upper_idx = {
780                let mut low = lo;
781                let mut high = hi;
782                while low < high {
783                    let mid = low + (high - low) / 2;
784                    if cum_freq < self.cdf_table[mid] {
785                        high = mid;
786                    } else {
787                        low = mid + 1;
788                    }
789                }
790                low
791            };
792            let start_idx = upper_idx - 1;
793
794            let s0 = self.cdf_table[start_idx];
795            let s1 = self.cdf_table[start_idx + 1];
796            let freq = s1 - s0;
797
798            if !decoder.advance_symbol(&Rans64DecSymbol::new(s0, freq), self.symbol_bits) {
799                return Err(EntropyError::InvalidStream);
800            }
801
802            let mut symbol = (start_idx - base_offset) as i32;
803            if symbol == desc.bypass_sentinel {
804                let bypass_count = self.decode_bypass_count_64(decoder)?;
805                let bypass_value = self.decode_bypass_value_payload_64(decoder, bypass_count)?;
806                // Issue 5: safe conversion with overflow checks using i64
807                let half = (bypass_value >> 1) as i64;
808                if bypass_value & 1 != 0 {
809                    // Negative: 2*(-value) - 1 -> value = -(bypassValue >> 1) - 1
810                    // = -(half as i64) - 1
811                    symbol = (-half)
812                        .checked_sub(1)
813                        .ok_or(EntropyError::InvalidStream)?
814                        .try_into()
815                        .map_err(|_| EntropyError::InvalidStream)?;
816                } else {
817                    // Positive: 2*(value - sentinel) -> value = (bypassValue >> 1) + sentinel
818                    symbol = half
819                        .checked_add(desc.bypass_sentinel as i64)
820                        .ok_or(EntropyError::InvalidStream)?
821                        .try_into()
822                        .map_err(|_| EntropyError::InvalidStream)?;
823                }
824            }
825
826            values[i] = (symbol as i64)
827                .checked_sub(desc.value_offset as i64)
828                .ok_or(EntropyError::InvalidStream)?
829                .try_into()
830                .map_err(|_| EntropyError::InvalidStream)?;
831        }
832        Ok(())
833    }
834}
835
836// ---------------------------------------------------------------------------
837// Helper traits to map RansParams to encoder/decoder types
838// ---------------------------------------------------------------------------
839
840/// Maps `RansParams` implementations to their encoder symbol types.
841///
842/// This trait is automatically implemented for both `RansByte` and `Rans64`
843/// and should not need to be implemented manually.
844pub trait EncoderVariantForS: RansParams {
845    /// The encoder symbol type for this variant.
846    type EncSymbol: EncSymbol;
847
848    /// The raw rANS encoder type for this variant.
849    type RawEnc: RawEncoder<Symbol = Self::EncSymbol>;
850
851    /// Maximum scale bits for this variant.
852    const MAX_SCALE_BITS: u32;
853
854    /// Convert raw encoder units to a byte vector.
855    fn units_to_bytes(units: Vec<<Self::RawEnc as RawEncoder>::Unit>) -> Vec<u8>;
856
857    /// Create a new encoder instance.
858    fn make_encoder() -> Self::RawEnc;
859}
860
861impl EncoderVariantForS for RansByte {
862    type EncSymbol = RansByteEncSymbol;
863    type RawEnc = RansByteEncoder<VecSink<u8>>;
864    const MAX_SCALE_BITS: u32 = 30;
865    fn units_to_bytes(units: Vec<u8>) -> Vec<u8> {
866        units
867    }
868    fn make_encoder() -> Self::RawEnc {
869        RansByteEncoder::new(VecSink::new(4096))
870    }
871}
872
873impl EncoderVariantForS for Rans64 {
874    type EncSymbol = Rans64EncSymbol;
875    type RawEnc = Rans64Encoder<VecSink<u32>>;
876    const MAX_SCALE_BITS: u32 = 32;
877    fn units_to_bytes(units: Vec<u32>) -> Vec<u8> {
878        let mut bytes = Vec::with_capacity(units.len() * 4);
879        for &u in &units {
880            bytes.extend_from_slice(&u.to_le_bytes());
881        }
882        bytes
883    }
884    fn make_encoder() -> Self::RawEnc {
885        Rans64Encoder::new(VecSink::new(4096))
886    }
887}
888
889// ---------------------------------------------------------------------------
890// Public EntropyEncoder
891// ---------------------------------------------------------------------------
892
893/// High-level entropy encoder using PMF distributions with bypass support.
894///
895/// Generic over `S: RansParams` (`RansByte` or `Rans64`).
896///
897/// # Example
898///
899/// ```ignore
900/// use msrtc_rans::entropy::EntropyEncoder;
901/// use msrtc_rans::RansByte;
902///
903/// let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
904/// enc.initialize(
905///     &[4, 6],       // pmf_lengths
906///     &[1, 2],       // pmf_offsets
907///     &[1, 3, 1, 1, 1, 3, 5, 3, 1, 1],  // pmf_table
908///     16,            // symbol_bits
909///     4,             // bypass_bits
910/// ).unwrap();
911///
912/// let mut buffer = Vec::new();
913/// enc.encode(&[0, 1], &[1, 1], &mut buffer).unwrap();
914/// ```
915pub struct EntropyEncoder<S: EncoderVariantForS> {
916    state: EncoderState<<S as EncoderVariantForS>::EncSymbol>,
917}
918
919impl<S: EncoderVariantForS> EntropyEncoder<S> {
920    /// Create a new uninitialized entropy encoder.
921    pub fn new() -> Self {
922        Self {
923            state: EncoderState::uninitialized(),
924        }
925    }
926
927    /// Initialize the encoder with PMF distribution data.
928    ///
929    /// * `pmf_lengths` — number of symbols per distribution (including bypass sentinel)
930    /// * `pmf_offsets` — value offsets per distribution
931    /// * `pmf_table` — flat array of symbol frequencies for all distributions
932    /// * `symbol_bits` — number of bits for symbol encoding (e.g. 16)
933    /// * `bypass_bits` — number of bits for bypass encoding (e.g. 4)
934    pub fn initialize(
935        &mut self,
936        pmf_lengths: &[i32],
937        pmf_offsets: &[i32],
938        pmf_table: &[i32],
939        symbol_bits: u32,
940        bypass_bits: u32,
941    ) -> Result<(), EntropyError> {
942        self.state.initialize(
943            pmf_lengths,
944            pmf_offsets,
945            pmf_table,
946            symbol_bits as i32,
947            bypass_bits as i32,
948            <S as EncoderVariantForS>::MAX_SCALE_BITS,
949        )
950    }
951
952    /// Encode symbols onto an existing raw encoder for persistent streaming.
953    ///
954    /// Unlike `encode()`, this does NOT flush the encoder or finalize output.
955    /// Call `encoder.flush()` and extract units after all batches are pushed.
956    pub fn encode_batch(
957        &self,
958        indices: &[i32],
959        values: &[i32],
960        encoder: &mut <S as EncoderVariantForS>::RawEnc,
961    ) -> Result<(), EntropyError> {
962        self.state.encode_batch(indices, values, encoder)
963    }
964
965    /// One-shot encode: encode `indices`/`values` into `buffer`.
966    ///
967    /// The encoded bytes are appended to `buffer`.
968    pub fn encode(
969        &self,
970        indices: &[i32],
971        values: &[i32],
972        buffer: &mut Vec<u8>,
973    ) -> Result<(), EntropyError> {
974        let units = self.state.encode_to_vec(indices, values, S::make_encoder)?;
975        let bytes = S::units_to_bytes(units);
976        buffer.extend_from_slice(&bytes);
977        Ok(())
978    }
979}
980
981impl<S: EncoderVariantForS> Default for EntropyEncoder<S> {
982    fn default() -> Self {
983        Self::new()
984    }
985}
986
987fn _assert_encoder_bounds() {
988    fn _is_encoder<S: EncoderVariantForS>() {}
989    _is_encoder::<RansByte>();
990    _is_encoder::<Rans64>();
991}
992
993// ---------------------------------------------------------------------------
994// Public EntropyDecoder
995// ---------------------------------------------------------------------------
996
997/// High-level entropy decoder using CDF tables for symbol lookup.
998///
999/// Generic over `S: RansParams` (`RansByte` or `Rans64`).
1000pub struct EntropyDecoder<S: RansParams> {
1001    state: DecoderState,
1002    _phantom: core::marker::PhantomData<S>,
1003}
1004
1005impl<S: RansParams> EntropyDecoder<S> {
1006    /// Create a new uninitialized entropy decoder.
1007    pub fn new() -> Self {
1008        Self {
1009            state: DecoderState::uninitialized(),
1010            _phantom: core::marker::PhantomData,
1011        }
1012    }
1013
1014    /// Initialize the decoder with PMF distribution data.
1015    ///
1016    /// * `pmf_lengths` — number of symbols per distribution (including bypass sentinel)
1017    /// * `pmf_offsets` — value offsets per distribution
1018    /// * `pmf_table` — flat array of symbol frequencies for all distributions
1019    /// * `symbol_bits` — number of bits for symbol encoding (e.g. 16)
1020    /// * `bypass_bits` — number of bits for bypass encoding (e.g. 4)
1021    pub fn initialize(
1022        &mut self,
1023        pmf_lengths: &[i32],
1024        pmf_offsets: &[i32],
1025        pmf_table: &[i32],
1026        symbol_bits: u32,
1027        bypass_bits: u32,
1028    ) -> Result<(), EntropyError> {
1029        let max_scale_bits = match S::NAME {
1030            "RansByte" => 30u32,
1031            "Rans64" => 32u32,
1032            _ => return Err(EntropyError::InvalidParams),
1033        };
1034        self.state.initialize(
1035            pmf_lengths,
1036            pmf_offsets,
1037            pmf_table,
1038            symbol_bits as i32,
1039            bypass_bits as i32,
1040            max_scale_bits,
1041        )
1042    }
1043
1044    /// One-shot decode: decode from `data` into `values`.
1045    ///
1046    /// * `values` — output buffer (must be same length as `indices`)
1047    /// * `indices` — distribution indices for each value to decode
1048    /// * `data` — encoded byte stream
1049    pub fn decode(
1050        &self,
1051        values: &mut [i32],
1052        indices: &[i32],
1053        data: &[u8],
1054    ) -> Result<(), EntropyError> {
1055        let is_byte = match S::NAME {
1056            "RansByte" => true,
1057            "Rans64" => false,
1058            _ => return Err(EntropyError::InvalidParams),
1059        };
1060        self.state.decode_from_slice(values, indices, data, is_byte)
1061    }
1062
1063    /// Decode from a slice but do NOT require the source to be fully exhausted.
1064    ///
1065    /// This is used when decoding from a `RansDecoderStream` where multiple encoded
1066    /// segments are concatenated. The method decodes `values`/`indices` from the
1067    /// beginning of `data` and returns the number of bytes consumed.
1068    ///
1069    /// * `values` — output buffer (must be same length as `indices`)
1070    /// * `indices` — distribution indices for each value to decode
1071    /// * `data` — encoded byte stream (may contain extra trailing data)
1072    ///
1073    /// Returns the number of bytes consumed from `data` on success.
1074    pub fn decode_partial(
1075        &self,
1076        values: &mut [i32],
1077        indices: &[i32],
1078        data: &[u8],
1079    ) -> Result<usize, EntropyError> {
1080        if self.state.symbol_bits == 0 {
1081            return Err(EntropyError::InvalidState);
1082        }
1083        if values.len() != indices.len() {
1084            return Err(EntropyError::InvalidParams);
1085        }
1086
1087        let consumed = match S::NAME {
1088            "RansByte" => {
1089                let units = data.to_vec();
1090                let source = SliceSource::new(&units);
1091                let mut decoder = msrtc_rans_core::RansByteDecoder::new(source);
1092                if !decoder.init() {
1093                    return Err(EntropyError::InvalidStream);
1094                }
1095                self.state
1096                    .decode_inner_byte(&mut decoder, values, indices)?;
1097                if !decoder.check_eof() {
1098                    return Err(EntropyError::InvalidStream);
1099                }
1100                decoder.source().position()
1101            }
1102            "Rans64" => {
1103                if data.len() % 4 != 0 {
1104                    return Err(EntropyError::InvalidStream);
1105                }
1106                let units = bytes_to_u32_units(data);
1107                let source = SliceSource::new(&units);
1108                let mut decoder = msrtc_rans_core::Rans64Decoder::new(source);
1109                if !decoder.init() {
1110                    return Err(EntropyError::InvalidStream);
1111                }
1112                self.state.decode_inner_64(&mut decoder, values, indices)?;
1113                if !decoder.check_eof() {
1114                    return Err(EntropyError::InvalidStream);
1115                }
1116                decoder.source().position() * 4
1117            }
1118            _ => return Err(EntropyError::InvalidParams),
1119        };
1120
1121        Ok(consumed)
1122    }
1123
1124    /// Decode a batch of symbols from data without requiring EOF.
1125    ///
1126    /// This is used for multipart stream decoding, where the stream contains
1127    /// multiple messages concatenated. The first decoder reads its symbols
1128    /// and returns the bytes consumed, leaving the rest for subsequent decoders.
1129    ///
1130    /// Returns the number of bytes consumed.
1131    pub fn decode_batch(
1132        &self,
1133        values: &mut [i32],
1134        indices: &[i32],
1135        data: &[u8],
1136    ) -> Result<usize, EntropyError> {
1137        if self.state.symbol_bits == 0 {
1138            return Err(EntropyError::InvalidState);
1139        }
1140        if values.len() != indices.len() {
1141            return Err(EntropyError::InvalidParams);
1142        }
1143
1144        let consumed = match S::NAME {
1145            "RansByte" => {
1146                let units = data.to_vec();
1147                let source = SliceSource::new(&units);
1148                let mut decoder = msrtc_rans_core::RansByteDecoder::new(source);
1149                if !decoder.init() {
1150                    return Err(EntropyError::InvalidStream);
1151                }
1152                self.state
1153                    .decode_inner_byte(&mut decoder, values, indices)?;
1154                decoder.source().position()
1155            }
1156            "Rans64" => {
1157                if data.len() % 4 != 0 {
1158                    return Err(EntropyError::InvalidStream);
1159                }
1160                let units = bytes_to_u32_units(data);
1161                let source = SliceSource::new(&units);
1162                let mut decoder = msrtc_rans_core::Rans64Decoder::new(source);
1163                if !decoder.init() {
1164                    return Err(EntropyError::InvalidStream);
1165                }
1166                self.state.decode_inner_64(&mut decoder, values, indices)?;
1167                decoder.source().position() * 4
1168            }
1169            _ => return Err(EntropyError::InvalidParams),
1170        };
1171
1172        Ok(consumed)
1173    }
1174
1175    /// Decode a batch of symbols from a sub-slice of data, returning bytes consumed.
1176    ///
1177    /// This is an alias for `decode_batch` used by the Python stream decoder.
1178    pub fn decode_stream(
1179        &self,
1180        values: &mut [i32],
1181        indices: &[i32],
1182        data: &[u8],
1183    ) -> Result<usize, EntropyError> {
1184        self.decode_batch(values, indices, data)
1185    }
1186
1187    /// Continue decoding from a persistent RansByte decoder (stream mode).
1188    ///
1189    /// The decoder must already be initialized (via `init()` on the first
1190    /// call). The caller owns the decoder and its source cursor.
1191    pub fn decode_byte_continue(
1192        &self,
1193        raw: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
1194        values: &mut [i32],
1195        indices: &[i32],
1196    ) -> Result<(), EntropyError> {
1197        self.state.decode_inner_byte(raw, values, indices)
1198    }
1199
1200    /// Continue decoding from a persistent Rans64 decoder (stream mode).
1201    pub fn decode_64_continue(
1202        &self,
1203        raw: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
1204        values: &mut [i32],
1205        indices: &[i32],
1206    ) -> Result<(), EntropyError> {
1207        self.state.decode_inner_64(raw, values, indices)
1208    }
1209}
1210
1211impl<S: RansParams> Default for EntropyDecoder<S> {
1212    fn default() -> Self {
1213        Self::new()
1214    }
1215}
1216
1217// ---------------------------------------------------------------------------
1218// Tests
1219// ---------------------------------------------------------------------------
1220
1221#[cfg(test)]
1222mod tests {
1223    use super::*;
1224
1225    // Reference test case from test_msrtc_rans.py:
1226    //   PMF_LENGTHS = [4, 6]
1227    //   PMF_OFFSETS = [1, 2]
1228    //   PMF_TABLE   = [1, 3, 1, 1, 1, 3, 5, 3, 1, 1]
1229    //   INDICES     = [0, 1, 0, 1]
1230    //   VALUES      = [-2, 1, 0, 1]
1231    //   SYMBOL_BITS = 16
1232    //   BYPASS_BITS = 4
1233    //
1234    // Reference bitstreams from upstream oracle (EntropyCoder.cpp):
1235    //   RansByte: hex = "0500bd040001a10003000b00"
1236    //   Rans64:   hex = "0500a1bd04000000110a002f03000300"
1237
1238    const PMF_LENGTHS: [i32; 2] = [4, 6];
1239    const PMF_OFFSETS: [i32; 2] = [1, 2];
1240    const PMF_TABLE: [i32; 10] = [1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
1241    const INDICES: [i32; 4] = [0, 1, 0, 1];
1242    const VALUES: [i32; 4] = [-2, 1, 0, 1];
1243    const SYMBOL_BITS: u32 = 16;
1244    const BYPASS_BITS: u32 = 4;
1245
1246    const REF_HEX_BYTE: &str = "0500bd040001a10003000b00";
1247    const REF_HEX_64: &str = "0500a1bd04000000110a002f03000300";
1248
1249    fn hex_decode(hex: &str) -> Vec<u8> {
1250        (0..hex.len())
1251            .step_by(2)
1252            .map(|i| u8::from_str_radix(&hex[i..i + 2], 16).unwrap())
1253            .collect()
1254    }
1255
1256    #[test]
1257    fn test_encoder_byte_initialize() {
1258        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1259        assert!(
1260            enc.initialize(
1261                &PMF_LENGTHS,
1262                &PMF_OFFSETS,
1263                &PMF_TABLE,
1264                SYMBOL_BITS,
1265                BYPASS_BITS
1266            )
1267            .is_ok()
1268        );
1269    }
1270
1271    #[test]
1272    fn test_encoder_64_initialize() {
1273        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1274        assert!(
1275            enc.initialize(
1276                &PMF_LENGTHS,
1277                &PMF_OFFSETS,
1278                &PMF_TABLE,
1279                SYMBOL_BITS,
1280                BYPASS_BITS
1281            )
1282            .is_ok()
1283        );
1284    }
1285
1286    #[test]
1287    fn test_encoder_rejects_invalid_pmf() {
1288        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1289        // Mismatched lengths/offsets
1290        assert_eq!(
1291            enc.initialize(&[4], &[1, 2], &PMF_TABLE, SYMBOL_BITS, BYPASS_BITS),
1292            Err(EntropyError::InvalidPmf)
1293        );
1294    }
1295
1296    #[test]
1297    fn test_encoder_rejects_invalid_params() {
1298        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1299        // symbol_bits < 2
1300        assert_eq!(
1301            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 1, BYPASS_BITS),
1302            Err(EntropyError::InvalidParams)
1303        );
1304    }
1305
1306    #[test]
1307    fn test_encoder_byte_rejects_length_leq_one() {
1308        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1309        // Length <= 1 is invalid (need tail mass for bypass)
1310        assert_eq!(
1311            enc.initialize(&[1, 6], &[1, 2], &PMF_TABLE, SYMBOL_BITS, BYPASS_BITS),
1312            Err(EntropyError::InvalidPmf)
1313        );
1314    }
1315
1316    #[test]
1317    fn test_encode_byte_matches_reference() {
1318        // Encodes values=[-2, 1, 0, 1] with RansByte.
1319        // Value -2 (index 0, offset 1 => adjusted=-1) triggers bypass.
1320        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1321        enc.initialize(
1322            &PMF_LENGTHS,
1323            &PMF_OFFSETS,
1324            &PMF_TABLE,
1325            SYMBOL_BITS,
1326            BYPASS_BITS,
1327        )
1328        .unwrap();
1329
1330        let mut buffer = Vec::new();
1331        enc.encode(&INDICES, &VALUES, &mut buffer).unwrap();
1332
1333        let expected = hex_decode(REF_HEX_BYTE);
1334        assert_eq!(
1335            buffer, expected,
1336            "RansByte encode output does not match reference hex"
1337        );
1338    }
1339
1340    #[test]
1341    fn test_encode_64_matches_reference() {
1342        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1343        enc.initialize(
1344            &PMF_LENGTHS,
1345            &PMF_OFFSETS,
1346            &PMF_TABLE,
1347            SYMBOL_BITS,
1348            BYPASS_BITS,
1349        )
1350        .unwrap();
1351
1352        let mut buffer = Vec::new();
1353        enc.encode(&INDICES, &VALUES, &mut buffer).unwrap();
1354
1355        let expected = hex_decode(REF_HEX_64);
1356        assert_eq!(
1357            buffer, expected,
1358            "Rans64 encode output does not match reference hex"
1359        );
1360    }
1361
1362    #[test]
1363    fn test_encode_in_range_values_no_bypass() {
1364        // Values that are all in-range (no bypass):
1365        // Dist 0: offset=1, sentinel=3, valid adjusted: [0,2] => value: [-1, 1]
1366        // Dist 1: offset=2, sentinel=5, valid adjusted: [0,4] => value: [-2, 2]
1367        let in_range_values = [1i32, 1, 0, 1];
1368        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1369        enc.initialize(
1370            &PMF_LENGTHS,
1371            &PMF_OFFSETS,
1372            &PMF_TABLE,
1373            SYMBOL_BITS,
1374            BYPASS_BITS,
1375        )
1376        .unwrap();
1377
1378        let mut buffer = Vec::new();
1379        let result = enc.encode(&INDICES, &in_range_values, &mut buffer);
1380        assert!(result.is_ok(), "encode should succeed: {:?}", result);
1381        assert!(!buffer.is_empty(), "encoded buffer should not be empty");
1382    }
1383
1384    #[test]
1385    fn test_decode_byte_roundtrip_in_range() {
1386        // Verify roundtrip encode-decode with in-range values (no bypass).
1387        let values = [1i32, 1, 0, 1];
1388
1389        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1390        enc.initialize(
1391            &PMF_LENGTHS,
1392            &PMF_OFFSETS,
1393            &PMF_TABLE,
1394            SYMBOL_BITS,
1395            BYPASS_BITS,
1396        )
1397        .unwrap();
1398
1399        let mut encoded = Vec::new();
1400        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1401
1402        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1403        dec.initialize(
1404            &PMF_LENGTHS,
1405            &PMF_OFFSETS,
1406            &PMF_TABLE,
1407            SYMBOL_BITS,
1408            BYPASS_BITS,
1409        )
1410        .unwrap();
1411
1412        let mut decoded = vec![0i32; values.len()];
1413        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1414
1415        assert_eq!(
1416            decoded, values,
1417            "roundtrip decode should match original values"
1418        );
1419    }
1420
1421    #[test]
1422    fn test_decode_64_roundtrip_in_range() {
1423        let values = [1i32, 1, 0, 1];
1424
1425        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1426        enc.initialize(
1427            &PMF_LENGTHS,
1428            &PMF_OFFSETS,
1429            &PMF_TABLE,
1430            SYMBOL_BITS,
1431            BYPASS_BITS,
1432        )
1433        .unwrap();
1434
1435        let mut encoded = Vec::new();
1436        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1437
1438        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1439        dec.initialize(
1440            &PMF_LENGTHS,
1441            &PMF_OFFSETS,
1442            &PMF_TABLE,
1443            SYMBOL_BITS,
1444            BYPASS_BITS,
1445        )
1446        .unwrap();
1447
1448        let mut decoded = vec![0i32; values.len()];
1449        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1450
1451        assert_eq!(
1452            decoded, values,
1453            "Rans64 roundtrip decode should match original values"
1454        );
1455    }
1456
1457    #[test]
1458    fn test_decode_byte_roundtrip_bypass() {
1459        // Roundtrip with values that require bypass (value=-2 is out-of-range).
1460        let values = [-2i32, 1, 0, 1];
1461
1462        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1463        enc.initialize(
1464            &PMF_LENGTHS,
1465            &PMF_OFFSETS,
1466            &PMF_TABLE,
1467            SYMBOL_BITS,
1468            BYPASS_BITS,
1469        )
1470        .unwrap();
1471
1472        let mut encoded = Vec::new();
1473        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1474
1475        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1476        dec.initialize(
1477            &PMF_LENGTHS,
1478            &PMF_OFFSETS,
1479            &PMF_TABLE,
1480            SYMBOL_BITS,
1481            BYPASS_BITS,
1482        )
1483        .unwrap();
1484
1485        let mut decoded = vec![0i32; values.len()];
1486        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1487
1488        assert_eq!(
1489            decoded, values,
1490            "bypass roundtrip decode should match original values"
1491        );
1492    }
1493
1494    #[test]
1495    fn test_decode_64_roundtrip_bypass() {
1496        let values = [-2i32, 1, 0, 1];
1497
1498        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1499        enc.initialize(
1500            &PMF_LENGTHS,
1501            &PMF_OFFSETS,
1502            &PMF_TABLE,
1503            SYMBOL_BITS,
1504            BYPASS_BITS,
1505        )
1506        .unwrap();
1507
1508        let mut encoded = Vec::new();
1509        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1510
1511        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1512        dec.initialize(
1513            &PMF_LENGTHS,
1514            &PMF_OFFSETS,
1515            &PMF_TABLE,
1516            SYMBOL_BITS,
1517            BYPASS_BITS,
1518        )
1519        .unwrap();
1520
1521        let mut decoded = vec![0i32; values.len()];
1522        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1523
1524        assert_eq!(
1525            decoded, values,
1526            "Rans64 bypass roundtrip decode should match original values"
1527        );
1528    }
1529
1530    // -----------------------------------------------------------------------
1531    // Issue 2d: scale-32 safe / reject tests
1532    // -----------------------------------------------------------------------
1533
1534    #[test]
1535    fn test_encoder_64_symbol_bits_31_accepted() {
1536        // Rans64 symbol_bits=31 is within max_safe_bits (31)
1537        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1538        assert!(
1539            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 31, BYPASS_BITS)
1540                .is_ok()
1541        );
1542    }
1543
1544    #[test]
1545    fn test_encoder_64_symbol_bits_32_rejected() {
1546        // Rans64 symbol_bits=32 exceeds max_safe_bits (31), must not panic
1547        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1548        assert_eq!(
1549            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 32, BYPASS_BITS),
1550            Err(EntropyError::InvalidParams)
1551        );
1552    }
1553
1554    #[test]
1555    fn test_encoder_64_bypass_bits_32_rejected() {
1556        // Rans64 bypass_bits=32 exceeds max_safe_bits (31), must not panic
1557        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1558        assert_eq!(
1559            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 32),
1560            Err(EntropyError::InvalidParams)
1561        );
1562    }
1563
1564    // -----------------------------------------------------------------------
1565    // Issue 3: misaligned Rans64 streams rejected
1566    // -----------------------------------------------------------------------
1567
1568    #[test]
1569    fn test_decode_64_rejects_misaligned_1_extra_byte() {
1570        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1571        enc.initialize(
1572            &PMF_LENGTHS,
1573            &PMF_OFFSETS,
1574            &PMF_TABLE,
1575            SYMBOL_BITS,
1576            BYPASS_BITS,
1577        )
1578        .unwrap();
1579        let mut encoded = Vec::new();
1580        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1581            .unwrap();
1582
1583        // Append 1 extra byte to make it misaligned
1584        let mut misaligned = encoded.clone();
1585        misaligned.push(0xAB);
1586
1587        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1588        dec.initialize(
1589            &PMF_LENGTHS,
1590            &PMF_OFFSETS,
1591            &PMF_TABLE,
1592            SYMBOL_BITS,
1593            BYPASS_BITS,
1594        )
1595        .unwrap();
1596        let mut decoded = vec![0i32; 4];
1597        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1598        assert_eq!(result, Err(EntropyError::InvalidStream));
1599    }
1600
1601    #[test]
1602    fn test_decode_64_rejects_misaligned_2_extra_bytes() {
1603        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1604        enc.initialize(
1605            &PMF_LENGTHS,
1606            &PMF_OFFSETS,
1607            &PMF_TABLE,
1608            SYMBOL_BITS,
1609            BYPASS_BITS,
1610        )
1611        .unwrap();
1612        let mut encoded = Vec::new();
1613        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1614            .unwrap();
1615
1616        let mut misaligned = encoded.clone();
1617        misaligned.extend_from_slice(&[0xAB, 0xCD]);
1618
1619        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1620        dec.initialize(
1621            &PMF_LENGTHS,
1622            &PMF_OFFSETS,
1623            &PMF_TABLE,
1624            SYMBOL_BITS,
1625            BYPASS_BITS,
1626        )
1627        .unwrap();
1628        let mut decoded = vec![0i32; 4];
1629        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1630        assert_eq!(result, Err(EntropyError::InvalidStream));
1631    }
1632
1633    #[test]
1634    fn test_decode_64_rejects_misaligned_3_extra_bytes() {
1635        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1636        enc.initialize(
1637            &PMF_LENGTHS,
1638            &PMF_OFFSETS,
1639            &PMF_TABLE,
1640            SYMBOL_BITS,
1641            BYPASS_BITS,
1642        )
1643        .unwrap();
1644        let mut encoded = Vec::new();
1645        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1646            .unwrap();
1647
1648        let mut misaligned = encoded.clone();
1649        misaligned.extend_from_slice(&[0xAB, 0xCD, 0xEF]);
1650
1651        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1652        dec.initialize(
1653            &PMF_LENGTHS,
1654            &PMF_OFFSETS,
1655            &PMF_TABLE,
1656            SYMBOL_BITS,
1657            BYPASS_BITS,
1658        )
1659        .unwrap();
1660        let mut decoded = vec![0i32; 4];
1661        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1662        assert_eq!(result, Err(EntropyError::InvalidStream));
1663    }
1664
1665    #[test]
1666    fn test_decode_byte_accepts_extra_bytes() {
1667        // RansByte has byte-level alignment, extra bytes should not be rejected
1668        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1669        enc.initialize(
1670            &PMF_LENGTHS,
1671            &PMF_OFFSETS,
1672            &PMF_TABLE,
1673            SYMBOL_BITS,
1674            BYPASS_BITS,
1675        )
1676        .unwrap();
1677        let mut encoded = Vec::new();
1678        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1679            .unwrap();
1680
1681        // Append extra bytes
1682        let mut extended = encoded.clone();
1683        extended.extend_from_slice(&[0xAB, 0xCD]);
1684
1685        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1686        dec.initialize(
1687            &PMF_LENGTHS,
1688            &PMF_OFFSETS,
1689            &PMF_TABLE,
1690            SYMBOL_BITS,
1691            BYPASS_BITS,
1692        )
1693        .unwrap();
1694        let mut decoded = vec![0i32; 4];
1695        // This may fail because the decoder checks EOF, but it should NOT
1696        // be rejected for misalignment
1697        let _ = dec.decode(&mut decoded, &INDICES, &extended);
1698        // We don't assert success or failure — we just assert no panic
1699    }
1700
1701    // -----------------------------------------------------------------------
1702    // Issue 4: Expanded bypass coverage
1703    // -----------------------------------------------------------------------
1704
1705    #[test]
1706    fn test_encode_bypass_positive_outlier() {
1707        // Value > sentinel (positive outlier): dist 0 sentinel=3, offset=1,
1708        // value=10 => adjusted=11 > 3 => bypass (8/2=4 above sentinel => value 4+3=7-1=6...
1709        // actually: adjusted=11, sentinel=3, bypass_value = 2*(11-3) = 16
1710        // decode: 16>>1=8, 8+3=11-1=10 ✓
1711        let values = [10i32, 1, 0, 1];
1712        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1713        enc.initialize(
1714            &PMF_LENGTHS,
1715            &PMF_OFFSETS,
1716            &PMF_TABLE,
1717            SYMBOL_BITS,
1718            BYPASS_BITS,
1719        )
1720        .unwrap();
1721        let mut encoded = Vec::new();
1722        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1723
1724        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1725        dec.initialize(
1726            &PMF_LENGTHS,
1727            &PMF_OFFSETS,
1728            &PMF_TABLE,
1729            SYMBOL_BITS,
1730            BYPASS_BITS,
1731        )
1732        .unwrap();
1733        let mut decoded = vec![0i32; 4];
1734        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1735        assert_eq!(decoded, values);
1736    }
1737
1738    #[test]
1739    fn test_encode_bypass_multi_digit_value() {
1740        // Value requiring multiple bypassBits-sized chunks (bypass_bits=4)
1741        // Large bypass value: dist 0 sentinel=3, offset=1, value=200 => adjusted=201
1742        // bypass_value = 2*(201-3) = 396 = 0x18C, needs multiple 4-bit chunks
1743        let values = [200i32, 1, 0, 1];
1744        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1745        enc.initialize(
1746            &PMF_LENGTHS,
1747            &PMF_OFFSETS,
1748            &PMF_TABLE,
1749            SYMBOL_BITS,
1750            BYPASS_BITS,
1751        )
1752        .unwrap();
1753        let mut encoded = Vec::new();
1754        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1755
1756        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1757        dec.initialize(
1758            &PMF_LENGTHS,
1759            &PMF_OFFSETS,
1760            &PMF_TABLE,
1761            SYMBOL_BITS,
1762            BYPASS_BITS,
1763        )
1764        .unwrap();
1765        let mut decoded = vec![0i32; 4];
1766        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1767        assert_eq!(decoded, values);
1768    }
1769
1770    #[test]
1771    fn test_encode_bypass_bits_2() {
1772        // Minimum bypass_bits = 2
1773        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1774        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 2)
1775            .unwrap();
1776        let values = [10i32, 1, 0, 1]; // value 10 => bypass
1777        let mut encoded = Vec::new();
1778        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1779
1780        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1781        dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 2)
1782            .unwrap();
1783        let mut decoded = vec![0i32; 4];
1784        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1785        assert_eq!(decoded, values);
1786    }
1787
1788    #[test]
1789    fn test_encode_bypass_bits_3() {
1790        // Odd bypass_bits = 3
1791        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1792        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 3)
1793            .unwrap();
1794        let values = [10i32, 1, 0, 1];
1795        let mut encoded = Vec::new();
1796        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1797
1798        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1799        dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 3)
1800            .unwrap();
1801        let mut decoded = vec![0i32; 4];
1802        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1803        assert_eq!(decoded, values);
1804    }
1805
1806    #[test]
1807    fn test_encode_bypass_bits_8() {
1808        // Larger bypass_bits = 8
1809        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1810        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 8)
1811            .unwrap();
1812        let values = [10i32, 1, 0, 1];
1813        let mut encoded = Vec::new();
1814        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1815
1816        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1817        dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 8)
1818            .unwrap();
1819        let mut decoded = vec![0i32; 4];
1820        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1821        assert_eq!(decoded, values);
1822    }
1823
1824    #[test]
1825    fn test_encode_bypass_multiple_bypasses() {
1826        // Multiple bypass values in one stream: both -2 and 10 need bypass
1827        // dist 0: offset=1 sentinel=3, dist 1: offset=2 sentinel=5
1828        // value -2 (dist 0) => adjusted=-1 => bypass (negative)
1829        // value 10 (dist 1) => adjusted=12 => bypass (positive)
1830        let values = [-2i32, 10, 0, 1];
1831        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1832        enc.initialize(
1833            &PMF_LENGTHS,
1834            &PMF_OFFSETS,
1835            &PMF_TABLE,
1836            SYMBOL_BITS,
1837            BYPASS_BITS,
1838        )
1839        .unwrap();
1840        let mut encoded = Vec::new();
1841        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1842
1843        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1844        dec.initialize(
1845            &PMF_LENGTHS,
1846            &PMF_OFFSETS,
1847            &PMF_TABLE,
1848            SYMBOL_BITS,
1849            BYPASS_BITS,
1850        )
1851        .unwrap();
1852        let mut decoded = vec![0i32; 4];
1853        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1854        assert_eq!(decoded, values);
1855    }
1856
1857    #[test]
1858    fn test_encode_bypass_mixed_in_range_and_bypass() {
1859        // Mix of in-range and bypass values
1860        // dist 0: offset=1 sentinel=3, valid adjusted [0,2] => values [-1, 1]
1861        // dist 1: offset=2 sentinel=5, valid adjusted [0,4] => values [-2, 2]
1862        // value 0 in dist 0 => in-range, value 1 in dist 0 => in-range
1863        // value 5 in dist 1 => bypass (adjusted=7 > 4), value -3 in dist 1 => bypass (adjusted=-1 < 0)
1864        let values = [0i32, 5, 1, -3];
1865        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1866        enc.initialize(
1867            &PMF_LENGTHS,
1868            &PMF_OFFSETS,
1869            &PMF_TABLE,
1870            SYMBOL_BITS,
1871            BYPASS_BITS,
1872        )
1873        .unwrap();
1874        let mut encoded = Vec::new();
1875        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1876
1877        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1878        dec.initialize(
1879            &PMF_LENGTHS,
1880            &PMF_OFFSETS,
1881            &PMF_TABLE,
1882            SYMBOL_BITS,
1883            BYPASS_BITS,
1884        )
1885        .unwrap();
1886        let mut decoded = vec![0i32; 4];
1887        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1888        assert_eq!(decoded, values);
1889    }
1890
1891    #[test]
1892    fn test_encode_bypass_negative_outlier_at_boundary() {
1893        // Negative outlier at boundary: -1 - sentinel (very negative)
1894        // dist 0: offset=1, sentinel=3, value=-10 => adjusted=-9 => bypass
1895        // (-9 < 0) => bypass_value = 2*9-1 = 17, decode: 17>>1=8, -(8+1) = -9, -9+1 = -8... wait
1896        // decode bypass: negative flag set, half=8, symbol = -(8+1) = -9, -9 = -9+1 = -8...
1897        // Actually: symbol = -9, values[i] = symbol - offset = (-9) - 1 = -10 ✓
1898        let values = [-10i32, 1, 0, 1];
1899        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1900        enc.initialize(
1901            &PMF_LENGTHS,
1902            &PMF_OFFSETS,
1903            &PMF_TABLE,
1904            SYMBOL_BITS,
1905            BYPASS_BITS,
1906        )
1907        .unwrap();
1908        let mut encoded = Vec::new();
1909        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1910
1911        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1912        dec.initialize(
1913            &PMF_LENGTHS,
1914            &PMF_OFFSETS,
1915            &PMF_TABLE,
1916            SYMBOL_BITS,
1917            BYPASS_BITS,
1918        )
1919        .unwrap();
1920        let mut decoded = vec![0i32; 4];
1921        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1922        assert_eq!(decoded, values);
1923    }
1924
1925    #[test]
1926    fn test_encode_bypass_large_positive_outlier() {
1927        // Large positive outlier
1928        let values = [10000i32, 1, 0, 1];
1929        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1930        enc.initialize(
1931            &PMF_LENGTHS,
1932            &PMF_OFFSETS,
1933            &PMF_TABLE,
1934            SYMBOL_BITS,
1935            BYPASS_BITS,
1936        )
1937        .unwrap();
1938        let mut encoded = Vec::new();
1939        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1940
1941        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1942        dec.initialize(
1943            &PMF_LENGTHS,
1944            &PMF_OFFSETS,
1945            &PMF_TABLE,
1946            SYMBOL_BITS,
1947            BYPASS_BITS,
1948        )
1949        .unwrap();
1950        let mut decoded = vec![0i32; 4];
1951        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1952        assert_eq!(decoded, values);
1953    }
1954
1955    // -----------------------------------------------------------------------
1956    // Issue 5: extreme value overflow protection
1957    // -----------------------------------------------------------------------
1958
1959    #[test]
1960    fn test_encode_bypass_extreme_negative_i32_min_plus_one() {
1961        // i32::MIN + 1 with offset 1 gives adjusted = i32::MIN + 2 = -2147483646
1962        // checked_neg of that gives 2147483646, no overflow.
1963        // This should succeed (no overflow).
1964        let values = [i32::MIN + 1, 1, 0, 1];
1965        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1966        enc.initialize(
1967            &PMF_LENGTHS,
1968            &PMF_OFFSETS,
1969            &PMF_TABLE,
1970            SYMBOL_BITS,
1971            BYPASS_BITS,
1972        )
1973        .unwrap();
1974        let mut encoded = Vec::new();
1975        let result = enc.encode(&INDICES, &values, &mut encoded);
1976        // Must not panic — should succeed or return InvalidParams gracefully
1977        assert!(result.is_ok() || result == Err(EntropyError::InvalidParams));
1978    }
1979
1980    #[test]
1981    fn test_encode_bypass_extreme_positive_i32_max() {
1982        // i32::MAX with offset could cause overflow in checked_add -> InvalidParams
1983        let values = [i32::MAX, 1, 0, 1];
1984        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1985        enc.initialize(
1986            &PMF_LENGTHS,
1987            &PMF_OFFSETS,
1988            &PMF_TABLE,
1989            SYMBOL_BITS,
1990            BYPASS_BITS,
1991        )
1992        .unwrap();
1993        let mut encoded = Vec::new();
1994        let result = enc.encode(&INDICES, &values, &mut encoded);
1995        assert_eq!(result, Err(EntropyError::InvalidParams));
1996    }
1997}