use oxideav_core::{Error, Result};
pub struct BitReader<'a> {
buf: &'a [u8],
bit_pos: usize, }
impl<'a> BitReader<'a> {
pub fn new(buf: &'a [u8]) -> Self {
Self { buf, bit_pos: 0 }
}
pub fn read_bit(&mut self) -> Result<u32> {
let byte_idx = self.bit_pos / 8;
if byte_idx >= self.buf.len() {
return Err(Error::invalid("prores bitstream: EOF"));
}
let bit_idx = 7 - (self.bit_pos % 8);
self.bit_pos += 1;
Ok(((self.buf[byte_idx] >> bit_idx) & 1) as u32)
}
pub fn read_bits(&mut self, n: u32) -> Result<u32> {
debug_assert!((1..=32).contains(&n));
let mut v = 0u32;
for _ in 0..n {
v = (v << 1) | self.read_bit()?;
}
Ok(v)
}
pub fn end_of_data(&self) -> bool {
let total_bits = self.buf.len() * 8;
if self.bit_pos >= total_bits {
return true;
}
let remaining = total_bits - self.bit_pos;
if remaining > 31 {
return false;
}
let mut pos = self.bit_pos;
while pos < total_bits {
let byte_idx = pos / 8;
let bit_idx = 7 - (pos % 8);
if ((self.buf[byte_idx] >> bit_idx) & 1) != 0 {
return false;
}
pos += 1;
}
true
}
pub fn byte_pos(&self) -> usize {
self.bit_pos.div_ceil(8)
}
pub fn bit_pos(&self) -> usize {
self.bit_pos
}
}
pub struct BitWriter {
buf: Vec<u8>,
cur_bits_used: u32,
}
impl Default for BitWriter {
fn default() -> Self {
Self::new()
}
}
impl BitWriter {
pub fn new() -> Self {
Self {
buf: Vec::new(),
cur_bits_used: 8,
}
}
pub fn write_bit(&mut self, b: u32) {
if self.cur_bits_used == 8 {
self.buf.push(0);
self.cur_bits_used = 0;
}
let last = self.buf.last_mut().unwrap();
let shift = 7 - self.cur_bits_used;
*last |= ((b & 1) as u8) << shift;
self.cur_bits_used += 1;
}
pub fn write_bits(&mut self, v: u32, n: u32) {
debug_assert!((1..=32).contains(&n));
for i in (0..n).rev() {
self.write_bit((v >> i) & 1);
}
}
pub fn align_byte(&mut self) {
while self.cur_bits_used != 8 {
self.write_bit(0);
}
}
pub fn finish(mut self) -> Vec<u8> {
self.align_byte();
self.buf
}
pub fn byte_len(&self) -> usize {
self.buf.len()
}
pub fn bit_len(&self) -> usize {
if self.cur_bits_used == 8 {
self.buf.len() * 8
} else {
(self.buf.len() - 1) * 8 + self.cur_bits_used as usize
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bits_roundtrip() {
let mut w = BitWriter::new();
w.write_bits(0b1010_1100, 8);
w.write_bits(0x5A_5A, 16);
w.write_bit(1);
w.write_bit(0);
let data = w.finish();
let mut r = BitReader::new(&data);
assert_eq!(r.read_bits(8).unwrap(), 0b1010_1100);
assert_eq!(r.read_bits(16).unwrap(), 0x5A_5A);
assert_eq!(r.read_bit().unwrap(), 1);
assert_eq!(r.read_bit().unwrap(), 0);
}
#[test]
fn end_of_data_with_zero_padding() {
let buf = [0b1000_0000u8];
let mut r = BitReader::new(&buf);
assert!(!r.end_of_data()); let _ = r.read_bit().unwrap();
assert!(r.end_of_data());
}
#[test]
fn end_of_data_with_nonzero_padding() {
let buf = [0b1000_0001u8];
let mut r = BitReader::new(&buf);
let _ = r.read_bit().unwrap(); assert!(!r.end_of_data());
}
}