开源C项目重构:中型
一、案例选择:zlib核心模块
zlib是世界最广泛使用的压缩库,约4000行核心代码,被PNG、HTTP、SSH、git等无数项目依赖。解压器处理不受信任输入,是重写的最高价值目标。
二、zlib位读取宏的不安全分析
#define NEEDBITS(n) do { \
while (bits < (unsigned)(n)) { \
hold += (unsigned long)(*next++) << bits; \
bits += 8; \
} \
} while (0)
#define BITS(n) ((unsigned)hold & ((1U << (n)) - 1))
#define DROPBITS(n) do { hold >>= (n); bits -= (n); } while (0)| 宏 | 风险 |
|---|---|
| NEEDBITS(n) | *next++可能越界;n>=32时1U<<n是UB |
| BITS(n) | 依赖调用者确保NEEDBITS已满足 |
| DROPBITS(n) | n>bits时bits变成负数(unsigned underflow) |
| 整体耦合 | hold, bits, next三变量必须保持一致性 |
三、Rust安全BitReader
pub struct BitReader<'a> {
data: &'a [u8], pos: usize, bit_buf: u64, bits_in_buf: u32,
}
impl<'a> BitReader<'a> {
pub fn need_bits(&mut self, n: u32) -> Result<(), DecompressError> {
while self.bits_in_buf < n {
if self.pos >= self.data.len() { return Err(DecompressError::UnexpectedEof); }
self.bit_buf |= (self.data[self.pos] as u64) << self.bits_in_buf;
self.bits_in_buf += 8; self.pos += 1;
}
Ok(())
}
pub fn peek_bits(&self, n: u32) -> u32 { /* 断言 n<=bits_in_buf */ }
pub fn read_bits(&mut self, n: u32) -> Result<u32, DecompressError> {
self.need_bits(n)?; let val = self.peek_bits(n); self.drop_bits(n); Ok(val)
}
pub fn drop_bits(&mut self, n: u32) { /* 断言 n<=bits_in_buf */ }
}消除C版本中NEEDBITS/BITS/DROPBITS的全部不安全因素:越界访问返回Err而非UB;位操作带断言检查;三个字段封装在一个结构体中。
四、Huffman解码器
pub struct HuffmanTable {
table: Vec<HuffmanEntry>,
max_bits: u32,
}
struct HuffmanEntry { symbol: u16, bits: u8 }
impl HuffmanTable {
pub fn fixed_literal_length() -> Self {
let mut table = Vec::with_capacity(288);
for i in 0..=143 { table.push(HuffmanEntry { symbol: i, bits: 8 }); }
for i in 144..=255 { table.push(HuffmanEntry { symbol: i, bits: 9 }); }
for i in 256..=279 { table.push(HuffmanEntry { symbol: i, bits: 7 }); }
for i in 280..=287 { table.push(HuffmanEntry { symbol: i, bits: 8 }); }
HuffmanTable { table, max_bits: 9 }
}
pub fn from_code_lengths(lengths: &[u8]) -> Result<Self, DecompressError> {
// RFC 1951 第3.2.2节: 计数排序构建Huffman表
let mut bl_count = [0u16; 16];
for &len in lengths { if len > 15 { return Err(...); } if len > 0 { bl_count[len as usize] += 1; } }
let mut next_code = [0u16; 16];
let mut code: u16 = 0;
for bits in 1..=15 { code = (code + bl_count[bits-1]) << 1; next_code[bits] = code; }
// ... 分配编码值
Ok(HuffmanTable { table, max_bits })
}
}五、核心解压循环
const LENGTH_BASE: [u16; 29] = [3,4,5,6,7,8,9,10,11,13,15,17,19,23,27,31,35,43,51,59,67,83,99,115,131,163,195,227,258];
const DISTANCE_BASE: [u16; 30] = [1,2,3,4,5,7,9,13,17,25,33,49,65,97,129,193,257,385,513,769,1025,1537,2049,3073,4097,6145,8193,12289,16385,24577];
pub fn inflate_deflate(input: &[u8], output: &mut Vec<u8>) -> Result<(), DecompressError> {
let mut reader = BitReader::new(input);
let mut is_final = false;
while !is_final {
is_final = reader.read_bits(1)? != 0;
match reader.read_bits(2)? {
0 => inflate_stored(&mut reader, output)?,
1 => inflate_huffman_block(&mut reader, output, &fixed_lit, &fixed_dist)?,
2 => { let (lit, dist) = decode_dynamic_tables(&mut reader)?;
inflate_huffman_block(&mut reader, output, &lit, &dist)?; }
3 => return Err(DecompressError::InvalidBlockType),
_ => unreachable!(),
}
}
Ok(())
}六、LZ77回引复制(安全版)
C版本中用裸指针算术 output[copy_start + i] 实现回引复制,无边界检查。
fn inflate_huffman_block(reader: &mut BitReader, output: &mut Vec<u8>,
lit_table: &HuffmanTable, dist_table: &HuffmanTable) -> Result<(), DecompressError>
{
loop {
let symbol = lit_table.decode(reader)?;
match symbol {
0..=255 => output.push(symbol as u8),
256 => break,
257..=285 => {
let length = LENGTH_BASE[(symbol-257) as usize] as usize
+ reader.read_bits(LENGTH_EXTRA_BITS[(symbol-257) as usize] as u32)? as usize;
let dist_symbol = dist_table.decode(reader)?;
let distance = DISTANCE_BASE[dist_symbol as usize] as usize
+ reader.read_bits(DISTANCE_EXTRA_BITS[dist_symbol as usize] as u32)? as usize;
if distance > output.len() { return Err(DecompressError::InvalidDistance); }
let copy_start = output.len() - distance;
for i in 0..length { output.push(output[copy_start + i]); }
}
_ => return Err(DecompressError::InvalidLiteral),
}
}
Ok(())
}distance边界检查确保不会回引到输出缓冲区外部。
七、动态Huffman表解码
fn decode_dynamic_tables(reader: &mut BitReader) -> Result<(HuffmanTable, HuffmanTable), DecompressError> {
let hlit = reader.read_bits(5)? as usize + 257;
let hdist = reader.read_bits(5)? as usize + 1;
let hclen = reader.read_bits(4)? as usize + 4;
const CODE_LENGTH_ORDER: [usize; 19] = [16,17,18,0,8,7,9,6,10,5,11,4,12,3,13,2,14,1,15];
let mut code_length_lengths = [0u8; 19];
for i in 0..hclen { code_length_lengths[CODE_LENGTH_ORDER[i]] = reader.read_bits(3)? as u8; }
let code_table = HuffmanTable::from_code_lengths(&code_length_lengths)?;
let mut lengths = Vec::with_capacity(hlit + hdist);
while lengths.len() < hlit + hdist {
let symbol = code_table.decode(reader)?;
match symbol {
0..=15 => lengths.push(symbol as u8),
16 => { let repeat = reader.read_bits(2)? as usize + 3;
let last = *lengths.last().unwrap_or(&0);
lengths.resize(lengths.len() + repeat, last); }
17 => { let repeat = reader.read_bits(3)? as usize + 3;
lengths.resize(lengths.len() + repeat, 0); }
18 => { let repeat = reader.read_bits(7)? as usize + 11;
lengths.resize(lengths.len() + repeat, 0); }
_ => return Err(DecompressError::HuffmanError),
}
}
let lit_table = HuffmanTable::from_code_lengths(&lengths[..hlit])?;
let dist_table = HuffmanTable::from_code_lengths(&lengths[hlit..hlit+hdist])?;
Ok((lit_table, dist_table))
}八、保持C ABI兼容
#[no_mangle]
pub extern "C" fn inflate(strm: *mut ZStream, flush: c_int) -> c_int {
if strm.is_null() { return Z_STREAM_ERROR; }
let input = unsafe { std::slice::from_raw_parts(strm.next_in, strm.avail_in as usize) };
let mut output = Vec::new();
match inflate_zlib(input) {
Ok(decompressed) => { /* 复制到strm.next_out */ Z_STREAM_END }
Err(_) => Z_DATA_ERROR,
}
}九、安全性对比
| 安全属性 | zlib (C) | Rust重写 |
|---|---|---|
| 缓冲区溢出 | 无防护 | 编译器强制 |
| 位操作安全 | 宏无保护 | 方法带断言 |
| Huffman表访问 | 裸索引,可能越界 | 带边界检查 |
| 错误传播 | 状态码+errno | Result + ? |
| 整数溢出 | 无检查 | checked/wrapping语义 |