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]
348 pub fn source(&self) -> &Sr {
349 &self.source
350 }
351 #[inline]
352 pub fn source_mut(&mut self) -> &mut Sr {
353 &mut self.source
354 }
355 #[inline]
356 pub fn into_source(self) -> Sr {
357 self.source
358 }
359 #[inline]
360 pub fn state(&self) -> $state_ty {
361 self.state
362 }
363 #[inline]
364 pub fn check_eof(&self) -> bool {
365 self.state == $lower_bound
366 }
367
368 #[inline]
370 pub fn init(&mut self) -> bool {
371 let mut unit = <$unit_ty>::default();
372 if !self.source.read(&mut unit) {
373 return false;
374 }
375 let mut x = unit as $state_ty;
376
377 for i in 1..$units_per_state {
378 if !self.source.read(&mut unit) {
379 return false;
380 }
381 x = x.wrapping_add((unit as $state_ty) << (i * $unit_bits as usize));
382 }
383
384 if x < $lower_bound {
385 return false;
386 }
387 self.state = x;
388 true
389 }
390
391 #[inline]
396 pub fn get(&self, scale_bits: Freq) -> Freq {
397 assert!(
398 scale_bits < 32,
399 "scale_bits={} causes overflow (use try_get)",
400 scale_bits
401 );
402 self.get_unchecked(scale_bits)
403 }
404
405 #[inline]
408 pub(crate) fn get_unchecked(&self, scale_bits: Freq) -> Freq {
409 let mask = (1u32 << scale_bits) - 1;
410 (self.state as Freq) & mask
411 }
412
413 #[inline]
415 pub fn try_get(
416 &self,
417 scale_bits: Freq,
418 ) -> core::result::Result<Freq, crate::error::RawRansError> {
419 if scale_bits < 2 || scale_bits > $max_scale_bits || scale_bits >= 32 {
420 return Err(crate::error::RawRansError::InvalidScaleBits {
421 provided: scale_bits as u32,
422 max_safe: core::cmp::min($max_scale_bits, 31),
423 });
424 }
425 Ok(self.get_unchecked(scale_bits))
426 }
427
428 #[inline]
438 pub fn advance(&mut self, start: Freq, freq: Freq, scale_bits: Freq) -> bool {
439 self.try_advance(start, freq, scale_bits)
440 .expect("invalid raw rANS decoder parameters")
441 }
442
443 #[inline]
446 pub(crate) fn advance_unchecked(
447 &mut self,
448 start: Freq,
449 freq: Freq,
450 scale_bits: Freq,
451 ) -> bool {
452 debug_assert!(start < (1u32 << scale_bits));
453 debug_assert!(freq > 0 && freq <= (1u32 << scale_bits) - start);
454
455 let scale = 1u32 << scale_bits;
456
457 let x = self.state;
458 let mask = (scale - 1) as $state_ty;
459 let value = x & mask;
460 debug_assert!(value >= start as $state_ty);
461
462 let shift = scale_bits as u32;
464 let mut x_new = (freq as $state_ty) * (x >> shift) + (value - start as $state_ty);
465
466 let mut renorm_unit = <$unit_ty>::default();
468 while x_new < $lower_bound {
469 if !self.source.read(&mut renorm_unit) {
470 return false;
471 }
472 x_new = (x_new << $unit_bits) + renorm_unit as $state_ty;
473 if $max_scale_bits <= $unit_bits {
474 debug_assert!(x_new >= $lower_bound);
475 break;
476 }
477 }
478
479 self.state = x_new;
481 true
482 }
483
484 #[inline]
486 pub fn try_advance(
487 &mut self,
488 start: Freq,
489 freq: Freq,
490 scale_bits: Freq,
491 ) -> core::result::Result<bool, crate::error::RawRansError> {
492 if scale_bits < 2 || scale_bits > $max_scale_bits || scale_bits >= 32 {
493 return Err(crate::error::RawRansError::InvalidScaleBits {
494 provided: scale_bits as u32,
495 max_safe: core::cmp::min($max_scale_bits, 31),
496 });
497 }
498 let scale = 1u32 << scale_bits;
499 if start >= scale || freq == 0 || freq > scale - start {
500 return Err(crate::error::RawRansError::InvalidParameters);
501 }
502 Ok(self.advance_unchecked(start, freq, scale_bits))
503 }
504
505 #[inline]
507 pub fn advance_symbol(&mut self, symbol: &$dec_symbol, scale_bits: Freq) -> bool {
508 self.advance(symbol.start, symbol.freq, scale_bits)
509 }
510 }
511 };
512}
513
514generate_rans_impl! {
519 RansByteEncSymbol, RansByteDecSymbol, RansByteEncoder, RansByteDecoder, u32, u8, 31, 30, 1u32 << 23, 8, 4, }
530
531generate_rans_impl! {
532 Rans64EncSymbol, Rans64DecSymbol, Rans64Encoder, Rans64Decoder, u64, u32, 63, 32, 1u64 << 31, 32, 2, }
543
544#[cfg(test)]
549mod tests {
550 use super::*;
551 use crate::sink::VecSink;
552 use crate::source::SliceSource;
553
554 #[test]
557 fn test_ransbyte_initial_state() {
558 let sink = VecSink::<u8>::new(64);
559 let encoder = RansByteEncoder::new(sink);
560 assert_eq!(encoder.state(), 1u32 << 23);
561 }
562
563 #[test]
564 fn test_ransbyte_raw_roundtrip() {
565 let sink = VecSink::<u8>::new(64);
566 let mut encoder = RansByteEncoder::new(sink);
567 encoder.put_raw(0, 128, 8);
568 encoder.flush();
569 let encoded = encoder.into_sink().encoded().to_vec();
570 assert!(!encoded.is_empty());
571
572 let source = SliceSource::new(&encoded[..]);
573 let mut decoder = RansByteDecoder::new(source);
574 assert!(decoder.init());
575
576 let freq_val = decoder.get(8);
577 assert_eq!(freq_val, 0);
578 assert!(decoder.advance(0, 128, 8));
579 assert!(decoder.check_eof());
580 }
581
582 #[test]
583 fn test_ransbyte_prepared_symbol_matches_raw() {
584 let sink1 = VecSink::<u8>::new(64);
585 let mut enc1 = RansByteEncoder::new(sink1);
586 enc1.put_raw(0, 128, 8);
587 enc1.flush();
588 let out1 = enc1.into_sink().encoded().to_vec();
589
590 let sink2 = VecSink::<u8>::new(64);
591 let mut enc2 = RansByteEncoder::new(sink2);
592 let sym = RansByteEncSymbol::new(0, 128, 8);
593 enc2.put(&sym);
594 enc2.flush();
595 let out2 = enc2.into_sink().encoded().to_vec();
596
597 assert_eq!(out1, out2);
598 }
599
600 #[test]
601 fn test_ransbyte_decoder_rejects_empty() {
602 let data: [u8; 0] = [];
603 let source = SliceSource::new(&data[..]);
604 let mut decoder = RansByteDecoder::new(source);
605 assert!(!decoder.init());
606 }
607
608 #[test]
609 fn test_ransbyte_encoder_reset() {
610 let sink = VecSink::<u8>::new(64);
611 let mut encoder = RansByteEncoder::new(sink);
612 encoder.put_raw(0, 128, 8);
613 encoder.reset();
614 assert_eq!(encoder.state(), 1u32 << 23);
615 }
616
617 #[test]
620 fn test_rans64_initial_state() {
621 let sink = VecSink::<u32>::new(64);
622 let encoder = Rans64Encoder::new(sink);
623 assert_eq!(encoder.state(), 1u64 << 31);
624 }
625
626 #[test]
627 fn test_rans64_raw_roundtrip() {
628 let sink = VecSink::<u32>::new(64);
629 let mut encoder = Rans64Encoder::new(sink);
630 encoder.put_raw(0, 128, 8);
631 encoder.flush();
632 let encoded = encoder.into_sink().encoded().to_vec();
633 assert!(!encoded.is_empty());
634
635 let source = SliceSource::new(&encoded[..]);
636 let mut decoder = Rans64Decoder::new(source);
637 assert!(decoder.init());
638 assert!(decoder.advance(0, 128, 8));
639 assert!(decoder.check_eof());
640 }
641
642 #[test]
643 fn test_rans64_prepared_symbol_matches_raw() {
644 let sink1 = VecSink::<u32>::new(64);
645 let mut enc1 = Rans64Encoder::new(sink1);
646 enc1.put_raw(0, 128, 8);
647 enc1.flush();
648 let out1 = enc1.into_sink().encoded().to_vec();
649
650 let sink2 = VecSink::<u32>::new(64);
651 let mut enc2 = Rans64Encoder::new(sink2);
652 let sym = Rans64EncSymbol::new(0, 128, 8);
653 enc2.put(&sym);
654 enc2.flush();
655 let out2 = enc2.into_sink().encoded().to_vec();
656
657 assert_eq!(out1, out2);
658 }
659
660 #[test]
661 fn test_rans64_encoder_reset() {
662 let sink = VecSink::<u32>::new(64);
663 let mut encoder = Rans64Encoder::new(sink);
664 encoder.put_raw(0, 128, 8);
665 encoder.reset();
666 assert_eq!(encoder.state(), 1u64 << 31);
667 }
668
669 #[test]
670 fn test_scale_bits_32_rejected_for_try_new() {
671 let result = Rans64EncSymbol::try_new(0, 128, 32);
672 assert!(result.is_err(), "scale_bits=32 must be rejected");
673 if let Err(e) = result {
674 match e {
675 crate::error::RawRansError::InvalidScaleBits { provided, max_safe } => {
676 assert_eq!(provided, 32);
677 assert_eq!(max_safe, 31);
678 }
679 _ => panic!("wrong error type"),
680 }
681 }
682 }
683
684 #[test]
685 fn test_scale_bits_32_rejected_for_try_put_raw() {
686 let sink = VecSink::<u32>::new(64);
687 let mut encoder = Rans64Encoder::new(sink);
688 let result = encoder.try_put_raw(0, 128, 32);
689 assert!(result.is_err(), "scale_bits=32 must be rejected");
690 }
691
692 #[test]
693 fn test_scale_bits_31_accepted_for_rans64() {
694 let result = Rans64EncSymbol::try_new(0, 1, 31);
695 assert!(result.is_ok(), "scale_bits=31 must be accepted for Rans64");
696 }
697
698 #[test]
699 fn test_scale_bits_1_rejected() {
700 let result = RansByteEncSymbol::try_new(0, 128, 1);
701 assert!(
702 result.is_err(),
703 "scale_bits=1 must be rejected (minimum is 2)"
704 );
705 }
706
707 #[test]
708 fn test_try_new_rejects_freq_zero() {
709 let result = RansByteEncSymbol::try_new(0, 0, 8);
710 assert!(result.is_err(), "freq=0 must be rejected");
711 }
712
713 #[test]
714 fn test_try_new_accepts_valid_params() {
715 let result = RansByteEncSymbol::try_new(0, 128, 8);
716 assert!(result.is_ok(), "valid params must be accepted");
717 }
718}