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
12pub 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 #[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 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 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 #[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}