#![forbid(unsafe_code)]
use crate::DecodeError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum ObuType {
SequenceHeader,
TemporalDelimiter,
FrameHeader,
TileGroup,
Frame,
RedundantFrameHeader,
Other(u8),
}
impl ObuType {
const fn from_u8(value: u8) -> Self {
match value {
1 => Self::SequenceHeader,
2 => Self::TemporalDelimiter,
3 => Self::FrameHeader,
4 => Self::TileGroup,
6 => Self::Frame,
7 => Self::RedundantFrameHeader,
other => Self::Other(other),
}
}
}
#[derive(Debug, Clone, Copy)]
pub(super) struct Obu<'a> {
pub(super) obu_type: ObuType,
pub(super) payload: &'a [u8],
}
pub(super) fn read_leb128(data: &[u8]) -> Result<(u64, usize), DecodeError> {
let mut value: u64 = 0;
for i in 0..8usize {
let &byte = data.get(i).ok_or(DecodeError::InvalidInput)?;
let low7 = u64::from(byte & 0x7f);
let shift = u32::try_from(i * 7).unwrap_or(u32::MAX);
value |= low7.checked_shl(shift).ok_or(DecodeError::InvalidInput)?;
if byte & 0x80 == 0 {
return Ok((value, i + 1));
}
}
Err(DecodeError::InvalidInput)
}
const fn parse_obu_header(byte: u8) -> Result<(ObuType, bool), DecodeError> {
let forbidden_bit = (byte >> 7) & 1;
if forbidden_bit != 0 {
return Err(DecodeError::InvalidInput);
}
let obu_type = (byte >> 3) & 0b1111;
let extension_flag = (byte >> 2) & 1;
if extension_flag != 0 {
return Err(DecodeError::Unsupported);
}
let has_size_field = (byte >> 1) & 1;
Ok((ObuType::from_u8(obu_type), has_size_field != 0))
}
pub(super) fn split_obus(data: &[u8]) -> Result<Vec<Obu<'_>>, DecodeError> {
let mut out = Vec::new();
let mut pos = 0usize;
while pos < data.len() {
let &header_byte = data.get(pos).ok_or(DecodeError::InvalidInput)?;
let (obu_type, has_size_field) = parse_obu_header(header_byte)?;
if !has_size_field {
return Err(DecodeError::Unsupported);
}
pos += 1;
let (size, size_len) = read_leb128(data.get(pos..).ok_or(DecodeError::InvalidInput)?)?;
pos += size_len;
let size = usize::try_from(size).map_err(|_err| DecodeError::InvalidInput)?;
let end = pos.checked_add(size).ok_or(DecodeError::InvalidInput)?;
let payload = data.get(pos..end).ok_or(DecodeError::InvalidInput)?;
out.push(Obu { obu_type, payload });
pos = end;
}
Ok(out)
}
#[cfg(test)]
#[path = "av1_obu_tests.rs"]
mod tests;