use std::io::{Error, ErrorKind, Read, Result};
use byteorder::ReadBytesExt;
pub trait ReadBits {
fn get(&mut self, cbit: u32) -> Result<u32>;
}
pub struct BitReader<R> {
binary_reader: R,
bits_read: u32,
bit_count: u32,
}
impl<R: Read> ReadBits for BitReader<R> {
fn get(&mut self, cbit: u32) -> Result<u32> {
BitReader::get(self, cbit)
}
}
impl<R: Read> BitReader<R> {
pub fn new(binary_reader: R) -> Self {
BitReader {
binary_reader,
bits_read: 0,
bit_count: 0,
}
}
pub fn flush_buffer_to_byte_boundary(&mut self) {
self.bit_count = 0;
}
pub fn bit_position_in_current_byte(&self) -> u32 {
8 - self.bit_count
}
pub fn read_byte(&mut self) -> Result<u8> {
assert!(self.bit_count == 0, "BitReader Error: Attempt to read bytes without first calling FlushBufferToByteBoundary");
let result = self.binary_reader.read_u8()?;
Ok(result)
}
pub fn get(&mut self, cbit: u32) -> Result<u32> {
let mut wret: u32 = 0;
let mut cbits_added = 0;
if cbit == 0 {
return Ok(wret);
}
if cbit > 32 {
return Err(Error::new(
ErrorKind::InvalidInput,
"BitReader Error: Attempt to read more than 32 bits",
));
}
while cbits_added < cbit {
let cbits_needed = cbit - cbits_added;
if self.bit_count == 0 {
self.bits_read = self.binary_reader.read_u8()? as u32;
self.bit_count = 8;
}
let cbits_from_buffer = std::cmp::min(cbits_needed, self.bit_count);
wret |= (self.bits_read & !(u32::MAX << cbits_from_buffer)) << cbits_added;
self.bits_read >>= cbits_from_buffer;
self.bit_count -= cbits_from_buffer;
cbits_added += cbits_from_buffer;
}
Ok(wret)
}
}