1use std::io::{BufReader, Read};
9
10use crate::error::Result;
11
12const RANGE_MAX: u16 = 0xFFFF;
14
15const MSB_MASK: u16 = 0x8000;
17
18const UNDERFLOW_MASK: u16 = 0x4000;
20
21pub struct ArithmeticDecoder<R: Read> {
27 input: BufReader<R>,
29 high: u16,
31 low: u16,
33 code: u16,
35 byte_buffer: u8,
37 bits_remaining: u8,
39}
40
41impl<R: Read> ArithmeticDecoder<R> {
42 #[inline]
44 pub fn new(reader: R) -> Result<Self> {
45 let mut input = BufReader::new(reader);
46
47 let mut initial_bytes = [0u8; 2];
49 input.read_exact(&mut initial_bytes)?;
50 let code = u16::from_be_bytes(initial_bytes);
51
52 Ok(Self {
53 input,
54 high: RANGE_MAX,
55 low: 0,
56 code,
57 byte_buffer: 0,
58 bits_remaining: 0,
59 })
60 }
61
62 #[inline(always)]
64 fn read_bit(&mut self) -> u16 {
65 if self.bits_remaining == 0 {
66 let mut byte = [0u8; 1];
67 if self.input.read_exact(&mut byte).is_ok() {
69 self.byte_buffer = byte[0];
70 } else {
71 self.byte_buffer = 0;
72 }
73 self.bits_remaining = 8;
74 }
75
76 self.bits_remaining -= 1;
77 ((self.byte_buffer >> self.bits_remaining) & 1) as u16
78 }
79
80 #[inline]
85 pub fn threshold_val(&self, total: u16) -> u16 {
86 let range = (self.high - self.low) as u32 + 1;
88 let offset = (self.code - self.low) as u32 + 1;
90 ((offset * total as u32 - 1) / range) as u16
91 }
92
93 #[inline]
99 pub fn decode_update(&mut self, cum_low: u16, cum_high: u16, total: u16) -> Result<()> {
100 let range = (self.high - self.low) as u32 + 1;
101 let scale = total as u32;
102
103 let new_high = self.low.wrapping_add(((range * cum_high as u32 / scale) - 1) as u16);
105 let new_low = self.low.wrapping_add((range * cum_low as u32 / scale) as u16);
106
107 self.high = new_high;
108 self.low = new_low;
109
110 self.renormalize();
112
113 Ok(())
114 }
115
116 #[inline(always)]
118 fn renormalize(&mut self) {
119 loop {
120 if (self.high ^ self.low) & MSB_MASK == 0 {
121 self.shift_out_msb();
123 } else if (self.low & UNDERFLOW_MASK) != 0 && (self.high & UNDERFLOW_MASK) == 0 {
124 self.handle_underflow();
127 } else {
128 break;
130 }
131 }
132 }
133
134 #[inline(always)]
136 fn shift_out_msb(&mut self) {
137 self.low <<= 1;
138 self.high = (self.high << 1) | 1;
139 self.code = (self.code << 1) | self.read_bit();
140 }
141
142 #[inline(always)]
144 fn handle_underflow(&mut self) {
145 self.low = (self.low << 1) & 0x7FFF;
147 self.high = (self.high << 1) | 0x8001;
148 self.code = ((self.code << 1) ^ MSB_MASK) | self.read_bit();
150 }
151}
152
153#[cfg(test)]
154mod tests {
155 use super::*;
156 use std::io::Cursor;
157
158 #[test]
159 fn test_initialization() {
160 let data = vec![0xAB, 0xCD];
161 let decoder = ArithmeticDecoder::new(Cursor::new(data)).unwrap();
162
163 assert_eq!(decoder.low, 0);
164 assert_eq!(decoder.high, 0xFFFF);
165 assert_eq!(decoder.code, 0xABCD);
166 }
167
168 #[test]
169 fn test_threshold_midpoint() {
170 let data = vec![0x80, 0x00];
172 let decoder = ArithmeticDecoder::new(Cursor::new(data)).unwrap();
173
174 assert_eq!(decoder.threshold_val(256), 128);
176 }
177
178 #[test]
179 fn test_threshold_boundaries() {
180 let data = vec![0x00, 0x00];
182 let decoder = ArithmeticDecoder::new(Cursor::new(data)).unwrap();
183 assert_eq!(decoder.threshold_val(100), 0);
184
185 let data = vec![0xFF, 0xFF];
187 let decoder = ArithmeticDecoder::new(Cursor::new(data)).unwrap();
188 assert_eq!(decoder.threshold_val(100), 99);
189 }
190}