1use crate::Freq;
19use crate::arithmetic;
20#[allow(unused_imports)]
21use crate::sink::Sink;
22#[allow(unused_imports)]
23use crate::source::Source;
24
25macro_rules! generate_rans_impl {
31 (
32 $enc_symbol:ident, $dec_symbol:ident, $encoder:ident, $decoder:ident, $state_ty:ty, $unit_ty:ty, $state_bits:expr, $max_scale_bits:expr,$lower_bound:expr, $unit_bits:expr, $units_per_state:expr, ) => {
44 #[derive(Debug, Clone, Copy)]
46 pub struct $enc_symbol {
47 pub x_max_hi: Freq,
48 pub freq_rcp_shift: u32,
49 pub freq_rcp: $state_ty,
50 pub freq_cmpl: Freq,
51 pub bias: Freq,
52 }
53
54 impl $enc_symbol {
55 #[inline]
57 pub fn new(start: Freq, freq: Freq, scale_bits: Freq) -> Self {
58 let scale = 1u32 << scale_bits;
59 debug_assert!(
60 0 < scale_bits && scale_bits as u32 <= $max_scale_bits,
61 "invalid scale_bits"
62 );
63 debug_assert!(start < scale, "start out of range");
64 debug_assert!(freq > 0 && freq <= scale - start, "freq out of range");
65
66 let min_bits =
67 core::cmp::min($state_bits, (core::mem::size_of::<Freq>() * 8 - 1) as u32);
68 let x_max_hi = freq << (min_bits - scale_bits);
69
70 let (freq_rcp, mut freq_rcp_shift, bias) = if freq > 1 {
71 let shift = arithmetic::reciprocal_shift(freq);
72 let rcp: $state_ty = if core::mem::size_of::<$state_ty>() >= 8 {
73 arithmetic::compute_reciprocal_u64(freq) as $state_ty
74 } else {
75 arithmetic::compute_reciprocal_u32(freq) as $state_ty
76 };
77 (rcp, shift - 1, start)
78 } else {
79 let rcp: $state_ty = if core::mem::size_of::<$state_ty>() >= 8 {
80 !0u64 as $state_ty
81 } else {
82 !0u32 as $state_ty
83 };
84 let bias_adj = start + scale - 1;
85 (rcp, 0, bias_adj)
86 };
87
88 if core::mem::size_of::<$state_ty>() < 8 {
89 freq_rcp_shift += (core::mem::size_of::<$state_ty>() * 8) as u32;
90 }
91
92 let freq_cmpl = scale - freq;
93
94 Self {
95 x_max_hi,
96 freq_rcp_shift,
97 freq_rcp,
98 freq_cmpl,
99 bias,
100 }
101 }
102
103 #[inline]
109 pub fn try_new(
110 start: Freq,
111 freq: Freq,
112 scale_bits: Freq,
113 ) -> core::result::Result<Self, crate::error::RawRansError> {
114 if scale_bits < 2 || scale_bits > $max_scale_bits || scale_bits >= 32 {
117 return Err(crate::error::RawRansError::InvalidScaleBits {
118 provided: scale_bits as u32,
119 max_safe: core::cmp::min($max_scale_bits, 31),
120 });
121 }
122 let scale = 1u32 << scale_bits;
123 if start >= scale {
124 return Err(crate::error::RawRansError::InvalidParameters);
125 }
126 if freq == 0 || freq > scale - start {
127 return Err(crate::error::RawRansError::InvalidParameters);
128 }
129
130 let min_bits =
131 core::cmp::min($state_bits, (core::mem::size_of::<Freq>() * 8 - 1) as u32);
132 let x_max_hi = freq << (min_bits - scale_bits);
133
134 let (freq_rcp, mut freq_rcp_shift, bias) = if freq > 1 {
135 let shift = arithmetic::reciprocal_shift(freq);
136 let rcp: $state_ty = if core::mem::size_of::<$state_ty>() >= 8 {
137 arithmetic::compute_reciprocal_u64(freq) as $state_ty
138 } else {
139 arithmetic::compute_reciprocal_u32(freq) as $state_ty
140 };
141 (rcp, shift - 1, start)
142 } else {
143 let rcp: $state_ty = if core::mem::size_of::<$state_ty>() >= 8 {
144 !0u64 as $state_ty
145 } else {
146 !0u32 as $state_ty
147 };
148 let bias_adj = start + scale - 1;
149 (rcp, 0, bias_adj)
150 };
151
152 if core::mem::size_of::<$state_ty>() < 8 {
153 freq_rcp_shift += (core::mem::size_of::<$state_ty>() * 8) as u32;
154 }
155
156 let freq_cmpl = scale - freq;
157
158 Ok(Self {
159 x_max_hi,
160 freq_rcp_shift,
161 freq_rcp,
162 freq_cmpl,
163 bias,
164 })
165 }
166
167 #[inline]
169 pub fn quotient(&self, x: $state_ty) -> $state_ty {
170 if core::mem::size_of::<$state_ty>() >= 8 {
171 let x_u64 = x as u64;
172 let rcp_u64 = self.freq_rcp as u64;
173 arithmetic::fast_quotient_u64(x_u64, rcp_u64, self.freq_rcp_shift) as $state_ty
174 } else {
175 let x_u32 = x as u32;
176 let rcp_u32 = self.freq_rcp as u32;
177 arithmetic::fast_quotient_u32(x_u32, rcp_u32, self.freq_rcp_shift) as $state_ty
178 }
179 }
180 }
181
182 #[derive(Debug, Clone, Copy)]
184 pub struct $dec_symbol {
185 pub freq: Freq,
186 pub start: Freq,
187 }
188
189 impl $dec_symbol {
190 #[inline]
192 pub fn new(start: Freq, freq: Freq) -> Self {
193 debug_assert!(freq > 0, "frequency must be positive");
194 Self { freq, start }
195 }
196 }
197
198 #[derive(Debug)]
200 pub struct $encoder<Sk: Sink<$unit_ty>> {
201 sink: Sk,
202 state: $state_ty,
203 }
204
205 impl<Sk: Sink<$unit_ty>> $encoder<Sk> {
206 #[inline]
207 pub fn new(sink: Sk) -> Self {
208 Self {
209 sink,
210 state: $lower_bound,
211 }
212 }
213
214 #[inline]
215 pub fn sink(&self) -> &Sk {
216 &self.sink
217 }
218 #[inline]
219 pub fn sink_mut(&mut self) -> &mut Sk {
220 &mut self.sink
221 }
222 #[inline]
223 pub fn into_sink(self) -> Sk {
224 self.sink
225 }
226 #[inline]
227 pub fn state(&self) -> $state_ty {
228 self.state
229 }
230 #[inline]
231 pub fn reset(&mut self) {
232 self.state = $lower_bound;
233 }
234
235 #[inline]
240 pub fn put_raw(&mut self, start: Freq, freq: Freq, scale_bits: Freq) {
241 self.try_put_raw(start, freq, scale_bits)
242 .expect("invalid raw rANS encoder parameters");
243 }
244
245 #[inline]
248 pub(crate) fn put_raw_unchecked(&mut self, start: Freq, freq: Freq, scale_bits: Freq) {
249 debug_assert!(start < (1u32 << scale_bits));
250 debug_assert!(freq > 0 && freq <= (1u32 << scale_bits) - start);
251
252 let x_max: $state_ty = (freq as $state_ty) << ($state_bits - scale_bits as u32);
253 let x = self.renormalize(x_max);
254
255 let shift = scale_bits as u32;
256 self.state = ((x / freq as $state_ty) << shift)
257 + (start as $state_ty)
258 + (x % freq as $state_ty);
259 }
260
261 #[inline]
266 pub fn try_put_raw(
267 &mut self,
268 start: Freq,
269 freq: Freq,
270 scale_bits: Freq,
271 ) -> core::result::Result<(), crate::error::RawRansError> {
272 if scale_bits < 2 || scale_bits > $max_scale_bits || scale_bits >= 32 {
274 return Err(crate::error::RawRansError::InvalidScaleBits {
275 provided: scale_bits as u32,
276 max_safe: core::cmp::min($max_scale_bits, 31),
277 });
278 }
279 let scale = 1u32 << scale_bits;
280 if start >= scale {
281 return Err(crate::error::RawRansError::InvalidParameters);
282 }
283 if freq == 0 || freq > scale - start {
284 return Err(crate::error::RawRansError::InvalidParameters);
285 }
286 self.put_raw_unchecked(start, freq, scale_bits);
287 Ok(())
288 }
289
290 #[inline]
292 pub fn put(&mut self, symbol: &$enc_symbol) {
293 let mut x_max = symbol.x_max_hi as $state_ty;
294 if $state_bits > (core::mem::size_of::<Freq>() * 8 - 1) as u32 {
295 let shift =
296 ($state_bits - (core::mem::size_of::<Freq>() * 8 - 1) as u32) as usize;
297 x_max = x_max << shift;
298 }
299 let x = self.renormalize(x_max);
300
301 let q = symbol.quotient(x);
302 self.state = x + (q * symbol.freq_cmpl as $state_ty) + symbol.bias as $state_ty;
303 }
304
305 #[inline]
307 pub fn flush(&mut self) {
308 let x = self.state;
309 for i in (1..$units_per_state).rev() {
310 let shift = i * $unit_bits as usize;
311 self.sink.write((x >> shift) as $unit_ty);
312 }
313 self.sink.write(x as $unit_ty);
314 }
315
316 #[inline]
317 fn renormalize(&mut self, x_max: $state_ty) -> $state_ty {
318 let mut x = self.state;
319 while x >= x_max {
320 self.sink.write(x as $unit_ty);
321 x >>= $unit_bits;
322 if $max_scale_bits <= $unit_bits {
323 debug_assert!(x < x_max);
324 break;
325 }
326 }
327 x
328 }
329 }
330
331 #[derive(Debug)]
333 pub struct $decoder<Sr: Source<$unit_ty>> {
334 source: Sr,
335 state: $state_ty,
336 }
337
338 impl<Sr: Source<$unit_ty>> $decoder<Sr> {
339 #[inline]
340 pub fn new(source: Sr) -> Self {
341 Self {
342 source,
343 state: $lower_bound,
344 }
345 }
346
347 #[inline]
350 pub fn from_state(source: Sr, state: $state_ty) -> Self {
351 Self { source, state }
352 }
353
354 #[inline]
355 pub fn source(&self) -> &Sr {
356 &self.source
357 }
358 #[inline]
359 pub fn source_mut(&mut self) -> &mut Sr {
360 &mut self.source
361 }
362 #[inline]
363 pub fn into_source(self) -> Sr {
364 self.source
365 }
366 #[inline]
367 pub fn state(&self) -> $state_ty {
368 self.state
369 }
370 #[inline]
371 pub fn check_eof(&self) -> bool {
372 self.state == $lower_bound
373 }
374
375 #[inline]
377 pub fn init(&mut self) -> bool {
378 let mut unit = <$unit_ty>::default();
379 if !self.source.read(&mut unit) {
380 return false;
381 }
382 let mut x = unit as $state_ty;
383
384 for i in 1..$units_per_state {
385 if !self.source.read(&mut unit) {
386 return false;
387 }
388 x = x.wrapping_add((unit as $state_ty) << (i * $unit_bits as usize));
389 }
390
391 if x < $lower_bound {
392 return false;
393 }
394 self.state = x;
395 true
396 }
397
398 #[inline]
403 pub fn get(&self, scale_bits: Freq) -> Freq {
404 assert!(
405 scale_bits < 32,
406 "scale_bits={} causes overflow (use try_get)",
407 scale_bits
408 );
409 self.get_unchecked(scale_bits)
410 }
411
412 #[inline]
415 pub(crate) fn get_unchecked(&self, scale_bits: Freq) -> Freq {
416 let mask = (1u32 << scale_bits) - 1;
417 (self.state as Freq) & mask
418 }
419
420 #[inline]
422 pub fn try_get(
423 &self,
424 scale_bits: Freq,
425 ) -> core::result::Result<Freq, crate::error::RawRansError> {
426 if scale_bits < 2 || scale_bits > $max_scale_bits || scale_bits >= 32 {
427 return Err(crate::error::RawRansError::InvalidScaleBits {
428 provided: scale_bits as u32,
429 max_safe: core::cmp::min($max_scale_bits, 31),
430 });
431 }
432 Ok(self.get_unchecked(scale_bits))
433 }
434
435 #[inline]
445 pub fn advance(&mut self, start: Freq, freq: Freq, scale_bits: Freq) -> bool {
446 self.try_advance(start, freq, scale_bits)
447 .expect("invalid raw rANS decoder parameters")
448 }
449
450 #[inline]
453 pub(crate) fn advance_unchecked(
454 &mut self,
455 start: Freq,
456 freq: Freq,
457 scale_bits: Freq,
458 ) -> bool {
459 debug_assert!(start < (1u32 << scale_bits));
460 debug_assert!(freq > 0 && freq <= (1u32 << scale_bits) - start);
461
462 let scale = 1u32 << scale_bits;
463
464 let x = self.state;
465 let mask = (scale - 1) as $state_ty;
466 let value = x & mask;
467 if value < start as $state_ty {
475 return false;
476 }
477
478 let shift = scale_bits as u32;
480 let mut x_new = (freq as $state_ty) * (x >> shift) + (value - start as $state_ty);
481
482 let mut renorm_unit = <$unit_ty>::default();
484 while x_new < $lower_bound {
485 if !self.source.read(&mut renorm_unit) {
486 return false;
487 }
488 x_new = (x_new << $unit_bits) + renorm_unit as $state_ty;
489 if $max_scale_bits <= $unit_bits {
490 debug_assert!(x_new >= $lower_bound);
491 break;
492 }
493 }
494
495 self.state = x_new;
497 true
498 }
499
500 #[inline]
502 pub fn try_advance(
503 &mut self,
504 start: Freq,
505 freq: Freq,
506 scale_bits: Freq,
507 ) -> core::result::Result<bool, crate::error::RawRansError> {
508 if scale_bits < 2 || scale_bits > $max_scale_bits || scale_bits >= 32 {
509 return Err(crate::error::RawRansError::InvalidScaleBits {
510 provided: scale_bits as u32,
511 max_safe: core::cmp::min($max_scale_bits, 31),
512 });
513 }
514 let scale = 1u32 << scale_bits;
515 if start >= scale || freq == 0 || freq > scale - start {
516 return Err(crate::error::RawRansError::InvalidParameters);
517 }
518 Ok(self.advance_unchecked(start, freq, scale_bits))
519 }
520
521 #[inline]
523 pub fn advance_symbol(&mut self, symbol: &$dec_symbol, scale_bits: Freq) -> bool {
524 self.advance(symbol.start, symbol.freq, scale_bits)
525 }
526 }
527 };
528}
529
530generate_rans_impl! {
535 RansByteEncSymbol, RansByteDecSymbol, RansByteEncoder, RansByteDecoder, u32, u8, 31, 30, 1u32 << 23, 8, 4, }
546
547generate_rans_impl! {
548 Rans64EncSymbol, Rans64DecSymbol, Rans64Encoder, Rans64Decoder, u64, u32, 63, 32, 1u64 << 31, 32, 2, }
559
560#[cfg(test)]
565mod tests {
566 use super::*;
567 use crate::sink::VecSink;
568 use crate::source::SliceSource;
569
570 #[test]
573 fn test_ransbyte_initial_state() {
574 let sink = VecSink::<u8>::new(64);
575 let encoder = RansByteEncoder::new(sink);
576 assert_eq!(encoder.state(), 1u32 << 23);
577 }
578
579 #[test]
580 fn test_ransbyte_raw_roundtrip() {
581 let sink = VecSink::<u8>::new(64);
582 let mut encoder = RansByteEncoder::new(sink);
583 encoder.put_raw(0, 128, 8);
584 encoder.flush();
585 let encoded = encoder.into_sink().encoded().to_vec();
586 assert!(!encoded.is_empty());
587
588 let source = SliceSource::new(&encoded[..]);
589 let mut decoder = RansByteDecoder::new(source);
590 assert!(decoder.init());
591
592 let freq_val = decoder.get(8);
593 assert_eq!(freq_val, 0);
594 assert!(decoder.advance(0, 128, 8));
595 assert!(decoder.check_eof());
596 }
597
598 #[test]
599 fn test_ransbyte_prepared_symbol_matches_raw() {
600 let sink1 = VecSink::<u8>::new(64);
601 let mut enc1 = RansByteEncoder::new(sink1);
602 enc1.put_raw(0, 128, 8);
603 enc1.flush();
604 let out1 = enc1.into_sink().encoded().to_vec();
605
606 let sink2 = VecSink::<u8>::new(64);
607 let mut enc2 = RansByteEncoder::new(sink2);
608 let sym = RansByteEncSymbol::new(0, 128, 8);
609 enc2.put(&sym);
610 enc2.flush();
611 let out2 = enc2.into_sink().encoded().to_vec();
612
613 assert_eq!(out1, out2);
614 }
615
616 #[test]
617 fn test_ransbyte_decoder_rejects_empty() {
618 let data: [u8; 0] = [];
619 let source = SliceSource::new(&data[..]);
620 let mut decoder = RansByteDecoder::new(source);
621 assert!(!decoder.init());
622 }
623
624 #[test]
625 fn test_ransbyte_encoder_reset() {
626 let sink = VecSink::<u8>::new(64);
627 let mut encoder = RansByteEncoder::new(sink);
628 encoder.put_raw(0, 128, 8);
629 encoder.reset();
630 assert_eq!(encoder.state(), 1u32 << 23);
631 }
632
633 #[test]
636 fn test_rans64_initial_state() {
637 let sink = VecSink::<u32>::new(64);
638 let encoder = Rans64Encoder::new(sink);
639 assert_eq!(encoder.state(), 1u64 << 31);
640 }
641
642 #[test]
643 fn test_rans64_raw_roundtrip() {
644 let sink = VecSink::<u32>::new(64);
645 let mut encoder = Rans64Encoder::new(sink);
646 encoder.put_raw(0, 128, 8);
647 encoder.flush();
648 let encoded = encoder.into_sink().encoded().to_vec();
649 assert!(!encoded.is_empty());
650
651 let source = SliceSource::new(&encoded[..]);
652 let mut decoder = Rans64Decoder::new(source);
653 assert!(decoder.init());
654 assert!(decoder.advance(0, 128, 8));
655 assert!(decoder.check_eof());
656 }
657
658 #[test]
659 fn test_rans64_prepared_symbol_matches_raw() {
660 let sink1 = VecSink::<u32>::new(64);
661 let mut enc1 = Rans64Encoder::new(sink1);
662 enc1.put_raw(0, 128, 8);
663 enc1.flush();
664 let out1 = enc1.into_sink().encoded().to_vec();
665
666 let sink2 = VecSink::<u32>::new(64);
667 let mut enc2 = Rans64Encoder::new(sink2);
668 let sym = Rans64EncSymbol::new(0, 128, 8);
669 enc2.put(&sym);
670 enc2.flush();
671 let out2 = enc2.into_sink().encoded().to_vec();
672
673 assert_eq!(out1, out2);
674 }
675
676 #[test]
677 fn test_rans64_encoder_reset() {
678 let sink = VecSink::<u32>::new(64);
679 let mut encoder = Rans64Encoder::new(sink);
680 encoder.put_raw(0, 128, 8);
681 encoder.reset();
682 assert_eq!(encoder.state(), 1u64 << 31);
683 }
684
685 #[test]
686 fn test_scale_bits_32_rejected_for_try_new() {
687 let result = Rans64EncSymbol::try_new(0, 128, 32);
688 assert!(result.is_err(), "scale_bits=32 must be rejected");
689 if let Err(e) = result {
690 match e {
691 crate::error::RawRansError::InvalidScaleBits { provided, max_safe } => {
692 assert_eq!(provided, 32);
693 assert_eq!(max_safe, 31);
694 }
695 _ => panic!("wrong error type"),
696 }
697 }
698 }
699
700 #[test]
701 fn test_scale_bits_32_rejected_for_try_put_raw() {
702 let sink = VecSink::<u32>::new(64);
703 let mut encoder = Rans64Encoder::new(sink);
704 let result = encoder.try_put_raw(0, 128, 32);
705 assert!(result.is_err(), "scale_bits=32 must be rejected");
706 }
707
708 #[test]
709 fn test_scale_bits_31_accepted_for_rans64() {
710 let result = Rans64EncSymbol::try_new(0, 1, 31);
711 assert!(result.is_ok(), "scale_bits=31 must be accepted for Rans64");
712 }
713
714 #[test]
715 fn test_scale_bits_1_rejected() {
716 let result = RansByteEncSymbol::try_new(0, 128, 1);
717 assert!(
718 result.is_err(),
719 "scale_bits=1 must be rejected (minimum is 2)"
720 );
721 }
722
723 #[test]
724 fn test_try_new_rejects_freq_zero() {
725 let result = RansByteEncSymbol::try_new(0, 0, 8);
726 assert!(result.is_err(), "freq=0 must be rejected");
727 }
728
729 #[test]
730 fn test_try_new_accepts_valid_params() {
731 let result = RansByteEncSymbol::try_new(0, 128, 8);
732 assert!(result.is_ok(), "valid params must be accepted");
733 }
734}