1use bitstream_io::read::BitRead as _;
26use bitstream_io::write::BitWrite as _;
27use std::borrow::Cow;
28use std::io::BufRead;
29use std::io::Read;
30use std::io::Write;
31use std::num::NonZeroUsize;
32
33#[derive(Copy, Clone, Debug)]
34enum ParseState {
35 Start(u8),
38 Skip(NonZeroUsize),
39 Three,
40 PostThree,
41}
42
43const H264_HEADER_LEN: NonZeroUsize = match NonZeroUsize::new(1) {
44 Some(one) => one,
45 None => panic!("1 should be non-zero"),
46};
47
48fn zero_pair_finder() -> &'static memchr::memmem::Finder<'static> {
49 static FINDER: std::sync::OnceLock<memchr::memmem::Finder<'static>> =
50 std::sync::OnceLock::new();
51 FINDER.get_or_init(|| memchr::memmem::Finder::new(b"\x00\x00"))
52}
53
54#[derive(Clone)]
64pub struct ByteReader<R: BufRead> {
65 inner: R,
72 state: ParseState,
73 i: usize,
74
75 max_fill: usize,
79}
80impl<R: BufRead> ByteReader<R> {
81 pub fn without_skip(inner: R) -> Self {
83 Self {
84 inner,
85 state: ParseState::Start(0),
86 i: 0,
87 max_fill: 128,
88 }
89 }
90
91 pub fn skipping_h264_header(inner: R) -> Self {
93 Self {
94 inner,
95 state: ParseState::Skip(H264_HEADER_LEN),
96 i: 0,
97 max_fill: 128,
98 }
99 }
100
101 pub fn skipping_bytes(inner: R, skip: NonZeroUsize) -> Self {
106 Self {
107 inner,
108 state: ParseState::Skip(skip),
109 i: 0,
110 max_fill: 128,
111 }
112 }
113
114 fn try_fill_buf_slow(&mut self) -> std::io::Result<bool> {
118 debug_assert_eq!(self.i, 0);
119 let chunk = self.inner.fill_buf()?;
120 if chunk.is_empty() {
121 return Ok(false);
122 }
123
124 let limit = std::cmp::min(chunk.len(), self.max_fill);
125 while self.i < limit {
126 match self.state {
127 ParseState::Start(zero_count) => {
128 let after_pair = if zero_count >= 2 {
131 Some(self.i)
133 } else if zero_count == 1 && chunk[self.i] == 0x00 {
134 if self.i + 1 < limit {
136 Some(self.i + 1)
137 } else {
138 self.state = ParseState::Start(2);
139 self.i += 1;
140 None
141 }
142 } else {
143 match zero_pair_finder().find(&chunk[self.i..limit]) {
145 Some(offset) => {
146 let ap = self.i + offset + 2;
147 if ap < limit {
148 Some(ap)
149 } else {
150 self.state = ParseState::Start(2);
151 self.i = ap;
152 None
153 }
154 }
155 None => {
156 let trailing = if limit > self.i && chunk[limit - 1] == 0x00 {
157 1
158 } else {
159 0
160 };
161 self.state = ParseState::Start(trailing);
162 self.i = limit;
163 None
164 }
165 }
166 };
167 let Some(after_pair) = after_pair else { break };
168 match chunk[after_pair] {
170 0x03 => {
171 self.i = after_pair;
172 self.state = ParseState::Three;
173 break;
174 }
175
176 b @ 0x00..=0x02 => {
183 return Err(std::io::Error::new(
184 std::io::ErrorKind::InvalidData,
185 format!("invalid RBSP byte {:#x} in state {:?}", b, &self.state,),
186 ));
187 }
188 _ => {
189 self.i = after_pair + 1;
190 self.state = ParseState::Start(0);
191 continue;
192 }
193 }
194 }
195 ParseState::Skip(remaining) => {
196 debug_assert_eq!(self.i, 0);
197 let skip = std::cmp::min(chunk.len(), remaining.get());
198 self.inner.consume(skip);
199 self.state = NonZeroUsize::new(remaining.get() - skip)
200 .map(ParseState::Skip)
201 .unwrap_or(ParseState::Start(0));
202 break;
203 }
204 ParseState::Three => {
205 debug_assert_eq!(self.i, 0);
206 self.inner.consume(1);
207 self.state = ParseState::PostThree;
208 break;
209 }
210 ParseState::PostThree => {
211 match chunk[self.i] {
212 0x00 => self.state = ParseState::Start(1),
213 0x01 | 0x02 | 0x03 => self.state = ParseState::Start(0),
214 o => {
215 return Err(std::io::Error::new(
216 std::io::ErrorKind::InvalidData,
217 format!("invalid RBSP byte {:#x} in state {:?}", o, &self.state),
218 ))
219 }
220 }
221 self.i += 1;
222 }
223 }
224 }
225 Ok(true)
226 }
227
228 pub fn reader(&mut self) -> &mut R {
230 &mut self.inner
231 }
232}
233impl<R: BufRead> Read for ByteReader<R> {
234 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
235 let chunk = self.fill_buf()?;
236 let amt = std::cmp::min(buf.len(), chunk.len());
237 if amt == 1 {
238 buf[0] = chunk[0];
242 } else {
243 buf[..amt].copy_from_slice(&chunk[..amt]);
244 }
245 self.consume(amt);
246 Ok(amt)
247 }
248}
249impl<R: BufRead> BufRead for ByteReader<R> {
250 fn fill_buf(&mut self) -> std::io::Result<&[u8]> {
251 while self.i == 0 && self.try_fill_buf_slow()? {}
252 Ok(&self.inner.fill_buf()?[0..self.i])
253 }
254
255 fn consume(&mut self, amt: usize) {
256 self.i = self.i.checked_sub(amt).unwrap();
257 self.inner.consume(amt);
258 }
259}
260
261pub fn decode_nal<'a>(nal_unit: &'a [u8]) -> Result<Cow<'a, [u8]>, std::io::Error> {
283 let mut reader = ByteReader {
284 inner: nal_unit,
285 state: ParseState::Skip(H264_HEADER_LEN),
286 i: 0,
287 max_fill: usize::MAX, };
289 let buf = reader.fill_buf()?;
290 if buf.len() + 1 == nal_unit.len() {
291 return Ok(Cow::Borrowed(&nal_unit[1..]));
292 }
293 let mut dst = Vec::with_capacity(nal_unit.len() - 2);
295 loop {
296 let buf = reader.fill_buf()?;
297 if buf.is_empty() {
298 break;
299 }
300 dst.extend_from_slice(buf);
301 let len = buf.len();
302 reader.consume(len);
303 }
304 Ok(Cow::Owned(dst))
305}
306
307#[derive(Debug)]
308pub enum BitReaderError {
309 ReaderError(&'static str, std::io::Error),
311
312 ExpGolombTooLarge(&'static str),
314
315 RemainingData,
317}
318
319pub trait Integer: bitstream_io::Integer + std::fmt::Debug {}
320impl<I: bitstream_io::Integer + std::fmt::Debug> Integer for I {}
321
322pub trait Primitive: bitstream_io::Primitive + std::fmt::Debug {}
323impl<P: bitstream_io::Primitive + std::fmt::Debug> Primitive for P {}
324
325pub trait BitWrite {
329 fn write_ue(&mut self, value: u32) -> std::io::Result<()>;
331
332 fn write_se(&mut self, value: i32) -> std::io::Result<()>;
334
335 fn write_bit(&mut self, bit: bool) -> std::io::Result<()>;
337
338 fn write<const BITS: u32, I: Integer>(&mut self, value: I) -> std::io::Result<()>;
340
341 fn write_var<I: Integer>(&mut self, bit_count: u32, value: I) -> std::io::Result<()>;
343
344 fn write_rbsp_trailing_bits(&mut self) -> std::io::Result<()>;
346}
347
348pub trait BitRead {
350 fn read_ue(&mut self, name: &'static str) -> Result<u32, BitReaderError>;
352
353 fn read_se(&mut self, name: &'static str) -> Result<i32, BitReaderError>;
355
356 fn read_bit(&mut self, name: &'static str) -> Result<bool, BitReaderError>;
358
359 fn read<const BITS: u32, I: Integer>(
363 &mut self,
364 name: &'static str,
365 ) -> Result<I, BitReaderError>;
366
367 fn read_var<I: Integer>(
371 &mut self,
372 bit_count: u32,
373 name: &'static str,
374 ) -> Result<I, BitReaderError>;
375
376 fn read_to<V: Primitive>(&mut self, name: &'static str) -> Result<V, BitReaderError>;
379
380 fn skip(&mut self, bit_count: u32, name: &'static str) -> Result<(), BitReaderError>;
383
384 fn byte_aligned(&self) -> bool;
386
387 fn has_more_rbsp_data(&mut self, name: &'static str) -> Result<bool, BitReaderError>;
392
393 fn finish_rbsp(self) -> Result<(), BitReaderError>;
395
396 fn finish_sei_payload(self) -> Result<(), BitReaderError>;
401}
402
403pub struct BitReader<R: std::io::BufRead + Clone> {
409 reader: bitstream_io::read::BitReader<R, bitstream_io::BigEndian>,
410}
411impl<R: std::io::BufRead + Clone> BitReader<R> {
412 pub fn new(inner: R) -> Self {
413 Self {
414 reader: bitstream_io::read::BitReader::new(inner),
415 }
416 }
417
418 pub fn reader(&mut self) -> Option<&mut R> {
420 self.reader.reader()
421 }
422
423 pub fn into_reader(self) -> R {
429 self.reader.into_reader()
430 }
431}
432
433impl<R: std::io::BufRead + Clone> BitRead for BitReader<R> {
434 fn read_ue(&mut self, name: &'static str) -> Result<u32, BitReaderError> {
435 let count = self
436 .reader
437 .read_unary::<1>()
438 .map_err(|e| BitReaderError::ReaderError(name, e))?;
439 if count > 31 {
440 return Err(BitReaderError::ExpGolombTooLarge(name));
441 } else if count > 0 {
442 let val: u32 = self.read_var(count, name)?;
443 Ok((1 << count) - 1 + val)
444 } else {
445 Ok(0)
446 }
447 }
448
449 fn read_se(&mut self, name: &'static str) -> Result<i32, BitReaderError> {
450 Ok(golomb_to_signed(self.read_ue(name)?))
451 }
452
453 fn read_bit(&mut self, name: &'static str) -> Result<bool, BitReaderError> {
454 self.reader
455 .read_bit()
456 .map_err(|e| BitReaderError::ReaderError(name, e))
457 }
458
459 fn read<const BITS: u32, I: Integer>(
460 &mut self,
461 name: &'static str,
462 ) -> Result<I, BitReaderError> {
463 self.reader
464 .read::<BITS, I>()
465 .map_err(|e| BitReaderError::ReaderError(name, e))
466 }
467
468 fn read_var<I: Integer>(
469 &mut self,
470 bit_count: u32,
471 name: &'static str,
472 ) -> Result<I, BitReaderError> {
473 self.reader
474 .read_var(bit_count)
475 .map_err(|e| BitReaderError::ReaderError(name, e))
476 }
477
478 fn read_to<V: Primitive>(&mut self, name: &'static str) -> Result<V, BitReaderError> {
479 self.reader
480 .read_to()
481 .map_err(|e| BitReaderError::ReaderError(name, e))
482 }
483
484 fn skip(&mut self, bit_count: u32, name: &'static str) -> Result<(), BitReaderError> {
485 self.reader
486 .skip(bit_count)
487 .map_err(|e| BitReaderError::ReaderError(name, e))
488 }
489
490 fn byte_aligned(&self) -> bool {
491 self.reader.byte_aligned()
492 }
493
494 fn has_more_rbsp_data(&mut self, name: &'static str) -> Result<bool, BitReaderError> {
495 let mut throwaway = self.reader.clone();
496 let r = (move || {
497 throwaway.skip(1)?;
498 throwaway.read_unary::<1>()?;
499 Ok::<_, std::io::Error>(())
500 })();
501 match r {
502 Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => Ok(false),
503 Err(e) => Err(BitReaderError::ReaderError(name, e)),
504 Ok(_) => Ok(true),
505 }
506 }
507
508 fn finish_rbsp(mut self) -> Result<(), BitReaderError> {
509 if !self
511 .reader
512 .read_bit()
513 .map_err(|e| BitReaderError::ReaderError("finish", e))?
514 {
515 match self.reader.read_unary::<1>() {
517 Err(e) => return Err(BitReaderError::ReaderError("finish", e)),
518 Ok(_) => return Err(BitReaderError::RemainingData),
519 }
520 }
521 match self.reader.read_unary::<1>() {
523 Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => Ok(()),
524 Err(e) => Err(BitReaderError::ReaderError("finish", e)),
525 Ok(_) => Err(BitReaderError::RemainingData),
526 }
527 }
528
529 fn finish_sei_payload(mut self) -> Result<(), BitReaderError> {
530 match self.reader.read_bit() {
531 Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(()),
532 Err(e) => return Err(BitReaderError::ReaderError("finish", e)),
533 Ok(false) => return Err(BitReaderError::RemainingData),
534 Ok(true) => {}
535 }
536 match self.reader.read_unary::<1>() {
537 Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => Ok(()),
538 Err(e) => Err(BitReaderError::ReaderError("finish", e)),
539 Ok(_) => Err(BitReaderError::RemainingData),
540 }
541 }
542}
543fn golomb_to_signed(val: u32) -> i32 {
544 let sign = (((val & 0x1) as i32) << 1) - 1;
545 ((val >> 1) as i32 + (val & 0x1) as i32) * sign
546}
547
548pub struct BitWriter<W: std::io::Write> {
550 inner: bitstream_io::write::BitWriter<W, bitstream_io::BigEndian>,
551}
552
553impl<W: std::io::Write> BitWriter<W> {
554 pub fn new(writer: W) -> Self {
556 Self {
557 inner: bitstream_io::write::BitWriter::new(writer),
558 }
559 }
560
561 pub fn writer(&mut self) -> Option<&mut W> {
563 self.inner.writer()
564 }
565
566 pub fn into_writer(self) -> W {
572 self.inner.into_writer()
573 }
574}
575
576impl<W: std::io::Write> BitWrite for BitWriter<W> {
577 fn write_ue(&mut self, value: u32) -> std::io::Result<()> {
578 if value == 0 {
582 self.inner.write_bit(true)
583 } else {
584 let code_num = value + 1;
585 let bits = 32 - code_num.leading_zeros(); let zeros = bits - 1;
587 for _ in 0..zeros {
589 self.inner.write_bit(false)?;
590 }
591 self.inner.write_var(bits, code_num)
593 }
594 }
595
596 fn write_se(&mut self, value: i32) -> std::io::Result<()> {
597 let code_num = if value > 0 {
599 (value as u32) * 2 - 1
600 } else {
601 (-value as u32) * 2
602 };
603 self.write_ue(code_num)
604 }
605
606 fn write_bit(&mut self, bit: bool) -> std::io::Result<()> {
607 self.inner.write_bit(bit)
608 }
609
610 fn write<const BITS: u32, I: Integer>(&mut self, value: I) -> std::io::Result<()> {
611 self.inner.write::<BITS, I>(value)
612 }
613
614 fn write_var<I: Integer>(&mut self, bit_count: u32, value: I) -> std::io::Result<()> {
615 self.inner.write_var::<I>(bit_count, value)
616 }
617
618 fn write_rbsp_trailing_bits(&mut self) -> std::io::Result<()> {
619 self.inner.write_bit(true)?; self.inner.byte_align()?; Ok(())
622 }
623}
624
625pub struct ByteWriter<W: Write> {
632 inner: W,
633 zero_count: u8,
636}
637
638impl<W: Write> ByteWriter<W> {
639 pub fn new(inner: W) -> Self {
641 Self {
642 inner,
643 zero_count: 0,
644 }
645 }
646
647 pub fn into_writer(self) -> W {
649 self.inner
650 }
651}
652
653impl<W: Write> Write for ByteWriter<W> {
654 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
655 let mut i = 0;
656 let mut chunk_start = 0;
657 while i < buf.len() {
658 if self.zero_count == 2 {
661 let b = buf[i];
662 if b <= 3 {
663 self.inner.write_all(&buf[chunk_start..i])?;
664 chunk_start = i;
665 self.inner.write_all(&[0x03])?;
666 }
667 self.zero_count = 0;
668 }
670 match memchr::memchr(0x00, &buf[i..]) {
672 None => {
673 self.zero_count = 0;
674 break;
675 }
676 Some(rel) => {
677 self.zero_count = if rel > 0 { 1 } else { self.zero_count + 1 };
680 i += rel + 1;
681 }
682 }
683 }
684 self.inner.write_all(&buf[chunk_start..])?;
685 Ok(buf.len())
686 }
687
688 fn flush(&mut self) -> std::io::Result<()> {
689 self.inner.flush()
690 }
691}
692
693#[cfg(test)]
694mod tests {
695 use super::*;
696 use hex_literal::*;
697 use hex_slice::AsHex;
698
699 #[test]
700 fn byte_reader() {
701 let data = hex!(
702 "67 64 00 0A AC 72 84 44 26 84 00 00 03
703 00 04 00 00 03 00 CA 3C 48 96 11 80"
704 );
705 for i in 1..data.len() - 1 {
706 let (head, tail) = data.split_at(i);
707 let r = head.chain(tail);
708 let mut r = ByteReader::skipping_h264_header(r);
709 let mut rbsp = Vec::new();
710 r.read_to_end(&mut rbsp).unwrap();
711 let expected = hex!(
712 "64 00 0A AC 72 84 44 26 84 00 00
713 00 04 00 00 00 CA 3C 48 96 11 80"
714 );
715 assert!(
716 rbsp == &expected[..],
717 "Mismatch with on split_at({}):\nrbsp {:02x}\nexpected {:02x}",
718 i,
719 rbsp.as_hex(),
720 expected.as_hex()
721 );
722 }
723 }
724
725 #[test]
726 fn bitreader_has_more_data() {
727 let mut reader = BitReader::new(&[0x12, 0x80][..]);
729 assert!(reader.has_more_rbsp_data("call 1").unwrap());
730 assert_eq!(reader.read::<8, u8>("u8 1").unwrap(), 0x12);
731 assert!(!reader.has_more_rbsp_data("call 2").unwrap());
732
733 let mut reader = BitReader::new(&[0x18][..]);
735 assert!(reader.has_more_rbsp_data("call 3").unwrap());
736 assert_eq!(reader.read::<4, u8>("u8 2").unwrap(), 0x1);
737 assert!(!reader.has_more_rbsp_data("call 4").unwrap());
738
739 let mut reader = BitReader::new(&[0x80, 0x00, 0x00][..]);
741 assert!(!reader
742 .has_more_rbsp_data("at end with cabac-zero-words")
743 .unwrap());
744 }
745
746 #[test]
747 fn byte_reader_emulation_prevention_beyond_max_fill() {
748 let mut input = vec![0xFF; 129];
754 input.extend_from_slice(&[0x00, 0x00, 0x03, 0x01]);
755 let mut r = ByteReader::without_skip(&input[..]);
756 let mut rbsp = Vec::new();
757 r.read_to_end(&mut rbsp).unwrap();
758 let mut expected = vec![0xFF; 129];
759 expected.extend_from_slice(&[0x00, 0x00, 0x01]);
760 assert_eq!(rbsp, expected, "emulation prevention byte was not stripped");
761 }
762
763 #[test]
764 fn read_ue_overflow() {
765 let mut reader = BitReader::new(&[0, 0, 0, 0, 255, 255, 255, 255, 255][..]);
766 assert!(matches!(
767 reader.read_ue("test"),
768 Err(BitReaderError::ExpGolombTooLarge("test"))
769 ));
770 }
771
772 fn byte_writer_encode(rbsp: &[u8]) -> Vec<u8> {
774 let mut out = Vec::new();
775 ByteWriter::new(&mut out).write_all(rbsp).unwrap();
776 out
777 }
778
779 #[test]
780 fn byte_writer_no_escaping_needed() {
781 assert_eq!(byte_writer_encode(b"hello"), b"hello");
783 assert_eq!(
784 byte_writer_encode(&[0xFF, 0xFE, 0x01, 0x02]),
785 &[0xFF, 0xFE, 0x01, 0x02]
786 );
787 assert_eq!(byte_writer_encode(&[0x00, 0x04]), &[0x00, 0x04]);
789 assert_eq!(byte_writer_encode(&[0x00, 0x00, 0x04]), &[0x00, 0x00, 0x04]);
791 }
792
793 #[test]
794 fn byte_writer_escaping() {
795 assert_eq!(
797 byte_writer_encode(&[0x00, 0x00, 0x00]),
798 &[0x00, 0x00, 0x03, 0x00]
799 );
800 assert_eq!(
801 byte_writer_encode(&[0x00, 0x00, 0x01]),
802 &[0x00, 0x00, 0x03, 0x01]
803 );
804 assert_eq!(
805 byte_writer_encode(&[0x00, 0x00, 0x02]),
806 &[0x00, 0x00, 0x03, 0x02]
807 );
808 assert_eq!(
809 byte_writer_encode(&[0x00, 0x00, 0x03]),
810 &[0x00, 0x00, 0x03, 0x03]
811 );
812 }
813
814 #[test]
815 fn byte_writer_multiple_escapes() {
816 assert_eq!(
820 byte_writer_encode(&[0x00, 0x00, 0x00, 0x00, 0x00, 0x01]),
821 &[0x00, 0x00, 0x03, 0x00, 0x00, 0x03, 0x00, 0x01],
822 );
823 }
824
825 #[test]
826 fn byte_writer_split_writes() {
827 let mut out = Vec::new();
829 let mut w = ByteWriter::new(&mut out);
830 w.write_all(&[0x00, 0x00]).unwrap();
831 w.write_all(&[0x03]).unwrap(); drop(w);
833 assert_eq!(out, &[0x00, 0x00, 0x03, 0x03]);
834
835 let mut out2 = Vec::new();
836 let mut w2 = ByteWriter::new(&mut out2);
837 w2.write_all(&[0x00]).unwrap();
838 w2.write_all(&[0x00]).unwrap();
839 w2.write_all(&[0x01]).unwrap(); drop(w2);
841 assert_eq!(out2, &[0x00, 0x00, 0x03, 0x01]);
842 }
843
844 fn make_nal(hdr: u8, rbsp: &[u8]) -> Vec<u8> {
846 let mut out = Vec::with_capacity(1 + rbsp.len() + rbsp.len() / 3);
848 out.push(hdr);
849 ByteWriter::new(&mut out).write_all(rbsp).unwrap();
850 out
851 }
852
853 #[test]
854 fn byte_writer_roundtrip() {
855 let rbsp = hex!(
857 "64 00 0A AC 72 84 44 26 84 00 00
858 00 04 00 00 00 CA 3C 48 96 11 80"
859 );
860 let nal = make_nal(0x67, &rbsp);
861 let decoded = decode_nal(&nal).unwrap();
862 assert_eq!(&*decoded, &rbsp[..]);
863 }
864
865 #[test]
866 fn byte_reader_rejects_forbidden_sequences() {
867 for forbidden in [0x00u8, 0x01, 0x02] {
871 let nal = [0x67, 0x12, 0x00, 0x00, forbidden, 0x34];
872 let mut r = ByteReader::skipping_h264_header(&nal[..]);
873 let mut buf = Vec::new();
874 let err = r.read_to_end(&mut buf).unwrap_err();
875 assert_eq!(
876 err.kind(),
877 std::io::ErrorKind::InvalidData,
878 "expected InvalidData for 0x00 0x00 {:#04x}, got {:?}",
879 forbidden,
880 err.kind(),
881 );
882 }
883 }
884
885 #[test]
886 fn byte_writer_escape_inserted_in_nal() {
887 assert_eq!(
889 make_nal(0x68, &hex!("12 34 00 00 00 86")),
890 hex!("68 12 34 00 00 03 00 86"),
891 );
892 }
893}