开源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表访问裸索引,可能越界带边界检查
错误传播状态码+errnoResult + ?
整数溢出无检查checked/wrapping语义

相关链接