use crate::error::Error;
#[inline]
pub(crate) fn mask64(n: u32) -> u64 {
if n >= 64 { u64::MAX } else { (1u64 << n) - 1 }
}
pub(crate) struct ForwardBitReader<'a> {
data: &'a [u8],
bit_pos: usize,
}
impl<'a> ForwardBitReader<'a> {
pub(crate) fn new(data: &'a [u8]) -> Self {
Self { data, bit_pos: 0 }
}
pub(crate) fn peek(&self, n: u32) -> u32 {
debug_assert!(n <= 32);
if n == 0 {
return 0;
}
let byte = self.bit_pos / 8;
let shift = self.bit_pos % 8;
let mut buf = [0u8; 8];
if byte < self.data.len() {
let take = (self.data.len() - byte).min(8);
buf[..take].copy_from_slice(&self.data[byte..byte + take]);
}
((u64::from_le_bytes(buf) >> shift) & mask64(n)) as u32
}
pub(crate) fn consume(&mut self, n: u32) {
self.bit_pos += n as usize;
}
pub(crate) fn read(&mut self, n: u32) -> u32 {
let v = self.peek(n);
self.consume(n);
v
}
pub(crate) fn bits_consumed(&self) -> usize {
self.bit_pos
}
}
pub(crate) struct ReverseBitReader<'a> {
data: &'a [u8],
bits_remaining: i64,
cache: u64,
top: i64,
}
impl<'a> ReverseBitReader<'a> {
pub(crate) fn new(data: &'a [u8]) -> Result<Self, Error> {
let last = *data.last().ok_or(Error::Corrupted("empty bitstream"))?;
if last == 0 {
return Err(Error::Corrupted("bitstream has no initialization bit"));
}
let padding = i64::from(last.leading_zeros()) + 1;
Ok(Self {
data,
bits_remaining: data.len() as i64 * 8 - padding,
cache: 0,
top: 0,
})
}
pub(crate) fn bits_remaining(&self) -> i64 {
self.bits_remaining
}
pub(crate) fn finished_exactly(&self) -> bool {
self.bits_remaining == 0
}
#[inline]
pub(crate) fn peek(&mut self, n: u32) -> u64 {
debug_assert!(n <= 56);
if n == 0 {
return 0;
}
if self.bits_remaining < i64::from(n) {
if self.bits_remaining <= 0 {
return 0;
}
let avail = self.bits_remaining as u32;
return self.extract(0, avail) << (n - avail);
}
if self.top < i64::from(n) {
self.reload_cache();
}
let shift = (self.top - i64::from(n)) as u32;
(self.cache >> shift) & mask64(n)
}
#[inline]
pub(crate) fn consume(&mut self, n: u32) {
self.bits_remaining -= i64::from(n);
self.top -= i64::from(n);
}
#[inline]
pub(crate) fn read(&mut self, n: u32) -> u64 {
let v = self.peek(n);
self.consume(n);
v
}
#[cold]
fn reload_cache(&mut self) {
let hi = self.bits_remaining as usize;
let byte = hi.saturating_sub(64).div_ceil(8);
self.cache = self.load8(byte);
self.top = (hi - 8 * byte) as i64;
}
fn load8(&self, byte: usize) -> u64 {
if byte + 8 <= self.data.len() {
u64::from_le_bytes(self.data[byte..byte + 8].try_into().unwrap())
} else {
let mut buf = [0u8; 8];
if byte < self.data.len() {
let tail = &self.data[byte..];
buf[..tail.len()].copy_from_slice(tail);
}
u64::from_le_bytes(buf)
}
}
fn extract(&self, from_bit: usize, n: u32) -> u64 {
let byte = from_bit / 8;
let shift = from_bit % 8;
(self.load8(byte) >> shift) & mask64(n)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn forward_reader_reads_lsb_first() {
let mut br = ForwardBitReader::new(&[0xD6, 0x01]);
assert_eq!(br.read(3), 0b110);
assert_eq!(br.read(5), 0b11010);
assert_eq!(br.read(4), 0b0001);
assert_eq!(br.bits_consumed(), 12);
assert_eq!(br.read(8), 0);
}
#[test]
fn reverse_reader_strips_padding_and_reads_backward() {
let mut br = ReverseBitReader::new(&[0x71, 0x01]).unwrap();
assert_eq!(br.bits_remaining(), 8);
assert_eq!(br.read(3), 0b011);
assert_eq!(br.read(5), 0b10001);
assert!(br.finished_exactly());
}
#[test]
fn reverse_reader_zero_fills_past_start() {
let mut br = ReverseBitReader::new(&[0x05]).unwrap();
assert_eq!(br.bits_remaining(), 2);
assert_eq!(br.peek(4), 0b0100);
br.consume(4);
assert_eq!(br.bits_remaining(), -2);
assert!(!br.finished_exactly());
}
#[test]
fn reverse_reader_rejects_empty_and_zero_padding() {
assert!(ReverseBitReader::new(&[]).is_err());
assert!(ReverseBitReader::new(&[0x00]).is_err());
}
}