#[derive(Debug, Default, Clone)]
pub struct BitWriter {
bytes: Vec<u8>,
cache: u64,
nbits: u32,
}
impl BitWriter {
pub fn new() -> Self {
Self::default()
}
pub fn bit_len(&self) -> usize {
self.bytes.len() * 8 + self.nbits as usize
}
pub fn is_byte_aligned(&self) -> bool {
self.nbits % 8 == 0
}
#[inline]
pub fn write_bits(&mut self, value: u32, n: u32) {
debug_assert!(n <= 32, "write_bits supports up to 32 bits");
if n == 0 {
return;
}
let mask = (1u64 << n) - 1; self.cache = (self.cache << n) | (value as u64 & mask);
self.nbits += n;
if self.nbits >= 32 {
self.nbits -= 32;
let word = (self.cache >> self.nbits) as u32;
self.bytes.extend_from_slice(&word.to_be_bytes());
self.cache &= (1u64 << self.nbits) - 1; }
}
#[inline]
pub fn write_bit(&mut self, bit: bool) {
self.write_bits(bit as u32, 1);
}
#[inline]
fn put_golomb(&mut self, x: u64) {
let n = 63 - x.leading_zeros(); self.write_bits(0, n);
if n < 32 {
self.write_bits(x as u32, n + 1);
} else {
self.write_bits((x >> 32) as u32, n - 31);
self.write_bits(x as u32, 32);
}
}
pub fn write_ue(&mut self, value: u32) {
self.put_golomb(value as u64 + 1);
}
pub fn write_se(&mut self, value: i32) {
let code_num = if value <= 0 {
(-(value as i64) as u64) * 2
} else {
(value as u64) * 2 - 1
};
self.put_golomb(code_num + 1);
}
pub fn rbsp_trailing_bits(&mut self) {
self.write_bits(1, 1);
self.align_zero();
}
pub fn align_zero(&mut self) {
let pad = (8 - self.nbits % 8) % 8;
if pad != 0 {
self.write_bits(0, pad);
}
while self.nbits >= 8 {
self.nbits -= 8;
self.bytes.push((self.cache >> self.nbits) as u8);
}
self.cache = 0;
}
pub fn into_bytes(mut self) -> Vec<u8> {
while self.nbits >= 8 {
self.nbits -= 8;
self.bytes.push((self.cache >> self.nbits) as u8);
}
assert!(
self.nbits == 0,
"BitWriter::into_bytes called with {} dangling bits",
self.nbits
);
self.bytes
}
pub fn as_bytes(&self) -> &[u8] {
&self.bytes
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn write_bits_is_msb_first() {
let mut w = BitWriter::new();
w.write_bits(0b101, 3);
w.write_bits(0b1, 1);
w.align_zero();
assert_eq!(w.into_bytes(), vec![0b1011_0000]);
}
#[test]
fn ue_known_values() {
let cases: &[(u32, &str)] = &[
(0, "1"),
(1, "010"),
(2, "011"),
(3, "00100"),
(4, "00101"),
(5, "00110"),
(6, "00111"),
(7, "0001000"),
(8, "0001001"),
];
for &(v, bits) in cases {
let mut w = BitWriter::new();
w.write_ue(v);
assert_eq!(bitstring(&w), bits, "ue({v})");
}
}
#[test]
fn se_known_values() {
let cases: &[(i32, &str)] = &[
(0, "1"), (1, "010"), (-1, "011"), (2, "00100"),
(-2, "00101"),
(3, "00110"),
(-3, "00111"),
];
for &(v, bits) in cases {
let mut w = BitWriter::new();
w.write_se(v);
assert_eq!(bitstring(&w), bits, "se({v})");
}
}
#[test]
fn ue_max_does_not_overflow() {
let mut w = BitWriter::new();
w.write_ue(u32::MAX);
assert_eq!(w.bit_len(), 65);
}
#[test]
fn rbsp_trailing_aligns_to_byte() {
let mut w = BitWriter::new();
w.write_bits(0b101, 3);
w.rbsp_trailing_bits();
assert_eq!(w.into_bytes(), vec![0b1011_0000]);
}
fn bitstring(w: &BitWriter) -> String {
let mut s = String::new();
for &b in w.as_bytes() {
for i in (0..8).rev() {
s.push(if (b >> i) & 1 == 1 { '1' } else { '0' });
}
}
for i in (0..w.nbits).rev() {
s.push(if (w.cache >> i) & 1 == 1 { '1' } else { '0' });
}
s
}
}