Skip to main content

opus_rs/
range_coder.rs

1use crate::fixedvec::FixedVec;
2
3pub const EC_SYM_BITS: u32 = 8;
4pub const EC_CODE_BITS: u32 = 32;
5pub const EC_SYM_MAX: u32 = (1 << EC_SYM_BITS) - 1;
6pub const EC_CODE_SHIFT: u32 = EC_CODE_BITS - EC_SYM_BITS - 1;
7pub const EC_CODE_TOP: u32 = 1 << (EC_CODE_BITS - 1);
8pub const EC_CODE_BOT: u32 = EC_CODE_TOP >> EC_SYM_BITS;
9pub const EC_CODE_EXTRA: u32 = (EC_CODE_BITS - 2) % EC_SYM_BITS + 1;
10pub const BITRES: i32 = 3;
11
12/// Maximum range-coder buffer capacity. Opus packets are bounded by 1276 bytes
13/// (RFC 6716 §3.1); standalone CELT usage (e.g. tests) may use up to 2048, so the
14/// heap-free buffer is sized for that. The Opus encoder additionally clamps its
15/// per-frame budget to `OPUS_MAX_PACKET_BYTES = 1276` (see `lib.rs`).
16pub const RANGE_BUF_MAX: usize = 2048;
17
18#[macro_export]
19macro_rules! tell_frac_inline {
20    ($rc:expr) => {{
21        static CORRECTION: [u32; 8] = [35733, 38967, 42495, 46340, 50535, 55109, 60097, 65535];
22        let nbits = $rc.nbits_total << BITRES;
23        let l = 32 - $rc.rng.leading_zeros() as i32;
24        let r = $rc.rng >> (l - 16);
25        let b = (r >> 12).wrapping_sub(8);
26
27        let correction = unsafe { *CORRECTION.get_unchecked(b as usize) };
28        let b = b + (r > correction) as u32;
29        nbits - (l << 3) - b as i32
30    }};
31}
32
33#[derive(Clone)]
34pub struct RangeCoder {
35    pub buf: FixedVec<u8, RANGE_BUF_MAX>,
36    pub storage: u32,
37    pub end_offs: u32,
38    pub end_window: u32,
39    pub nend_bits: i32,
40    pub nbits_total: i32,
41    pub offs: u32,
42    pub rng: u32,
43    pub val: u32,
44    pub ext: u32,
45    pub rem: i32,
46    pub error: i32,
47}
48
49impl RangeCoder {
50    pub fn new_encoder(size: u32) -> Self {
51        let size = (size as usize).min(RANGE_BUF_MAX).max(1);
52        let buf = FixedVec::from_value(0u8, size);
53        RangeCoder {
54            buf,
55            storage: size as u32,
56            end_offs: 0,
57            end_window: 0,
58            nend_bits: 0,
59            nbits_total: 33,
60            offs: 0,
61            rng: 1 << 31,
62            val: 0,
63            ext: 0,
64            rem: -1,
65            error: 0,
66        }
67    }
68
69    #[inline]
70    pub fn reset_for_encode(&mut self, size: u32) {
71        let size = (size as usize).min(RANGE_BUF_MAX).max(1);
72        self.buf.resize(size, 0);
73        self.storage = size as u32;
74        self.end_offs = 0;
75        self.end_window = 0;
76        self.nend_bits = 0;
77        self.nbits_total = 33;
78        self.offs = 0;
79        self.rng = 1 << 31;
80        self.val = 0;
81        self.ext = 0;
82        self.rem = -1;
83        self.error = 0;
84    }
85
86    pub fn new_decoder(data: &[u8]) -> Self {
87        let n = data.len().min(RANGE_BUF_MAX);
88        let storage = n as u32;
89        let buf = FixedVec::from_slice(&data[..n]);
90        let mut rc = RangeCoder {
91            buf,
92            storage,
93            end_offs: 0,
94            end_window: 0,
95            nend_bits: 0,
96            nbits_total: (EC_CODE_BITS + 1
97                - ((EC_CODE_BITS - EC_CODE_EXTRA) / EC_SYM_BITS) * EC_SYM_BITS)
98                as i32,
99            offs: 0,
100            rng: 1 << EC_CODE_EXTRA,
101            val: 0,
102            ext: 0,
103            rem: 0,
104            error: 0,
105        };
106
107        rc.rem = rc.read_byte() as i32;
108        rc.val = rc
109            .rng
110            .wrapping_sub(1)
111            .wrapping_sub(rc.rem as u32 >> (EC_SYM_BITS - EC_CODE_EXTRA));
112
113        rc.normalize_decoder();
114        rc
115    }
116
117    #[inline(always)]
118    fn normalize_decoder(&mut self) {
119        let mut guard = 0u32;
120        while self.rng <= EC_CODE_BOT {
121            guard += 1;
122            if guard > 100 {
123                self.error = 1;
124                self.rng = EC_CODE_BOT + 1;
125                break;
126            }
127            self.nbits_total += EC_SYM_BITS as i32;
128            self.rng <<= EC_SYM_BITS;
129
130            let sym = self.rem;
131            self.rem = self.read_byte() as i32;
132
133            let combined_sym = ((sym << EC_SYM_BITS) | self.rem) >> (EC_SYM_BITS - EC_CODE_EXTRA);
134            self.val = (self.val << EC_SYM_BITS).wrapping_add(EC_SYM_MAX & !combined_sym as u32)
135                & (EC_CODE_TOP - 1);
136        }
137    }
138
139    fn read_byte(&mut self) -> u8 {
140        if self.offs < self.storage {
141            let b = self.buf[self.offs as usize];
142            self.offs += 1;
143            b
144        } else {
145            0
146        }
147    }
148
149    #[inline(always)]
150    pub fn enc_uint(&mut self, fl: u32, ft: u32) {
151        if ft > 1 {
152            let ft_minus_1 = ft - 1;
153            let ftb = 32 - ft_minus_1.leading_zeros() as i32;
154            if ftb > 8 {
155                let s = ftb - 8;
156                let fl_low = fl & ((1u32 << s) - 1);
157                let fl = fl >> s;
158                let ft = (ft_minus_1 >> s) + 1;
159                self.encode(fl, fl.wrapping_add(1), ft);
160                self.enc_bits(fl_low, s as u32);
161            } else {
162                self.encode(fl, fl.wrapping_add(1), ft);
163            }
164        }
165    }
166
167    #[inline(always)]
168    pub fn dec_uint(&mut self, ft: u32) -> u32 {
169        if ft > 1 {
170            let ft_minus_1 = ft - 1;
171            let ftb = 32 - ft_minus_1.leading_zeros() as i32;
172            if ftb > 8 {
173                let s = ftb - 8;
174                let ft = (ft_minus_1 >> s) + 1;
175                let fs = self.decode(ft);
176                self.update(fs, fs.wrapping_add(1), ft);
177                let r = self.dec_bits(s as u32);
178                (fs << s) | r
179            } else {
180                let fs = self.decode(ft);
181                self.update(fs, fs.wrapping_add(1), ft);
182                fs
183            }
184        } else {
185            0
186        }
187    }
188
189    #[inline(always)]
190    pub fn enc_bits(&mut self, val: u32, bits: u32) {
191        if bits == 0 {
192            return;
193        }
194        let mut window = self.end_window;
195        let mut used = self.nend_bits;
196        if (used as u32) + bits > EC_CODE_BITS {
197            while used >= EC_SYM_BITS as i32 {
198                self.write_byte_at_end((window & EC_SYM_MAX) as u8);
199                window >>= EC_SYM_BITS;
200                used -= EC_SYM_BITS as i32;
201            }
202        }
203        window |= (val & ((1 << bits) - 1)) << used;
204        used += bits as i32;
205        self.end_window = window;
206        self.nend_bits = used;
207        self.nbits_total += bits as i32;
208    }
209
210    pub fn pad_to_bits(&mut self, target_bits: i32) {
211        let remaining = target_bits - self.nbits_total;
212        if remaining <= 0 {
213            return;
214        }
215        let mut remaining = remaining as u32;
216
217        let partial =
218            (EC_SYM_BITS - (self.nend_bits as u32 & (EC_SYM_BITS - 1))) & (EC_SYM_BITS - 1);
219        if partial > 0 && remaining >= partial {
220            self.enc_bits(0, partial.min(remaining));
221            remaining -= partial.min(remaining);
222        }
223
224        let full_bytes = remaining / EC_SYM_BITS;
225        if full_bytes > 0 {
226            let available = self.storage - self.offs - self.end_offs;
227            let write_count = full_bytes.min(available);
228            if write_count > 0 {
229                let start = (self.storage - self.end_offs - write_count) as usize;
230                unsafe {
231                    core::ptr::write_bytes(
232                        self.buf.as_mut_ptr().add(start),
233                        0,
234                        write_count as usize,
235                    );
236                }
237                self.end_offs += write_count;
238            }
239            if write_count < full_bytes {
240                self.error = 1;
241            }
242            self.nbits_total += (full_bytes * EC_SYM_BITS) as i32;
243            remaining -= full_bytes * EC_SYM_BITS;
244        }
245
246        if remaining > 0 {
247            self.enc_bits(0, remaining);
248        }
249    }
250
251    pub fn dec_bits(&mut self, bits: u32) -> u32 {
252        if bits == 0 {
253            return 0;
254        }
255        let mut window = self.end_window;
256        let mut used = self.nend_bits;
257        if used < bits as i32 {
258            loop {
259                let byte = if self.end_offs < self.storage {
260                    self.end_offs += 1;
261                    self.buf[(self.storage - self.end_offs) as usize]
262                } else {
263                    0
264                };
265                window |= (byte as u32) << used;
266                used += 8;
267                if used > 32 - 8 {
268                    break;
269                }
270            }
271        }
272        let ret = window & ((1 << bits) - 1);
273        self.end_window = window >> bits;
274        self.nend_bits = used - bits as i32;
275        self.nbits_total += bits as i32;
276        ret
277    }
278
279    #[inline(always)]
280    pub fn tell_frac(&self) -> i32 {
281        static CORRECTION: [u32; 8] = [35733, 38967, 42495, 46340, 50535, 55109, 60097, 65535];
282        let nbits = self.nbits_total << BITRES;
283        let l = 32 - self.rng.leading_zeros() as i32;
284        let r = self.rng >> (l - 16);
285        let b = (r >> 12).wrapping_sub(8);
286        let b = b + (r > CORRECTION[b as usize]) as u32;
287        nbits - (l << 3) - b as i32
288    }
289
290    /// Integer bit count matching C's `ec_tell()`: `nbits_total - EC_ILOG(rng)`.
291    #[inline(always)]
292    pub fn tell(&self) -> i32 {
293        self.nbits_total - (32 - self.rng.leading_zeros() as i32)
294    }
295
296    #[inline(always)]
297    pub fn tell_fast(&self) -> i32 {
298        self.nbits_total
299    }
300
301    /// Shrink the range coder buffer, moving end-coded bytes to the new end.
302    /// Equivalent to C's `ec_enc_shrink`.
303    pub fn shrink(&mut self, new_size: u32) {
304        debug_assert!(self.offs + self.end_offs <= new_size);
305        if self.end_offs > 0 {
306            let old_end_start = (self.storage - self.end_offs) as usize;
307            let old_end_end = self.storage as usize;
308            let new_end_start = (new_size - self.end_offs) as usize;
309            self.buf
310                .copy_within(old_end_start..old_end_end, new_end_start);
311        }
312        self.storage = new_size;
313    }
314
315    #[inline(always)]
316    fn write_byte(&mut self, value: u8) {
317        if self.offs + self.end_offs < self.storage {
318            unsafe {
319                *self.buf.get_unchecked_mut(self.offs as usize) = value;
320            }
321            self.offs += 1;
322        } else {
323            self.error = 1;
324        }
325    }
326
327    #[inline(always)]
328    fn carry_out(&mut self, c: i32) {
329        if c != EC_SYM_MAX as i32 {
330            let carry = c >> EC_SYM_BITS;
331            if self.rem >= 0 {
332                self.write_byte((self.rem + carry) as u8);
333            }
334            if self.ext > 0 {
335                let sym = (EC_SYM_MAX as i32 + carry) & EC_SYM_MAX as i32;
336
337                let ext = self.ext as usize;
338                for _j in 0..ext {
339                    self.write_byte(sym as u8);
340                }
341                self.ext = 0;
342            }
343            self.rem = c & EC_SYM_MAX as i32;
344        } else {
345            self.ext += 1;
346        }
347    }
348
349    #[inline(always)]
350    fn celt_udiv(n: u32, d: u32) -> u32 {
351        debug_assert!(d > 0);
352        n / d
353    }
354
355    #[inline(always)]
356    pub fn encode(&mut self, fl: u32, fh: u32, ft: u32) {
357        debug_assert!(ft > 0, "encode: ft must be > 0");
358        let r = Self::celt_udiv(self.rng, ft);
359        if fl > 0 {
360            self.val = self
361                .val
362                .wrapping_add(self.rng.wrapping_sub(r.wrapping_mul(ft.wrapping_sub(fl))));
363            self.rng = r.wrapping_mul(fh.wrapping_sub(fl));
364        } else {
365            self.rng = self.rng.wrapping_sub(r.wrapping_mul(ft.wrapping_sub(fh)));
366        }
367        self.normalize_encoder();
368    }
369
370    #[inline(always)]
371    fn normalize_encoder(&mut self) {
372        while self.rng <= EC_CODE_BOT {
373            let c = (self.val >> EC_CODE_SHIFT) as i32;
374            if c != EC_SYM_MAX as i32 {
375                let carry = c >> EC_SYM_BITS;
376                if self.rem >= 0 {
377                    if self.offs + self.end_offs < self.storage {
378                        unsafe {
379                            *self.buf.get_unchecked_mut(self.offs as usize) =
380                                (self.rem + carry) as u8;
381                        }
382                        self.offs += 1;
383                    } else {
384                        self.error = 1;
385                    }
386                }
387                if self.ext > 0 {
388                    let sym = (EC_SYM_MAX as i32 + carry) & EC_SYM_MAX as i32;
389                    let ext = self.ext as usize;
390                    for _j in 0..ext {
391                        if self.offs + self.end_offs < self.storage {
392                            unsafe {
393                                *self.buf.get_unchecked_mut(self.offs as usize) = sym as u8;
394                            }
395                            self.offs += 1;
396                        } else {
397                            self.error = 1;
398                        }
399                    }
400                    self.ext = 0;
401                }
402                self.rem = c & EC_SYM_MAX as i32;
403            } else {
404                self.ext += 1;
405            }
406            self.val = (self.val << EC_SYM_BITS) & (EC_CODE_TOP - 1);
407            self.rng <<= EC_SYM_BITS;
408            self.nbits_total = self.nbits_total.wrapping_add(EC_SYM_BITS as i32);
409        }
410    }
411
412    #[inline(always)]
413    pub fn encode_bit_logp(&mut self, val: bool, logp: u32) {
414        let s = self.rng >> logp;
415        let r = self.rng.wrapping_sub(s);
416        if val {
417            self.val = self.val.wrapping_add(r);
418            self.rng = s;
419        } else {
420            self.rng = r;
421        }
422        self.normalize_encoder();
423    }
424
425    #[inline(always)]
426    pub fn encode_icdf(&mut self, s: i32, icdf: &[u8], ftb: u32) {
427        let r = self.rng >> ftb;
428        if s > 0 {
429            let val = unsafe { *icdf.get_unchecked((s - 1) as usize) as u32 };
430            self.val = self
431                .val
432                .wrapping_add(self.rng.wrapping_sub(r.wrapping_mul(val)));
433            let lower = unsafe { *icdf.get_unchecked(s as usize) };
434            self.rng = r.wrapping_mul(val.wrapping_sub(lower as u32));
435        } else {
436            let val = unsafe { *icdf.get_unchecked(s as usize) as u32 };
437            self.rng = self.rng.wrapping_sub(r.wrapping_mul(val));
438        }
439        self.normalize_encoder();
440    }
441
442    #[inline(always)]
443    pub fn decode_bit_logp(&mut self, logp: u32) -> bool {
444        let s = self.rng >> logp;
445        let ret = self.val < s;
446        if !ret {
447            self.val = self.val.wrapping_sub(s);
448            self.rng = self.rng.wrapping_sub(s);
449        } else {
450            self.rng = s;
451        }
452        self.normalize_decoder();
453        ret
454    }
455
456    #[inline(always)]
457    pub fn decode_icdf(&mut self, icdf: &[u8], ftb: u32) -> i32 {
458        let mut s = self.rng;
459        let d = self.val;
460        let r = s >> ftb;
461        let mut ret = 0;
462        let mut t;
463
464        loop {
465            t = s;
466            s = r.wrapping_mul(icdf[ret] as u32);
467            ret += 1;
468            if d >= s {
469                break;
470            }
471        }
472
473        self.val = d.wrapping_sub(s);
474        self.rng = t.wrapping_sub(s);
475        self.normalize_decoder();
476        (ret - 1) as i32
477    }
478
479    #[inline(always)]
480    pub fn decode(&mut self, ft: u32) -> u32 {
481        let r = Self::celt_udiv(self.rng, ft);
482        self.ext = r;
483        let s = self.val / r;
484        ft - ft.min(s.wrapping_add(1))
485    }
486
487    #[inline(always)]
488    pub fn update(&mut self, fl: u32, fh: u32, ft: u32) {
489        let s = self.ext.wrapping_mul(ft.wrapping_sub(fh));
490        self.val = self.val.wrapping_sub(s);
491        self.rng = if fl > 0 {
492            self.ext.wrapping_mul(fh.wrapping_sub(fl))
493        } else {
494            self.rng.wrapping_sub(s)
495        };
496        self.normalize_decoder();
497    }
498
499    pub fn laplace_encode(&mut self, value: &mut i32, fs: u32, decay: i32) {
500        let mut val = *value;
501        let mut fl = 0;
502        let mut fs_val = fs;
503
504        if val != 0 {
505            let s = if val < 0 { -1 } else { 0 };
506            val = (val + s) ^ s;
507            fl = fs_val;
508            fs_val = self.laplace_get_freq1(fs_val, decay);
509
510            let mut i = 1;
511            while fs_val > 0 && i < val {
512                fs_val *= 2;
513                fl += fs_val + 2;
514                fs_val = ((fs_val as i32 * decay) >> 15) as u32;
515                i += 1;
516            }
517
518            if fs_val == 0 {
519                let ndi_max = 32768 - fl + 1 - 1;
520                let ndi_max = (ndi_max as i32 - s) >> 1;
521                let di = (val - i).min(ndi_max - 1);
522                fl += (2 * di + 1 + s) as u32;
523                fs_val = 1u32.min(32768 - fl);
524                *value = (i + di + s) ^ s;
525            } else {
526                fs_val += 1;
527                fl += fs_val & (!s as u32);
528            }
529        }
530        self.encode(fl, fl.wrapping_add(fs_val), 1 << 15);
531    }
532
533    fn laplace_get_freq1(&self, fs0: u32, decay: i32) -> u32 {
534        let ft = 32768 - (2 * 16) - fs0;
535        ((ft as i32 * (16384 - decay)) >> 15) as u32
536    }
537
538    pub fn laplace_decode(&mut self, fs: u32, decay: i32) -> i32 {
539        let fm = self.decode(1 << 15);
540        let mut fl = 0;
541        let mut fs_val = fs;
542        let mut val = 0;
543
544        if fm >= fs_val {
545            val += 1;
546            fl = fs_val;
547            fs_val = self.laplace_get_freq1(fs_val, decay) + 1;
548
549            while fs_val > 1 && fm >= fl + 2 * fs_val {
550                fs_val *= 2;
551                fl += fs_val;
552                fs_val = (((fs_val as i32 - 2) * decay) >> 15) as u32 + 1;
553                val += 1;
554            }
555
556            if fs_val <= 1 {
557                let di = (fm - fl) >> 1;
558                val += di as i32;
559                fl += 2 * di;
560            }
561
562            if fm < fl + fs_val {
563                val = -val;
564            } else {
565                fl += fs_val;
566            }
567        }
568
569        self.update(fl, fl.wrapping_add(fs_val.min(32768 - fl)), 1 << 15);
570        val
571    }
572
573    #[inline(always)]
574    fn write_byte_at_end(&mut self, value: u8) {
575        if self.offs + self.end_offs < self.storage {
576            self.end_offs += 1;
577            let idx = (self.storage - self.end_offs) as usize;
578            unsafe {
579                *self.buf.get_unchecked_mut(idx) = value;
580            }
581        } else {
582            self.error = 1;
583        }
584    }
585
586    pub fn patch_initial_bits(&mut self, val: u32, nbits: u32) {
587        let shift = EC_SYM_BITS - nbits;
588        let mask = ((1u32 << nbits) - 1) << shift;
589        if self.offs > 0 {
590            self.buf[0] = ((self.buf[0] as u32 & !mask) | (val << shift)) as u8;
591        } else if self.rem >= 0 {
592            self.rem = ((self.rem as u32 & !mask) | (val << shift)) as i32;
593        } else if self.rng <= (EC_CODE_TOP >> nbits) {
594            let mask_shifted = mask << EC_CODE_SHIFT;
595            self.val = (self.val & !mask_shifted) | (val << (EC_CODE_SHIFT + shift));
596        } else {
597            self.error = -1;
598        }
599    }
600
601    pub fn done(&mut self) {
602        let ilog = 32 - self.rng.leading_zeros();
603        let mut l = (EC_CODE_BITS - ilog) as i32;
604        let mut msk = (EC_CODE_TOP - 1) >> l;
605        let mut end = (self.val.wrapping_add(msk)) & !msk;
606
607        if (end | msk) >= self.val.wrapping_add(self.rng) {
608            l += 1;
609            msk >>= 1;
610            end = (self.val.wrapping_add(msk)) & !msk;
611        }
612
613        while l > 0 {
614            self.carry_out((end >> EC_CODE_SHIFT) as i32);
615            end = (end << EC_SYM_BITS) & (EC_CODE_TOP - 1);
616            l -= EC_SYM_BITS as i32;
617        }
618
619        if self.rem >= 0 || self.ext > 0 {
620            self.carry_out(0);
621        }
622
623        let mut window = self.end_window;
624        let mut used = self.nend_bits;
625        while used >= EC_SYM_BITS as i32 {
626            self.write_byte_at_end((window & EC_SYM_MAX) as u8);
627            window >>= EC_SYM_BITS;
628            used -= EC_SYM_BITS as i32;
629        }
630
631        if self.error == 0 {
632            for i in self.offs..(self.storage - self.end_offs) {
633                self.buf[i as usize] = 0;
634            }
635
636            if used > 0 {
637                if self.end_offs >= self.storage {
638                    self.error = -1;
639                } else {
640                    // If we've busted, don't add too many extra bits to the
641                    // last byte; it would corrupt the range coder data.
642                    if self.offs + self.end_offs >= self.storage && -l < used {
643                        window &= (1u32 << (-l as u32)) - 1;
644                        self.error = -1;
645                    }
646                    let idx = (self.storage - self.end_offs - 1) as usize;
647                    self.buf[idx] |= window as u8;
648                }
649            }
650        }
651    }
652
653    /// Flush pending output and return the encoded bytes.
654    ///
655    /// Allocates and is therefore only available with the `std` feature. The
656    /// `#![no_std]` encoder path reads `self.buf[..self.offs]` directly instead.
657    #[cfg(feature = "std")]
658    pub fn finish(&mut self) -> Vec<u8> {
659        self.done();
660
661        let extra_end = if self.nend_bits > 0 && self.end_offs == 0 {
662            1
663        } else {
664            0
665        };
666        let mut result = Vec::with_capacity((self.offs + self.end_offs + extra_end) as usize);
667        result.extend_from_slice(&self.buf[0..self.offs as usize]);
668        if extra_end > 0 {
669            result.push(self.buf[(self.storage - 1) as usize]);
670        }
671        result.extend_from_slice(
672            &self.buf[(self.storage - self.end_offs) as usize..self.storage as usize],
673        );
674        result
675    }
676}
677
678#[cfg(all(test, feature = "std"))]
679mod tests {
680    use super::*;
681
682    #[test]
683    fn test_laplace() {
684        let mut enc = RangeCoder::new_encoder(100);
685        let mut val = -3;
686        let fs = 100 << 7;
687        let decay = 120 << 6;
688        enc.laplace_encode(&mut val, fs, decay);
689        enc.done();
690
691        assert_eq!(enc.offs, 1);
692        assert_eq!(enc.buf[0], 224);
693
694        let mut dec = RangeCoder::new_decoder(&enc.buf[..enc.offs as usize]);
695        let decoded_val = dec.laplace_decode(fs, decay);
696        assert_eq!(decoded_val, -3);
697    }
698
699    #[test]
700    fn test_icdf_consistency() {
701        let mut enc = RangeCoder::new_encoder(1024);
702        let icdf = [2, 1, 0];
703        enc.encode_icdf(0, &icdf, 2);
704        enc.encode_icdf(1, &icdf, 2);
705        enc.encode_icdf(2, &icdf, 2);
706        enc.done();
707        let data = enc.buf[..enc.offs as usize].to_vec();
708
709        let mut dec = RangeCoder::new_decoder(&data);
710        let s0 = dec.decode_icdf(&icdf, 2);
711        let s1 = dec.decode_icdf(&icdf, 2);
712        let s2 = dec.decode_icdf(&icdf, 2);
713
714        assert_eq!(s0, 0);
715        assert_eq!(s1, 1);
716        assert_eq!(s2, 2);
717    }
718
719    #[test]
720    fn test_icdf_last_symbol_no_oob() {
721        let icdf: &[u8] = &[170, 85, 0];
722        let ftb = 8u32;
723
724        for sym in 0..3i32 {
725            let mut enc = RangeCoder::new_encoder(256);
726            enc.encode_icdf(sym, icdf, ftb);
727            enc.done();
728            let data = enc.buf[..enc.offs as usize].to_vec();
729
730            let mut dec = RangeCoder::new_decoder(&data);
731            let decoded = dec.decode_icdf(icdf, ftb);
732            assert_eq!(decoded, sym, "往返失败: 编码 symbol={sym} 解码得 {decoded}");
733        }
734    }
735
736    #[test]
737    fn test_icdf_decode_terminates() {
738        let icdf: &[u8] = &[192, 128, 64, 0];
739        let ftb = 8u32;
740
741        let symbols = [0i32, 1, 2, 3];
742        let mut enc = RangeCoder::new_encoder(256);
743        for &s in &symbols {
744            enc.encode_icdf(s, icdf, ftb);
745        }
746        enc.done();
747        let data = enc.buf[..enc.offs as usize].to_vec();
748
749        let mut dec = RangeCoder::new_decoder(&data);
750        for &expected in &symbols {
751            let got = dec.decode_icdf(icdf, ftb);
752            assert_eq!(got, expected, "解码器输出 {got},期望 {expected}");
753        }
754    }
755
756    #[test]
757    fn test_bits_only() {
758        let mut enc = RangeCoder::new_encoder(1024);
759
760        enc.enc_bits(1, 1);
761        enc.enc_bits(5, 3);
762        enc.enc_bits(7, 3);
763        enc.enc_bits(0, 2);
764
765        let data = enc.finish();
766        let mut dec = RangeCoder::new_decoder(&data);
767
768        let b1 = dec.dec_bits(1);
769        let b2 = dec.dec_bits(3);
770        let b3 = dec.dec_bits(3);
771        let b4 = dec.dec_bits(2);
772
773        assert_eq!(b1, 1);
774        assert_eq!(b2, 5);
775        assert_eq!(b3, 7);
776        assert_eq!(b4, 0);
777    }
778
779    #[test]
780    fn test_interleaved_bits_entropy() {
781        let mut enc = RangeCoder::new_encoder(1024);
782
783        enc.enc_bits(1, 1);
784
785        enc.encode(10, 20, 100);
786
787        enc.enc_bits(5, 3);
788
789        enc.encode(50, 60, 100);
790
791        let data = enc.finish();
792
793        let mut dec = RangeCoder::new_decoder(&data);
794
795        let b1 = dec.dec_bits(1);
796        let d1 = dec.decode(100);
797        dec.update(10, 20, 100);
798        let b2 = dec.dec_bits(3);
799        let d2 = dec.decode(100);
800        dec.update(50, 60, 100);
801
802        assert_eq!(b1, 1);
803        assert!((10..20).contains(&d1), "d1={} expected in [10, 20)", d1);
804        assert_eq!(b2, 5);
805        assert!((50..60).contains(&d2), "d2={} expected in [50, 60)", d2);
806    }
807}