pub mod aud;
pub mod pps;
pub mod prefix;
pub mod sei;
pub mod slice;
pub mod sps;
pub mod sps_extension;
pub mod subset_sps;
use crate::rbsp;
use hex_slice::AsHex;
use std::fmt;
use std::io::Read;
use std::num::NonZeroUsize;
#[derive(PartialEq, Hash, Debug, Copy, Clone)]
pub enum UnitType {
Unspecified(u8),
SliceLayerWithoutPartitioningNonIdr,
SliceDataPartitionALayer,
SliceDataPartitionBLayer,
SliceDataPartitionCLayer,
SliceLayerWithoutPartitioningIdr,
SEI,
SeqParameterSet,
PicParameterSet,
AccessUnitDelimiter,
EndOfSeq,
EndOfStream,
FillerData,
SeqParameterSetExtension,
PrefixNALUnit,
SubsetSeqParameterSet,
DepthParameterSet,
SliceLayerWithoutPartitioningAux,
SliceExtension,
SliceExtensionViewComponent,
Reserved(u8),
}
impl UnitType {
pub fn for_id(id: u8) -> Result<UnitType, UnitTypeError> {
if id > 31 {
Err(UnitTypeError::ValueOutOfRange(id))
} else {
let t = match id {
0 => UnitType::Unspecified(0),
1 => UnitType::SliceLayerWithoutPartitioningNonIdr,
2 => UnitType::SliceDataPartitionALayer,
3 => UnitType::SliceDataPartitionBLayer,
4 => UnitType::SliceDataPartitionCLayer,
5 => UnitType::SliceLayerWithoutPartitioningIdr,
6 => UnitType::SEI,
7 => UnitType::SeqParameterSet,
8 => UnitType::PicParameterSet,
9 => UnitType::AccessUnitDelimiter,
10 => UnitType::EndOfSeq,
11 => UnitType::EndOfStream,
12 => UnitType::FillerData,
13 => UnitType::SeqParameterSetExtension,
14 => UnitType::PrefixNALUnit,
15 => UnitType::SubsetSeqParameterSet,
16 => UnitType::DepthParameterSet,
17..=18 => UnitType::Reserved(id),
19 => UnitType::SliceLayerWithoutPartitioningAux,
20 => UnitType::SliceExtension,
21 => UnitType::SliceExtensionViewComponent,
22..=23 => UnitType::Reserved(id),
24..=31 => UnitType::Unspecified(id),
_ => panic!("unexpected {}", id), };
Ok(t)
}
}
pub fn id(self) -> u8 {
match self {
UnitType::Unspecified(v) => v,
UnitType::SliceLayerWithoutPartitioningNonIdr => 1,
UnitType::SliceDataPartitionALayer => 2,
UnitType::SliceDataPartitionBLayer => 3,
UnitType::SliceDataPartitionCLayer => 4,
UnitType::SliceLayerWithoutPartitioningIdr => 5,
UnitType::SEI => 6,
UnitType::SeqParameterSet => 7,
UnitType::PicParameterSet => 8,
UnitType::AccessUnitDelimiter => 9,
UnitType::EndOfSeq => 10,
UnitType::EndOfStream => 11,
UnitType::FillerData => 12,
UnitType::SeqParameterSetExtension => 13,
UnitType::PrefixNALUnit => 14,
UnitType::SubsetSeqParameterSet => 15,
UnitType::DepthParameterSet => 16,
UnitType::SliceLayerWithoutPartitioningAux => 19,
UnitType::SliceExtension => 20,
UnitType::SliceExtensionViewComponent => 21,
UnitType::Reserved(v) => v,
}
}
}
#[derive(Debug)]
pub enum UnitTypeError {
ValueOutOfRange(u8),
}
#[derive(Copy, Clone, PartialEq, Eq)]
pub struct NalHeader(u8);
#[derive(Debug)]
pub enum NalHeaderError {
ForbiddenZeroBit,
}
impl NalHeader {
pub fn new(header_value: u8) -> Result<NalHeader, NalHeaderError> {
if header_value & 0b1000_0000 != 0 {
Err(NalHeaderError::ForbiddenZeroBit)
} else {
Ok(NalHeader(header_value))
}
}
pub fn nal_ref_idc(self) -> u8 {
(self.0 & 0b0110_0000) >> 5
}
pub fn nal_unit_type(self) -> UnitType {
UnitType::for_id(self.0 & 0b0001_1111).unwrap()
}
}
impl From<NalHeader> for u8 {
fn from(v: NalHeader) -> Self {
v.0
}
}
impl fmt::Debug for NalHeader {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
f.debug_struct("NalHeader")
.field("nal_ref_idc", &self.nal_ref_idc())
.field("nal_unit_type", &self.nal_unit_type())
.finish()
}
}
#[derive(Copy, Clone, PartialEq, Eq)]
pub struct NalHeaderMvcExtension([u8; 3]);
impl NalHeaderMvcExtension {
pub fn non_idr_flag(&self) -> bool {
self.0[0] & 0x40 != 0
}
pub fn priority_id(&self) -> u8 {
self.0[0] & 0x3F
}
pub fn view_id(&self) -> u16 {
((self.0[1] as u16) << 2) | ((self.0[2] as u16) >> 6)
}
pub fn temporal_id(&self) -> u8 {
(self.0[2] >> 3) & 0x07
}
pub fn anchor_pic_flag(&self) -> bool {
self.0[2] & 0x04 != 0
}
pub fn inter_view_flag(&self) -> bool {
self.0[2] & 0x02 != 0
}
}
impl fmt::Debug for NalHeaderMvcExtension {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NalHeaderMvcExtension")
.field("non_idr_flag", &self.non_idr_flag())
.field("priority_id", &self.priority_id())
.field("view_id", &self.view_id())
.field("temporal_id", &self.temporal_id())
.field("anchor_pic_flag", &self.anchor_pic_flag())
.field("inter_view_flag", &self.inter_view_flag())
.finish()
}
}
#[derive(Copy, Clone, PartialEq, Eq)]
pub struct NalHeaderSvcExtension([u8; 3]);
impl NalHeaderSvcExtension {
pub fn idr_flag(&self) -> bool {
self.0[0] & 0x40 != 0
}
pub fn priority_id(&self) -> u8 {
self.0[0] & 0x3F
}
pub fn no_inter_layer_pred_flag(&self) -> bool {
self.0[1] & 0x80 != 0
}
pub fn dependency_id(&self) -> u8 {
(self.0[1] >> 4) & 0x07
}
pub fn quality_id(&self) -> u8 {
self.0[1] & 0x0F
}
pub fn temporal_id(&self) -> u8 {
(self.0[2] >> 5) & 0x07
}
pub fn use_ref_base_pic_flag(&self) -> bool {
self.0[2] & 0x10 != 0
}
pub fn discardable_flag(&self) -> bool {
self.0[2] & 0x08 != 0
}
pub fn output_flag(&self) -> bool {
self.0[2] & 0x04 != 0
}
}
impl fmt::Debug for NalHeaderSvcExtension {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NalHeaderSvcExtension")
.field("idr_flag", &self.idr_flag())
.field("priority_id", &self.priority_id())
.field("no_inter_layer_pred_flag", &self.no_inter_layer_pred_flag())
.field("dependency_id", &self.dependency_id())
.field("quality_id", &self.quality_id())
.field("temporal_id", &self.temporal_id())
.field("use_ref_base_pic_flag", &self.use_ref_base_pic_flag())
.field("discardable_flag", &self.discardable_flag())
.field("output_flag", &self.output_flag())
.finish()
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum NalHeaderExtension {
Mvc(NalHeaderMvcExtension),
Svc(NalHeaderSvcExtension),
}
impl NalHeaderExtension {
pub fn from_bytes(bytes: [u8; 3]) -> Self {
if bytes[0] & 0x80 != 0 {
NalHeaderExtension::Svc(NalHeaderSvcExtension(bytes))
} else {
NalHeaderExtension::Mvc(NalHeaderMvcExtension(bytes))
}
}
}
pub fn parse_nal_header_extension<N: Nal>(
nal: &N,
) -> Result<(NalHeaderExtension, rbsp::ByteReader<N::BufRead>), std::io::Error> {
let mut reader = nal.reader();
let mut buf = [0u8; 4];
reader.read_exact(&mut buf)?;
let ext = NalHeaderExtension::from_bytes([buf[1], buf[2], buf[3]]);
let rbsp = rbsp::ByteReader::without_skip(reader);
Ok((ext, rbsp))
}
pub fn extended_rbsp_bytes<N: Nal>(nal: &N) -> rbsp::ByteReader<N::BufRead> {
let skip = NonZeroUsize::new(4).unwrap();
rbsp::ByteReader::skipping_bytes(nal.reader(), skip)
}
pub trait Nal {
type BufRead: std::io::BufRead + Clone;
fn is_complete(&self) -> bool;
fn header(&self) -> Result<NalHeader, NalHeaderError>;
fn reader(&self) -> Self::BufRead;
#[inline]
fn rbsp_bytes(&self) -> rbsp::ByteReader<Self::BufRead> {
rbsp::ByteReader::skipping_h264_header(self.reader())
}
#[inline]
fn rbsp_bits(&self) -> rbsp::BitReader<rbsp::ByteReader<Self::BufRead>> {
rbsp::BitReader::new(self.rbsp_bytes())
}
}
#[derive(Clone, Eq, PartialEq)]
pub struct RefNal<'a> {
header: u8,
complete: bool,
head: &'a [u8],
tail: &'a [&'a [u8]],
}
impl<'a> RefNal<'a> {
#[inline]
pub fn new(head: &'a [u8], tail: &'a [&'a [u8]], complete: bool) -> Self {
for buf in tail {
debug_assert!(!buf.is_empty());
}
Self {
header: *head.first().expect("RefNal must be non-empty"),
head,
tail,
complete,
}
}
}
impl<'a> Nal for RefNal<'a> {
type BufRead = RefNalReader<'a>;
#[inline]
fn is_complete(&self) -> bool {
self.complete
}
#[inline]
fn header(&self) -> Result<NalHeader, NalHeaderError> {
NalHeader::new(self.header)
}
#[inline]
fn reader(&self) -> Self::BufRead {
RefNalReader {
cur: self.head,
tail: self.tail,
complete: self.complete,
}
}
}
impl<'a> std::fmt::Debug for RefNal<'a> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RefNal")
.field("header", &self.header())
.field(
"data",
&RefNalReader {
cur: self.head,
tail: self.tail,
complete: self.complete,
},
)
.finish()
}
}
#[derive(Clone)]
pub struct RefNalReader<'a> {
cur: &'a [u8],
tail: &'a [&'a [u8]],
complete: bool,
}
impl<'a> RefNalReader<'a> {
fn next_chunk(&mut self) {
match self.tail {
[first, tail @ ..] => {
self.cur = first;
self.tail = tail;
}
_ => self.cur = &[], }
}
}
impl<'a> std::io::Read for RefNalReader<'a> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
let len;
if buf.is_empty() {
len = 0;
} else if self.cur.is_empty() && !self.complete {
return Err(std::io::Error::new(
std::io::ErrorKind::WouldBlock,
"reached end of partially-buffered NAL",
));
} else if buf.len() < self.cur.len() {
len = buf.len();
let (copy, keep) = self.cur.split_at(len);
buf.copy_from_slice(copy);
self.cur = keep;
} else {
len = self.cur.len();
buf[..len].copy_from_slice(self.cur);
self.next_chunk();
}
Ok(len)
}
}
impl<'a> std::io::BufRead for RefNalReader<'a> {
fn fill_buf(&mut self) -> std::io::Result<&[u8]> {
if self.cur.is_empty() && !self.complete {
return Err(std::io::Error::new(
std::io::ErrorKind::WouldBlock,
"reached end of partially-buffered NAL",
));
}
Ok(self.cur)
}
fn consume(&mut self, amt: usize) {
self.cur = &self.cur[amt..];
if self.cur.is_empty() {
self.next_chunk();
}
}
}
impl<'a> std::fmt::Debug for RefNalReader<'a> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{:02x}", self.cur.plain_hex(true))?;
for buf in self.tail {
write!(f, " {:02x}", buf.plain_hex(true))?;
}
if !self.complete {
f.write_str(" ...")?;
}
Ok(())
}
}
pub trait WritableNal {
fn write_bits<W: crate::rbsp::BitWrite>(&self, w: &mut W) -> std::io::Result<()>;
fn write_with_header<W: std::io::Write>(
&self,
hdr: NalHeader,
w: &mut W,
) -> std::io::Result<()> {
w.write_all(&[hdr.into()])?;
let mut w = crate::rbsp::BitWriter::new(crate::rbsp::ByteWriter::new(w));
self.write_bits(&mut w)?;
Ok(())
}
fn to_vec_with_header(&self, hdr: NalHeader) -> Vec<u8> {
let mut v = Vec::new();
self.write_with_header(hdr, &mut v)
.expect("writing to Vec<u8> should not fail");
v
}
}
#[cfg(test)]
mod test {
use std::io::{BufRead, Read};
use super::*;
#[test]
fn header() {
let h = NalHeader::new(0b0101_0001).unwrap();
assert_eq!(0b10, h.nal_ref_idc());
assert_eq!(UnitType::Reserved(17), h.nal_unit_type());
}
#[test]
fn ref_nal() {
fn common<'a>(head: &'a [u8], tail: &'a [&'a [u8]], complete: bool) -> RefNal<'a> {
let nal = RefNal::new(head, tail, complete);
assert_eq!(NalHeader::new(0b0101_0001).unwrap(), nal.header().unwrap());
let mut r = nal.reader();
let mut buf = [0u8; 5];
r.read_exact(&mut buf).unwrap();
assert_eq!(&buf[..], &[0b0101_0001, 1, 2, 3, 4]);
if complete {
assert_eq!(r.read(&mut buf[..]).unwrap(), 0);
let mut buf = Vec::new();
nal.reader().read_to_end(&mut buf).unwrap();
assert_eq!(buf, &[0b0101_0001, 1, 2, 3, 4]);
} else {
assert_eq!(
r.read(&mut buf[..]).unwrap_err().kind(),
std::io::ErrorKind::WouldBlock
);
}
nal
}
let nal = common(&[0b0101_0001, 1, 2, 3, 4], &[], false);
let mut r = nal.reader();
assert_eq!(r.fill_buf().unwrap(), &[0b0101_0001, 1, 2, 3, 4]);
r.consume(1);
assert_eq!(r.fill_buf().unwrap(), &[1, 2, 3, 4]);
r.consume(4);
assert_eq!(
r.fill_buf().unwrap_err().kind(),
std::io::ErrorKind::WouldBlock
);
let nal = common(&[0b0101_0001], &[&[1, 2], &[3, 4]], false);
let mut r = nal.reader();
assert_eq!(r.fill_buf().unwrap(), &[0b0101_0001]);
r.consume(1);
assert_eq!(r.fill_buf().unwrap(), &[1, 2]);
r.consume(2);
assert_eq!(r.fill_buf().unwrap(), &[3, 4]);
r.consume(1);
assert_eq!(r.fill_buf().unwrap(), &[4]);
r.consume(1);
assert_eq!(
r.fill_buf().unwrap_err().kind(),
std::io::ErrorKind::WouldBlock
);
let nal = common(&[0b0101_0001, 1, 2, 3, 4], &[], true);
let mut r = nal.reader();
assert_eq!(r.fill_buf().unwrap(), &[0b0101_0001, 1, 2, 3, 4]);
r.consume(1);
assert_eq!(r.fill_buf().unwrap(), &[1, 2, 3, 4]);
r.consume(4);
assert!(r.fill_buf().unwrap().is_empty());
}
#[test]
fn mvc_header_extension() {
let bytes: [u8; 3] = [
0b0100_0011, 0b01100000, 0b1010_1101, ];
let ext = NalHeaderExtension::from_bytes(bytes);
match ext {
NalHeaderExtension::Mvc(mvc) => {
assert!(mvc.non_idr_flag());
assert_eq!(mvc.priority_id(), 3);
assert_eq!(mvc.view_id(), 386);
assert_eq!(mvc.temporal_id(), 5);
assert!(mvc.anchor_pic_flag());
assert!(!mvc.inter_view_flag());
}
_ => panic!("expected MVC extension"),
}
}
#[test]
fn mvc_header_extension_view_id_zero() {
let bytes: [u8; 3] = [0x00, 0x00, 0x01];
let ext = NalHeaderExtension::from_bytes(bytes);
match ext {
NalHeaderExtension::Mvc(mvc) => {
assert!(!mvc.non_idr_flag());
assert_eq!(mvc.priority_id(), 0);
assert_eq!(mvc.view_id(), 0);
assert_eq!(mvc.temporal_id(), 0);
assert!(!mvc.anchor_pic_flag());
assert!(!mvc.inter_view_flag());
}
_ => panic!("expected MVC extension"),
}
}
#[test]
fn mvc_header_extension_max_view_id() {
let bytes: [u8; 3] = [0x00, 0xFF, 0b1100_0001];
let ext = NalHeaderExtension::from_bytes(bytes);
match ext {
NalHeaderExtension::Mvc(mvc) => {
assert_eq!(mvc.view_id(), 1023);
}
_ => panic!("expected MVC extension"),
}
}
#[test]
fn svc_header_extension() {
let bytes: [u8; 3] = [
0b1010_1010, 0b1110_0011, 0b0101_0111, ];
let ext = NalHeaderExtension::from_bytes(bytes);
match ext {
NalHeaderExtension::Svc(svc) => {
assert!(!svc.idr_flag());
assert_eq!(svc.priority_id(), 42);
assert!(svc.no_inter_layer_pred_flag());
assert_eq!(svc.dependency_id(), 6);
assert_eq!(svc.quality_id(), 3);
assert_eq!(svc.temporal_id(), 2);
assert!(svc.use_ref_base_pic_flag());
assert!(!svc.discardable_flag());
assert!(svc.output_flag());
}
_ => panic!("expected SVC extension"),
}
}
#[test]
fn svc_header_extension_idr() {
let bytes: [u8; 3] = [0b1100_0000, 0b0000_0000, 0b0000_0011];
let ext = NalHeaderExtension::from_bytes(bytes);
match ext {
NalHeaderExtension::Svc(svc) => {
assert!(svc.idr_flag());
assert_eq!(svc.priority_id(), 0);
assert!(!svc.no_inter_layer_pred_flag());
assert_eq!(svc.dependency_id(), 0);
assert_eq!(svc.quality_id(), 0);
assert_eq!(svc.temporal_id(), 0);
assert!(!svc.use_ref_base_pic_flag());
assert!(!svc.discardable_flag());
assert!(!svc.output_flag());
}
_ => panic!("expected SVC extension"),
}
}
#[test]
fn parse_nal_header_extension_from_refnal() {
let nal_bytes: &[u8] = &[
0x6E, 0x00, 0x00, 0b0100_0001, 0xAA,
0xBB, ];
let nal = RefNal::new(nal_bytes, &[], true);
let (ext, _rbsp) = parse_nal_header_extension(&nal).unwrap();
match ext {
NalHeaderExtension::Mvc(mvc) => {
assert_eq!(mvc.view_id(), 1);
assert!(!mvc.non_idr_flag());
}
_ => panic!("expected MVC extension"),
}
}
#[test]
fn reader_debug() {
assert_eq!(
format!(
"{:?}",
RefNalReader {
cur: &b"\x00"[..],
tail: &[&b"\x01"[..], &b"\x02\x03"[..]],
complete: false,
}
),
"00 01 02 03 ..."
);
}
}