use crate::error::DecodeError;
pub(crate) struct BitReader<'a> {
data: &'a [u8],
byte_pos: usize,
bit_pos: u8, }
impl<'a> BitReader<'a> {
pub(crate) fn new(data: &'a [u8]) -> Self {
BitReader {
data,
byte_pos: 0,
bit_pos: 0,
}
}
#[inline]
pub(crate) fn read_bit(&mut self) -> Result<u32, DecodeError> {
if self.byte_pos >= self.data.len() {
return Err(DecodeError::Bitstream("unexpected end of stream".into()));
}
let bit = (self.data[self.byte_pos] >> (7 - self.bit_pos)) & 1;
self.bit_pos += 1;
if self.bit_pos == 8 {
self.byte_pos += 1;
self.bit_pos = 0;
}
Ok(bit as u32)
}
#[inline]
pub(crate) fn read_bits(&mut self, n: u32) -> Result<u32, DecodeError> {
let mut v = 0u32;
for _ in 0..n {
v = (v << 1) | self.read_bit()?;
}
Ok(v)
}
#[inline]
pub(crate) fn read_flag(&mut self) -> Result<bool, DecodeError> {
Ok(self.read_bit()? != 0)
}
pub(crate) fn read_ue(&mut self) -> Result<u32, DecodeError> {
let mut leading_zeros = 0u32;
while self.read_bit()? == 0 {
leading_zeros += 1;
if leading_zeros > 31 {
return Err(DecodeError::Bitstream(
"ue(v) leading zeros exceed 31".into(),
));
}
}
if leading_zeros == 0 {
return Ok(0);
}
let suffix = self.read_bits(leading_zeros)?;
Ok((1 << leading_zeros) - 1 + suffix)
}
pub(crate) fn read_se(&mut self) -> Result<i32, DecodeError> {
let ue = self.read_ue()?;
let v = if ue & 1 == 1 {
((ue + 1) >> 1) as i32
} else {
-((ue >> 1) as i32)
};
Ok(v)
}
pub(crate) fn bit_pos(&self) -> usize {
self.byte_pos * 8 + self.bit_pos as usize
}
}
pub(crate) fn unescape_rbsp(src: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(src.len());
let mut i = 0usize;
while i < src.len() {
if i + 2 < src.len() && src[i] == 0x00 && src[i + 1] == 0x00 && src[i + 2] == 0x03 {
out.push(0x00);
out.push(0x00);
i += 3; } else {
out.push(src[i]);
i += 1;
}
}
out
}
#[cfg(test)]
pub(crate) fn unescape_rbsp_with_map(src: &[u8]) -> (Vec<u8>, Vec<usize>) {
let mut out = Vec::with_capacity(src.len());
let mut src_of = Vec::with_capacity(src.len());
let mut i = 0usize;
while i < src.len() {
if i + 2 < src.len() && src[i] == 0x00 && src[i + 1] == 0x00 && src[i + 2] == 0x03 {
out.push(0x00);
src_of.push(i);
out.push(0x00);
src_of.push(i + 1);
i += 3; } else {
out.push(src[i]);
src_of.push(i);
i += 1;
}
}
(out, src_of)
}
pub(crate) fn rbsp_src_map(src: &[u8]) -> Vec<usize> {
let mut src_of = Vec::with_capacity(src.len());
let mut i = 0usize;
while i < src.len() {
if i + 2 < src.len() && src[i] == 0x00 && src[i + 1] == 0x00 && src[i + 2] == 0x03 {
src_of.push(i);
src_of.push(i + 1);
i += 3;
} else {
src_of.push(i);
i += 1;
}
}
src_of
}
pub(crate) fn nal_to_rbsp_offset(src_of: &[usize], nal_off: usize) -> usize {
src_of.partition_point(|&s| s < nal_off)
}
#[cfg(test)]
mod tests {
use crate::bitreader::{BitReader, rbsp_src_map, unescape_rbsp, unescape_rbsp_with_map};
#[test]
fn src_map_matches_with_map_variant() {
let nal = [
0x00u8, 0x00, 0x03, 0x00, 0xAA, 0x00, 0x00, 0x03, 0x01, 0xBB, 0xCC,
];
let (_rbsp, map_ref) = unescape_rbsp_with_map(&nal);
let map = rbsp_src_map(&nal);
assert_eq!(map, map_ref);
}
#[test]
fn round_trip_ue() {
let bits: &[u8] = &[
0b1_010_011_0, 0b0110_0000, ];
let mut r = BitReader::new(bits);
assert_eq!(r.read_ue().unwrap(), 0);
assert_eq!(r.read_ue().unwrap(), 1);
assert_eq!(r.read_ue().unwrap(), 2);
assert_eq!(r.read_ue().unwrap(), 5);
}
#[test]
fn unescape_removes_emulation_prevention() {
let input = vec![0x00, 0x00, 0x03, 0x01, 0xFF];
let out = unescape_rbsp(&input);
assert_eq!(out, vec![0x00, 0x00, 0x01, 0xFF]);
}
}