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
1170impl<S: RansParams> Default for EntropyDecoder<S> {
1171    fn default() -> Self {
1172        Self::new()
1173    }
1174}
1175
1176// ---------------------------------------------------------------------------
1177// Tests
1178// ---------------------------------------------------------------------------
1179
1180#[cfg(test)]
1181mod tests {
1182    use super::*;
1183
1184    // Reference test case from test_msrtc_rans.py:
1185    //   PMF_LENGTHS = [4, 6]
1186    //   PMF_OFFSETS = [1, 2]
1187    //   PMF_TABLE   = [1, 3, 1, 1, 1, 3, 5, 3, 1, 1]
1188    //   INDICES     = [0, 1, 0, 1]
1189    //   VALUES      = [-2, 1, 0, 1]
1190    //   SYMBOL_BITS = 16
1191    //   BYPASS_BITS = 4
1192    //
1193    // Reference bitstreams from upstream oracle (EntropyCoder.cpp):
1194    //   RansByte: hex = "0500bd040001a10003000b00"
1195    //   Rans64:   hex = "0500a1bd04000000110a002f03000300"
1196
1197    const PMF_LENGTHS: [i32; 2] = [4, 6];
1198    const PMF_OFFSETS: [i32; 2] = [1, 2];
1199    const PMF_TABLE: [i32; 10] = [1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
1200    const INDICES: [i32; 4] = [0, 1, 0, 1];
1201    const VALUES: [i32; 4] = [-2, 1, 0, 1];
1202    const SYMBOL_BITS: u32 = 16;
1203    const BYPASS_BITS: u32 = 4;
1204
1205    const REF_HEX_BYTE: &str = "0500bd040001a10003000b00";
1206    const REF_HEX_64: &str = "0500a1bd04000000110a002f03000300";
1207
1208    fn hex_decode(hex: &str) -> Vec<u8> {
1209        (0..hex.len())
1210            .step_by(2)
1211            .map(|i| u8::from_str_radix(&hex[i..i + 2], 16).unwrap())
1212            .collect()
1213    }
1214
1215    #[test]
1216    fn test_encoder_byte_initialize() {
1217        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1218        assert!(
1219            enc.initialize(
1220                &PMF_LENGTHS,
1221                &PMF_OFFSETS,
1222                &PMF_TABLE,
1223                SYMBOL_BITS,
1224                BYPASS_BITS
1225            )
1226            .is_ok()
1227        );
1228    }
1229
1230    #[test]
1231    fn test_encoder_64_initialize() {
1232        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1233        assert!(
1234            enc.initialize(
1235                &PMF_LENGTHS,
1236                &PMF_OFFSETS,
1237                &PMF_TABLE,
1238                SYMBOL_BITS,
1239                BYPASS_BITS
1240            )
1241            .is_ok()
1242        );
1243    }
1244
1245    #[test]
1246    fn test_encoder_rejects_invalid_pmf() {
1247        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1248        // Mismatched lengths/offsets
1249        assert_eq!(
1250            enc.initialize(&[4], &[1, 2], &PMF_TABLE, SYMBOL_BITS, BYPASS_BITS),
1251            Err(EntropyError::InvalidPmf)
1252        );
1253    }
1254
1255    #[test]
1256    fn test_encoder_rejects_invalid_params() {
1257        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1258        // symbol_bits < 2
1259        assert_eq!(
1260            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 1, BYPASS_BITS),
1261            Err(EntropyError::InvalidParams)
1262        );
1263    }
1264
1265    #[test]
1266    fn test_encoder_byte_rejects_length_leq_one() {
1267        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1268        // Length <= 1 is invalid (need tail mass for bypass)
1269        assert_eq!(
1270            enc.initialize(&[1, 6], &[1, 2], &PMF_TABLE, SYMBOL_BITS, BYPASS_BITS),
1271            Err(EntropyError::InvalidPmf)
1272        );
1273    }
1274
1275    #[test]
1276    fn test_encode_byte_matches_reference() {
1277        // Encodes values=[-2, 1, 0, 1] with RansByte.
1278        // Value -2 (index 0, offset 1 => adjusted=-1) triggers bypass.
1279        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1280        enc.initialize(
1281            &PMF_LENGTHS,
1282            &PMF_OFFSETS,
1283            &PMF_TABLE,
1284            SYMBOL_BITS,
1285            BYPASS_BITS,
1286        )
1287        .unwrap();
1288
1289        let mut buffer = Vec::new();
1290        enc.encode(&INDICES, &VALUES, &mut buffer).unwrap();
1291
1292        let expected = hex_decode(REF_HEX_BYTE);
1293        assert_eq!(
1294            buffer, expected,
1295            "RansByte encode output does not match reference hex"
1296        );
1297    }
1298
1299    #[test]
1300    fn test_encode_64_matches_reference() {
1301        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1302        enc.initialize(
1303            &PMF_LENGTHS,
1304            &PMF_OFFSETS,
1305            &PMF_TABLE,
1306            SYMBOL_BITS,
1307            BYPASS_BITS,
1308        )
1309        .unwrap();
1310
1311        let mut buffer = Vec::new();
1312        enc.encode(&INDICES, &VALUES, &mut buffer).unwrap();
1313
1314        let expected = hex_decode(REF_HEX_64);
1315        assert_eq!(
1316            buffer, expected,
1317            "Rans64 encode output does not match reference hex"
1318        );
1319    }
1320
1321    #[test]
1322    fn test_encode_in_range_values_no_bypass() {
1323        // Values that are all in-range (no bypass):
1324        // Dist 0: offset=1, sentinel=3, valid adjusted: [0,2] => value: [-1, 1]
1325        // Dist 1: offset=2, sentinel=5, valid adjusted: [0,4] => value: [-2, 2]
1326        let in_range_values = [1i32, 1, 0, 1];
1327        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1328        enc.initialize(
1329            &PMF_LENGTHS,
1330            &PMF_OFFSETS,
1331            &PMF_TABLE,
1332            SYMBOL_BITS,
1333            BYPASS_BITS,
1334        )
1335        .unwrap();
1336
1337        let mut buffer = Vec::new();
1338        let result = enc.encode(&INDICES, &in_range_values, &mut buffer);
1339        assert!(result.is_ok(), "encode should succeed: {:?}", result);
1340        assert!(!buffer.is_empty(), "encoded buffer should not be empty");
1341    }
1342
1343    #[test]
1344    fn test_decode_byte_roundtrip_in_range() {
1345        // Verify roundtrip encode-decode with in-range values (no bypass).
1346        let values = [1i32, 1, 0, 1];
1347
1348        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1349        enc.initialize(
1350            &PMF_LENGTHS,
1351            &PMF_OFFSETS,
1352            &PMF_TABLE,
1353            SYMBOL_BITS,
1354            BYPASS_BITS,
1355        )
1356        .unwrap();
1357
1358        let mut encoded = Vec::new();
1359        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1360
1361        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1362        dec.initialize(
1363            &PMF_LENGTHS,
1364            &PMF_OFFSETS,
1365            &PMF_TABLE,
1366            SYMBOL_BITS,
1367            BYPASS_BITS,
1368        )
1369        .unwrap();
1370
1371        let mut decoded = vec![0i32; values.len()];
1372        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1373
1374        assert_eq!(
1375            decoded, values,
1376            "roundtrip decode should match original values"
1377        );
1378    }
1379
1380    #[test]
1381    fn test_decode_64_roundtrip_in_range() {
1382        let values = [1i32, 1, 0, 1];
1383
1384        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1385        enc.initialize(
1386            &PMF_LENGTHS,
1387            &PMF_OFFSETS,
1388            &PMF_TABLE,
1389            SYMBOL_BITS,
1390            BYPASS_BITS,
1391        )
1392        .unwrap();
1393
1394        let mut encoded = Vec::new();
1395        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1396
1397        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1398        dec.initialize(
1399            &PMF_LENGTHS,
1400            &PMF_OFFSETS,
1401            &PMF_TABLE,
1402            SYMBOL_BITS,
1403            BYPASS_BITS,
1404        )
1405        .unwrap();
1406
1407        let mut decoded = vec![0i32; values.len()];
1408        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1409
1410        assert_eq!(
1411            decoded, values,
1412            "Rans64 roundtrip decode should match original values"
1413        );
1414    }
1415
1416    #[test]
1417    fn test_decode_byte_roundtrip_bypass() {
1418        // Roundtrip with values that require bypass (value=-2 is out-of-range).
1419        let values = [-2i32, 1, 0, 1];
1420
1421        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1422        enc.initialize(
1423            &PMF_LENGTHS,
1424            &PMF_OFFSETS,
1425            &PMF_TABLE,
1426            SYMBOL_BITS,
1427            BYPASS_BITS,
1428        )
1429        .unwrap();
1430
1431        let mut encoded = Vec::new();
1432        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1433
1434        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1435        dec.initialize(
1436            &PMF_LENGTHS,
1437            &PMF_OFFSETS,
1438            &PMF_TABLE,
1439            SYMBOL_BITS,
1440            BYPASS_BITS,
1441        )
1442        .unwrap();
1443
1444        let mut decoded = vec![0i32; values.len()];
1445        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1446
1447        assert_eq!(
1448            decoded, values,
1449            "bypass roundtrip decode should match original values"
1450        );
1451    }
1452
1453    #[test]
1454    fn test_decode_64_roundtrip_bypass() {
1455        let values = [-2i32, 1, 0, 1];
1456
1457        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1458        enc.initialize(
1459            &PMF_LENGTHS,
1460            &PMF_OFFSETS,
1461            &PMF_TABLE,
1462            SYMBOL_BITS,
1463            BYPASS_BITS,
1464        )
1465        .unwrap();
1466
1467        let mut encoded = Vec::new();
1468        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1469
1470        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1471        dec.initialize(
1472            &PMF_LENGTHS,
1473            &PMF_OFFSETS,
1474            &PMF_TABLE,
1475            SYMBOL_BITS,
1476            BYPASS_BITS,
1477        )
1478        .unwrap();
1479
1480        let mut decoded = vec![0i32; values.len()];
1481        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1482
1483        assert_eq!(
1484            decoded, values,
1485            "Rans64 bypass roundtrip decode should match original values"
1486        );
1487    }
1488
1489    // -----------------------------------------------------------------------
1490    // Issue 2d: scale-32 safe / reject tests
1491    // -----------------------------------------------------------------------
1492
1493    #[test]
1494    fn test_encoder_64_symbol_bits_31_accepted() {
1495        // Rans64 symbol_bits=31 is within max_safe_bits (31)
1496        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1497        assert!(
1498            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 31, BYPASS_BITS)
1499                .is_ok()
1500        );
1501    }
1502
1503    #[test]
1504    fn test_encoder_64_symbol_bits_32_rejected() {
1505        // Rans64 symbol_bits=32 exceeds max_safe_bits (31), must not panic
1506        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1507        assert_eq!(
1508            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, 32, BYPASS_BITS),
1509            Err(EntropyError::InvalidParams)
1510        );
1511    }
1512
1513    #[test]
1514    fn test_encoder_64_bypass_bits_32_rejected() {
1515        // Rans64 bypass_bits=32 exceeds max_safe_bits (31), must not panic
1516        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1517        assert_eq!(
1518            enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 32),
1519            Err(EntropyError::InvalidParams)
1520        );
1521    }
1522
1523    // -----------------------------------------------------------------------
1524    // Issue 3: misaligned Rans64 streams rejected
1525    // -----------------------------------------------------------------------
1526
1527    #[test]
1528    fn test_decode_64_rejects_misaligned_1_extra_byte() {
1529        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1530        enc.initialize(
1531            &PMF_LENGTHS,
1532            &PMF_OFFSETS,
1533            &PMF_TABLE,
1534            SYMBOL_BITS,
1535            BYPASS_BITS,
1536        )
1537        .unwrap();
1538        let mut encoded = Vec::new();
1539        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1540            .unwrap();
1541
1542        // Append 1 extra byte to make it misaligned
1543        let mut misaligned = encoded.clone();
1544        misaligned.push(0xAB);
1545
1546        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1547        dec.initialize(
1548            &PMF_LENGTHS,
1549            &PMF_OFFSETS,
1550            &PMF_TABLE,
1551            SYMBOL_BITS,
1552            BYPASS_BITS,
1553        )
1554        .unwrap();
1555        let mut decoded = vec![0i32; 4];
1556        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1557        assert_eq!(result, Err(EntropyError::InvalidStream));
1558    }
1559
1560    #[test]
1561    fn test_decode_64_rejects_misaligned_2_extra_bytes() {
1562        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1563        enc.initialize(
1564            &PMF_LENGTHS,
1565            &PMF_OFFSETS,
1566            &PMF_TABLE,
1567            SYMBOL_BITS,
1568            BYPASS_BITS,
1569        )
1570        .unwrap();
1571        let mut encoded = Vec::new();
1572        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1573            .unwrap();
1574
1575        let mut misaligned = encoded.clone();
1576        misaligned.extend_from_slice(&[0xAB, 0xCD]);
1577
1578        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1579        dec.initialize(
1580            &PMF_LENGTHS,
1581            &PMF_OFFSETS,
1582            &PMF_TABLE,
1583            SYMBOL_BITS,
1584            BYPASS_BITS,
1585        )
1586        .unwrap();
1587        let mut decoded = vec![0i32; 4];
1588        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1589        assert_eq!(result, Err(EntropyError::InvalidStream));
1590    }
1591
1592    #[test]
1593    fn test_decode_64_rejects_misaligned_3_extra_bytes() {
1594        let mut enc: EntropyEncoder<Rans64> = EntropyEncoder::new();
1595        enc.initialize(
1596            &PMF_LENGTHS,
1597            &PMF_OFFSETS,
1598            &PMF_TABLE,
1599            SYMBOL_BITS,
1600            BYPASS_BITS,
1601        )
1602        .unwrap();
1603        let mut encoded = Vec::new();
1604        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1605            .unwrap();
1606
1607        let mut misaligned = encoded.clone();
1608        misaligned.extend_from_slice(&[0xAB, 0xCD, 0xEF]);
1609
1610        let mut dec: EntropyDecoder<Rans64> = EntropyDecoder::new();
1611        dec.initialize(
1612            &PMF_LENGTHS,
1613            &PMF_OFFSETS,
1614            &PMF_TABLE,
1615            SYMBOL_BITS,
1616            BYPASS_BITS,
1617        )
1618        .unwrap();
1619        let mut decoded = vec![0i32; 4];
1620        let result = dec.decode(&mut decoded, &INDICES, &misaligned);
1621        assert_eq!(result, Err(EntropyError::InvalidStream));
1622    }
1623
1624    #[test]
1625    fn test_decode_byte_accepts_extra_bytes() {
1626        // RansByte has byte-level alignment, extra bytes should not be rejected
1627        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1628        enc.initialize(
1629            &PMF_LENGTHS,
1630            &PMF_OFFSETS,
1631            &PMF_TABLE,
1632            SYMBOL_BITS,
1633            BYPASS_BITS,
1634        )
1635        .unwrap();
1636        let mut encoded = Vec::new();
1637        enc.encode(&INDICES, &[1i32, 1, 0, 1], &mut encoded)
1638            .unwrap();
1639
1640        // Append extra bytes
1641        let mut extended = encoded.clone();
1642        extended.extend_from_slice(&[0xAB, 0xCD]);
1643
1644        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1645        dec.initialize(
1646            &PMF_LENGTHS,
1647            &PMF_OFFSETS,
1648            &PMF_TABLE,
1649            SYMBOL_BITS,
1650            BYPASS_BITS,
1651        )
1652        .unwrap();
1653        let mut decoded = vec![0i32; 4];
1654        // This may fail because the decoder checks EOF, but it should NOT
1655        // be rejected for misalignment
1656        let _ = dec.decode(&mut decoded, &INDICES, &extended);
1657        // We don't assert success or failure — we just assert no panic
1658    }
1659
1660    // -----------------------------------------------------------------------
1661    // Issue 4: Expanded bypass coverage
1662    // -----------------------------------------------------------------------
1663
1664    #[test]
1665    fn test_encode_bypass_positive_outlier() {
1666        // Value > sentinel (positive outlier): dist 0 sentinel=3, offset=1,
1667        // value=10 => adjusted=11 > 3 => bypass (8/2=4 above sentinel => value 4+3=7-1=6...
1668        // actually: adjusted=11, sentinel=3, bypass_value = 2*(11-3) = 16
1669        // decode: 16>>1=8, 8+3=11-1=10 ✓
1670        let values = [10i32, 1, 0, 1];
1671        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1672        enc.initialize(
1673            &PMF_LENGTHS,
1674            &PMF_OFFSETS,
1675            &PMF_TABLE,
1676            SYMBOL_BITS,
1677            BYPASS_BITS,
1678        )
1679        .unwrap();
1680        let mut encoded = Vec::new();
1681        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1682
1683        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1684        dec.initialize(
1685            &PMF_LENGTHS,
1686            &PMF_OFFSETS,
1687            &PMF_TABLE,
1688            SYMBOL_BITS,
1689            BYPASS_BITS,
1690        )
1691        .unwrap();
1692        let mut decoded = vec![0i32; 4];
1693        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1694        assert_eq!(decoded, values);
1695    }
1696
1697    #[test]
1698    fn test_encode_bypass_multi_digit_value() {
1699        // Value requiring multiple bypassBits-sized chunks (bypass_bits=4)
1700        // Large bypass value: dist 0 sentinel=3, offset=1, value=200 => adjusted=201
1701        // bypass_value = 2*(201-3) = 396 = 0x18C, needs multiple 4-bit chunks
1702        let values = [200i32, 1, 0, 1];
1703        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1704        enc.initialize(
1705            &PMF_LENGTHS,
1706            &PMF_OFFSETS,
1707            &PMF_TABLE,
1708            SYMBOL_BITS,
1709            BYPASS_BITS,
1710        )
1711        .unwrap();
1712        let mut encoded = Vec::new();
1713        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1714
1715        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1716        dec.initialize(
1717            &PMF_LENGTHS,
1718            &PMF_OFFSETS,
1719            &PMF_TABLE,
1720            SYMBOL_BITS,
1721            BYPASS_BITS,
1722        )
1723        .unwrap();
1724        let mut decoded = vec![0i32; 4];
1725        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1726        assert_eq!(decoded, values);
1727    }
1728
1729    #[test]
1730    fn test_encode_bypass_bits_2() {
1731        // Minimum bypass_bits = 2
1732        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1733        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 2)
1734            .unwrap();
1735        let values = [10i32, 1, 0, 1]; // value 10 => bypass
1736        let mut encoded = Vec::new();
1737        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1738
1739        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1740        dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 2)
1741            .unwrap();
1742        let mut decoded = vec![0i32; 4];
1743        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1744        assert_eq!(decoded, values);
1745    }
1746
1747    #[test]
1748    fn test_encode_bypass_bits_3() {
1749        // Odd bypass_bits = 3
1750        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1751        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 3)
1752            .unwrap();
1753        let values = [10i32, 1, 0, 1];
1754        let mut encoded = Vec::new();
1755        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1756
1757        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1758        dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 3)
1759            .unwrap();
1760        let mut decoded = vec![0i32; 4];
1761        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1762        assert_eq!(decoded, values);
1763    }
1764
1765    #[test]
1766    fn test_encode_bypass_bits_8() {
1767        // Larger bypass_bits = 8
1768        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1769        enc.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 8)
1770            .unwrap();
1771        let values = [10i32, 1, 0, 1];
1772        let mut encoded = Vec::new();
1773        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1774
1775        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1776        dec.initialize(&PMF_LENGTHS, &PMF_OFFSETS, &PMF_TABLE, SYMBOL_BITS, 8)
1777            .unwrap();
1778        let mut decoded = vec![0i32; 4];
1779        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1780        assert_eq!(decoded, values);
1781    }
1782
1783    #[test]
1784    fn test_encode_bypass_multiple_bypasses() {
1785        // Multiple bypass values in one stream: both -2 and 10 need bypass
1786        // dist 0: offset=1 sentinel=3, dist 1: offset=2 sentinel=5
1787        // value -2 (dist 0) => adjusted=-1 => bypass (negative)
1788        // value 10 (dist 1) => adjusted=12 => bypass (positive)
1789        let values = [-2i32, 10, 0, 1];
1790        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1791        enc.initialize(
1792            &PMF_LENGTHS,
1793            &PMF_OFFSETS,
1794            &PMF_TABLE,
1795            SYMBOL_BITS,
1796            BYPASS_BITS,
1797        )
1798        .unwrap();
1799        let mut encoded = Vec::new();
1800        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1801
1802        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1803        dec.initialize(
1804            &PMF_LENGTHS,
1805            &PMF_OFFSETS,
1806            &PMF_TABLE,
1807            SYMBOL_BITS,
1808            BYPASS_BITS,
1809        )
1810        .unwrap();
1811        let mut decoded = vec![0i32; 4];
1812        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1813        assert_eq!(decoded, values);
1814    }
1815
1816    #[test]
1817    fn test_encode_bypass_mixed_in_range_and_bypass() {
1818        // Mix of in-range and bypass values
1819        // dist 0: offset=1 sentinel=3, valid adjusted [0,2] => values [-1, 1]
1820        // dist 1: offset=2 sentinel=5, valid adjusted [0,4] => values [-2, 2]
1821        // value 0 in dist 0 => in-range, value 1 in dist 0 => in-range
1822        // value 5 in dist 1 => bypass (adjusted=7 > 4), value -3 in dist 1 => bypass (adjusted=-1 < 0)
1823        let values = [0i32, 5, 1, -3];
1824        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1825        enc.initialize(
1826            &PMF_LENGTHS,
1827            &PMF_OFFSETS,
1828            &PMF_TABLE,
1829            SYMBOL_BITS,
1830            BYPASS_BITS,
1831        )
1832        .unwrap();
1833        let mut encoded = Vec::new();
1834        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1835
1836        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1837        dec.initialize(
1838            &PMF_LENGTHS,
1839            &PMF_OFFSETS,
1840            &PMF_TABLE,
1841            SYMBOL_BITS,
1842            BYPASS_BITS,
1843        )
1844        .unwrap();
1845        let mut decoded = vec![0i32; 4];
1846        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1847        assert_eq!(decoded, values);
1848    }
1849
1850    #[test]
1851    fn test_encode_bypass_negative_outlier_at_boundary() {
1852        // Negative outlier at boundary: -1 - sentinel (very negative)
1853        // dist 0: offset=1, sentinel=3, value=-10 => adjusted=-9 => bypass
1854        // (-9 < 0) => bypass_value = 2*9-1 = 17, decode: 17>>1=8, -(8+1) = -9, -9+1 = -8... wait
1855        // decode bypass: negative flag set, half=8, symbol = -(8+1) = -9, -9 = -9+1 = -8...
1856        // Actually: symbol = -9, values[i] = symbol - offset = (-9) - 1 = -10 ✓
1857        let values = [-10i32, 1, 0, 1];
1858        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1859        enc.initialize(
1860            &PMF_LENGTHS,
1861            &PMF_OFFSETS,
1862            &PMF_TABLE,
1863            SYMBOL_BITS,
1864            BYPASS_BITS,
1865        )
1866        .unwrap();
1867        let mut encoded = Vec::new();
1868        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1869
1870        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1871        dec.initialize(
1872            &PMF_LENGTHS,
1873            &PMF_OFFSETS,
1874            &PMF_TABLE,
1875            SYMBOL_BITS,
1876            BYPASS_BITS,
1877        )
1878        .unwrap();
1879        let mut decoded = vec![0i32; 4];
1880        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1881        assert_eq!(decoded, values);
1882    }
1883
1884    #[test]
1885    fn test_encode_bypass_large_positive_outlier() {
1886        // Large positive outlier
1887        let values = [10000i32, 1, 0, 1];
1888        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1889        enc.initialize(
1890            &PMF_LENGTHS,
1891            &PMF_OFFSETS,
1892            &PMF_TABLE,
1893            SYMBOL_BITS,
1894            BYPASS_BITS,
1895        )
1896        .unwrap();
1897        let mut encoded = Vec::new();
1898        enc.encode(&INDICES, &values, &mut encoded).unwrap();
1899
1900        let mut dec: EntropyDecoder<RansByte> = EntropyDecoder::new();
1901        dec.initialize(
1902            &PMF_LENGTHS,
1903            &PMF_OFFSETS,
1904            &PMF_TABLE,
1905            SYMBOL_BITS,
1906            BYPASS_BITS,
1907        )
1908        .unwrap();
1909        let mut decoded = vec![0i32; 4];
1910        dec.decode(&mut decoded, &INDICES, &encoded).unwrap();
1911        assert_eq!(decoded, values);
1912    }
1913
1914    // -----------------------------------------------------------------------
1915    // Issue 5: extreme value overflow protection
1916    // -----------------------------------------------------------------------
1917
1918    #[test]
1919    fn test_encode_bypass_extreme_negative_i32_min_plus_one() {
1920        // i32::MIN + 1 with offset 1 gives adjusted = i32::MIN + 2 = -2147483646
1921        // checked_neg of that gives 2147483646, no overflow.
1922        // This should succeed (no overflow).
1923        let values = [i32::MIN + 1, 1, 0, 1];
1924        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1925        enc.initialize(
1926            &PMF_LENGTHS,
1927            &PMF_OFFSETS,
1928            &PMF_TABLE,
1929            SYMBOL_BITS,
1930            BYPASS_BITS,
1931        )
1932        .unwrap();
1933        let mut encoded = Vec::new();
1934        let result = enc.encode(&INDICES, &values, &mut encoded);
1935        // Must not panic — should succeed or return InvalidParams gracefully
1936        assert!(result.is_ok() || result == Err(EntropyError::InvalidParams));
1937    }
1938
1939    #[test]
1940    fn test_encode_bypass_extreme_positive_i32_max() {
1941        // i32::MAX with offset could cause overflow in checked_add -> InvalidParams
1942        let values = [i32::MAX, 1, 0, 1];
1943        let mut enc: EntropyEncoder<RansByte> = EntropyEncoder::new();
1944        enc.initialize(
1945            &PMF_LENGTHS,
1946            &PMF_OFFSETS,
1947            &PMF_TABLE,
1948            SYMBOL_BITS,
1949            BYPASS_BITS,
1950        )
1951        .unwrap();
1952        let mut encoded = Vec::new();
1953        let result = enc.encode(&INDICES, &values, &mut encoded);
1954        assert_eq!(result, Err(EntropyError::InvalidParams));
1955    }
1956}