1pub mod aud;
8pub mod pps;
9pub mod prefix;
10pub mod sei;
11pub mod slice;
12pub mod sps;
13pub mod sps_extension;
14pub mod subset_sps;
15
16use crate::rbsp;
17use hex_slice::AsHex;
18use std::fmt;
19use std::io::Read;
20use std::num::NonZeroUsize;
21
22#[derive(PartialEq, Hash, Debug, Copy, Clone)]
23pub enum UnitType {
24 Unspecified(u8),
26 SliceLayerWithoutPartitioningNonIdr,
27 SliceDataPartitionALayer,
28 SliceDataPartitionBLayer,
29 SliceDataPartitionCLayer,
30 SliceLayerWithoutPartitioningIdr,
31 SEI,
33 SeqParameterSet,
34 PicParameterSet,
35 AccessUnitDelimiter,
36 EndOfSeq,
37 EndOfStream,
38 FillerData,
39 SeqParameterSetExtension,
40 PrefixNALUnit,
41 SubsetSeqParameterSet,
42 DepthParameterSet,
43 SliceLayerWithoutPartitioningAux,
44 SliceExtension,
45 SliceExtensionViewComponent,
46 Reserved(u8),
48}
49impl UnitType {
50 pub fn for_id(id: u8) -> Result<UnitType, UnitTypeError> {
51 if id > 31 {
52 Err(UnitTypeError::ValueOutOfRange(id))
53 } else {
54 let t = match id {
55 0 => UnitType::Unspecified(0),
56 1 => UnitType::SliceLayerWithoutPartitioningNonIdr,
57 2 => UnitType::SliceDataPartitionALayer,
58 3 => UnitType::SliceDataPartitionBLayer,
59 4 => UnitType::SliceDataPartitionCLayer,
60 5 => UnitType::SliceLayerWithoutPartitioningIdr,
61 6 => UnitType::SEI,
62 7 => UnitType::SeqParameterSet,
63 8 => UnitType::PicParameterSet,
64 9 => UnitType::AccessUnitDelimiter,
65 10 => UnitType::EndOfSeq,
66 11 => UnitType::EndOfStream,
67 12 => UnitType::FillerData,
68 13 => UnitType::SeqParameterSetExtension,
69 14 => UnitType::PrefixNALUnit,
70 15 => UnitType::SubsetSeqParameterSet,
71 16 => UnitType::DepthParameterSet,
72 17..=18 => UnitType::Reserved(id),
73 19 => UnitType::SliceLayerWithoutPartitioningAux,
74 20 => UnitType::SliceExtension,
75 21 => UnitType::SliceExtensionViewComponent,
76 22..=23 => UnitType::Reserved(id),
77 24..=31 => UnitType::Unspecified(id),
78 _ => panic!("unexpected {}", id), };
80 Ok(t)
81 }
82 }
83
84 pub fn id(self) -> u8 {
85 match self {
86 UnitType::Unspecified(v) => v,
87 UnitType::SliceLayerWithoutPartitioningNonIdr => 1,
88 UnitType::SliceDataPartitionALayer => 2,
89 UnitType::SliceDataPartitionBLayer => 3,
90 UnitType::SliceDataPartitionCLayer => 4,
91 UnitType::SliceLayerWithoutPartitioningIdr => 5,
92 UnitType::SEI => 6,
93 UnitType::SeqParameterSet => 7,
94 UnitType::PicParameterSet => 8,
95 UnitType::AccessUnitDelimiter => 9,
96 UnitType::EndOfSeq => 10,
97 UnitType::EndOfStream => 11,
98 UnitType::FillerData => 12,
99 UnitType::SeqParameterSetExtension => 13,
100 UnitType::PrefixNALUnit => 14,
101 UnitType::SubsetSeqParameterSet => 15,
102 UnitType::DepthParameterSet => 16,
103 UnitType::SliceLayerWithoutPartitioningAux => 19,
104 UnitType::SliceExtension => 20,
105 UnitType::SliceExtensionViewComponent => 21,
106 UnitType::Reserved(v) => v,
107 }
108 }
109}
110
111#[derive(Debug)]
112pub enum UnitTypeError {
113 ValueOutOfRange(u8),
115}
116
117#[derive(Copy, Clone, PartialEq, Eq)]
118pub struct NalHeader(u8);
119
120#[derive(Debug)]
121pub enum NalHeaderError {
122 ForbiddenZeroBit,
124}
125impl NalHeader {
126 pub fn new(header_value: u8) -> Result<NalHeader, NalHeaderError> {
127 if header_value & 0b1000_0000 != 0 {
128 Err(NalHeaderError::ForbiddenZeroBit)
129 } else {
130 Ok(NalHeader(header_value))
131 }
132 }
133
134 pub fn nal_ref_idc(self) -> u8 {
135 (self.0 & 0b0110_0000) >> 5
136 }
137
138 pub fn nal_unit_type(self) -> UnitType {
139 UnitType::for_id(self.0 & 0b0001_1111).unwrap()
140 }
141}
142impl From<NalHeader> for u8 {
143 fn from(v: NalHeader) -> Self {
144 v.0
145 }
146}
147
148impl fmt::Debug for NalHeader {
149 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
150 f.debug_struct("NalHeader")
151 .field("nal_ref_idc", &self.nal_ref_idc())
152 .field("nal_unit_type", &self.nal_unit_type())
153 .finish()
154 }
155}
156
157#[derive(Copy, Clone, PartialEq, Eq)]
168pub struct NalHeaderMvcExtension([u8; 3]);
169
170impl NalHeaderMvcExtension {
171 pub fn non_idr_flag(&self) -> bool {
172 self.0[0] & 0x40 != 0
173 }
174 pub fn priority_id(&self) -> u8 {
175 self.0[0] & 0x3F
176 }
177 pub fn view_id(&self) -> u16 {
178 ((self.0[1] as u16) << 2) | ((self.0[2] as u16) >> 6)
179 }
180 pub fn temporal_id(&self) -> u8 {
181 (self.0[2] >> 3) & 0x07
182 }
183 pub fn anchor_pic_flag(&self) -> bool {
184 self.0[2] & 0x04 != 0
185 }
186 pub fn inter_view_flag(&self) -> bool {
187 self.0[2] & 0x02 != 0
188 }
189}
190impl fmt::Debug for NalHeaderMvcExtension {
191 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
192 f.debug_struct("NalHeaderMvcExtension")
193 .field("non_idr_flag", &self.non_idr_flag())
194 .field("priority_id", &self.priority_id())
195 .field("view_id", &self.view_id())
196 .field("temporal_id", &self.temporal_id())
197 .field("anchor_pic_flag", &self.anchor_pic_flag())
198 .field("inter_view_flag", &self.inter_view_flag())
199 .finish()
200 }
201}
202
203#[derive(Copy, Clone, PartialEq, Eq)]
213pub struct NalHeaderSvcExtension([u8; 3]);
214
215impl NalHeaderSvcExtension {
216 pub fn idr_flag(&self) -> bool {
217 self.0[0] & 0x40 != 0
218 }
219 pub fn priority_id(&self) -> u8 {
220 self.0[0] & 0x3F
221 }
222 pub fn no_inter_layer_pred_flag(&self) -> bool {
223 self.0[1] & 0x80 != 0
224 }
225 pub fn dependency_id(&self) -> u8 {
226 (self.0[1] >> 4) & 0x07
227 }
228 pub fn quality_id(&self) -> u8 {
229 self.0[1] & 0x0F
230 }
231 pub fn temporal_id(&self) -> u8 {
232 (self.0[2] >> 5) & 0x07
233 }
234 pub fn use_ref_base_pic_flag(&self) -> bool {
235 self.0[2] & 0x10 != 0
236 }
237 pub fn discardable_flag(&self) -> bool {
238 self.0[2] & 0x08 != 0
239 }
240 pub fn output_flag(&self) -> bool {
241 self.0[2] & 0x04 != 0
242 }
243}
244impl fmt::Debug for NalHeaderSvcExtension {
245 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
246 f.debug_struct("NalHeaderSvcExtension")
247 .field("idr_flag", &self.idr_flag())
248 .field("priority_id", &self.priority_id())
249 .field("no_inter_layer_pred_flag", &self.no_inter_layer_pred_flag())
250 .field("dependency_id", &self.dependency_id())
251 .field("quality_id", &self.quality_id())
252 .field("temporal_id", &self.temporal_id())
253 .field("use_ref_base_pic_flag", &self.use_ref_base_pic_flag())
254 .field("discardable_flag", &self.discardable_flag())
255 .field("output_flag", &self.output_flag())
256 .finish()
257 }
258}
259
260#[derive(Copy, Clone, Debug, PartialEq, Eq)]
266pub enum NalHeaderExtension {
267 Mvc(NalHeaderMvcExtension),
268 Svc(NalHeaderSvcExtension),
269}
270
271impl NalHeaderExtension {
272 pub fn from_bytes(bytes: [u8; 3]) -> Self {
274 if bytes[0] & 0x80 != 0 {
275 NalHeaderExtension::Svc(NalHeaderSvcExtension(bytes))
276 } else {
277 NalHeaderExtension::Mvc(NalHeaderMvcExtension(bytes))
278 }
279 }
280}
281
282pub fn parse_nal_header_extension<N: Nal>(
288 nal: &N,
289) -> Result<(NalHeaderExtension, rbsp::ByteReader<N::BufRead>), std::io::Error> {
290 let mut reader = nal.reader();
291 let mut buf = [0u8; 4];
292 reader.read_exact(&mut buf)?;
293 let ext = NalHeaderExtension::from_bytes([buf[1], buf[2], buf[3]]);
294 let rbsp = rbsp::ByteReader::without_skip(reader);
295 Ok((ext, rbsp))
296}
297
298pub fn extended_rbsp_bytes<N: Nal>(nal: &N) -> rbsp::ByteReader<N::BufRead> {
302 let skip = NonZeroUsize::new(4).unwrap();
304 rbsp::ByteReader::skipping_bytes(nal.reader(), skip)
305}
306
307pub trait Nal {
359 type BufRead: std::io::BufRead + Clone;
360
361 fn is_complete(&self) -> bool;
363
364 fn header(&self) -> Result<NalHeader, NalHeaderError>;
366
367 fn reader(&self) -> Self::BufRead;
371
372 #[inline]
375 fn rbsp_bytes(&self) -> rbsp::ByteReader<Self::BufRead> {
376 rbsp::ByteReader::skipping_h264_header(self.reader())
377 }
378
379 #[inline]
381 fn rbsp_bits(&self) -> rbsp::BitReader<rbsp::ByteReader<Self::BufRead>> {
382 rbsp::BitReader::new(self.rbsp_bytes())
383 }
384}
385
386#[derive(Clone, Eq, PartialEq)]
388pub struct RefNal<'a> {
389 header: u8,
390 complete: bool,
391
392 head: &'a [u8],
394 tail: &'a [&'a [u8]],
395}
396impl<'a> RefNal<'a> {
397 #[inline]
399 pub fn new(head: &'a [u8], tail: &'a [&'a [u8]], complete: bool) -> Self {
400 for buf in tail {
401 debug_assert!(!buf.is_empty());
402 }
403 Self {
404 header: *head.first().expect("RefNal must be non-empty"),
405 head,
406 tail,
407 complete,
408 }
409 }
410}
411impl<'a> Nal for RefNal<'a> {
412 type BufRead = RefNalReader<'a>;
413
414 #[inline]
415 fn is_complete(&self) -> bool {
416 self.complete
417 }
418
419 #[inline]
420 fn header(&self) -> Result<NalHeader, NalHeaderError> {
421 NalHeader::new(self.header)
422 }
423
424 #[inline]
425 fn reader(&self) -> Self::BufRead {
426 RefNalReader {
427 cur: self.head,
428 tail: self.tail,
429 complete: self.complete,
430 }
431 }
432}
433impl<'a> std::fmt::Debug for RefNal<'a> {
434 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
435 f.debug_struct("RefNal")
437 .field("header", &self.header())
438 .field(
439 "data",
440 &RefNalReader {
441 cur: self.head,
442 tail: self.tail,
443 complete: self.complete,
444 },
445 )
446 .finish()
447 }
448}
449
450#[derive(Clone)]
456pub struct RefNalReader<'a> {
457 cur: &'a [u8],
459 tail: &'a [&'a [u8]],
460 complete: bool,
461}
462impl<'a> RefNalReader<'a> {
463 fn next_chunk(&mut self) {
464 match self.tail {
465 [first, tail @ ..] => {
466 self.cur = first;
467 self.tail = tail;
468 }
469 _ => self.cur = &[], }
471 }
472}
473impl<'a> std::io::Read for RefNalReader<'a> {
474 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
475 let len;
476 if buf.is_empty() {
477 len = 0;
478 } else if self.cur.is_empty() && !self.complete {
479 return Err(std::io::Error::new(
480 std::io::ErrorKind::WouldBlock,
481 "reached end of partially-buffered NAL",
482 ));
483 } else if buf.len() < self.cur.len() {
484 len = buf.len();
485 let (copy, keep) = self.cur.split_at(len);
486 buf.copy_from_slice(copy);
487 self.cur = keep;
488 } else {
489 len = self.cur.len();
490 buf[..len].copy_from_slice(self.cur);
491 self.next_chunk();
492 }
493 Ok(len)
494 }
495}
496impl<'a> std::io::BufRead for RefNalReader<'a> {
497 fn fill_buf(&mut self) -> std::io::Result<&[u8]> {
498 if self.cur.is_empty() && !self.complete {
499 return Err(std::io::Error::new(
500 std::io::ErrorKind::WouldBlock,
501 "reached end of partially-buffered NAL",
502 ));
503 }
504 Ok(self.cur)
505 }
506 fn consume(&mut self, amt: usize) {
507 self.cur = &self.cur[amt..];
508 if self.cur.is_empty() {
509 self.next_chunk();
510 }
511 }
512}
513impl<'a> std::fmt::Debug for RefNalReader<'a> {
514 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
515 write!(f, "{:02x}", self.cur.plain_hex(true))?;
516 for buf in self.tail {
517 write!(f, " {:02x}", buf.plain_hex(true))?;
518 }
519 if !self.complete {
520 f.write_str(" ...")?;
521 }
522 Ok(())
523 }
524}
525
526pub trait WritableNal {
527 fn write_bits<W: crate::rbsp::BitWrite>(&self, w: &mut W) -> std::io::Result<()>;
529
530 fn write_with_header<W: std::io::Write>(
533 &self,
534 hdr: NalHeader,
535 w: &mut W,
536 ) -> std::io::Result<()> {
537 w.write_all(&[hdr.into()])?;
538 let mut w = crate::rbsp::BitWriter::new(crate::rbsp::ByteWriter::new(w));
539 self.write_bits(&mut w)?;
540 Ok(())
541 }
542
543 fn to_vec_with_header(&self, hdr: NalHeader) -> Vec<u8> {
545 let mut v = Vec::new();
546 self.write_with_header(hdr, &mut v)
547 .expect("writing to Vec<u8> should not fail");
548 v
549 }
550}
551
552#[cfg(test)]
553mod test {
554 use std::io::{BufRead, Read};
555
556 use super::*;
557
558 #[test]
559 fn header() {
560 let h = NalHeader::new(0b0101_0001).unwrap();
561 assert_eq!(0b10, h.nal_ref_idc());
562 assert_eq!(UnitType::Reserved(17), h.nal_unit_type());
563 }
564
565 #[test]
566 fn ref_nal() {
567 fn common<'a>(head: &'a [u8], tail: &'a [&'a [u8]], complete: bool) -> RefNal<'a> {
568 let nal = RefNal::new(head, tail, complete);
569 assert_eq!(NalHeader::new(0b0101_0001).unwrap(), nal.header().unwrap());
570
571 let mut r = nal.reader();
573 let mut buf = [0u8; 5];
574 r.read_exact(&mut buf).unwrap();
575 assert_eq!(&buf[..], &[0b0101_0001, 1, 2, 3, 4]);
576 if complete {
577 assert_eq!(r.read(&mut buf[..]).unwrap(), 0);
578
579 let mut buf = Vec::new();
581 nal.reader().read_to_end(&mut buf).unwrap();
582 assert_eq!(buf, &[0b0101_0001, 1, 2, 3, 4]);
583 } else {
584 assert_eq!(
585 r.read(&mut buf[..]).unwrap_err().kind(),
586 std::io::ErrorKind::WouldBlock
587 );
588 }
589
590 nal
592 }
593
594 let nal = common(&[0b0101_0001, 1, 2, 3, 4], &[], false);
596 let mut r = nal.reader();
597 assert_eq!(r.fill_buf().unwrap(), &[0b0101_0001, 1, 2, 3, 4]);
598 r.consume(1);
599 assert_eq!(r.fill_buf().unwrap(), &[1, 2, 3, 4]);
600 r.consume(4);
601 assert_eq!(
602 r.fill_buf().unwrap_err().kind(),
603 std::io::ErrorKind::WouldBlock
604 );
605
606 let nal = common(&[0b0101_0001], &[&[1, 2], &[3, 4]], false);
608 let mut r = nal.reader();
609 assert_eq!(r.fill_buf().unwrap(), &[0b0101_0001]);
610 r.consume(1);
611 assert_eq!(r.fill_buf().unwrap(), &[1, 2]);
612 r.consume(2);
613 assert_eq!(r.fill_buf().unwrap(), &[3, 4]);
614 r.consume(1);
615 assert_eq!(r.fill_buf().unwrap(), &[4]);
616 r.consume(1);
617 assert_eq!(
618 r.fill_buf().unwrap_err().kind(),
619 std::io::ErrorKind::WouldBlock
620 );
621
622 let nal = common(&[0b0101_0001, 1, 2, 3, 4], &[], true);
624 let mut r = nal.reader();
625 assert_eq!(r.fill_buf().unwrap(), &[0b0101_0001, 1, 2, 3, 4]);
626 r.consume(1);
627 assert_eq!(r.fill_buf().unwrap(), &[1, 2, 3, 4]);
628 r.consume(4);
629 assert!(r.fill_buf().unwrap().is_empty());
630 }
631
632 #[test]
633 fn mvc_header_extension() {
634 let bytes: [u8; 3] = [
639 0b0100_0011, 0b01100000, 0b1010_1101, ];
643 let ext = NalHeaderExtension::from_bytes(bytes);
644 match ext {
645 NalHeaderExtension::Mvc(mvc) => {
646 assert!(mvc.non_idr_flag());
647 assert_eq!(mvc.priority_id(), 3);
648 assert_eq!(mvc.view_id(), 386);
650 assert_eq!(mvc.temporal_id(), 5);
651 assert!(mvc.anchor_pic_flag());
652 assert!(!mvc.inter_view_flag());
653 }
654 _ => panic!("expected MVC extension"),
655 }
656 }
657
658 #[test]
659 fn mvc_header_extension_view_id_zero() {
660 let bytes: [u8; 3] = [0x00, 0x00, 0x01];
662 let ext = NalHeaderExtension::from_bytes(bytes);
663 match ext {
664 NalHeaderExtension::Mvc(mvc) => {
665 assert!(!mvc.non_idr_flag());
666 assert_eq!(mvc.priority_id(), 0);
667 assert_eq!(mvc.view_id(), 0);
668 assert_eq!(mvc.temporal_id(), 0);
669 assert!(!mvc.anchor_pic_flag());
670 assert!(!mvc.inter_view_flag());
671 }
672 _ => panic!("expected MVC extension"),
673 }
674 }
675
676 #[test]
677 fn mvc_header_extension_max_view_id() {
678 let bytes: [u8; 3] = [0x00, 0xFF, 0b1100_0001];
681 let ext = NalHeaderExtension::from_bytes(bytes);
682 match ext {
683 NalHeaderExtension::Mvc(mvc) => {
684 assert_eq!(mvc.view_id(), 1023);
685 }
686 _ => panic!("expected MVC extension"),
687 }
688 }
689
690 #[test]
691 fn svc_header_extension() {
692 let bytes: [u8; 3] = [
697 0b1010_1010, 0b1110_0011, 0b0101_0111, ];
701 let ext = NalHeaderExtension::from_bytes(bytes);
702 match ext {
703 NalHeaderExtension::Svc(svc) => {
704 assert!(!svc.idr_flag());
705 assert_eq!(svc.priority_id(), 42);
706 assert!(svc.no_inter_layer_pred_flag());
707 assert_eq!(svc.dependency_id(), 6);
708 assert_eq!(svc.quality_id(), 3);
709 assert_eq!(svc.temporal_id(), 2);
710 assert!(svc.use_ref_base_pic_flag());
711 assert!(!svc.discardable_flag());
712 assert!(svc.output_flag());
713 }
714 _ => panic!("expected SVC extension"),
715 }
716 }
717
718 #[test]
719 fn svc_header_extension_idr() {
720 let bytes: [u8; 3] = [0b1100_0000, 0b0000_0000, 0b0000_0011];
722 let ext = NalHeaderExtension::from_bytes(bytes);
723 match ext {
724 NalHeaderExtension::Svc(svc) => {
725 assert!(svc.idr_flag());
726 assert_eq!(svc.priority_id(), 0);
727 assert!(!svc.no_inter_layer_pred_flag());
728 assert_eq!(svc.dependency_id(), 0);
729 assert_eq!(svc.quality_id(), 0);
730 assert_eq!(svc.temporal_id(), 0);
731 assert!(!svc.use_ref_base_pic_flag());
732 assert!(!svc.discardable_flag());
733 assert!(!svc.output_flag());
734 }
735 _ => panic!("expected SVC extension"),
736 }
737 }
738
739 #[test]
740 fn parse_nal_header_extension_from_refnal() {
741 let nal_bytes: &[u8] = &[
745 0x6E, 0x00, 0x00, 0b0100_0001, 0xAA,
750 0xBB, ];
752 let nal = RefNal::new(nal_bytes, &[], true);
753 let (ext, _rbsp) = parse_nal_header_extension(&nal).unwrap();
754 match ext {
755 NalHeaderExtension::Mvc(mvc) => {
756 assert_eq!(mvc.view_id(), 1);
757 assert!(!mvc.non_idr_flag());
758 }
759 _ => panic!("expected MVC extension"),
760 }
761 }
762
763 #[test]
764 fn reader_debug() {
765 assert_eq!(
766 format!(
767 "{:?}",
768 RefNalReader {
769 cur: &b"\x00"[..],
770 tail: &[&b"\x01"[..], &b"\x02\x03"[..]],
771 complete: false,
772 }
773 ),
774 "00 01 02 03 ..."
775 );
776 }
777}