use crate::error::BitError;
#[derive(Debug)]
pub struct BitWriter<'a> {
out: &'a mut [u8],
bit_pos: usize,
}
impl<'a> BitWriter<'a> {
pub fn new(out: &'a mut [u8]) -> Self {
out.fill(0);
Self { out, bit_pos: 0 }
}
pub fn write_bool(&mut self, v: bool) -> Result<(), BitError> {
self.write_bits(u64::from(v), 1)
}
pub fn write_bits(&mut self, value: u64, nbits: u32) -> Result<(), BitError> {
if nbits == 0 || nbits > 64 {
return Err(BitError::OutOfRange {
field: "write_bits",
});
}
let nbits = nbits as usize;
if self.bit_pos + nbits > self.out.len() * 8 {
return Err(BitError::BufferFull);
}
for i in 0..nbits {
let bit = ((value >> (nbits - 1 - i)) & 1) as u8;
let byte_idx = self.bit_pos / 8;
let shift = 7 - (self.bit_pos % 8);
self.out[byte_idx] |= bit << shift;
self.bit_pos += 1;
}
Ok(())
}
pub fn skip(&mut self, nbits: u32) -> Result<(), BitError> {
let n = nbits as usize;
if self.bit_pos + n > self.out.len() * 8 {
return Err(BitError::BufferFull);
}
self.bit_pos += n;
Ok(())
}
#[must_use]
pub const fn bytes_written(&self) -> usize {
self.bit_pos.div_ceil(8)
}
}
#[cfg(test)]
mod tests {
use super::super::reader::BitReader;
use super::*;
#[test]
fn writes_and_reads_back_known_value() {
let mut buf = [0u8; 2];
let mut w = BitWriter::new(&mut buf);
w.write_bits(0b1010, 4).unwrap();
w.write_bits(0b1100, 4).unwrap();
w.write_bits(0xFF, 8).unwrap();
assert_eq!(w.bytes_written(), 2);
let mut r = BitReader::new(&buf);
assert_eq!(r.read_u8(4).unwrap(), 0b1010);
assert_eq!(r.read_u8(4).unwrap(), 0b1100);
assert_eq!(r.read_u8(8).unwrap(), 0xFF);
}
#[test]
fn buffer_full_is_reported_not_panicked() {
let mut buf = [0u8; 1];
let mut w = BitWriter::new(&mut buf);
w.write_bits(0, 8).unwrap();
assert_eq!(w.write_bits(0, 1).unwrap_err(), BitError::BufferFull);
}
#[test]
fn skip_leaves_bits_zero() {
let mut buf = [0xFFu8; 1];
let mut w = BitWriter::new(&mut buf);
w.skip(4).unwrap();
w.write_bits(0b1111, 4).unwrap();
assert_eq!(buf[0], 0b0000_1111);
}
}