Skip to main content

bitio_rs/
reader.rs

1use crate::byte_order::ByteOrder;
2use crate::error::BitReadWriteError;
3use crate::traits::{BitPeek, BitRead};
4use std::io::{BufReader, Read};
5
6// ------------------------------- BitReader ------------------------------- //
7
8pub struct BitReader<R: Read> {
9    byte_order: ByteOrder,
10    inner: BufReader<R>,
11
12    bits_buffer: u64, // 比特缓冲区:rust 中并没有表达 "一系列比特" 的具名数据结构,但是事实上 u64 就可以表达一系列比特
13    bits_in_buffer: usize, // 当前比特缓冲区中持有的比特数
14}
15
16impl<R: Read> BitReader<R> {
17    pub fn new(inner: R) -> Self {
18        Self::with_byte_order(ByteOrder::BigEndian, inner)
19    }
20
21    pub fn with_byte_order(byte_order: ByteOrder, inner: R) -> Self {
22        Self {
23            byte_order,
24            inner: BufReader::new(inner),
25            bits_buffer: 0,
26            bits_in_buffer: 0,
27        }
28    }
29}
30
31impl<R: Read> BitReader<R> {
32    fn put_into_bits_buffer(&mut self, n: usize) -> std::io::Result<()> {
33        let bits_needed = n.saturating_sub(self.bits_in_buffer); // 使用 saturating_sub 防止下溢
34        let mut bytes_needed = (bits_needed + 7) / 8; // 这是一种常见的 向上取整除法技巧(Ceiling Division Trick),用于计算容纳指定位数所需的最小字节数(当`bits_needed`不是8的倍数时,加上7就会使得总和至少达到下一个8的倍数,从而在除以8时得到正确地向上取整的结果)
35        let max_bytes_needed = (64 - self.bits_in_buffer) / 8;
36        if bytes_needed > max_bytes_needed {
37            bytes_needed = max_bytes_needed;
38        }
39        if bytes_needed > 0 {
40            let mut buf = [0u8; 8]; // 注意这里没有用 vector(堆上分配) 而是使用了栈上分配数组,这是一个性能优化
41            let slice = &mut buf[..bytes_needed];
42            if self.inner.read(slice)? < bytes_needed {
43                return Err(BitReadWriteError::UnexpectedEof.into());
44            };
45            for &mut b in slice {
46                // 所谓低地址就是如果顺序的将一块字流读取出来,首个字节索引是 0,第二个字节索引是 1,以此类推,0 就是低地址,也就是最读到的(索引最大的那个)必然是高地址
47                // 大端序时来的数据越晚,左移的位数就越少,这样最后一个数据(最高地址数据)就在最右边(最低位)
48                // 小端序时来的数据越晚,左移的位数就越多,这样最后一个数据(最高地址数据)就在最左边(最高位)
49                let shift = match self.byte_order {
50                    ByteOrder::BigEndian => {
51                        // 大端序的低位字节存储在高地址,高位字节存储在低地址
52                        // 大端序读取时,新读到数据(高地址数据)总是放置在比特缓冲区剩余空间的最低位(最右边)
53                        let s = 64u32 - 8u32 - self.bits_in_buffer as u32; // shift = 64 - 8 - available_bits
54                        s
55                    }
56                    ByteOrder::LittleEndian => {
57                        // 小端序的低位字节存储在低地址,高位字节存储在高地址
58                        // 小端序读取时,新读到数据(高地址数据)总是要放置在比特缓冲区的最高位(最左边)
59                        let s = self.bits_in_buffer as u32;
60                        s
61                    }
62                };
63                // 将新读到数据(高地址数据)左移 shift 位,然后与比特缓冲区进行或运算,这样就是将新数据放到了比特缓冲区的最高位(最左边)
64                self.bits_buffer |= u64::from(b).wrapping_shl(shift);
65                // 更新比特缓冲区可用比特数
66                self.bits_in_buffer = (self.bits_in_buffer + 8).min(64);
67            }
68        }
69        Ok(())
70    }
71
72    fn get_from_bits_buffer(&mut self, n: usize, take: bool) -> std::io::Result<u64> {
73        let bit_value = match self.byte_order {
74            ByteOrder::BigEndian => {
75                // 提取比特缓冲区高位 n 位(从左数的 n 位)
76                let value = self.bits_buffer >> (64 - n);
77                value
78            }
79            ByteOrder::LittleEndian => {
80                // 用位掩码提取低 n 位
81                let mask = if n == 64 { u64::MAX } else { (1u64 << n) - 1 };
82                let value = self.bits_buffer & mask;
83                value
84            }
85        };
86        if take {
87            if n == 64 {
88                self.bits_buffer = 0;
89            } else {
90                match self.byte_order {
91                    ByteOrder::BigEndian => {
92                        self.bits_buffer <<= n;
93                    }
94                    ByteOrder::LittleEndian => {
95                        self.bits_buffer >>= n;
96                    }
97                }
98            }
99
100            self.bits_in_buffer -= n;
101        }
102        Ok(bit_value)
103    }
104}
105
106impl<R: Read> BitReader<R> {
107    /// Returns `true` if at byte boundary (no pending bits)
108    ///
109    /// When true:
110    /// - `read()` operations are permitted
111    /// - Next bit read will start from a fresh byte
112    pub fn is_byte_aligned(&self) -> bool {
113        self.bits_in_buffer % 8 == 0
114    }
115}
116
117impl<R: Read> BitRead for BitReader<R> {
118    type Output = u64;
119
120    /// Reads exactly `n` bits from the stream (1-64 bits)
121    ///
122    /// # Arguments
123    /// * `n` - Number of bits to read (1 to 64 inclusive)
124    ///
125    /// # Returns
126    /// Bits read
127    ///
128    /// # Errors
129    /// Returns error if `n` is not between 1-64 or not enough bits are available
130    fn read_bits(&mut self, n: usize) -> std::io::Result<Self::Output> {
131        // 校验 n
132        if n == 0 || n > 64 {
133            return Err(BitReadWriteError::InvalidBitCount(n).into());
134        }
135
136        // 填充比特缓冲区
137        self.put_into_bits_buffer(n)?;
138
139        // 从比特缓冲区取 n 比特,并且消费掉
140        self.get_from_bits_buffer(n, true)
141    }
142}
143
144impl<R: Read> Read for BitReader<R> {
145    /// Reads bytes from the underlying bit stream.
146    ///
147    /// This method behaves differently depending on the bit buffer state:
148    /// - When the bit buffer is **empty** (byte-aligned state), it delegates directly to the inner reader
149    /// - When the bit buffer contains **unconsumed bits** (non-byte-aligned state), it returns an
150    ///   [`UnalignedAccess`](BitReadWriteError::UnalignedAccess) error
151    ///
152    /// # Byte Alignment Requirement
153    /// The fundamental reason for this behavior is **bit stream integrity**. When partially consumed
154    /// bits exist in the buffer:
155    /// 1. Direct byte access would bypass the bit buffer's state tracking
156    /// 2. Reading bytes would consume underlying bytes that contain the *remaining portions* of
157    ///    partially read bit sequences
158    /// 3. This would irreversibly corrupt the bit-level parsing state
159    ///
160    /// # Correct Usage
161    /// To read byte data:
162    /// 1. Use `is_byte_aligned()` to check if you're in byte-aligned state
163    /// 2. For mixed bit/byte reading, always consume all bits in the current byte before reading bytes
164    ///
165    /// # Errors
166    /// Returns `BitReadWriteError::UnalignedAccess` (wrapped in `io::Error`) when called with
167    /// non-empty bit buffer. This prevents:
168    /// - Undefined state transitions
169    /// - Silent data corruption
170    /// - Loss of partially buffered bits
171    fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
172        let mut written = 0;
173
174        // 1) 如果完全空,直接读取
175        if self.bits_in_buffer == 0 {
176            return self.inner.read(buf);
177        }
178
179        // 2) 如果有残留,但已经是整字节边界(8 的倍数),先拆 buffer
180        if self.bits_in_buffer % 8 == 0 {
181            // 一直拆,直到 buffer 中不再有整字节或 buf 写满
182            while self.bits_in_buffer >= 8 && written < buf.len() {
183                // 每次取 8 位并消费
184                let byte = self.get_from_bits_buffer(8, true)? as u8;
185                buf[written] = byte;
186                written += 1;
187            }
188
189            // 拆完后,buffer 要么空,要么剩 <8 位(此处一定是空,因为 bits_in_buffer%8==0)
190            // 剩余 buf 空间,再走一次底层读以获取后续字节
191            if written < buf.len() {
192                let n = self.inner.read(&mut buf[written..])?;
193                written += n;
194            }
195
196            return Ok(written);
197        }
198
199        // 3) 剩余 bits 不是 8 的倍数 —— 非字节对齐,禁止直接读
200        Err(BitReadWriteError::UnalignedAccess.into())
201    }
202}
203
204// ------------------------------- PeekableBitReader ------------------------------- //
205
206pub struct PeekableBitReader<R: Read> {
207    inner: BitReader<R>,
208}
209
210impl<R: Read> PeekableBitReader<R> {
211    pub fn new(inner: R) -> Self {
212        Self {
213            inner: BitReader::new(inner),
214        }
215    }
216
217    pub fn with_byte_order(inner: R) -> Self {
218        Self {
219            inner: BitReader::with_byte_order(ByteOrder::LittleEndian, inner),
220        }
221    }
222}
223
224impl<R: Read> BitRead for PeekableBitReader<R> {
225    type Output = u64;
226
227    fn read_bits(&mut self, n: usize) -> std::io::Result<Self::Output> {
228        self.inner.read_bits(n)
229    }
230}
231
232impl<R: Read> BitPeek for PeekableBitReader<R> {
233    type Output = u64;
234
235    fn peek_bits(&mut self, n: usize) -> std::io::Result<Self::Output> {
236        if n == 0 || n > 64 {
237            return Err(BitReadWriteError::InvalidBitCount(n).into());
238        }
239
240        // 填充比特缓冲区
241        self.inner.put_into_bits_buffer(n)?;
242
243        // 从比特缓冲区取 n 比特,但是并不消费掉
244        self.inner.get_from_bits_buffer(n, false)
245    }
246}
247
248// ------------------------------- BulkBitReader ------------------------------- //
249
250pub struct BulkBitReader<R: Read> {
251    inner: BitReader<R>,
252}
253
254impl<R: Read> BulkBitReader<R> {
255    pub fn new(inner: R) -> Self {
256        Self {
257            inner: BitReader::new(inner),
258        }
259    }
260
261    pub fn with_endianness(endianness: ByteOrder, inner: R) -> Self {
262        Self {
263            inner: BitReader::with_byte_order(endianness, inner),
264        }
265    }
266}
267
268impl<R: Read> BitRead for BulkBitReader<R> {
269    type Output = Vec<u64>;
270
271    fn read_bits(&mut self, n: usize) -> std::io::Result<Self::Output> {
272        if n == 0 {
273            return Err(BitReadWriteError::InvalidBitCount(n).into());
274        }
275        let mut remaining = n;
276        let mut chunks = Vec::with_capacity((n + 63) / 64);
277        while remaining > 0 {
278            let take = remaining.min(64);
279            chunks.push(self.inner.read_bits(take)?);
280            remaining -= take;
281        }
282        Ok(chunks)
283    }
284}