use crate::error::{CodecError, Result};
#[derive(Debug, Default)]
pub struct BitWriter {
buf: Vec<u8>,
bit_len: usize,
}
impl BitWriter {
pub const fn new() -> Self {
Self {
buf: Vec::new(),
bit_len: 0,
}
}
#[cfg(test)]
pub const fn bit_len(&self) -> usize {
self.bit_len
}
pub fn write_bits(&mut self, value: u32, count: u8) {
assert!(count <= 32, "cannot write {count} bits from a u32");
for i in (0..count).rev() {
self.write_bit(value >> i & 1 == 1);
}
}
pub fn write_bit(&mut self, set: bool) {
let offset = self.bit_len % 8;
if offset == 0 {
self.buf.push(0);
}
if set {
let last = self.buf.len() - 1;
self.buf[last] |= 0x80 >> offset;
}
self.bit_len += 1;
}
pub fn write_slice_bits(&mut self, data: &[u8], bit_count: usize) -> Result<()> {
if data.len() < bit_count.div_ceil(8) {
return Err(CodecError::BufferTooSmall {
needed: bit_count.div_ceil(8),
actual: data.len(),
});
}
for index in 0..bit_count {
let byte = data[index / 8];
let set = byte >> (7 - index % 8) & 1 == 1;
self.write_bit(set);
}
Ok(())
}
pub fn align_to_octet(&mut self) {
while !self.bit_len.is_multiple_of(8) {
self.write_bit(false);
}
}
pub fn finish(mut self) -> Vec<u8> {
self.align_to_octet();
self.buf
}
}
#[derive(Debug)]
pub struct BitReader<'a> {
buf: &'a [u8],
bit_pos: usize,
}
impl<'a> BitReader<'a> {
pub const fn new(buf: &'a [u8]) -> Self {
Self { buf, bit_pos: 0 }
}
pub const fn remaining_bits(&self) -> usize {
self.buf.len() * 8 - self.bit_pos
}
pub fn read_bits(&mut self, count: u8) -> Result<u32> {
assert!(count <= 32, "cannot read {count} bits into a u32");
if self.remaining_bits() < count as usize {
return Err(CodecError::InvalidPayload {
details: format!(
"truncated AMR payload: needed {count} bits, {} remain",
self.remaining_bits()
),
});
}
let mut value = 0u32;
for _ in 0..count {
value = value << 1 | u32::from(self.read_bit_unchecked());
}
Ok(value)
}
fn read_bit_unchecked(&mut self) -> u8 {
let byte = self.buf[self.bit_pos / 8];
let bit = byte >> (7 - self.bit_pos % 8) & 1;
self.bit_pos += 1;
bit
}
pub fn read_slice_bits(&mut self, bit_count: usize) -> Result<Vec<u8>> {
if self.remaining_bits() < bit_count {
return Err(CodecError::InvalidPayload {
details: format!(
"truncated AMR payload: needed {bit_count} bits, {} remain",
self.remaining_bits()
),
});
}
let mut out = vec![0u8; bit_count.div_ceil(8)];
for index in 0..bit_count {
if self.read_bit_unchecked() == 1 {
out[index / 8] |= 0x80 >> (index % 8);
}
}
Ok(out)
}
pub const fn align_to_octet(&mut self) {
self.bit_pos = self.bit_pos.div_ceil(8) * 8;
}
#[cfg(test)]
pub fn remaining_bits_are_zero(&self) -> bool {
let mut probe = Self {
buf: self.buf,
bit_pos: self.bit_pos,
};
(0..probe.remaining_bits()).all(|_| probe.read_bit_unchecked() == 0)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn writes_bits_most_significant_first() {
let mut w = BitWriter::new();
w.write_bits(0b1010, 4);
w.write_bits(0b11, 2);
assert_eq!(w.bit_len(), 6);
assert_eq!(w.finish(), vec![0b1010_1100]);
}
#[test]
fn writes_across_byte_boundaries() {
let mut w = BitWriter::new();
w.write_bits(0b1111_1111_1111, 12);
assert_eq!(w.bit_len(), 12);
assert_eq!(w.finish(), vec![0xFF, 0xF0]);
}
#[test]
fn only_the_low_count_bits_are_written() {
let mut w = BitWriter::new();
w.write_bits(0xFFFF_FFF5, 4);
assert_eq!(w.finish(), vec![0b0101_0000]);
}
#[test]
fn reads_back_what_was_written() {
let mut w = BitWriter::new();
w.write_bits(0b1010, 4);
w.write_bits(0b110_0110, 7);
w.write_bits(1, 1);
let bytes = w.finish();
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(4).unwrap(), 0b1010);
assert_eq!(r.read_bits(7).unwrap(), 0b110_0110);
assert_eq!(r.read_bits(1).unwrap(), 1);
assert!(r.remaining_bits_are_zero());
}
#[test]
fn reading_past_the_end_is_an_error_not_a_panic() {
let bytes = [0xFFu8];
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(8).unwrap(), 0xFF);
assert_eq!(r.remaining_bits(), 0);
let err = r.read_bits(1).unwrap_err();
assert!(matches!(err, CodecError::InvalidPayload { .. }));
assert!(r.read_slice_bits(1).is_err());
}
#[test]
fn slice_bits_round_trip_at_non_octet_lengths() {
for bit_count in [1usize, 7, 8, 9, 39, 40, 95, 244, 477] {
let byte_len = bit_count.div_ceil(8);
let mut src = vec![0u8; byte_len];
for (i, byte) in src.iter_mut().enumerate() {
*byte = u8::try_from(i % 256).unwrap_or(0).wrapping_mul(37) | 0x81;
}
let tail = bit_count % 8;
if tail != 0 {
let mask = 0xFFu8 << (8 - tail);
let last = byte_len - 1;
src[last] &= mask;
}
let mut w = BitWriter::new();
w.write_bits(0b1011, 4);
w.write_slice_bits(&src, bit_count).unwrap();
let bytes = w.finish();
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(4).unwrap(), 0b1011);
let got = r.read_slice_bits(bit_count).unwrap();
assert_eq!(got, src, "round trip failed at {bit_count} bits");
assert!(
r.remaining_bits_are_zero(),
"padding not zero at {bit_count} bits"
);
}
}
#[test]
fn write_slice_bits_rejects_a_short_source() {
let mut w = BitWriter::new();
let err = w.write_slice_bits(&[0xFF], 9).unwrap_err();
assert!(matches!(err, CodecError::BufferTooSmall { .. }));
}
#[test]
fn align_skips_to_the_octet_boundary() {
let mut w = BitWriter::new();
w.write_bits(0b101, 3);
w.align_to_octet();
assert_eq!(w.bit_len(), 8);
w.write_bits(0xAB, 8);
assert_eq!(w.finish(), vec![0b1010_0000, 0xAB]);
let bytes = [0b1010_0000u8, 0xAB];
let mut r = BitReader::new(&bytes);
assert_eq!(r.read_bits(3).unwrap(), 0b101);
r.align_to_octet();
assert_eq!(r.read_bits(8).unwrap(), 0xAB);
let mut r = BitReader::new(&bytes);
r.align_to_octet();
assert_eq!(r.remaining_bits(), 16);
}
#[test]
fn non_zero_padding_is_detectable() {
let bytes = [0b1010_0001u8];
let mut r = BitReader::new(&bytes);
r.read_bits(4).unwrap();
assert!(!r.remaining_bits_are_zero());
let bytes = [0b1010_0000u8];
let mut r = BitReader::new(&bytes);
r.read_bits(4).unwrap();
assert!(r.remaining_bits_are_zero());
}
#[test]
fn probing_padding_does_not_advance_the_reader() {
let bytes = [0b1010_0000u8];
let mut r = BitReader::new(&bytes);
r.read_bits(4).unwrap();
let before = r.remaining_bits();
assert!(r.remaining_bits_are_zero());
assert_eq!(r.remaining_bits(), before);
assert_eq!(r.read_bits(4).unwrap(), 0);
}
}