dsi_bitstream/impls/
bit_reader.rs1use core::convert::Infallible;
10#[cfg(feature = "mem_dbg")]
11use mem_dbg::{MemDbg, MemSize};
12
13use crate::codes::params::{DefaultReadParams, ReadParams};
14use crate::traits::*;
15
16#[derive(Debug, Clone)]
39#[cfg_attr(feature = "mem_dbg", derive(MemDbg, MemSize))]
40pub struct BitReader<E: Endianness, WR, RP: ReadParams = DefaultReadParams> {
41 backend: WR,
43 bit_index: u64,
45 _marker: core::marker::PhantomData<(E, RP)>,
46}
47
48impl<E: Endianness, WR, RP: ReadParams> BitReader<E, WR, RP> {
49 #[must_use]
51 pub const fn new(backend: WR) -> Self {
52 Self {
53 backend,
54 bit_index: 0,
55 _marker: core::marker::PhantomData,
56 }
57 }
58}
59
60impl<WR: WordRead<Word = u64> + WordSeek<Error = <WR as WordRead>::Error>, RP: ReadParams>
61 BitRead<BE> for BitReader<BE, WR, RP>
62{
63 type Error = <WR as WordRead>::Error;
64 type PeekWord = u32;
65 const PEEK_BITS: usize = 32;
66
67 #[inline]
68 fn skip_bits(&mut self, n_bits: usize) -> Result<(), Self::Error> {
69 self.bit_index += n_bits as u64;
70 Ok(())
71 }
72
73 #[inline]
74 fn read_bits(&mut self, num_bits: usize) -> Result<u64, Self::Error> {
75 debug_assert!(num_bits <= 64);
76 #[cfg(feature = "checks")]
77 assert!(num_bits <= 64);
78
79 if num_bits == 0 {
80 return Ok(0);
81 }
82
83 self.backend.set_word_pos(self.bit_index / 64)?;
84 let in_word_offset = (self.bit_index % 64) as usize;
85
86 let res = if (in_word_offset + num_bits) <= 64 {
87 let word = self.backend.read_word()?.to_be();
89 (word << in_word_offset) >> (64 - num_bits)
90 } else {
91 let high_word = self.backend.read_word()?.to_be();
93 let low_word = self.backend.read_word()?.to_be();
94 let shamt1 = 64 - num_bits;
95 let shamt2 = 128 - in_word_offset - num_bits;
96 ((high_word << in_word_offset) >> shamt1) | (low_word >> shamt2)
97 };
98 self.bit_index += num_bits as u64;
99 Ok(res)
100 }
101
102 #[inline]
103 fn peek_bits(&mut self, n_bits: usize) -> Result<u32, Self::Error> {
104 if n_bits == 0 {
105 return Ok(0);
106 }
107
108 #[cfg(feature = "checks")]
109 assert!(n_bits <= 32);
110
111 self.backend.set_word_pos(self.bit_index / 64)?;
112 let in_word_offset = (self.bit_index % 64) as usize;
113
114 let res = if (in_word_offset + n_bits) <= 64 {
115 let word = self.backend.read_word()?.to_be();
117 (word << in_word_offset) >> (64 - n_bits)
118 } else {
119 let high_word = self.backend.read_word()?.to_be();
121 let low_word = self.backend.read_word()?.to_be();
122 let shamt1 = 64 - n_bits;
123 let shamt2 = 128 - in_word_offset - n_bits;
124 ((high_word << in_word_offset) >> shamt1) | (low_word >> shamt2)
125 };
126 Ok(res as u32)
127 }
128
129 #[inline]
130 fn read_unary(&mut self) -> Result<u64, Self::Error> {
131 self.backend.set_word_pos(self.bit_index / 64)?;
132 let in_word_offset = self.bit_index % 64;
133 let mut bits_in_word = 64 - in_word_offset;
134 let mut total = 0;
135
136 let mut word = self.backend.read_word()?.to_be();
137 word <<= in_word_offset;
138 loop {
139 let zeros = word.leading_zeros() as u64;
140 if zeros < bits_in_word {
142 self.bit_index += total + zeros + 1;
143 return Ok(total + zeros);
144 }
145 total += bits_in_word;
146 bits_in_word = 64;
147 word = self.backend.read_word()?.to_be();
148 }
149 }
150
151 #[inline(always)]
152 fn skip_bits_after_peek(&mut self, n: usize) {
153 self.bit_index += n as u64;
154 }
155}
156
157impl<E: Endianness, WR: WordSeek, RP: ReadParams> BitSeek for BitReader<E, WR, RP> {
158 type Error = Infallible;
159
160 fn bit_pos(&mut self) -> Result<u64, Self::Error> {
161 Ok(self.bit_index)
162 }
163
164 fn set_bit_pos(&mut self, bit_index: u64) -> Result<(), Self::Error> {
165 self.bit_index = bit_index;
166 Ok(())
167 }
168}
169
170impl<WR: WordRead<Word = u64> + WordSeek<Error = <WR as WordRead>::Error>, RP: ReadParams>
171 BitRead<LE> for BitReader<LE, WR, RP>
172{
173 type Error = <WR as WordRead>::Error;
174 type PeekWord = u32;
175 const PEEK_BITS: usize = 32;
176
177 #[inline]
178 fn skip_bits(&mut self, n_bits: usize) -> Result<(), Self::Error> {
179 self.bit_index += n_bits as u64;
180 Ok(())
181 }
182
183 #[inline]
184 fn read_bits(&mut self, num_bits: usize) -> Result<u64, Self::Error> {
185 #[cfg(feature = "checks")]
186 assert!(num_bits <= 64);
187
188 if num_bits == 0 {
189 return Ok(0);
190 }
191
192 self.backend.set_word_pos(self.bit_index / 64)?;
193 let in_word_offset = (self.bit_index % 64) as usize;
194
195 let res = if (in_word_offset + num_bits) <= 64 {
196 let word = self.backend.read_word()?.to_le();
198 let shamt = 64 - num_bits;
199 (word << (shamt - in_word_offset)) >> shamt
200 } else {
201 let low_word = self.backend.read_word()?.to_le();
203 let high_word = self.backend.read_word()?.to_le();
204 let shamt1 = 128 - in_word_offset - num_bits;
205 let shamt2 = 64 - num_bits;
206 ((high_word << shamt1) >> shamt2) | (low_word >> in_word_offset)
207 };
208 self.bit_index += num_bits as u64;
209 Ok(res)
210 }
211
212 #[inline]
213 fn peek_bits(&mut self, n_bits: usize) -> Result<u32, Self::Error> {
214 if n_bits == 0 {
215 return Ok(0);
216 }
217
218 #[cfg(feature = "checks")]
219 assert!(n_bits <= 32);
220
221 self.backend.set_word_pos(self.bit_index / 64)?;
222 let in_word_offset = (self.bit_index % 64) as usize;
223
224 let res = if (in_word_offset + n_bits) <= 64 {
225 let word = self.backend.read_word()?.to_le();
227 let shamt = 64 - n_bits;
228 (word << (shamt - in_word_offset)) >> shamt
229 } else {
230 let low_word = self.backend.read_word()?.to_le();
232 let high_word = self.backend.read_word()?.to_le();
233 let shamt1 = 128 - in_word_offset - n_bits;
234 let shamt2 = 64 - n_bits;
235 ((high_word << shamt1) >> shamt2) | (low_word >> in_word_offset)
236 };
237 Ok(res as u32)
238 }
239
240 #[inline]
241 fn read_unary(&mut self) -> Result<u64, Self::Error> {
242 self.backend.set_word_pos(self.bit_index / 64)?;
243 let in_word_offset = self.bit_index % 64;
244 let mut bits_in_word = 64 - in_word_offset;
245 let mut total = 0;
246
247 let mut word = self.backend.read_word()?.to_le();
248 word >>= in_word_offset;
249 loop {
250 let zeros = word.trailing_zeros() as u64;
251 if zeros < bits_in_word {
253 self.bit_index += total + zeros + 1;
254 return Ok(total + zeros);
255 }
256 total += bits_in_word;
257 bits_in_word = 64;
258 word = self.backend.read_word()?.to_le();
259 }
260 }
261
262 #[inline(always)]
263 fn skip_bits_after_peek(&mut self, n: usize) {
264 self.bit_index += n as u64;
265 }
266}
267
268#[cfg(feature = "std")]
269impl<WR: WordRead<Word = u64> + WordSeek<Error = <WR as WordRead>::Error>, RP: ReadParams>
270 std::io::Read for BitReader<LE, WR, RP>
271{
272 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
280 let mut read = 0;
281 let mut iter = buf.chunks_exact_mut(8);
282
283 for chunk in &mut iter {
284 match self.read_bits(64) {
285 Ok(word) => {
286 chunk.copy_from_slice(&word.to_le_bytes());
287 read += 8;
288 }
289 Err(_) if read > 0 => return Ok(read),
292 Err(e) => {
293 return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, e));
294 }
295 }
296 }
297
298 let rem = iter.into_remainder();
299 if !rem.is_empty() {
300 match self.read_bits(rem.len() * 8) {
301 Ok(word) => {
302 rem.copy_from_slice(&word.to_le_bytes()[..rem.len()]);
303 read += rem.len();
304 }
305 Err(_) if read > 0 => return Ok(read),
306 Err(e) => {
307 return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, e));
308 }
309 }
310 }
311
312 Ok(read)
313 }
314}
315
316#[cfg(feature = "std")]
317impl<WR: WordRead<Word = u64> + WordSeek<Error = <WR as WordRead>::Error>, RP: ReadParams>
318 std::io::Read for BitReader<BE, WR, RP>
319{
320 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
328 let mut read = 0;
329 let mut iter = buf.chunks_exact_mut(8);
330
331 for chunk in &mut iter {
332 match self.read_bits(64) {
333 Ok(word) => {
334 chunk.copy_from_slice(&word.to_be_bytes());
335 read += 8;
336 }
337 Err(_) if read > 0 => return Ok(read),
340 Err(e) => {
341 return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, e));
342 }
343 }
344 }
345
346 let rem = iter.into_remainder();
347 if !rem.is_empty() {
348 match self.read_bits(rem.len() * 8) {
349 Ok(word) => {
350 rem.copy_from_slice(&word.to_be_bytes()[8 - rem.len()..]);
351 read += rem.len();
352 }
353 Err(_) if read > 0 => return Ok(read),
354 Err(e) => {
355 return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, e));
356 }
357 }
358 }
359
360 Ok(read)
361 }
362}