pub(crate) struct BitReader<'a> {
buf: &'a [u8],
val: u64,
pos: usize,
bit_pos: u32,
bits_left: u64,
eos: bool,
abs_bits: u64,
}
impl<'a> BitReader<'a> {
pub(crate) fn new(buf: &'a [u8]) -> Self {
let prime = buf.len().min(8);
let mut val = 0u64;
for (i, &byte) in buf.iter().take(prime).enumerate() {
val |= u64::from(byte) << (8 * i);
}
Self {
buf,
val,
pos: prime,
bit_pos: 0,
bits_left: buf.len() as u64 * 8,
eos: false,
abs_bits: 0,
}
}
pub(crate) fn new_at(buf: &'a [u8], bit_offset: u64) -> Self {
let start = usize::try_from(bit_offset / 8)
.unwrap_or(usize::MAX)
.min(buf.len());
let bit_pos = u32::try_from(bit_offset % 8).unwrap_or(0);
let loaded = (buf.len() - start).min(8);
let mut val = 0u64;
for (i, &byte) in buf[start..].iter().take(loaded).enumerate() {
val |= u64::from(byte) << (8 * i);
}
let total_bits = buf.len() as u64 * 8;
Self {
buf,
val,
pos: start + loaded,
bit_pos,
bits_left: total_bits.saturating_sub(bit_offset),
eos: false,
abs_bits: bit_offset,
}
}
#[allow(
clippy::cast_possible_truncation,
reason = "the accumulator is masked to n <= 24 bits; the high bits are unconsumed future bits"
)]
pub(crate) fn peek_bits(&self, n: u32) -> u32 {
debug_assert!(n <= 24, "peek_bits supports up to 24 bits per call");
let mask = (1u32 << n).wrapping_sub(1);
(self.val >> self.bit_pos) as u32 & mask
}
pub(crate) fn consume(&mut self, n: u32) {
self.abs_bits += u64::from(n);
if u64::from(n) > self.bits_left {
self.eos = true;
self.bits_left = 0;
} else {
self.bits_left -= u64::from(n);
}
self.bit_pos += n;
while self.bit_pos >= 8 {
self.val >>= 8;
if self.pos < self.buf.len() {
self.val |= u64::from(self.buf[self.pos]) << 56;
self.pos += 1;
}
self.bit_pos -= 8;
}
}
pub(crate) fn read_bits(&mut self, n: u32) -> u32 {
let result = self.peek_bits(n);
self.consume(n);
result
}
pub(crate) fn read_bit(&mut self) -> u32 {
self.read_bits(1)
}
pub(crate) const fn is_eos(&self) -> bool {
self.eos
}
pub(crate) const fn bit_position(&self) -> u64 {
self.abs_bits
}
}
#[cfg(test)]
mod tests {
use super::BitReader;
#[test]
fn reads_a_whole_byte() {
let mut r = BitReader::new(&[0x2F]);
assert_eq!(r.read_bits(8), 0x2F);
assert!(!r.is_eos());
}
#[test]
fn assembles_bits_lsb_first_within_a_byte() {
let mut r = BitReader::new(&[0xAC]);
assert_eq!(r.read_bits(2), 0b00); assert_eq!(r.read_bits(3), 0b011); assert_eq!(r.read_bits(3), 0b101); }
#[test]
fn multi_byte_values_are_little_endian() {
let mut r = BitReader::new(&[0x34, 0x12]);
assert_eq!(r.read_bits(16), 0x1234);
}
#[test]
fn reads_across_a_byte_boundary() {
let mut r = BitReader::new(&[0xF0, 0x0F]);
assert_eq!(r.read_bits(4), 0x0); assert_eq!(r.read_bits(8), 0xFF); assert_eq!(r.read_bits(4), 0x0); }
#[test]
fn decodes_a_vp8l_header_prefix() {
let bytes = [0x2F, 0x00, 0x00, 0x00, 0x00];
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(8), 0x2F);
assert_eq!(r.read_bits(14) + 1, 1); assert_eq!(r.read_bits(14) + 1, 1); assert_eq!(r.read_bit(), 0); assert_eq!(r.read_bits(3), 0); assert!(!r.is_eos());
}
#[test]
fn past_end_yields_zero_and_latches_eos() {
let mut r = BitReader::new(&[0xFF]);
assert_eq!(r.read_bits(8), 0xFF);
assert!(!r.is_eos());
assert_eq!(r.read_bits(8), 0x00); assert!(r.is_eos());
}
#[test]
fn works_on_an_empty_buffer() {
let mut r = BitReader::new(&[]);
assert_eq!(r.read_bits(4), 0);
assert!(r.is_eos());
}
fn snapshot(r: &BitReader<'_>) -> (u64, usize, u32, u64, bool, u64) {
(r.val, r.pos, r.bit_pos, r.bits_left, r.eos, r.abs_bits)
}
#[test]
fn new_at_matches_new_then_consume_exhaustively() {
for len in [0usize, 1, 2, 3, 7, 8, 9, 16, 20] {
let buf: Vec<u8> = (0..len)
.map(|i| u8::try_from((i * 37 + 5) & 0xff).unwrap_or(0))
.collect();
for off in 0..=(len as u64 * 8) {
let at_offset = BitReader::new_at(&buf, off);
let mut walked = BitReader::new(&buf);
walked.consume(u32::try_from(off).unwrap());
assert_eq!(
snapshot(&at_offset),
snapshot(&walked),
"field mismatch at len={len} off={off}"
);
let mut a = at_offset;
let mut b = walked;
for n in [1u32, 3, 8, 13, 24, 7, 5] {
assert_eq!(
a.read_bits(n),
b.read_bits(n),
"read mismatch at len={len} off={off} n={n}"
);
assert_eq!(a.bit_position(), b.bit_position());
assert_eq!(a.is_eos(), b.is_eos());
}
}
}
}
}
#[cfg(test)]
mod proptests {
use super::BitReader;
use proptest::prelude::*;
fn buf_and_offset() -> impl Strategy<Value = (Vec<u8>, u64)> {
proptest::collection::vec(any::<u8>(), 0..=64).prop_flat_map(|buf| {
let max = buf.len() as u64 * 8;
(Just(buf), 0..=max)
})
}
proptest! {
#[test]
fn new_at_equals_new_then_consume(
(buf, off) in buf_and_offset(),
reads in proptest::collection::vec(1u32..=24, 0..64),
) {
let at_offset = BitReader::new_at(&buf, off);
let mut walked = BitReader::new(&buf);
walked.consume(u32::try_from(off).unwrap());
prop_assert_eq!(at_offset.val, walked.val, "val");
prop_assert_eq!(at_offset.pos, walked.pos, "pos");
prop_assert_eq!(at_offset.bit_pos, walked.bit_pos, "bit_pos");
prop_assert_eq!(at_offset.bits_left, walked.bits_left, "bits_left");
prop_assert_eq!(at_offset.eos, walked.eos, "eos");
prop_assert_eq!(at_offset.abs_bits, walked.abs_bits, "abs_bits");
let mut a = at_offset;
let mut b = walked;
for &n in &reads {
prop_assert_eq!(a.read_bits(n), b.read_bits(n));
prop_assert_eq!(a.bit_position(), b.bit_position());
prop_assert_eq!(a.is_eos(), b.is_eos());
}
}
}
}