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 mut idx = data_size as isize - 1;
319        while idx >= 0 {
320            let index = indices[idx as usize];
321            let value = values[idx as usize];
322
323            if index < 0 {
324                // Skip encoding, decoder returns 0 for skipped indices
325                idx -= 1;
326                continue;
327            }
328
329            // Clamp distribution index to a valid range
330            let dist_len = self.distribution_descs.len();
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)
380        let max_parts = (FREQ_BITS as usize / self.bypass_bits as usize).max(2);
381        let mut bypass_buffer = Vec::with_capacity(max_parts);
382
383        let mut bv = bypass_value;
384        while bv != 0 {
385            bypass_buffer.push(bv & self.bypass_max_value);
386            bv >>= self.bypass_bits;
387        }
388
389        let mut bypass_count = bypass_buffer.len() as Freq;
390
391        // Put digits in reverse order (MSB first in the bitstream)
392        // since the rANS encoder writes from end to start
393        for &digit in bypass_buffer.iter().rev() {
394            encoder.put_raw(digit, 1, self.bypass_bits);
395        }
396
397        // Encode bypass count as remainder-coded prefix
398        // (each maxValue digit means "more to follow")
399        let mut bypass_prefix_count: Freq = 0;
400        while bypass_count >= self.bypass_max_value {
401            bypass_count -= self.bypass_max_value;
402            bypass_prefix_count += 1;
403        }
404        // Put bypassCount remainder (terminal digit)
405        encoder.put_raw(bypass_count, 1, self.bypass_bits);
406        // Put bypassCount prefix markers
407        for _ in 0..bypass_prefix_count {
408            encoder.put_raw(self.bypass_max_value, 1, self.bypass_bits);
409        }
410    }
411}
412
413// ---------------------------------------------------------------------------
414// Internal decoder state — generic across variants
415// ---------------------------------------------------------------------------
416
417struct DecoderState {
418    symbol_bits: Freq,
419    distribution_descs: Vec<DistributionDesc>,
420    cdf_table: Vec<Freq>,
421    bypass_bits: Freq,
422    bypass_max_value: Freq,
423}
424
425impl DecoderState {
426    fn uninitialized() -> Self {
427        Self {
428            symbol_bits: 0,
429            distribution_descs: Vec::new(),
430            cdf_table: Vec::new(),
431            bypass_bits: 0,
432            bypass_max_value: 0,
433        }
434    }
435
436    fn initialize(
437        &mut self,
438        pmf_lengths: &[i32],
439        pmf_offsets: &[i32],
440        pmf_table: &[i32],
441        symbol_bits: i32,
442        bypass_bits: i32,
443        max_scale_bits: u32,
444    ) -> Result<(), EntropyError> {
445        let sb = symbol_bits as Freq;
446        let bb = bypass_bits as Freq;
447        check_bits(sb, max_scale_bits)?;
448        check_bits(bb, max_scale_bits)?;
449
450        // Issue 2a: safe maximum
451        let is_byte_variant = max_scale_bits < 32;
452        let max_safe_bits = if is_byte_variant { 30u32 } else { 31u32 };
453        if sb > max_safe_bits || bb > max_safe_bits {
454            return Err(EntropyError::InvalidParams);
455        }
456
457        let mut distribution_descs = Vec::new();
458        initialize_distribution_desc(
459            &mut distribution_descs,
460            pmf_lengths,
461            pmf_offsets,
462            pmf_table.len(),
463        )?;
464
465        // Build CDF table: for each distribution, store cumulative starts
466        // plus one extra entry per distribution for the total sum
467        let num_dist = distribution_descs.len();
468        let mut cdf_table = vec![0u32; pmf_table.len() + num_dist];
469        // Issue 2b: use u64 for max_freq
470        let max_freq = 1u64 << symbol_bits;
471
472        let mut cursor: usize = 0;
473        for dist_idx in 0..num_dist {
474            // Update symbol_offset to point into the CDF table (not the PMF table)
475            distribution_descs[dist_idx].symbol_offset = cursor + dist_idx;
476
477            let mut start: u64 = 0;
478            for _i in 0..=distribution_descs[dist_idx].bypass_sentinel {
479                let freq = pmf_table[cursor] as u64;
480                if !(freq > 0 && freq <= max_freq - start) {
481                    return Err(EntropyError::InvalidPmf);
482                }
483                cdf_table[cursor + dist_idx] = start as Freq;
484                start += freq;
485                cursor += 1;
486            }
487            cdf_table[cursor + dist_idx] = start as Freq; // total sum
488        }
489
490        self.distribution_descs = distribution_descs;
491        self.cdf_table = cdf_table;
492        self.symbol_bits = sb;
493        self.bypass_bits = bb;
494        // Issue 2c: compute from u64
495        self.bypass_max_value = ((1u64 << bb) - 1) as Freq;
496        Ok(())
497    }
498
499    fn decode_from_slice(
500        &self,
501        values: &mut [i32],
502        indices: &[i32],
503        data: &[u8],
504        is_byte_variant: bool,
505    ) -> Result<(), EntropyError> {
506        if self.symbol_bits == 0 {
507            return Err(EntropyError::InvalidState);
508        }
509        if values.len() != indices.len() {
510            return Err(EntropyError::InvalidParams);
511        }
512
513        if is_byte_variant {
514            let units = data.to_vec();
515            let source = SliceSource::new(&units);
516            let mut decoder = msrtc_rans_core::RansByteDecoder::new(source);
517            if !decoder.init() {
518                return Err(EntropyError::InvalidStream);
519            }
520            self.decode_inner_byte(&mut decoder, values, indices)?;
521            if !decoder.source().is_exhausted() || !decoder.check_eof() {
522                return Err(EntropyError::InvalidStream);
523            }
524        } else {
525            // Issue 3: Reject misaligned Rans64 streams (must be 4-byte aligned)
526            if data.len() % 4 != 0 {
527                return Err(EntropyError::InvalidStream);
528            }
529            let units = bytes_to_u32_units(data);
530            let source = SliceSource::new(&units);
531            let mut decoder = msrtc_rans_core::Rans64Decoder::new(source);
532            if !decoder.init() {
533                return Err(EntropyError::InvalidStream);
534            }
535            self.decode_inner_64(&mut decoder, values, indices)?;
536            if !decoder.source().is_exhausted() || !decoder.check_eof() {
537                return Err(EntropyError::InvalidStream);
538            }
539        }
540        Ok(())
541    }
542
543    /// Helper: decode bypass count using remainder-coded prefix for RansByte decoder.
544    #[inline]
545    fn decode_bypass_count_byte(
546        &self,
547        decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
548    ) -> Result<Freq, EntropyError> {
549        let mut total: Freq = 0;
550        loop {
551            let value = decoder.get(self.bypass_bits);
552            if !decoder.advance(value, 1, self.bypass_bits) {
553                return Err(EntropyError::InvalidStream);
554            }
555            total += value;
556            if value != self.bypass_max_value {
557                break;
558            }
559            if total > FREQ_BITS {
560                return Err(EntropyError::InvalidStream);
561            }
562        }
563        Ok(total)
564    }
565
566    /// Helper: decode bypass count for Rans64 decoder.
567    #[inline]
568    fn decode_bypass_count_64(
569        &self,
570        decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
571    ) -> Result<Freq, EntropyError> {
572        let mut total: Freq = 0;
573        loop {
574            let value = decoder.get(self.bypass_bits);
575            if !decoder.advance(value, 1, self.bypass_bits) {
576                return Err(EntropyError::InvalidStream);
577            }
578            total += value;
579            if value != self.bypass_max_value {
580                break;
581            }
582            if total > FREQ_BITS {
583                return Err(EntropyError::InvalidStream);
584            }
585        }
586        Ok(total)
587    }
588
589    /// Helper: decode bypass value for RansByte decoder.
590    #[inline]
591    fn decode_bypass_value_payload_byte(
592        &self,
593        decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
594        bypass_count: Freq,
595    ) -> Result<Freq, EntropyError> {
596        // Issue 2c: use u64 for intermediate value to avoid overflow on shift
597        let mut encoded_value: u64 = 0;
598        let total_bits = bypass_count as u64 * self.bypass_bits as u64;
599        let mut shift: u64 = 0;
600        while shift < total_bits {
601            let v = decoder.get(self.bypass_bits);
602            if !decoder.advance(v, 1, self.bypass_bits) {
603                return Err(EntropyError::InvalidStream);
604            }
605            encoded_value |= (v as u64) << shift;
606            shift += self.bypass_bits as u64;
607        }
608        Ok(encoded_value as Freq)
609    }
610
611    /// Helper: decode bypass value for Rans64 decoder.
612    #[inline]
613    fn decode_bypass_value_payload_64(
614        &self,
615        decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
616        bypass_count: Freq,
617    ) -> Result<Freq, EntropyError> {
618        // Issue 2c: use u64 for intermediate value to avoid overflow on shift
619        let mut encoded_value: u64 = 0;
620        let total_bits = bypass_count as u64 * self.bypass_bits as u64;
621        let mut shift: u64 = 0;
622        while shift < total_bits {
623            let v = decoder.get(self.bypass_bits);
624            if !decoder.advance(v, 1, self.bypass_bits) {
625                return Err(EntropyError::InvalidStream);
626            }
627            encoded_value |= (v as u64) << shift;
628            shift += self.bypass_bits as u64;
629        }
630        Ok(encoded_value as Freq)
631    }
632
633    pub(crate) fn decode_inner_byte(
634        &self,
635        decoder: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
636        values: &mut [i32],
637        indices: &[i32],
638    ) -> Result<(), EntropyError> {
639        if self.symbol_bits == 0 {
640            return Err(EntropyError::InvalidState);
641        }
642        if values.len() != indices.len() {
643            return Err(EntropyError::InvalidParams);
644        }
645
646        for (i, &index) in indices.iter().enumerate() {
647            if index < 0 {
648                values[i] = 0;
649                continue;
650            }
651
652            let dist_len = self.distribution_descs.len();
653            let ui = if (index as usize) < dist_len {
654                index as usize
655            } else {
656                dist_len - 1
657            };
658            let desc = &self.distribution_descs[ui];
659
660            // Get cumulative frequency from state (low symbolBits bits)
661            let cum_freq = decoder.get(self.symbol_bits);
662            debug_assert!(cum_freq < (1u32 << self.symbol_bits));
663
664            // Binary search in CDF table to find the symbol
665            let base_offset = desc.symbol_offset;
666            let lo = base_offset + 1;
667            let hi = base_offset + desc.bypass_sentinel as usize + 1;
668
669            // upper_bound: first element > cum_freq
670            let upper_idx = {
671                let mut low = lo;
672                let mut high = hi;
673                while low < high {
674                    let mid = low + (high - low) / 2;
675                    if cum_freq < self.cdf_table[mid] {
676                        high = mid;
677                    } else {
678                        low = mid + 1;
679                    }
680                }
681                low
682            };
683            // upper_bound - 1 gives the last element ≤ cum_freq
684            let start_idx = upper_idx - 1;
685
686            let s0 = self.cdf_table[start_idx];
687            let s1 = self.cdf_table[start_idx + 1];
688            let freq = s1 - s0;
689
690            if !decoder.advance_symbol(&RansByteDecSymbol::new(s0, freq), self.symbol_bits) {
691                return Err(EntropyError::InvalidStream);
692            }
693
694            let mut symbol = (start_idx - base_offset) as i32;
695            if symbol == desc.bypass_sentinel {
696                let bypass_count = self.decode_bypass_count_byte(decoder)?;
697                let bypass_value = self.decode_bypass_value_payload_byte(decoder, bypass_count)?;
698                // Issue 5: safe conversion with overflow checks using i64
699                let half = (bypass_value >> 1) as i64;
700                if bypass_value & 1 != 0 {
701                    // Negative: 2*(-value) - 1 -> value = -(bypassValue >> 1) - 1
702                    // = -(half as i64) - 1
703                    symbol = (-half)
704                        .checked_sub(1)
705                        .ok_or(EntropyError::InvalidStream)?
706                        .try_into()
707                        .map_err(|_| EntropyError::InvalidStream)?;
708                } else {
709                    // Positive: 2*(value - sentinel) -> value = (bypassValue >> 1) + sentinel
710                    symbol = half
711                        .checked_add(desc.bypass_sentinel as i64)
712                        .ok_or(EntropyError::InvalidStream)?
713                        .try_into()
714                        .map_err(|_| EntropyError::InvalidStream)?;
715                }
716            }
717
718            values[i] = (symbol as i64)
719                .checked_sub(desc.value_offset as i64)
720                .ok_or(EntropyError::InvalidStream)?
721                .try_into()
722                .map_err(|_| EntropyError::InvalidStream)?;
723        }
724        Ok(())
725    }
726
727    pub(crate) fn decode_inner_64(
728        &self,
729        decoder: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
730        values: &mut [i32],
731        indices: &[i32],
732    ) -> Result<(), EntropyError> {
733        if self.symbol_bits == 0 {
734            return Err(EntropyError::InvalidState);
735        }
736        if values.len() != indices.len() {
737            return Err(EntropyError::InvalidParams);
738        }
739
740        for (i, &index) in indices.iter().enumerate() {
741            if index < 0 {
742                values[i] = 0;
743                continue;
744            }
745
746            let dist_len = self.distribution_descs.len();
747            let ui = if (index as usize) < dist_len {
748                index as usize
749            } else {
750                dist_len - 1
751            };
752            let desc = &self.distribution_descs[ui];
753
754            let cum_freq = decoder.get(self.symbol_bits);
755            debug_assert!(cum_freq < (1u32 << self.symbol_bits));
756
757            let base_offset = desc.symbol_offset;
758            let lo = base_offset + 1;
759            let hi = base_offset + desc.bypass_sentinel as usize + 1;
760
761            let upper_idx = {
762                let mut low = lo;
763                let mut high = hi;
764                while low < high {
765                    let mid = low + (high - low) / 2;
766                    if cum_freq < self.cdf_table[mid] {
767                        high = mid;
768                    } else {
769                        low = mid + 1;
770                    }
771                }
772                low
773            };
774            let start_idx = upper_idx - 1;
775
776            let s0 = self.cdf_table[start_idx];
777            let s1 = self.cdf_table[start_idx + 1];
778            let freq = s1 - s0;
779
780            if !decoder.advance_symbol(&Rans64DecSymbol::new(s0, freq), self.symbol_bits) {
781                return Err(EntropyError::InvalidStream);
782            }
783
784            let mut symbol = (start_idx - base_offset) as i32;
785            if symbol == desc.bypass_sentinel {
786                let bypass_count = self.decode_bypass_count_64(decoder)?;
787                let bypass_value = self.decode_bypass_value_payload_64(decoder, bypass_count)?;
788                // Issue 5: safe conversion with overflow checks using i64
789                let half = (bypass_value >> 1) as i64;
790                if bypass_value & 1 != 0 {
791                    // Negative: 2*(-value) - 1 -> value = -(bypassValue >> 1) - 1
792                    // = -(half as i64) - 1
793                    symbol = (-half)
794                        .checked_sub(1)
795                        .ok_or(EntropyError::InvalidStream)?
796                        .try_into()
797                        .map_err(|_| EntropyError::InvalidStream)?;
798                } else {
799                    // Positive: 2*(value - sentinel) -> value = (bypassValue >> 1) + sentinel
800                    symbol = half
801                        .checked_add(desc.bypass_sentinel as i64)
802                        .ok_or(EntropyError::InvalidStream)?
803                        .try_into()
804                        .map_err(|_| EntropyError::InvalidStream)?;
805                }
806            }
807
808            values[i] = (symbol as i64)
809                .checked_sub(desc.value_offset as i64)
810                .ok_or(EntropyError::InvalidStream)?
811                .try_into()
812                .map_err(|_| EntropyError::InvalidStream)?;
813        }
814        Ok(())
815    }
816}
817
818// ---------------------------------------------------------------------------
819// Helper traits to map RansParams to encoder/decoder types
820// ---------------------------------------------------------------------------
821
822/// Maps `RansParams` implementations to their encoder symbol types.
823///
824/// This trait is automatically implemented for both `RansByte` and `Rans64`
825/// and should not need to be implemented manually.
826pub trait EncoderVariantForS: RansParams {
827    /// The encoder symbol type for this variant.
828    type EncSymbol: EncSymbol;
829
830    /// The raw rANS encoder type for this variant.
831    type RawEnc: RawEncoder<Symbol = Self::EncSymbol>;
832
833    /// Maximum scale bits for this variant.
834    const MAX_SCALE_BITS: u32;
835
836    /// Convert raw encoder units to a byte vector.
837    fn units_to_bytes(units: Vec<<Self::RawEnc as RawEncoder>::Unit>) -> Vec<u8>;
838
839    /// Create a new encoder instance.
840    fn make_encoder() -> Self::RawEnc;
841}
842
843impl EncoderVariantForS for RansByte {
844    type EncSymbol = RansByteEncSymbol;
845    type RawEnc = RansByteEncoder<VecSink<u8>>;
846    const MAX_SCALE_BITS: u32 = 30;
847    fn units_to_bytes(units: Vec<u8>) -> Vec<u8> {
848        units
849    }
850    fn make_encoder() -> Self::RawEnc {
851        RansByteEncoder::new(VecSink::new(4096))
852    }
853}
854
855impl EncoderVariantForS for Rans64 {
856    type EncSymbol = Rans64EncSymbol;
857    type RawEnc = Rans64Encoder<VecSink<u32>>;
858    const MAX_SCALE_BITS: u32 = 32;
859    fn units_to_bytes(units: Vec<u32>) -> Vec<u8> {
860        let mut bytes = Vec::with_capacity(units.len() * 4);
861        for &u in &units {
862            bytes.extend_from_slice(&u.to_le_bytes());
863        }
864        bytes
865    }
866    fn make_encoder() -> Self::RawEnc {
867        Rans64Encoder::new(VecSink::new(4096))
868    }
869}
870
871// ---------------------------------------------------------------------------
872// Public EntropyEncoder
873// ---------------------------------------------------------------------------
874
875/// High-level entropy encoder using PMF distributions with bypass support.
876///
877/// Generic over `S: RansParams` (`RansByte` or `Rans64`).
878///
879/// # Example
880///
881/// ```ignore
882/// use msrtc_rans::entropy::EntropyEncoder;
883/// use msrtc_rans::RansByte;
884///
885/// let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
886/// enc.initialize(
887///     &[4, 6],       // pmf_lengths
888///     &[1, 2],       // pmf_offsets
889///     &[1, 3, 1, 1, 1, 3, 5, 3, 1, 1],  // pmf_table
890///     16,            // symbol_bits
891///     4,             // bypass_bits
892/// ).unwrap();
893///
894/// let mut buffer = Vec::new();
895/// enc.encode(&[0, 1], &[1, 1], &mut buffer).unwrap();
896/// ```
897pub struct EntropyEncoder<S: EncoderVariantForS> {
898    state: EncoderState<<S as EncoderVariantForS>::EncSymbol>,
899}
900
901impl<S: EncoderVariantForS> EntropyEncoder<S> {
902    /// Create a new uninitialized entropy encoder.
903    pub fn new() -> Self {
904        Self {
905            state: EncoderState::uninitialized(),
906        }
907    }
908
909    /// Initialize the encoder with PMF distribution data.
910    ///
911    /// * `pmf_lengths` — number of symbols per distribution (including bypass sentinel)
912    /// * `pmf_offsets` — value offsets per distribution
913    /// * `pmf_table` — flat array of symbol frequencies for all distributions
914    /// * `symbol_bits` — number of bits for symbol encoding (e.g. 16)
915    /// * `bypass_bits` — number of bits for bypass encoding (e.g. 4)
916    pub fn initialize(
917        &mut self,
918        pmf_lengths: &[i32],
919        pmf_offsets: &[i32],
920        pmf_table: &[i32],
921        symbol_bits: u32,
922        bypass_bits: u32,
923    ) -> Result<(), EntropyError> {
924        self.state.initialize(
925            pmf_lengths,
926            pmf_offsets,
927            pmf_table,
928            symbol_bits as i32,
929            bypass_bits as i32,
930            <S as EncoderVariantForS>::MAX_SCALE_BITS,
931        )
932    }
933
934    /// Encode symbols onto an existing raw encoder for persistent streaming.
935    ///
936    /// Unlike `encode()`, this does NOT flush the encoder or finalize output.
937    /// Call `encoder.flush()` and extract units after all batches are pushed.
938    pub fn encode_batch(
939        &self,
940        indices: &[i32],
941        values: &[i32],
942        encoder: &mut <S as EncoderVariantForS>::RawEnc,
943    ) -> Result<(), EntropyError> {
944        self.state.encode_batch(indices, values, encoder)
945    }
946
947    /// One-shot encode: encode `indices`/`values` into `buffer`.
948    ///
949    /// The encoded bytes are appended to `buffer`.
950    pub fn encode(
951        &self,
952        indices: &[i32],
953        values: &[i32],
954        buffer: &mut Vec<u8>,
955    ) -> Result<(), EntropyError> {
956        let units = self.state.encode_to_vec(indices, values, S::make_encoder)?;
957        let bytes = S::units_to_bytes(units);
958        buffer.extend_from_slice(&bytes);
959        Ok(())
960    }
961}
962
963impl<S: EncoderVariantForS> Default for EntropyEncoder<S> {
964    fn default() -> Self {
965        Self::new()
966    }
967}
968
969fn _assert_encoder_bounds() {
970    fn _is_encoder<S: EncoderVariantForS>() {}
971    _is_encoder::<RansByte>();
972    _is_encoder::<Rans64>();
973}
974
975// ---------------------------------------------------------------------------
976// Public EntropyDecoder
977// ---------------------------------------------------------------------------
978
979/// High-level entropy decoder using CDF tables for symbol lookup.
980///
981/// Generic over `S: RansParams` (`RansByte` or `Rans64`).
982pub struct EntropyDecoder<S: RansParams> {
983    state: DecoderState,
984    _phantom: core::marker::PhantomData<S>,
985}
986
987impl<S: RansParams> EntropyDecoder<S> {
988    /// Create a new uninitialized entropy decoder.
989    pub fn new() -> Self {
990        Self {
991            state: DecoderState::uninitialized(),
992            _phantom: core::marker::PhantomData,
993        }
994    }
995
996    /// Initialize the decoder with PMF distribution data.
997    ///
998    /// * `pmf_lengths` — number of symbols per distribution (including bypass sentinel)
999    /// * `pmf_offsets` — value offsets per distribution
1000    /// * `pmf_table` — flat array of symbol frequencies for all distributions
1001    /// * `symbol_bits` — number of bits for symbol encoding (e.g. 16)
1002    /// * `bypass_bits` — number of bits for bypass encoding (e.g. 4)
1003    pub fn initialize(
1004        &mut self,
1005        pmf_lengths: &[i32],
1006        pmf_offsets: &[i32],
1007        pmf_table: &[i32],
1008        symbol_bits: u32,
1009        bypass_bits: u32,
1010    ) -> Result<(), EntropyError> {
1011        let max_scale_bits = match S::NAME {
1012            "RansByte" => 30u32,
1013            "Rans64" => 32u32,
1014            _ => return Err(EntropyError::InvalidParams),
1015        };
1016        self.state.initialize(
1017            pmf_lengths,
1018            pmf_offsets,
1019            pmf_table,
1020            symbol_bits as i32,
1021            bypass_bits as i32,
1022            max_scale_bits,
1023        )
1024    }
1025
1026    /// One-shot decode: decode from `data` into `values`.
1027    ///
1028    /// * `values` — output buffer (must be same length as `indices`)
1029    /// * `indices` — distribution indices for each value to decode
1030    /// * `data` — encoded byte stream
1031    pub fn decode(
1032        &self,
1033        values: &mut [i32],
1034        indices: &[i32],
1035        data: &[u8],
1036    ) -> Result<(), EntropyError> {
1037        let is_byte = match S::NAME {
1038            "RansByte" => true,
1039            "Rans64" => false,
1040            _ => return Err(EntropyError::InvalidParams),
1041        };
1042        self.state.decode_from_slice(values, indices, data, is_byte)
1043    }
1044
1045    /// Decode from a slice but do NOT require the source to be fully exhausted.
1046    ///
1047    /// This is used when decoding from a `RansDecoderStream` where multiple encoded
1048    /// segments are concatenated. The method decodes `values`/`indices` from the
1049    /// beginning of `data` and returns the number of bytes consumed.
1050    ///
1051    /// * `values` — output buffer (must be same length as `indices`)
1052    /// * `indices` — distribution indices for each value to decode
1053    /// * `data` — encoded byte stream (may contain extra trailing data)
1054    ///
1055    /// Returns the number of bytes consumed from `data` on success.
1056    pub fn decode_partial(
1057        &self,
1058        values: &mut [i32],
1059        indices: &[i32],
1060        data: &[u8],
1061    ) -> Result<usize, EntropyError> {
1062        if self.state.symbol_bits == 0 {
1063            return Err(EntropyError::InvalidState);
1064        }
1065        if values.len() != indices.len() {
1066            return Err(EntropyError::InvalidParams);
1067        }
1068
1069        let consumed = match S::NAME {
1070            "RansByte" => {
1071                let units = data.to_vec();
1072                let source = SliceSource::new(&units);
1073                let mut decoder = msrtc_rans_core::RansByteDecoder::new(source);
1074                if !decoder.init() {
1075                    return Err(EntropyError::InvalidStream);
1076                }
1077                self.state
1078                    .decode_inner_byte(&mut decoder, values, indices)?;
1079                if !decoder.check_eof() {
1080                    return Err(EntropyError::InvalidStream);
1081                }
1082                decoder.source().position()
1083            }
1084            "Rans64" => {
1085                if data.len() % 4 != 0 {
1086                    return Err(EntropyError::InvalidStream);
1087                }
1088                let units = bytes_to_u32_units(data);
1089                let source = SliceSource::new(&units);
1090                let mut decoder = msrtc_rans_core::Rans64Decoder::new(source);
1091                if !decoder.init() {
1092                    return Err(EntropyError::InvalidStream);
1093                }
1094                self.state.decode_inner_64(&mut decoder, values, indices)?;
1095                if !decoder.check_eof() {
1096                    return Err(EntropyError::InvalidStream);
1097                }
1098                decoder.source().position() * 4
1099            }
1100            _ => return Err(EntropyError::InvalidParams),
1101        };
1102
1103        Ok(consumed)
1104    }
1105
1106    /// Decode a batch of symbols from data without requiring EOF.
1107    ///
1108    /// This is used for multipart stream decoding, where the stream contains
1109    /// multiple messages concatenated. The first decoder reads its symbols
1110    /// and returns the bytes consumed, leaving the rest for subsequent decoders.
1111    ///
1112    /// Returns the number of bytes consumed.
1113    pub fn decode_batch(
1114        &self,
1115        values: &mut [i32],
1116        indices: &[i32],
1117        data: &[u8],
1118    ) -> Result<usize, EntropyError> {
1119        if self.state.symbol_bits == 0 {
1120            return Err(EntropyError::InvalidState);
1121        }
1122        if values.len() != indices.len() {
1123            return Err(EntropyError::InvalidParams);
1124        }
1125
1126        let consumed = match S::NAME {
1127            "RansByte" => {
1128                let units = data.to_vec();
1129                let source = SliceSource::new(&units);
1130                let mut decoder = msrtc_rans_core::RansByteDecoder::new(source);
1131                if !decoder.init() {
1132                    return Err(EntropyError::InvalidStream);
1133                }
1134                self.state
1135                    .decode_inner_byte(&mut decoder, values, indices)?;
1136                decoder.source().position()
1137            }
1138            "Rans64" => {
1139                if data.len() % 4 != 0 {
1140                    return Err(EntropyError::InvalidStream);
1141                }
1142                let units = bytes_to_u32_units(data);
1143                let source = SliceSource::new(&units);
1144                let mut decoder = msrtc_rans_core::Rans64Decoder::new(source);
1145                if !decoder.init() {
1146                    return Err(EntropyError::InvalidStream);
1147                }
1148                self.state.decode_inner_64(&mut decoder, values, indices)?;
1149                decoder.source().position() * 4
1150            }
1151            _ => return Err(EntropyError::InvalidParams),
1152        };
1153
1154        Ok(consumed)
1155    }
1156
1157    /// Decode a batch of symbols from a sub-slice of data, returning bytes consumed.
1158    ///
1159    /// This is an alias for `decode_batch` used by the Python stream decoder.
1160    pub fn decode_stream(
1161        &self,
1162        values: &mut [i32],
1163        indices: &[i32],
1164        data: &[u8],
1165    ) -> Result<usize, EntropyError> {
1166        self.decode_batch(values, indices, data)
1167    }
1168
1169    /// Continue decoding from a persistent RansByte decoder (stream mode).
1170    ///
1171    /// The decoder must already be initialized (via `init()` on the first
1172    /// call). The caller owns the decoder and its source cursor.
1173    pub fn decode_byte_continue(
1174        &self,
1175        raw: &mut msrtc_rans_core::RansByteDecoder<SliceSource<'_, u8>>,
1176        values: &mut [i32],
1177        indices: &[i32],
1178    ) -> Result<(), EntropyError> {
1179        self.state.decode_inner_byte(raw, values, indices)
1180    }
1181
1182    /// Continue decoding from a persistent Rans64 decoder (stream mode).
1183    pub fn decode_64_continue(
1184        &self,
1185        raw: &mut msrtc_rans_core::Rans64Decoder<SliceSource<'_, u32>>,
1186        values: &mut [i32],
1187        indices: &[i32],
1188    ) -> Result<(), EntropyError> {
1189        self.state.decode_inner_64(raw, values, indices)
1190    }
1191}
1192
1193impl<S: RansParams> Default for EntropyDecoder<S> {
1194    fn default() -> Self {
1195        Self::new()
1196    }
1197}
1198
1199// ---------------------------------------------------------------------------
1200// Tests
1201// ---------------------------------------------------------------------------
1202
1203#[cfg(test)]
1204mod tests {
1205    use super::*;
1206
1207    // Reference test case from test_msrtc_rans.py:
1208    //   PMF_LENGTHS = [4, 6]
1209    //   PMF_OFFSETS = [1, 2]
1210    //   PMF_TABLE   = [1, 3, 1, 1, 1, 3, 5, 3, 1, 1]
1211    //   INDICES     = [0, 1, 0, 1]
1212    //   VALUES      = [-2, 1, 0, 1]
1213    //   SYMBOL_BITS = 16
1214    //   BYPASS_BITS = 4
1215    //
1216    // Reference bitstreams from upstream oracle (EntropyCoder.cpp):
1217    //   RansByte: hex = "0500bd040001a10003000b00"
1218    //   Rans64:   hex = "0500a1bd04000000110a002f03000300"
1219
1220    const PMF_LENGTHS: [i32; 2] = [4, 6];
1221    const PMF_OFFSETS: [i32; 2] = [1, 2];
1222    const PMF_TABLE: [i32; 10] = [1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
1223    const INDICES: [i32; 4] = [0, 1, 0, 1];
1224    const VALUES: [i32; 4] = [-2, 1, 0, 1];
1225    const SYMBOL_BITS: u32 = 16;
1226    const BYPASS_BITS: u32 = 4;
1227
1228    const REF_HEX_BYTE: &str = "0500bd040001a10003000b00";
1229    const REF_HEX_64: &str = "0500a1bd04000000110a002f03000300";
1230
1231    fn hex_decode(hex: &str) -> Vec<u8> {
1232        (0..hex.len())
1233            .step_by(2)
1234            .map(|i| u8::from_str_radix(&hex[i..i + 2], 16).unwrap())
1235            .collect()
1236    }
1237
1238    #[test]
1239    fn test_encoder_byte_initialize() {
1240        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1241        assert!(
1242            enc.initialize(
1243                &PMF_LENGTHS,
1244                &PMF_OFFSETS,
1245                &PMF_TABLE,
1246                SYMBOL_BITS,
1247                BYPASS_BITS
1248            )
1249            .is_ok()
1250        );
1251    }
1252
1253    #[test]
1254    fn test_encoder_64_initialize() {
1255        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1256        assert!(
1257            enc.initialize(
1258                &PMF_LENGTHS,
1259                &PMF_OFFSETS,
1260                &PMF_TABLE,
1261                SYMBOL_BITS,
1262                BYPASS_BITS
1263            )
1264            .is_ok()
1265        );
1266    }
1267
1268    #[test]
1269    fn test_encoder_rejects_invalid_pmf() {
1270        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1271        // Mismatched lengths/offsets
1272        assert_eq!(
1273            enc.initialize(&[4], &[1, 2], &PMF_TABLE, SYMBOL_BITS, BYPASS_BITS),
1274            Err(EntropyError::InvalidPmf)
1275        );
1276    }
1277
1278    #[test]
1279    fn test_encoder_rejects_invalid_params() {
1280        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1281        // symbol_bits < 2
1282        assert_eq!(
1283            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 1, BYPASS_BITS),
1284            Err(EntropyError::InvalidParams)
1285        );
1286    }
1287
1288    #[test]
1289    fn test_encoder_byte_rejects_length_leq_one() {
1290        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1291        // Length <= 1 is invalid (need tail mass for bypass)
1292        assert_eq!(
1293            enc.initialize(&[1, 6], &[1, 2], &PMF_TABLE, SYMBOL_BITS, BYPASS_BITS),
1294            Err(EntropyError::InvalidPmf)
1295        );
1296    }
1297
1298    #[test]
1299    fn test_encode_byte_matches_reference() {
1300        // Encodes values=[-2, 1, 0, 1] with RansByte.
1301        // Value -2 (index 0, offset 1 => adjusted=-1) triggers bypass.
1302        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1303        enc.initialize(
1304            &PMF_LENGTHS,
1305            &PMF_OFFSETS,
1306            &PMF_TABLE,
1307            SYMBOL_BITS,
1308            BYPASS_BITS,
1309        )
1310        .unwrap();
1311
1312        let mut buffer = Vec::new();
1313        enc.encode(&INDICES, &VALUES, &mut buffer).unwrap();
1314
1315        let expected = hex_decode(REF_HEX_BYTE);
1316        assert_eq!(
1317            buffer, expected,
1318            "RansByte encode output does not match reference hex"
1319        );
1320    }
1321
1322    #[test]
1323    fn test_encode_64_matches_reference() {
1324        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1325        enc.initialize(
1326            &PMF_LENGTHS,
1327            &PMF_OFFSETS,
1328            &PMF_TABLE,
1329            SYMBOL_BITS,
1330            BYPASS_BITS,
1331        )
1332        .unwrap();
1333
1334        let mut buffer = Vec::new();
1335        enc.encode(&INDICES, &VALUES, &mut buffer).unwrap();
1336
1337        let expected = hex_decode(REF_HEX_64);
1338        assert_eq!(
1339            buffer, expected,
1340            "Rans64 encode output does not match reference hex"
1341        );
1342    }
1343
1344    #[test]
1345    fn test_encode_in_range_values_no_bypass() {
1346        // Values that are all in-range (no bypass):
1347        // Dist 0: offset=1, sentinel=3, valid adjusted: [0,2] => value: [-1, 1]
1348        // Dist 1: offset=2, sentinel=5, valid adjusted: [0,4] => value: [-2, 2]
1349        let in_range_values = [1i32, 1, 0, 1];
1350        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1351        enc.initialize(
1352            &PMF_LENGTHS,
1353            &PMF_OFFSETS,
1354            &PMF_TABLE,
1355            SYMBOL_BITS,
1356            BYPASS_BITS,
1357        )
1358        .unwrap();
1359
1360        let mut buffer = Vec::new();
1361        let result = enc.encode(&INDICES, &in_range_values, &mut buffer);
1362        assert!(result.is_ok(), "encode should succeed: {:?}", result);
1363        assert!(!buffer.is_empty(), "encoded buffer should not be empty");
1364    }
1365
1366    #[test]
1367    fn test_decode_byte_roundtrip_in_range() {
1368        // Verify roundtrip encode-decode with in-range values (no bypass).
1369        let values = [1i32, 1, 0, 1];
1370
1371        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1372        enc.initialize(
1373            &PMF_LENGTHS,
1374            &PMF_OFFSETS,
1375            &PMF_TABLE,
1376            SYMBOL_BITS,
1377            BYPASS_BITS,
1378        )
1379        .unwrap();
1380
1381        let mut encoded = Vec::new();
1382        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1383
1384        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1385        dec.initialize(
1386            &PMF_LENGTHS,
1387            &PMF_OFFSETS,
1388            &PMF_TABLE,
1389            SYMBOL_BITS,
1390            BYPASS_BITS,
1391        )
1392        .unwrap();
1393
1394        let mut decoded = vec![0i32; values.len()];
1395        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1396
1397        assert_eq!(
1398            decoded, values,
1399            "roundtrip decode should match original values"
1400        );
1401    }
1402
1403    #[test]
1404    fn test_decode_64_roundtrip_in_range() {
1405        let values = [1i32, 1, 0, 1];
1406
1407        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1408        enc.initialize(
1409            &PMF_LENGTHS,
1410            &PMF_OFFSETS,
1411            &PMF_TABLE,
1412            SYMBOL_BITS,
1413            BYPASS_BITS,
1414        )
1415        .unwrap();
1416
1417        let mut encoded = Vec::new();
1418        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1419
1420        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1421        dec.initialize(
1422            &PMF_LENGTHS,
1423            &PMF_OFFSETS,
1424            &PMF_TABLE,
1425            SYMBOL_BITS,
1426            BYPASS_BITS,
1427        )
1428        .unwrap();
1429
1430        let mut decoded = vec![0i32; values.len()];
1431        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1432
1433        assert_eq!(
1434            decoded, values,
1435            "Rans64 roundtrip decode should match original values"
1436        );
1437    }
1438
1439    #[test]
1440    fn test_decode_byte_roundtrip_bypass() {
1441        // Roundtrip with values that require bypass (value=-2 is out-of-range).
1442        let values = [-2i32, 1, 0, 1];
1443
1444        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1445        enc.initialize(
1446            &PMF_LENGTHS,
1447            &PMF_OFFSETS,
1448            &PMF_TABLE,
1449            SYMBOL_BITS,
1450            BYPASS_BITS,
1451        )
1452        .unwrap();
1453
1454        let mut encoded = Vec::new();
1455        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1456
1457        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1458        dec.initialize(
1459            &PMF_LENGTHS,
1460            &PMF_OFFSETS,
1461            &PMF_TABLE,
1462            SYMBOL_BITS,
1463            BYPASS_BITS,
1464        )
1465        .unwrap();
1466
1467        let mut decoded = vec![0i32; values.len()];
1468        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1469
1470        assert_eq!(
1471            decoded, values,
1472            "bypass roundtrip decode should match original values"
1473        );
1474    }
1475
1476    #[test]
1477    fn test_decode_64_roundtrip_bypass() {
1478        let values = [-2i32, 1, 0, 1];
1479
1480        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1481        enc.initialize(
1482            &PMF_LENGTHS,
1483            &PMF_OFFSETS,
1484            &PMF_TABLE,
1485            SYMBOL_BITS,
1486            BYPASS_BITS,
1487        )
1488        .unwrap();
1489
1490        let mut encoded = Vec::new();
1491        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1492
1493        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1494        dec.initialize(
1495            &PMF_LENGTHS,
1496            &PMF_OFFSETS,
1497            &PMF_TABLE,
1498            SYMBOL_BITS,
1499            BYPASS_BITS,
1500        )
1501        .unwrap();
1502
1503        let mut decoded = vec![0i32; values.len()];
1504        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1505
1506        assert_eq!(
1507            decoded, values,
1508            "Rans64 bypass roundtrip decode should match original values"
1509        );
1510    }
1511
1512    // -----------------------------------------------------------------------
1513    // Issue 2d: scale-32 safe / reject tests
1514    // -----------------------------------------------------------------------
1515
1516    #[test]
1517    fn test_encoder_64_symbol_bits_31_accepted() {
1518        // Rans64 symbol_bits=31 is within max_safe_bits (31)
1519        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1520        assert!(
1521            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 31, BYPASS_BITS)
1522                .is_ok()
1523        );
1524    }
1525
1526    #[test]
1527    fn test_encoder_64_symbol_bits_32_rejected() {
1528        // Rans64 symbol_bits=32 exceeds max_safe_bits (31), must not panic
1529        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1530        assert_eq!(
1531            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 32, BYPASS_BITS),
1532            Err(EntropyError::InvalidParams)
1533        );
1534    }
1535
1536    #[test]
1537    fn test_encoder_64_bypass_bits_32_rejected() {
1538        // Rans64 bypass_bits=32 exceeds max_safe_bits (31), must not panic
1539        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1540        assert_eq!(
1541            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 32),
1542            Err(EntropyError::InvalidParams)
1543        );
1544    }
1545
1546    // -----------------------------------------------------------------------
1547    // Issue 3: misaligned Rans64 streams rejected
1548    // -----------------------------------------------------------------------
1549
1550    #[test]
1551    fn test_decode_64_rejects_misaligned_1_extra_byte() {
1552        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1553        enc.initialize(
1554            &PMF_LENGTHS,
1555            &PMF_OFFSETS,
1556            &PMF_TABLE,
1557            SYMBOL_BITS,
1558            BYPASS_BITS,
1559        )
1560        .unwrap();
1561        let mut encoded = Vec::new();
1562        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1563            .unwrap();
1564
1565        // Append 1 extra byte to make it misaligned
1566        let mut misaligned = encoded.clone();
1567        misaligned.push(0xAB);
1568
1569        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1570        dec.initialize(
1571            &PMF_LENGTHS,
1572            &PMF_OFFSETS,
1573            &PMF_TABLE,
1574            SYMBOL_BITS,
1575            BYPASS_BITS,
1576        )
1577        .unwrap();
1578        let mut decoded = vec![0i32; 4];
1579        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1580        assert_eq!(result, Err(EntropyError::InvalidStream));
1581    }
1582
1583    #[test]
1584    fn test_decode_64_rejects_misaligned_2_extra_bytes() {
1585        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1586        enc.initialize(
1587            &PMF_LENGTHS,
1588            &PMF_OFFSETS,
1589            &PMF_TABLE,
1590            SYMBOL_BITS,
1591            BYPASS_BITS,
1592        )
1593        .unwrap();
1594        let mut encoded = Vec::new();
1595        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1596            .unwrap();
1597
1598        let mut misaligned = encoded.clone();
1599        misaligned.extend_from_slice(&[0xAB, 0xCD]);
1600
1601        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1602        dec.initialize(
1603            &PMF_LENGTHS,
1604            &PMF_OFFSETS,
1605            &PMF_TABLE,
1606            SYMBOL_BITS,
1607            BYPASS_BITS,
1608        )
1609        .unwrap();
1610        let mut decoded = vec![0i32; 4];
1611        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1612        assert_eq!(result, Err(EntropyError::InvalidStream));
1613    }
1614
1615    #[test]
1616    fn test_decode_64_rejects_misaligned_3_extra_bytes() {
1617        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1618        enc.initialize(
1619            &PMF_LENGTHS,
1620            &PMF_OFFSETS,
1621            &PMF_TABLE,
1622            SYMBOL_BITS,
1623            BYPASS_BITS,
1624        )
1625        .unwrap();
1626        let mut encoded = Vec::new();
1627        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1628            .unwrap();
1629
1630        let mut misaligned = encoded.clone();
1631        misaligned.extend_from_slice(&[0xAB, 0xCD, 0xEF]);
1632
1633        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1634        dec.initialize(
1635            &PMF_LENGTHS,
1636            &PMF_OFFSETS,
1637            &PMF_TABLE,
1638            SYMBOL_BITS,
1639            BYPASS_BITS,
1640        )
1641        .unwrap();
1642        let mut decoded = vec![0i32; 4];
1643        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1644        assert_eq!(result, Err(EntropyError::InvalidStream));
1645    }
1646
1647    #[test]
1648    fn test_decode_byte_accepts_extra_bytes() {
1649        // RansByte has byte-level alignment, extra bytes should not be rejected
1650        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1651        enc.initialize(
1652            &PMF_LENGTHS,
1653            &PMF_OFFSETS,
1654            &PMF_TABLE,
1655            SYMBOL_BITS,
1656            BYPASS_BITS,
1657        )
1658        .unwrap();
1659        let mut encoded = Vec::new();
1660        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1661            .unwrap();
1662
1663        // Append extra bytes
1664        let mut extended = encoded.clone();
1665        extended.extend_from_slice(&[0xAB, 0xCD]);
1666
1667        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1668        dec.initialize(
1669            &PMF_LENGTHS,
1670            &PMF_OFFSETS,
1671            &PMF_TABLE,
1672            SYMBOL_BITS,
1673            BYPASS_BITS,
1674        )
1675        .unwrap();
1676        let mut decoded = vec![0i32; 4];
1677        // This may fail because the decoder checks EOF, but it should NOT
1678        // be rejected for misalignment
1679        let _ = dec.decode(&mut decoded, &INDICES, &extended);
1680        // We don't assert success or failure — we just assert no panic
1681    }
1682
1683    // -----------------------------------------------------------------------
1684    // Issue 4: Expanded bypass coverage
1685    // -----------------------------------------------------------------------
1686
1687    #[test]
1688    fn test_encode_bypass_positive_outlier() {
1689        // Value > sentinel (positive outlier): dist 0 sentinel=3, offset=1,
1690        // value=10 => adjusted=11 > 3 => bypass (8/2=4 above sentinel => value 4+3=7-1=6...
1691        // actually: adjusted=11, sentinel=3, bypass_value = 2*(11-3) = 16
1692        // decode: 16>>1=8, 8+3=11-1=10 ✓
1693        let values = [10i32, 1, 0, 1];
1694        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1695        enc.initialize(
1696            &PMF_LENGTHS,
1697            &PMF_OFFSETS,
1698            &PMF_TABLE,
1699            SYMBOL_BITS,
1700            BYPASS_BITS,
1701        )
1702        .unwrap();
1703        let mut encoded = Vec::new();
1704        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1705
1706        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1707        dec.initialize(
1708            &PMF_LENGTHS,
1709            &PMF_OFFSETS,
1710            &PMF_TABLE,
1711            SYMBOL_BITS,
1712            BYPASS_BITS,
1713        )
1714        .unwrap();
1715        let mut decoded = vec![0i32; 4];
1716        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1717        assert_eq!(decoded, values);
1718    }
1719
1720    #[test]
1721    fn test_encode_bypass_multi_digit_value() {
1722        // Value requiring multiple bypassBits-sized chunks (bypass_bits=4)
1723        // Large bypass value: dist 0 sentinel=3, offset=1, value=200 => adjusted=201
1724        // bypass_value = 2*(201-3) = 396 = 0x18C, needs multiple 4-bit chunks
1725        let values = [200i32, 1, 0, 1];
1726        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1727        enc.initialize(
1728            &PMF_LENGTHS,
1729            &PMF_OFFSETS,
1730            &PMF_TABLE,
1731            SYMBOL_BITS,
1732            BYPASS_BITS,
1733        )
1734        .unwrap();
1735        let mut encoded = Vec::new();
1736        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1737
1738        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1739        dec.initialize(
1740            &PMF_LENGTHS,
1741            &PMF_OFFSETS,
1742            &PMF_TABLE,
1743            SYMBOL_BITS,
1744            BYPASS_BITS,
1745        )
1746        .unwrap();
1747        let mut decoded = vec![0i32; 4];
1748        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1749        assert_eq!(decoded, values);
1750    }
1751
1752    #[test]
1753    fn test_encode_bypass_bits_2() {
1754        // Minimum bypass_bits = 2
1755        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1756        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 2)
1757            .unwrap();
1758        let values = [10i32, 1, 0, 1]; // value 10 => bypass
1759        let mut encoded = Vec::new();
1760        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1761
1762        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1763        dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 2)
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_3() {
1772        // Odd bypass_bits = 3
1773        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1774        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 3)
1775            .unwrap();
1776        let values = [10i32, 1, 0, 1];
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, 3)
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_8() {
1790        // Larger bypass_bits = 8
1791        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1792        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 8)
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, 8)
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_multiple_bypasses() {
1808        // Multiple bypass values in one stream: both -2 and 10 need bypass
1809        // dist 0: offset=1 sentinel=3, dist 1: offset=2 sentinel=5
1810        // value -2 (dist 0) => adjusted=-1 => bypass (negative)
1811        // value 10 (dist 1) => adjusted=12 => bypass (positive)
1812        let values = [-2i32, 10, 0, 1];
1813        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1814        enc.initialize(
1815            &PMF_LENGTHS,
1816            &PMF_OFFSETS,
1817            &PMF_TABLE,
1818            SYMBOL_BITS,
1819            BYPASS_BITS,
1820        )
1821        .unwrap();
1822        let mut encoded = Vec::new();
1823        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1824
1825        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1826        dec.initialize(
1827            &PMF_LENGTHS,
1828            &PMF_OFFSETS,
1829            &PMF_TABLE,
1830            SYMBOL_BITS,
1831            BYPASS_BITS,
1832        )
1833        .unwrap();
1834        let mut decoded = vec![0i32; 4];
1835        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1836        assert_eq!(decoded, values);
1837    }
1838
1839    #[test]
1840    fn test_encode_bypass_mixed_in_range_and_bypass() {
1841        // Mix of in-range and bypass values
1842        // dist 0: offset=1 sentinel=3, valid adjusted [0,2] => values [-1, 1]
1843        // dist 1: offset=2 sentinel=5, valid adjusted [0,4] => values [-2, 2]
1844        // value 0 in dist 0 => in-range, value 1 in dist 0 => in-range
1845        // value 5 in dist 1 => bypass (adjusted=7 > 4), value -3 in dist 1 => bypass (adjusted=-1 < 0)
1846        let values = [0i32, 5, 1, -3];
1847        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1848        enc.initialize(
1849            &PMF_LENGTHS,
1850            &PMF_OFFSETS,
1851            &PMF_TABLE,
1852            SYMBOL_BITS,
1853            BYPASS_BITS,
1854        )
1855        .unwrap();
1856        let mut encoded = Vec::new();
1857        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1858
1859        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1860        dec.initialize(
1861            &PMF_LENGTHS,
1862            &PMF_OFFSETS,
1863            &PMF_TABLE,
1864            SYMBOL_BITS,
1865            BYPASS_BITS,
1866        )
1867        .unwrap();
1868        let mut decoded = vec![0i32; 4];
1869        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1870        assert_eq!(decoded, values);
1871    }
1872
1873    #[test]
1874    fn test_encode_bypass_negative_outlier_at_boundary() {
1875        // Negative outlier at boundary: -1 - sentinel (very negative)
1876        // dist 0: offset=1, sentinel=3, value=-10 => adjusted=-9 => bypass
1877        // (-9 < 0) => bypass_value = 2*9-1 = 17, decode: 17>>1=8, -(8+1) = -9, -9+1 = -8... wait
1878        // decode bypass: negative flag set, half=8, symbol = -(8+1) = -9, -9 = -9+1 = -8...
1879        // Actually: symbol = -9, values[i] = symbol - offset = (-9) - 1 = -10 ✓
1880        let values = [-10i32, 1, 0, 1];
1881        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1882        enc.initialize(
1883            &PMF_LENGTHS,
1884            &PMF_OFFSETS,
1885            &PMF_TABLE,
1886            SYMBOL_BITS,
1887            BYPASS_BITS,
1888        )
1889        .unwrap();
1890        let mut encoded = Vec::new();
1891        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1892
1893        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1894        dec.initialize(
1895            &PMF_LENGTHS,
1896            &PMF_OFFSETS,
1897            &PMF_TABLE,
1898            SYMBOL_BITS,
1899            BYPASS_BITS,
1900        )
1901        .unwrap();
1902        let mut decoded = vec![0i32; 4];
1903        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1904        assert_eq!(decoded, values);
1905    }
1906
1907    #[test]
1908    fn test_encode_bypass_large_positive_outlier() {
1909        // Large positive outlier
1910        let values = [10000i32, 1, 0, 1];
1911        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1912        enc.initialize(
1913            &PMF_LENGTHS,
1914            &PMF_OFFSETS,
1915            &PMF_TABLE,
1916            SYMBOL_BITS,
1917            BYPASS_BITS,
1918        )
1919        .unwrap();
1920        let mut encoded = Vec::new();
1921        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1922
1923        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1924        dec.initialize(
1925            &PMF_LENGTHS,
1926            &PMF_OFFSETS,
1927            &PMF_TABLE,
1928            SYMBOL_BITS,
1929            BYPASS_BITS,
1930        )
1931        .unwrap();
1932        let mut decoded = vec![0i32; 4];
1933        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1934        assert_eq!(decoded, values);
1935    }
1936
1937    // -----------------------------------------------------------------------
1938    // Issue 5: extreme value overflow protection
1939    // -----------------------------------------------------------------------
1940
1941    #[test]
1942    fn test_encode_bypass_extreme_negative_i32_min_plus_one() {
1943        // i32::MIN + 1 with offset 1 gives adjusted = i32::MIN + 2 = -2147483646
1944        // checked_neg of that gives 2147483646, no overflow.
1945        // This should succeed (no overflow).
1946        let values = [i32::MIN + 1, 1, 0, 1];
1947        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1948        enc.initialize(
1949            &PMF_LENGTHS,
1950            &PMF_OFFSETS,
1951            &PMF_TABLE,
1952            SYMBOL_BITS,
1953            BYPASS_BITS,
1954        )
1955        .unwrap();
1956        let mut encoded = Vec::new();
1957        let result = enc.encode(&INDICES, &values, &mut encoded);
1958        // Must not panic — should succeed or return InvalidParams gracefully
1959        assert!(result.is_ok() || result == Err(EntropyError::InvalidParams));
1960    }
1961
1962    #[test]
1963    fn test_encode_bypass_extreme_positive_i32_max() {
1964        // i32::MAX with offset could cause overflow in checked_add -> InvalidParams
1965        let values = [i32::MAX, 1, 0, 1];
1966        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1967        enc.initialize(
1968            &PMF_LENGTHS,
1969            &PMF_OFFSETS,
1970            &PMF_TABLE,
1971            SYMBOL_BITS,
1972            BYPASS_BITS,
1973        )
1974        .unwrap();
1975        let mut encoded = Vec::new();
1976        let result = enc.encode(&INDICES, &values, &mut encoded);
1977        assert_eq!(result, Err(EntropyError::InvalidParams));
1978    }
1979}