use alloc::vec::Vec;
use crate::error::{Error, Result};
#[derive(Debug, Clone, Copy)]
pub(crate) struct BitReaderState {
byte_index: usize,
bit_buffer: u64,
bit_count: u8,
}
#[derive(Debug)]
pub(crate) struct BitReader<'a> {
input: &'a [u8],
byte_index: usize,
bit_buffer: u64,
bit_count: u8,
}
impl<'a> BitReader<'a> {
#[cfg(test)]
pub(crate) fn new(input: &'a [u8]) -> Self {
Self::new_seeded(input, 0, 0)
}
pub(crate) fn new_seeded(input: &'a [u8], bit_buffer: u64, bit_count: u8) -> Self {
Self {
input,
byte_index: 0,
bit_buffer,
bit_count,
}
}
pub(crate) fn available_bits(&self) -> usize {
self.bit_count as usize + (self.input.len() - self.byte_index) * 8
}
pub(crate) fn residual_bit_buffer(&self) -> u64 {
self.bit_buffer
}
pub(crate) fn residual_bit_count(&self) -> u8 {
self.bit_count
}
pub(crate) fn read_bit(&mut self) -> Result<bool> {
Ok(self.read_bits(1)? != 0)
}
pub(crate) fn read_bits(&mut self, bit_count: u8) -> Result<u16> {
let bits = self.peek_bits(bit_count)?;
self.skip_bits(bit_count)?;
Ok(bits)
}
pub(crate) fn peek_bits(&mut self, bit_count: u8) -> Result<u16> {
while self.bit_count < bit_count {
let Some(&next) = self.input.get(self.byte_index) else {
return Err(Error::InvalidData(
"unexpected end of deflate stream".into(),
));
};
self.bit_buffer |= u64::from(next) << self.bit_count;
self.bit_count += 8;
self.byte_index += 1;
}
Ok((self.bit_buffer & ((1u64 << bit_count) - 1)) as u16)
}
pub(crate) fn skip_bits(&mut self, bit_count: u8) -> Result<()> {
if self.bit_count < bit_count {
self.peek_bits(bit_count)?;
}
self.bit_buffer >>= bit_count;
self.bit_count -= bit_count;
Ok(())
}
pub(crate) fn align_to_byte(&mut self) {
let extra = self.bit_count % 8;
self.bit_buffer >>= extra;
self.bit_count -= extra;
}
pub(crate) fn read_bytes(&mut self, len: usize) -> Result<&'a [u8]> {
debug_assert_eq!(self.bit_count % 8, 0);
let buffered_bytes = (self.bit_count / 8) as usize;
let start = self.byte_index - buffered_bytes;
self.bit_buffer = 0;
self.bit_count = 0;
let end = start + len;
let Some(bytes) = self.input.get(start..end) else {
return Err(Error::InvalidData(
"unexpected end of deflate stream".into(),
));
};
self.byte_index = end;
Ok(bytes)
}
pub(crate) fn snapshot(&self) -> BitReaderState {
BitReaderState {
byte_index: self.byte_index,
bit_buffer: self.bit_buffer,
bit_count: self.bit_count,
}
}
pub(crate) fn restore(&mut self, state: BitReaderState) {
self.byte_index = state.byte_index;
self.bit_buffer = state.bit_buffer;
self.bit_count = state.bit_count;
}
pub(crate) fn committed_bytes(&self) -> usize {
self.byte_index
}
}
#[derive(Debug)]
pub(crate) struct BitWriter<'a> {
output: &'a mut Vec<u8>,
bit_buffer: u64,
bit_count: u8,
}
impl<'a> BitWriter<'a> {
#[cfg(test)]
pub(crate) fn new(output: &'a mut Vec<u8>) -> Self {
Self::new_seeded(output, 0, 0)
}
pub(crate) fn new_seeded(output: &'a mut Vec<u8>, bit_buffer: u64, bit_count: u8) -> Self {
Self {
output,
bit_buffer,
bit_count,
}
}
pub(crate) fn write_bit(&mut self, bit: bool) {
self.write_bits(1, u16::from(bit));
}
pub(crate) fn write_bits(&mut self, bit_count: u8, bits: u16) {
debug_assert!(bit_count <= 16);
self.bit_buffer |= u64::from(bits) << self.bit_count;
self.bit_count += bit_count;
while self.bit_count >= 8 {
self.output.push(self.bit_buffer as u8);
self.bit_buffer >>= 8;
self.bit_count -= 8;
}
}
pub(crate) fn align_to_byte(&mut self) {
if self.bit_count > 0 {
self.output.push(self.bit_buffer as u8);
self.bit_buffer = 0;
self.bit_count = 0;
}
}
pub(crate) fn bit_state(self) -> (u64, u8) {
(self.bit_buffer, self.bit_count)
}
pub(crate) fn finish(mut self) {
self.align_to_byte();
}
}
#[cfg(test)]
mod tests {
use alloc::vec;
use alloc::vec::Vec;
use super::{BitReader, BitWriter};
#[test]
fn writer_basic() {
let mut out = Vec::new();
let mut w = BitWriter::new(&mut out);
w.write_bits(3, 0b101);
w.write_bits(5, 0b10011);
w.finish();
assert_eq!(out, vec![0x9D]);
}
#[test]
fn writer_multibyte() {
let mut out = Vec::new();
let mut w = BitWriter::new(&mut out);
w.write_bits(16, 0xABCD);
w.finish();
assert_eq!(out, vec![0xCD, 0xAB]);
}
#[test]
fn reader_roundtrip_with_writer() {
let mut out = Vec::new();
let mut w = BitWriter::new(&mut out);
w.write_bits(1, 1);
w.write_bits(3, 0b010);
w.write_bits(5, 0b11001);
w.write_bits(7, 0b0110110);
w.finish();
let mut r = BitReader::new(&out);
assert_eq!(r.read_bits(1).unwrap(), 1);
assert_eq!(r.read_bits(3).unwrap(), 0b010);
assert_eq!(r.read_bits(5).unwrap(), 0b11001);
assert_eq!(r.read_bits(7).unwrap(), 0b0110110);
}
#[test]
fn reader_eof_errors() {
let mut r = BitReader::new(&[0x01]);
r.read_bits(8).unwrap();
assert!(r.read_bits(1).is_err());
}
#[test]
fn reader_snapshot_restore() {
let data = [0xAB, 0xCD, 0xEF];
let mut r = BitReader::new(&data);
let snap = r.snapshot();
assert_eq!(r.read_bits(4).unwrap(), 0xB);
assert_eq!(r.read_bits(4).unwrap(), 0xA);
r.restore(snap);
assert_eq!(r.read_bits(8).unwrap(), 0xAB);
assert_eq!(r.read_bits(8).unwrap(), 0xCD);
}
#[test]
fn align_to_byte_discards_fraction() {
let data = [0x34, 0x12];
let mut r = BitReader::new(&data);
r.read_bits(4).unwrap();
r.align_to_byte();
assert_eq!(r.read_bytes(1).unwrap(), &[0x12]);
}
#[test]
fn committed_bytes_counts_loaded_bytes() {
let data = [0xAA, 0xBB, 0xCC];
let mut r = BitReader::new(&data);
assert_eq!(r.committed_bytes(), 0);
r.read_bits(4).unwrap();
assert_eq!(r.committed_bytes(), 1);
r.read_bits(4).unwrap();
assert_eq!(r.committed_bytes(), 1);
r.read_bits(1).unwrap();
assert_eq!(r.committed_bytes(), 2);
}
#[test]
fn available_bits() {
let data = [0x00, 0x00];
let mut r = BitReader::new(&data);
assert_eq!(r.available_bits(), 16);
r.read_bits(3).unwrap();
assert_eq!(r.available_bits(), 13);
}
}