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}