#![forbid(unsafe_code)]
use crate::AudioError;
#[derive(Debug, Clone)]
pub struct BitPacker {
bytes: Vec<u8>,
current_byte: u8,
bit_position: u8,
}
impl BitPacker {
#[must_use]
pub fn new() -> Self {
Self {
bytes: Vec::new(),
current_byte: 0,
bit_position: 0,
}
}
pub fn write_bit(&mut self, bit: bool) {
if bit {
self.current_byte |= 1 << self.bit_position;
}
self.bit_position += 1;
if self.bit_position == 8 {
self.flush_byte();
}
}
pub fn write_bits(&mut self, value: u32, bits: u8) {
for i in 0..bits {
let bit = (value >> i) & 1;
self.write_bit(bit != 0);
}
}
pub fn write_byte(&mut self, byte: u8) {
self.write_bits(u32::from(byte), 8);
}
pub fn write_bytes(&mut self, bytes: &[u8]) {
for &byte in bytes {
self.write_byte(byte);
}
}
#[allow(clippy::cast_sign_loss)]
pub fn write_signed(&mut self, value: i32, bits: u8) {
let unsigned = value as u32;
self.write_bits(unsigned, bits);
}
fn flush_byte(&mut self) {
self.bytes.push(self.current_byte);
self.current_byte = 0;
self.bit_position = 0;
}
#[must_use]
pub fn position(&self) -> usize {
self.bytes.len() * 8 + self.bit_position as usize
}
#[must_use]
pub fn finish(mut self) -> Vec<u8> {
if self.bit_position > 0 {
self.flush_byte();
}
self.bytes
}
#[must_use]
pub fn as_bytes(&self) -> &[u8] {
&self.bytes
}
#[must_use]
pub fn size(&self) -> usize {
let mut size = self.bytes.len();
if self.bit_position > 0 {
size += 1;
}
size
}
}
impl Default for BitPacker {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct BitReader<'a> {
data: &'a [u8],
byte_pos: usize,
bit_pos: u8,
}
impl<'a> BitReader<'a> {
#[must_use]
pub const fn new(data: &'a [u8]) -> Self {
Self {
data,
byte_pos: 0,
bit_pos: 0,
}
}
#[must_use]
pub const fn position(&self) -> usize {
self.byte_pos * 8 + self.bit_pos as usize
}
#[must_use]
pub fn is_exhausted(&self) -> bool {
self.byte_pos >= self.data.len()
}
pub fn read_bit(&mut self) -> Result<bool, AudioError> {
if self.byte_pos >= self.data.len() {
return Err(AudioError::Eof);
}
let bit = (self.data[self.byte_pos] >> self.bit_pos) & 1;
self.bit_pos += 1;
if self.bit_pos == 8 {
self.bit_pos = 0;
self.byte_pos += 1;
}
Ok(bit != 0)
}
pub fn read_bits(&mut self, n: u8) -> Result<u32, AudioError> {
if n > 32 {
return Err(AudioError::InvalidData(
"Cannot read more than 32 bits at once".into(),
));
}
let mut result = 0u32;
for i in 0..n {
let bit = self.read_bit()? as u32;
result |= bit << i; }
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_write_bit() {
let mut packer = BitPacker::new();
packer.write_bit(true);
packer.write_bit(false);
packer.write_bit(true);
packer.write_bit(true);
packer.write_bit(false);
packer.write_bit(false);
packer.write_bit(false);
packer.write_bit(false);
let bytes = packer.finish();
assert_eq!(bytes.len(), 1);
assert_eq!(bytes[0], 0b0000_1101); }
#[test]
fn test_write_bits() {
let mut packer = BitPacker::new();
packer.write_bits(0b1101, 4);
packer.write_bits(0b0010, 4);
let bytes = packer.finish();
assert_eq!(bytes.len(), 1);
assert_eq!(bytes[0], 0b0010_1101);
}
#[test]
fn test_write_byte() {
let mut packer = BitPacker::new();
packer.write_byte(0xAB);
packer.write_byte(0xCD);
let bytes = packer.finish();
assert_eq!(bytes, vec![0xAB, 0xCD]);
}
#[test]
fn test_write_bytes() {
let mut packer = BitPacker::new();
packer.write_bytes(&[0x01, 0x02, 0x03]);
let bytes = packer.finish();
assert_eq!(bytes, vec![0x01, 0x02, 0x03]);
}
#[test]
fn test_write_across_bytes() {
let mut packer = BitPacker::new();
packer.write_bits(0xFF, 6); packer.write_bits(0xFF, 6);
let bytes = packer.finish();
assert_eq!(bytes.len(), 2);
assert_eq!(bytes[0], 0xFF); assert_eq!(bytes[1], 0x0F); }
#[test]
fn test_position() {
let mut packer = BitPacker::new();
assert_eq!(packer.position(), 0);
packer.write_bits(0, 3);
assert_eq!(packer.position(), 3);
packer.write_bits(0, 5);
assert_eq!(packer.position(), 8);
packer.write_bits(0, 4);
assert_eq!(packer.position(), 12);
}
#[test]
fn test_size() {
let mut packer = BitPacker::new();
assert_eq!(packer.size(), 0);
packer.write_bits(0, 3);
assert_eq!(packer.size(), 1);
packer.write_bits(0, 5);
assert_eq!(packer.size(), 1);
packer.write_bits(0, 4);
assert_eq!(packer.size(), 2); }
#[test]
fn test_signed() {
let mut packer = BitPacker::new();
packer.write_signed(-1, 8);
let bytes = packer.finish();
assert_eq!(bytes[0], 0xFF);
}
#[test]
fn test_bit_reader_single_bit() {
let data = [0x01u8];
let mut reader = BitReader::new(&data);
assert!(reader.read_bit().expect("read bit 0"));
assert!(!reader.read_bit().expect("read bit 1"));
}
#[test]
fn test_bit_reader_exhausted() {
let data = [0x00u8];
let mut reader = BitReader::new(&data);
for _ in 0..8 {
let _ = reader.read_bit();
}
assert!(reader.is_exhausted());
assert!(reader.read_bit().is_err());
}
#[test]
fn test_bit_reader_read_bits() {
let data = [0x0Du8];
let mut reader = BitReader::new(&data);
let val = reader.read_bits(4).expect("read 4 bits");
assert_eq!(val, 0b1101); }
#[test]
fn test_packer_reader_round_trip() {
let mut packer = BitPacker::new();
packer.write_bits(0b1011, 4); packer.write_bits(0b00110, 5); let bytes = packer.finish();
let mut reader = BitReader::new(&bytes);
assert_eq!(reader.read_bits(4).expect("4 bits"), 0b1011);
assert_eq!(reader.read_bits(5).expect("5 bits"), 0b00110);
}
#[test]
fn test_bit_reader_position() {
let data = [0xFFu8, 0xFFu8];
let mut reader = BitReader::new(&data);
assert_eq!(reader.position(), 0);
let _ = reader.read_bits(3);
assert_eq!(reader.position(), 3);
let _ = reader.read_bits(5);
assert_eq!(reader.position(), 8);
let _ = reader.read_bits(4);
assert_eq!(reader.position(), 12);
}
}