use crate::error::{Result, fmt};
pub struct BitReader<'a> {
bytes: &'a [u8],
byte_pos: usize,
window: u64,
bits_in_window: u32,
bits_read: u64,
bits_total: u64,
}
impl<'a> BitReader<'a> {
pub fn new(bytes: &'a [u8]) -> Self {
Self {
bytes,
byte_pos: 0,
window: 0,
bits_in_window: 0,
bits_read: 0,
bits_total: (bytes.len() as u64) * 8,
}
}
pub fn bit_position(&self) -> u64 {
self.bits_read
}
pub fn bytes_consumed(&self) -> usize {
self.bits_read.div_ceil(8) as usize
}
#[inline]
pub fn read_bit(&mut self) -> Result<u8> {
if self.bits_read >= self.bits_total {
return Err(fmt!(ProtocolError, "BitReader: read past end"));
}
if !self.ensure_bits(1) {
return Err(fmt!(ProtocolError, "BitReader: read past end"));
}
let bit = (self.window & 1) as u8;
self.window >>= 1;
self.bits_in_window -= 1;
self.bits_read += 1;
Ok(bit)
}
#[inline]
pub fn read_bits(&mut self, n: u32) -> Result<u64> {
if n == 0 {
return Ok(0);
}
if n > 64 {
return Err(fmt!(
ProtocolError,
"BitReader: cannot read {} bits into u64",
n
));
}
if self.bits_read + n as u64 > self.bits_total {
return Err(fmt!(ProtocolError, "BitReader: read past end"));
}
let mut result: u64 = 0;
let mut remaining = n;
let mut shift: u32 = 0;
while remaining > 0 {
if self.bits_in_window == 0 {
let want = remaining.min(64);
if !self.ensure_bits(want) {
return Err(fmt!(ProtocolError, "BitReader: read past end"));
}
}
let take = remaining.min(self.bits_in_window);
let mask = if take == 64 {
u64::MAX
} else {
(1u64 << take) - 1
};
result |= (self.window & mask) << shift;
if take == 64 {
self.window = 0;
} else {
self.window >>= take;
}
self.bits_in_window -= take;
remaining -= take;
shift += take;
}
self.bits_read += n as u64;
Ok(result)
}
#[inline]
pub fn read_signed(&mut self, n: u32) -> Result<i64> {
let unsigned = self.read_bits(n)?;
if n == 0 || n == 64 {
return Ok(unsigned as i64);
}
let sign_bit = 1u64 << (n - 1);
let extended = if unsigned & sign_bit != 0 {
unsigned | (u64::MAX << n)
} else {
unsigned
};
Ok(extended as i64)
}
#[inline]
fn ensure_bits(&mut self, want: u32) -> bool {
while self.bits_in_window < want
&& self.bits_in_window <= 56
&& self.byte_pos < self.bytes.len()
{
let b = self.bytes[self.byte_pos] as u64;
self.byte_pos += 1;
self.window |= b << self.bits_in_window;
self.bits_in_window += 8;
}
self.bits_in_window >= want
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::ErrorCode;
#[test]
fn single_bits_lsb_first() {
let bytes = [0b1010_0001u8];
let mut r = BitReader::new(&bytes);
let order = [1, 0, 0, 0, 0, 1, 0, 1];
for (i, expected) in order.iter().enumerate() {
assert_eq!(r.read_bit().unwrap(), *expected, "bit {}", i);
}
assert_eq!(r.read_bit().unwrap_err().code(), ErrorCode::ProtocolError);
}
#[test]
fn read_bits_groups_lsb_first() {
let bytes = [0xAC, 0x02];
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(8).unwrap(), 0xAC);
assert_eq!(r.read_bits(4).unwrap(), 0x02);
}
#[test]
fn read_bits_spans_byte_boundary() {
let bytes = [0xFF, 0x01];
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(12).unwrap(), 0x1FF);
}
#[test]
fn read_signed_sign_extends() {
let bytes = [0x40];
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_signed(7).unwrap(), -64);
let bytes = [0b0011_1111];
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_signed(7).unwrap(), 63);
}
#[test]
fn read_64_bits_works() {
let bytes = 0x0102_0304_0506_0708u64.to_le_bytes();
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(64).unwrap(), 0x0102_0304_0506_0708);
assert!(r.read_bit().is_err()); }
#[test]
fn bit_position_and_bytes_consumed() {
let bytes = [0xFFu8, 0xFF, 0xFF];
let mut r = BitReader::new(&bytes);
let _ = r.read_bits(13).unwrap();
assert_eq!(r.bit_position(), 13);
assert_eq!(r.bytes_consumed(), 2); }
#[test]
fn n_zero_returns_zero() {
let bytes = [0u8; 0];
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(0).unwrap(), 0);
assert_eq!(r.bit_position(), 0);
}
#[test]
fn over_64_bits_rejected() {
let bytes = [0u8; 16];
let mut r = BitReader::new(&bytes);
assert_eq!(
r.read_bits(65).unwrap_err().code(),
ErrorCode::ProtocolError
);
}
#[test]
fn read_past_end_in_read_bits_errors() {
let bytes = [0xFFu8];
let mut r = BitReader::new(&bytes);
let _ = r.read_bits(7).unwrap();
assert!(r.read_bits(2).is_err()); }
}