use bitstream_io::read::BitRead as _;
use bitstream_io::write::BitWrite as _;
use std::borrow::Cow;
use std::io::BufRead;
use std::io::Read;
use std::io::Write;
use std::num::NonZeroUsize;
#[derive(Copy, Clone, Debug)]
enum ParseState {
Start(u8),
Skip(NonZeroUsize),
Three,
PostThree,
}
const H264_HEADER_LEN: NonZeroUsize = match NonZeroUsize::new(1) {
Some(one) => one,
None => panic!("1 should be non-zero"),
};
fn zero_pair_finder() -> &'static memchr::memmem::Finder<'static> {
static FINDER: std::sync::OnceLock<memchr::memmem::Finder<'static>> =
std::sync::OnceLock::new();
FINDER.get_or_init(|| memchr::memmem::Finder::new(b"\x00\x00"))
}
#[derive(Clone)]
pub struct ByteReader<R: BufRead> {
inner: R,
state: ParseState,
i: usize,
max_fill: usize,
}
impl<R: BufRead> ByteReader<R> {
pub fn without_skip(inner: R) -> Self {
Self {
inner,
state: ParseState::Start(0),
i: 0,
max_fill: 128,
}
}
pub fn skipping_h264_header(inner: R) -> Self {
Self {
inner,
state: ParseState::Skip(H264_HEADER_LEN),
i: 0,
max_fill: 128,
}
}
pub fn skipping_bytes(inner: R, skip: NonZeroUsize) -> Self {
Self {
inner,
state: ParseState::Skip(skip),
i: 0,
max_fill: 128,
}
}
fn try_fill_buf_slow(&mut self) -> std::io::Result<bool> {
debug_assert_eq!(self.i, 0);
let chunk = self.inner.fill_buf()?;
if chunk.is_empty() {
return Ok(false);
}
let limit = std::cmp::min(chunk.len(), self.max_fill);
while self.i < limit {
match self.state {
ParseState::Start(zero_count) => {
let after_pair = if zero_count >= 2 {
Some(self.i)
} else if zero_count == 1 && chunk[self.i] == 0x00 {
if self.i + 1 < limit {
Some(self.i + 1)
} else {
self.state = ParseState::Start(2);
self.i += 1;
None
}
} else {
match zero_pair_finder().find(&chunk[self.i..limit]) {
Some(offset) => {
let ap = self.i + offset + 2;
if ap < limit {
Some(ap)
} else {
self.state = ParseState::Start(2);
self.i = ap;
None
}
}
None => {
let trailing = if limit > self.i && chunk[limit - 1] == 0x00 {
1
} else {
0
};
self.state = ParseState::Start(trailing);
self.i = limit;
None
}
}
};
let Some(after_pair) = after_pair else { break };
match chunk[after_pair] {
0x03 => {
self.i = after_pair;
self.state = ParseState::Three;
break;
}
b @ 0x00..=0x02 => {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("invalid RBSP byte {:#x} in state {:?}", b, &self.state,),
));
}
_ => {
self.i = after_pair + 1;
self.state = ParseState::Start(0);
continue;
}
}
}
ParseState::Skip(remaining) => {
debug_assert_eq!(self.i, 0);
let skip = std::cmp::min(chunk.len(), remaining.get());
self.inner.consume(skip);
self.state = NonZeroUsize::new(remaining.get() - skip)
.map(ParseState::Skip)
.unwrap_or(ParseState::Start(0));
break;
}
ParseState::Three => {
debug_assert_eq!(self.i, 0);
self.inner.consume(1);
self.state = ParseState::PostThree;
break;
}
ParseState::PostThree => {
match chunk[self.i] {
0x00 => self.state = ParseState::Start(1),
0x01 | 0x02 | 0x03 => self.state = ParseState::Start(0),
o => {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("invalid RBSP byte {:#x} in state {:?}", o, &self.state),
))
}
}
self.i += 1;
}
}
}
Ok(true)
}
pub fn reader(&mut self) -> &mut R {
&mut self.inner
}
}
impl<R: BufRead> Read for ByteReader<R> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
let chunk = self.fill_buf()?;
let amt = std::cmp::min(buf.len(), chunk.len());
if amt == 1 {
buf[0] = chunk[0];
} else {
buf[..amt].copy_from_slice(&chunk[..amt]);
}
self.consume(amt);
Ok(amt)
}
}
impl<R: BufRead> BufRead for ByteReader<R> {
fn fill_buf(&mut self) -> std::io::Result<&[u8]> {
while self.i == 0 && self.try_fill_buf_slow()? {}
Ok(&self.inner.fill_buf()?[0..self.i])
}
fn consume(&mut self, amt: usize) {
self.i = self.i.checked_sub(amt).unwrap();
self.inner.consume(amt);
}
}
pub fn decode_nal<'a>(nal_unit: &'a [u8]) -> Result<Cow<'a, [u8]>, std::io::Error> {
let mut reader = ByteReader {
inner: nal_unit,
state: ParseState::Skip(H264_HEADER_LEN),
i: 0,
max_fill: usize::MAX, };
let buf = reader.fill_buf()?;
if buf.len() + 1 == nal_unit.len() {
return Ok(Cow::Borrowed(&nal_unit[1..]));
}
let mut dst = Vec::with_capacity(nal_unit.len() - 2);
loop {
let buf = reader.fill_buf()?;
if buf.is_empty() {
break;
}
dst.extend_from_slice(buf);
let len = buf.len();
reader.consume(len);
}
Ok(Cow::Owned(dst))
}
#[derive(Debug)]
pub enum BitReaderError {
ReaderError(&'static str, std::io::Error),
ExpGolombTooLarge(&'static str),
RemainingData,
}
pub trait Integer: bitstream_io::Integer + std::fmt::Debug {}
impl<I: bitstream_io::Integer + std::fmt::Debug> Integer for I {}
pub trait Primitive: bitstream_io::Primitive + std::fmt::Debug {}
impl<P: bitstream_io::Primitive + std::fmt::Debug> Primitive for P {}
pub trait BitWrite {
fn write_ue(&mut self, value: u32) -> std::io::Result<()>;
fn write_se(&mut self, value: i32) -> std::io::Result<()>;
fn write_bit(&mut self, bit: bool) -> std::io::Result<()>;
fn write<const BITS: u32, I: Integer>(&mut self, value: I) -> std::io::Result<()>;
fn write_var<I: Integer>(&mut self, bit_count: u32, value: I) -> std::io::Result<()>;
fn write_rbsp_trailing_bits(&mut self) -> std::io::Result<()>;
}
pub trait BitRead {
fn read_ue(&mut self, name: &'static str) -> Result<u32, BitReaderError>;
fn read_se(&mut self, name: &'static str) -> Result<i32, BitReaderError>;
fn read_bit(&mut self, name: &'static str) -> Result<bool, BitReaderError>;
fn read<const BITS: u32, I: Integer>(
&mut self,
name: &'static str,
) -> Result<I, BitReaderError>;
fn read_var<I: Integer>(
&mut self,
bit_count: u32,
name: &'static str,
) -> Result<I, BitReaderError>;
fn read_to<V: Primitive>(&mut self, name: &'static str) -> Result<V, BitReaderError>;
fn skip(&mut self, bit_count: u32, name: &'static str) -> Result<(), BitReaderError>;
fn byte_aligned(&self) -> bool;
fn has_more_rbsp_data(&mut self, name: &'static str) -> Result<bool, BitReaderError>;
fn finish_rbsp(self) -> Result<(), BitReaderError>;
fn finish_sei_payload(self) -> Result<(), BitReaderError>;
}
pub struct BitReader<R: std::io::BufRead + Clone> {
reader: bitstream_io::read::BitReader<R, bitstream_io::BigEndian>,
}
impl<R: std::io::BufRead + Clone> BitReader<R> {
pub fn new(inner: R) -> Self {
Self {
reader: bitstream_io::read::BitReader::new(inner),
}
}
pub fn reader(&mut self) -> Option<&mut R> {
self.reader.reader()
}
pub fn into_reader(self) -> R {
self.reader.into_reader()
}
}
impl<R: std::io::BufRead + Clone> BitRead for BitReader<R> {
fn read_ue(&mut self, name: &'static str) -> Result<u32, BitReaderError> {
let count = self
.reader
.read_unary::<1>()
.map_err(|e| BitReaderError::ReaderError(name, e))?;
if count > 31 {
return Err(BitReaderError::ExpGolombTooLarge(name));
} else if count > 0 {
let val: u32 = self.read_var(count, name)?;
Ok((1 << count) - 1 + val)
} else {
Ok(0)
}
}
fn read_se(&mut self, name: &'static str) -> Result<i32, BitReaderError> {
Ok(golomb_to_signed(self.read_ue(name)?))
}
fn read_bit(&mut self, name: &'static str) -> Result<bool, BitReaderError> {
self.reader
.read_bit()
.map_err(|e| BitReaderError::ReaderError(name, e))
}
fn read<const BITS: u32, I: Integer>(
&mut self,
name: &'static str,
) -> Result<I, BitReaderError> {
self.reader
.read::<BITS, I>()
.map_err(|e| BitReaderError::ReaderError(name, e))
}
fn read_var<I: Integer>(
&mut self,
bit_count: u32,
name: &'static str,
) -> Result<I, BitReaderError> {
self.reader
.read_var(bit_count)
.map_err(|e| BitReaderError::ReaderError(name, e))
}
fn read_to<V: Primitive>(&mut self, name: &'static str) -> Result<V, BitReaderError> {
self.reader
.read_to()
.map_err(|e| BitReaderError::ReaderError(name, e))
}
fn skip(&mut self, bit_count: u32, name: &'static str) -> Result<(), BitReaderError> {
self.reader
.skip(bit_count)
.map_err(|e| BitReaderError::ReaderError(name, e))
}
fn byte_aligned(&self) -> bool {
self.reader.byte_aligned()
}
fn has_more_rbsp_data(&mut self, name: &'static str) -> Result<bool, BitReaderError> {
let mut throwaway = self.reader.clone();
let r = (move || {
throwaway.skip(1)?;
throwaway.read_unary::<1>()?;
Ok::<_, std::io::Error>(())
})();
match r {
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => Ok(false),
Err(e) => Err(BitReaderError::ReaderError(name, e)),
Ok(_) => Ok(true),
}
}
fn finish_rbsp(mut self) -> Result<(), BitReaderError> {
if !self
.reader
.read_bit()
.map_err(|e| BitReaderError::ReaderError("finish", e))?
{
match self.reader.read_unary::<1>() {
Err(e) => return Err(BitReaderError::ReaderError("finish", e)),
Ok(_) => return Err(BitReaderError::RemainingData),
}
}
match self.reader.read_unary::<1>() {
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => Ok(()),
Err(e) => Err(BitReaderError::ReaderError("finish", e)),
Ok(_) => Err(BitReaderError::RemainingData),
}
}
fn finish_sei_payload(mut self) -> Result<(), BitReaderError> {
match self.reader.read_bit() {
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(()),
Err(e) => return Err(BitReaderError::ReaderError("finish", e)),
Ok(false) => return Err(BitReaderError::RemainingData),
Ok(true) => {}
}
match self.reader.read_unary::<1>() {
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => Ok(()),
Err(e) => Err(BitReaderError::ReaderError("finish", e)),
Ok(_) => Err(BitReaderError::RemainingData),
}
}
}
fn golomb_to_signed(val: u32) -> i32 {
let sign = (((val & 0x1) as i32) << 1) - 1;
((val >> 1) as i32 + (val & 0x1) as i32) * sign
}
pub struct BitWriter<W: std::io::Write> {
inner: bitstream_io::write::BitWriter<W, bitstream_io::BigEndian>,
}
impl<W: std::io::Write> BitWriter<W> {
pub fn new(writer: W) -> Self {
Self {
inner: bitstream_io::write::BitWriter::new(writer),
}
}
pub fn writer(&mut self) -> Option<&mut W> {
self.inner.writer()
}
pub fn into_writer(self) -> W {
self.inner.into_writer()
}
}
impl<W: std::io::Write> BitWrite for BitWriter<W> {
fn write_ue(&mut self, value: u32) -> std::io::Result<()> {
if value == 0 {
self.inner.write_bit(true)
} else {
let code_num = value + 1;
let bits = 32 - code_num.leading_zeros(); let zeros = bits - 1;
for _ in 0..zeros {
self.inner.write_bit(false)?;
}
self.inner.write_var(bits, code_num)
}
}
fn write_se(&mut self, value: i32) -> std::io::Result<()> {
let code_num = if value > 0 {
(value as u32) * 2 - 1
} else {
(-value as u32) * 2
};
self.write_ue(code_num)
}
fn write_bit(&mut self, bit: bool) -> std::io::Result<()> {
self.inner.write_bit(bit)
}
fn write<const BITS: u32, I: Integer>(&mut self, value: I) -> std::io::Result<()> {
self.inner.write::<BITS, I>(value)
}
fn write_var<I: Integer>(&mut self, bit_count: u32, value: I) -> std::io::Result<()> {
self.inner.write_var::<I>(bit_count, value)
}
fn write_rbsp_trailing_bits(&mut self) -> std::io::Result<()> {
self.inner.write_bit(true)?; self.inner.byte_align()?; Ok(())
}
}
pub struct ByteWriter<W: Write> {
inner: W,
zero_count: u8,
}
impl<W: Write> ByteWriter<W> {
pub fn new(inner: W) -> Self {
Self {
inner,
zero_count: 0,
}
}
pub fn into_writer(self) -> W {
self.inner
}
}
impl<W: Write> Write for ByteWriter<W> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
let mut i = 0;
let mut chunk_start = 0;
while i < buf.len() {
if self.zero_count == 2 {
let b = buf[i];
if b <= 3 {
self.inner.write_all(&buf[chunk_start..i])?;
chunk_start = i;
self.inner.write_all(&[0x03])?;
}
self.zero_count = 0;
}
match memchr::memchr(0x00, &buf[i..]) {
None => {
self.zero_count = 0;
break;
}
Some(rel) => {
self.zero_count = if rel > 0 { 1 } else { self.zero_count + 1 };
i += rel + 1;
}
}
}
self.inner.write_all(&buf[chunk_start..])?;
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
self.inner.flush()
}
}
#[cfg(test)]
mod tests {
use super::*;
use hex_literal::*;
use hex_slice::AsHex;
#[test]
fn byte_reader() {
let data = hex!(
"67 64 00 0A AC 72 84 44 26 84 00 00 03
00 04 00 00 03 00 CA 3C 48 96 11 80"
);
for i in 1..data.len() - 1 {
let (head, tail) = data.split_at(i);
let r = head.chain(tail);
let mut r = ByteReader::skipping_h264_header(r);
let mut rbsp = Vec::new();
r.read_to_end(&mut rbsp).unwrap();
let expected = hex!(
"64 00 0A AC 72 84 44 26 84 00 00
00 04 00 00 00 CA 3C 48 96 11 80"
);
assert!(
rbsp == &expected[..],
"Mismatch with on split_at({}):\nrbsp {:02x}\nexpected {:02x}",
i,
rbsp.as_hex(),
expected.as_hex()
);
}
}
#[test]
fn bitreader_has_more_data() {
let mut reader = BitReader::new(&[0x12, 0x80][..]);
assert!(reader.has_more_rbsp_data("call 1").unwrap());
assert_eq!(reader.read::<8, u8>("u8 1").unwrap(), 0x12);
assert!(!reader.has_more_rbsp_data("call 2").unwrap());
let mut reader = BitReader::new(&[0x18][..]);
assert!(reader.has_more_rbsp_data("call 3").unwrap());
assert_eq!(reader.read::<4, u8>("u8 2").unwrap(), 0x1);
assert!(!reader.has_more_rbsp_data("call 4").unwrap());
let mut reader = BitReader::new(&[0x80, 0x00, 0x00][..]);
assert!(!reader
.has_more_rbsp_data("at end with cabac-zero-words")
.unwrap());
}
#[test]
fn byte_reader_emulation_prevention_beyond_max_fill() {
let mut input = vec![0xFF; 129];
input.extend_from_slice(&[0x00, 0x00, 0x03, 0x01]);
let mut r = ByteReader::without_skip(&input[..]);
let mut rbsp = Vec::new();
r.read_to_end(&mut rbsp).unwrap();
let mut expected = vec![0xFF; 129];
expected.extend_from_slice(&[0x00, 0x00, 0x01]);
assert_eq!(rbsp, expected, "emulation prevention byte was not stripped");
}
#[test]
fn read_ue_overflow() {
let mut reader = BitReader::new(&[0, 0, 0, 0, 255, 255, 255, 255, 255][..]);
assert!(matches!(
reader.read_ue("test"),
Err(BitReaderError::ExpGolombTooLarge("test"))
));
}
fn byte_writer_encode(rbsp: &[u8]) -> Vec<u8> {
let mut out = Vec::new();
ByteWriter::new(&mut out).write_all(rbsp).unwrap();
out
}
#[test]
fn byte_writer_no_escaping_needed() {
assert_eq!(byte_writer_encode(b"hello"), b"hello");
assert_eq!(
byte_writer_encode(&[0xFF, 0xFE, 0x01, 0x02]),
&[0xFF, 0xFE, 0x01, 0x02]
);
assert_eq!(byte_writer_encode(&[0x00, 0x04]), &[0x00, 0x04]);
assert_eq!(byte_writer_encode(&[0x00, 0x00, 0x04]), &[0x00, 0x00, 0x04]);
}
#[test]
fn byte_writer_escaping() {
assert_eq!(
byte_writer_encode(&[0x00, 0x00, 0x00]),
&[0x00, 0x00, 0x03, 0x00]
);
assert_eq!(
byte_writer_encode(&[0x00, 0x00, 0x01]),
&[0x00, 0x00, 0x03, 0x01]
);
assert_eq!(
byte_writer_encode(&[0x00, 0x00, 0x02]),
&[0x00, 0x00, 0x03, 0x02]
);
assert_eq!(
byte_writer_encode(&[0x00, 0x00, 0x03]),
&[0x00, 0x00, 0x03, 0x03]
);
}
#[test]
fn byte_writer_multiple_escapes() {
assert_eq!(
byte_writer_encode(&[0x00, 0x00, 0x00, 0x00, 0x00, 0x01]),
&[0x00, 0x00, 0x03, 0x00, 0x00, 0x03, 0x00, 0x01],
);
}
#[test]
fn byte_writer_split_writes() {
let mut out = Vec::new();
let mut w = ByteWriter::new(&mut out);
w.write_all(&[0x00, 0x00]).unwrap();
w.write_all(&[0x03]).unwrap(); drop(w);
assert_eq!(out, &[0x00, 0x00, 0x03, 0x03]);
let mut out2 = Vec::new();
let mut w2 = ByteWriter::new(&mut out2);
w2.write_all(&[0x00]).unwrap();
w2.write_all(&[0x00]).unwrap();
w2.write_all(&[0x01]).unwrap(); drop(w2);
assert_eq!(out2, &[0x00, 0x00, 0x03, 0x01]);
}
fn make_nal(hdr: u8, rbsp: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(1 + rbsp.len() + rbsp.len() / 3);
out.push(hdr);
ByteWriter::new(&mut out).write_all(rbsp).unwrap();
out
}
#[test]
fn byte_writer_roundtrip() {
let rbsp = hex!(
"64 00 0A AC 72 84 44 26 84 00 00
00 04 00 00 00 CA 3C 48 96 11 80"
);
let nal = make_nal(0x67, &rbsp);
let decoded = decode_nal(&nal).unwrap();
assert_eq!(&*decoded, &rbsp[..]);
}
#[test]
fn byte_reader_rejects_forbidden_sequences() {
for forbidden in [0x00u8, 0x01, 0x02] {
let nal = [0x67, 0x12, 0x00, 0x00, forbidden, 0x34];
let mut r = ByteReader::skipping_h264_header(&nal[..]);
let mut buf = Vec::new();
let err = r.read_to_end(&mut buf).unwrap_err();
assert_eq!(
err.kind(),
std::io::ErrorKind::InvalidData,
"expected InvalidData for 0x00 0x00 {:#04x}, got {:?}",
forbidden,
err.kind(),
);
}
}
#[test]
fn byte_writer_escape_inserted_in_nal() {
assert_eq!(
make_nal(0x68, &hex!("12 34 00 00 00 86")),
hex!("68 12 34 00 00 03 00 86"),
);
}
}