Skip to main content

bitio_rs/
writer.rs

1use crate::byte_order::ByteOrder;
2use crate::error::BitReadWriteError;
3use crate::traits::BitWrite;
4use std::io::{BufWriter, Result, Write};
5
6pub struct BitWriter<W: Write> {
7    byte_order: ByteOrder,
8    inner: Option<BufWriter<W>>, // 用 BufWriter<W> 能避免频繁的系统调用
9
10    bits_buffer: u64,
11    bits_in_buffer: usize,
12}
13
14impl<W: Write> BitWriter<W> {
15    pub fn new(inner: W) -> Self {
16        Self::with_byte_order(ByteOrder::BigEndian, inner)
17    }
18
19    pub fn with_byte_order(byte_order: ByteOrder, inner: W) -> Self {
20        Self {
21            byte_order,
22            inner: Option::from(BufWriter::new(inner)),
23            bits_buffer: 0,
24            bits_in_buffer: 0,
25        }
26    }
27}
28
29impl<W: Write> BitWriter<W> {
30    fn inner_mut(&mut self) -> Result<&mut BufWriter<W>> {
31        self.inner
32            .as_mut()
33            .ok_or_else(|| std::io::Error::new(std::io::ErrorKind::Other, "inner writer is gone"))
34    }
35}
36
37impl<W: Write> BitWriter<W> {
38    /// 将对齐的(完整的)字节写入底层的写入器
39    fn write_aligned_bytes_to_inner(&mut self) -> Result<()> {
40        // 先算出有多少对齐的字节待写入底层
41        // 注意本操作只会处理对齐的字节
42        let count = self.bits_in_buffer / 8;
43        if count == 0 {
44            return Ok(());
45        }
46
47        let mut buf = Vec::with_capacity(count);
48        for _ in 0..count {
49            let byte = match self.byte_order {
50                ByteOrder::BigEndian => (self.bits_buffer >> 56) as u8, // 大端序每次都从比特缓冲区左边取 8 位,也就是 1 字节,注意这里没有改变比特缓冲区本身
51                ByteOrder::LittleEndian => self.bits_buffer as u8, // 小端序每次都从比特缓冲区右边取 8 位,也就是 1 字节,注意这里没有改变比特缓冲区本身
52            };
53            buf.push(byte);
54
55            match self.byte_order {
56                ByteOrder::BigEndian => {
57                    self.bits_buffer <<= 8; // 大端序每次从左边取完比特缓冲区 1 字节后,要从左边消除掉已经取出的 8 位,注意这里改变了比特缓冲区本身
58                }
59                ByteOrder::LittleEndian => {
60                    self.bits_buffer >>= 8; // 小端序每次从右边取完比特缓冲区 1 字节后,要从右边消除掉已经取出的 8 位,注意这里改变了比特缓冲区本身
61                }
62            }
63            self.bits_in_buffer -= 8; // 更改比特缓冲区位计数
64        }
65        self.inner_mut()?.write_all(&buf)?; // 一次写多个字节能减少潜在的系统调用
66        Ok(())
67    }
68
69    /// 将比特缓冲区尾部的不足 1 字节的数据写入底层的写入器,注意,这个函数只能在比特缓冲区中剩余位不足 1 字节(8 比特)时调用才有意义
70    fn write_residual_partial_byte_to_inner(&mut self) -> Result<()> {
71        if self.bits_in_buffer > 0 && self.bits_in_buffer < 8 {
72            let byte = match self.byte_order {
73                ByteOrder::BigEndian => (self.bits_buffer >> 56) as u8, // 对于大端序,将比特缓冲区最左边剩余的不足 1 字节的位写入底层的写入器
74                ByteOrder::LittleEndian => self.bits_buffer as u8, // 对于小端序,将比特缓冲区最右边剩余的不足 1 字节的位写入底层的写入器
75            };
76            self.inner_mut()?.write_all(&[byte])?;
77            self.bits_buffer = 0; // 清零比特缓冲区
78            self.bits_in_buffer = 0; // 清零比特缓冲区计数
79        }
80        Ok(())
81    }
82}
83
84impl<W: Write> BitWriter<W> {
85    pub fn into_inner(mut self) -> Result<W> {
86        self.write_residual_partial_byte_to_inner()?;
87        if let Some(mut inner) = self.inner.take() {
88            inner.flush()?;
89            inner.into_inner().map_err(|e| e.into_error())
90        } else {
91            // 如果已经被取走了,返回错误或自定义错误
92            Err(std::io::Error::new(
93                std::io::ErrorKind::Other,
94                "inner already taken",
95            ))
96        }
97    }
98}
99
100impl<W: Write> Write for BitWriter<W> {
101    fn write(&mut self, buf: &[u8]) -> Result<usize> {
102        // 在写入新来的字节组到底层写入器之前,先确保比特缓冲区中对齐的字节被写入底层的写入器
103        self.write_aligned_bytes_to_inner()?;
104
105        if self.bits_in_buffer == 0 {
106            // 如果执行完将比特缓冲区中所有对齐字节都写入底层的写入器后,如果比特缓冲区已经清零(此时已是干净的状态),那么就可以将新来的字节组直接写入底层的写入器(高速)
107            self.inner_mut()?.write(buf)?;
108            return Ok(buf.len());
109        }
110
111        // 如果执行完将比特缓冲区中所有对齐字节都写入底层的写入器后,比特缓冲区中还有剩余的位(也就是未对齐为 1 字节的位,比如 3 比特),那么就需要将字节组的每个字节都执行 “比特写”(在这个过程中实际上是先将所有自己组的字节都写到比特缓冲区然后由后续逻辑从比特缓冲区写到底层写入器,也就是不允许绕过比特缓冲区) 这样才能保证底层写入器是无空隙的(这样速度较字节组直写要慢,但是我们的底层写入器保证是 BufWriter 因此不会慢太多)
112        for &b in buf {
113            self.write_bits(b as u64, 8)?;
114        }
115
116        Ok(buf.len())
117    }
118
119    fn flush(&mut self) -> Result<()> {
120        // 注意冲刷操作一定要把比特缓冲区的残尾字节写入底层写入器,否则底层写入器就少尾部数据了
121        self.write_residual_partial_byte_to_inner()?;
122        self.inner_mut()?.flush()
123    }
124}
125
126impl<W: Write> Drop for BitWriter<W> {
127    fn drop(&mut self) {
128        // 先尝试写入残余的比特数据,忽略错误
129        // 注意这里显式的忽略了错误因为 Rust 规定 Drop 里不允许 panic,同样的,不能直接 self.flush().unwrap(); 因为 .unwrap() 可能会 panic
130        let _ = self.write_residual_partial_byte_to_inner();
131
132        // 访问 inner,如果存在则 flush,忽略错误
133        if let Some(ref mut inner) = self.inner {
134            let _ = inner.flush();
135        }
136        // 如果 inner 是 None,说明已经被 take() 过了,直接跳过即可
137    }
138}
139
140impl<W: Write> BitWrite for BitWriter<W> {
141    fn write_bits(&mut self, value: u64, n: usize) -> Result<()> {
142        // 校验 n
143        if n == 0 || n > 64 {
144            return Err(BitReadWriteError::InvalidBitCount(n).into());
145        }
146
147        let mut remaining = n;
148        let mask = if n == 64 { u64::MAX } else { (1u64 << n) - 1 }; // (1u64 << n) - 1 就是低位连续 n 个 1,高位全是 0
149        let mut val = value & mask; // 用掩码取出 n 位有效位,无效的位被丢弃
150
151        while remaining > 0 {
152            let available = 64 - self.bits_in_buffer;
153            let to_insert = remaining.min(available);
154            let insert_at_next_round = remaining - to_insert;
155            let to_insert_val = val >> insert_at_next_round; // 注意这里没有改变 val 本身,而是用 val 的一部分建立了新值
156
157            match self.byte_order {
158                ByteOrder::BigEndian => {
159                    self.bits_buffer |= to_insert_val << (available - to_insert); // 大端序时是把值从比特缓冲区的左边往右边堆(可以想象比特缓冲区是一个能容纳 64 块砖的长条盒子,大端序就是来一块砖就从左开始码放)
160                }
161                ByteOrder::LittleEndian => {
162                    self.bits_buffer |= to_insert_val << self.bits_in_buffer; // 小端序时是把值从比特缓冲区的右边往左边堆(可以想象比特缓冲区是一个能容纳 64 块砖的长条盒子,小端序就是来一块砖就从右开始码放)
163                }
164            }
165
166            self.bits_in_buffer += to_insert; // 更新比特缓冲区中已有的位数
167            remaining -= to_insert; // 更新剩余的要插入的位数
168
169            if insert_at_next_round > 0 {
170                val &= (1u64 << insert_at_next_round) - 1; //  (1u64 << insert_at_next_round) - 1 又是一个掩码,用下一轮要插入的位数来更新 val,相当于丢弃了 val 中本轮已经插入过的位,注意这里是直接修改了 val 本身
171            }
172
173            // 每凑够(包括大于的情况)1 字节就触发一次写入底层写入器的操作
174            if self.bits_in_buffer >= 8 || remaining == 0 {
175                self.write_aligned_bytes_to_inner()?; // 注意只能将对其的部分写入底层写入器,如果将未对齐的也写入了,后续再有新的字节组过来时,底层写入器就会因为本次写入了部分字节后出现位的断档
176            }
177        }
178
179        Ok(())
180    }
181}