Skip to main content

oggopus_embedded/
container.rs

1/*
2 * Copyright (c) 2025 Tomi Leppänen
3 * SPDX-License-Identifier: BSD-3-Clause
4 */
5//! Ogg parsing code.
6
7use super::ErrorValues;
8use bitflags::bitflags;
9use core::num::NonZeroUsize;
10use nom::{
11    bytes::complete::{tag, take},
12    error::ErrorKind,
13    number, Parser,
14};
15
16bitflags! {
17    #[derive(Debug, PartialEq)]
18    struct HeaderFlags: u8 {
19        const Continuation = 0b001;
20        const BeginOfStream = 0b010;
21        const EndOfStream = 0b100;
22    }
23}
24
25/// Error from parsing ogg container.
26#[derive(Debug, PartialEq)]
27pub enum OggError {
28    /// Unsupported ogg version.
29    UnsupportedVersion(u8),
30    /// Parsing error from nom library.
31    ParsingError(ErrorKind),
32    /// Stream ended abruptly.
33    EndOfStreamError(Option<NonZeroUsize>),
34    /// Stream did not validate as ogg stream.
35    InvalidStream(ErrorValues),
36    /// Stream is not supported, e.g. it contains a grouped stream.
37    UnsupportedStream(&'static str),
38    /// Stream is not ogg stream.
39    NotOggStream,
40    /**
41     * Buffer was too small to contain packet.
42     *
43     * Contains size of the buffer and how many bytes would have been actually needed.
44     */
45    BufferTooSmallError(usize, usize),
46}
47
48impl core::fmt::Display for OggError {
49    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
50        use OggError::*;
51        match self {
52            UnsupportedVersion(version) => {
53                f.write_fmt(format_args!("unsupported ogg version: {}", version))?
54            }
55            ParsingError(kind) => f.write_fmt(format_args!(
56                "parsing error with ogg: {}",
57                kind.description()
58            ))?,
59            EndOfStreamError(Some(size)) => f.write_fmt(format_args!(
60                "ogg stream ended abruptly with {} more bytes needed",
61                size
62            ))?,
63            EndOfStreamError(None) => f.write_fmt(format_args!("ogg stream ended abruptly"))?,
64            InvalidStream(error) => {
65                f.write_str("invalid stream: ")?;
66                error.fmt(f)?;
67            }
68            UnsupportedStream(error) => {
69                f.write_fmt(format_args!("unsupported stream: {}", error))?
70            }
71            NotOggStream => f.write_str("this is not an ogg stream")?,
72            BufferTooSmallError(got, needed) => f.write_fmt(format_args!(
73                "buffer is too small: got {} but needed {}",
74                got, needed
75            ))?,
76        };
77        Ok(())
78    }
79}
80
81impl core::error::Error for OggError {
82    fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
83        None
84    }
85}
86
87impl<'data> From<nom::Err<(&'data [u8], ErrorKind)>> for OggError {
88    fn from(error: nom::Err<(&'data [u8], ErrorKind)>) -> OggError {
89        use OggError::*;
90        fn convert(kind: ErrorKind) -> OggError {
91            if kind == ErrorKind::Eof {
92                EndOfStreamError(None)
93            } else {
94                ParsingError(kind)
95            }
96        }
97        match error {
98            nom::Err::Failure((_, kind)) => convert(kind),
99            nom::Err::Error((_, kind)) => convert(kind),
100            nom::Err::Incomplete(nom::Needed::Size(size)) => EndOfStreamError(Some(size)),
101            nom::Err::Incomplete(nom::Needed::Unknown) => EndOfStreamError(None),
102        }
103    }
104}
105
106pub(crate) type Result<'data, O> = core::result::Result<(&'data [u8], O), OggError>;
107
108#[derive(Debug, PartialEq)]
109struct Segment {
110    before: usize,
111    size: usize,
112    complete: bool,
113}
114
115#[derive(Debug, PartialEq)]
116struct SegmentTableIterator<'data> {
117    table: &'data [u8],
118    cumulated: usize,
119}
120
121impl Iterator for SegmentTableIterator<'_> {
122    type Item = Segment;
123
124    fn next(&mut self) -> Option<Self::Item> {
125        if self.table.is_empty() {
126            None
127        } else {
128            let mut index = 0;
129            let mut size = 0;
130            while index < self.table.len() && self.table[index] == 255 {
131                size += usize::from(self.table[index]);
132                index += 1;
133            }
134            let complete;
135            if index < self.table.len() {
136                assert!(self.table[index] != 255);
137                size += usize::from(self.table[index]);
138                self.table = &self.table[index + 1..];
139                complete = true;
140            } else {
141                self.table = &self.table[0..0];
142                complete = false;
143            }
144            let before = self.cumulated;
145            self.cumulated += size;
146            Some(Segment {
147                before,
148                size,
149                complete,
150            })
151        }
152    }
153}
154
155#[derive(Debug, PartialEq)]
156struct PageHeader<'data> {
157    version: u8,
158    header_type: HeaderFlags,
159    _granule_position: u64,
160    bitstream_serial_number: u32,
161    page_sequence_number: u32,
162    segment_table: &'data [u8],
163}
164
165impl PageHeader<'_> {
166    fn parse(input: &[u8]) -> Result<PageHeader<'_>> {
167        use OggError::*;
168        let (input, _) = tag(b"OggS".as_slice())(input)
169            .map_err(|_: nom::Err<(&[u8], ErrorKind)>| NotOggStream)?;
170        let (input, version) = number::u8().parse(input)?;
171        let (input, header_type) = number::u8()
172            .parse(input)
173            .map(|(input, flags)| (input, HeaderFlags::from_bits_retain(flags)))?;
174        let (input, granule_position) = number::le_u64().parse(input)?;
175        let (input, bitstream_serial_number) = number::le_u32().parse(input)?;
176        let (input, page_sequence_number) = number::le_u32().parse(input)?;
177        let (input, _crc_checksum) = number::le_u32().parse(input)?;
178        let (input, count) = number::u8().parse(input)?;
179        let (input, segment_table) = take(count)(input)?;
180        Ok((
181            input,
182            PageHeader {
183                version,
184                header_type,
185                _granule_position: granule_position,
186                bitstream_serial_number,
187                page_sequence_number,
188                segment_table,
189            },
190        ))
191    }
192
193    fn iter_segment_table(data: &[u8]) -> SegmentTableIterator<'_> {
194        let (_, header) = PageHeader::parse(data).unwrap();
195        SegmentTableIterator {
196            table: header.segment_table,
197            cumulated: 0,
198        }
199    }
200}
201
202#[derive(Debug, PartialEq)]
203pub(crate) struct Page<'data> {
204    header: PageHeader<'data>,
205    pub data: &'data [u8],
206}
207
208impl Page<'_> {
209    fn parse(input: &[u8]) -> Result<'_, Page<'_>> {
210        use OggError::*;
211        let (data, header) = PageHeader::parse(input)?;
212        if header.version != 0 {
213            return Err(UnsupportedVersion(header.version));
214        }
215        let size: usize = header.segment_table.iter().map(|x| usize::from(*x)).sum();
216        let (remaining, data) = take(size)(data)?;
217        Ok((remaining, Page { header, data }))
218    }
219
220    fn last_packet_continues(&self) -> bool {
221        *self.header.segment_table.last().unwrap() == 255
222    }
223
224    fn max_segment_size(&self, old_max: usize, accumulated: usize) -> (usize, usize) {
225        let (max, last_max) = self.header.segment_table.iter().fold(
226            (old_max, accumulated),
227            |(all_max, mut current_max), current| {
228                current_max += usize::from(*current);
229                if *current < 255 {
230                    (all_max.max(current_max), 0)
231                } else {
232                    (all_max, current_max)
233                }
234            },
235        );
236        if self.last_packet_continues() {
237            (max.max(last_max), last_max)
238        } else {
239            (max.max(last_max), 0)
240        }
241    }
242
243    /// Bitstream serial number for the page.
244    pub fn bitstream_serial_number(&self) -> u32 {
245        self.header.bitstream_serial_number
246    }
247
248    /// Page sequence number for the page.
249    pub fn page_sequence_number(&self) -> u32 {
250        self.header.page_sequence_number
251    }
252
253    /**
254     * Parse pages from data until end of page at packet boundary.
255     *
256     * Useful for skipping comment headers. Returns the last page which is useful for validating
257     * the stream.
258     */
259    pub(crate) fn skip(data: &[u8]) -> Result<'_, Page> {
260        use OggError::*;
261        let (mut remaining, mut page) = Self::parse(data)?;
262        let mut page_sequence_number = page.page_sequence_number();
263        let bitstream_serial_number = page.bitstream_serial_number();
264        while page.last_packet_continues() {
265            (remaining, page) = Self::parse(remaining)?;
266            if page.page_sequence_number() != page_sequence_number + 1 {
267                return Err(InvalidStream(ErrorValues::SequenceNumberMismatch(
268                    page_sequence_number,
269                    page.page_sequence_number(),
270                )));
271            }
272            page_sequence_number = page.page_sequence_number();
273            if page.bitstream_serial_number() != bitstream_serial_number {
274                return Err(UnsupportedStream(
275                    "bitstream serial number changed unexpectedly",
276                ));
277            }
278        }
279        Ok((remaining, page))
280    }
281}
282
283/**
284 * Iterator for ogg packets.
285 *
286 * Note that this does not implement [`Iterator`] trait because it is not possible to borrow from
287 * iterator in [`Item`][`Iterator::Item`].
288 */
289#[derive(Debug, PartialEq)]
290pub struct Packets<'data, const BUFFER_SIZE: usize> {
291    data: &'data [u8],
292    page: Page<'data>,
293    segments: SegmentTableIterator<'data>,
294    buffer: [u8; BUFFER_SIZE],
295}
296
297/// Ogg packet.
298pub struct Packet<'buffer> {
299    /// Data in ogg packet.
300    pub data: &'buffer [u8],
301}
302
303impl<const BUFFER_SIZE: usize> Packets<'_, BUFFER_SIZE> {
304    /// Parses input data for pages until a page that ends at packet boundary.
305    pub(crate) fn parse(data: &[u8]) -> Result<'_, Packets<'_, BUFFER_SIZE>> {
306        use OggError::*;
307        let (mut remaining, mut page) = Page::parse(data)?;
308        let (mut max_segment, mut acc) = page.max_segment_size(0, 0);
309        let mut page_sequence_number = page.page_sequence_number();
310        let bitstream_serial_number = page.bitstream_serial_number();
311        while page.last_packet_continues() {
312            (remaining, page) = Page::parse(remaining)?;
313            (max_segment, acc) = page.max_segment_size(max_segment, acc);
314            if page.page_sequence_number() != page_sequence_number + 1 {
315                return Err(InvalidStream(ErrorValues::SequenceNumberMismatch(
316                    page_sequence_number,
317                    page.page_sequence_number(),
318                )));
319            }
320            page_sequence_number = page.page_sequence_number();
321            if page.bitstream_serial_number() != bitstream_serial_number {
322                return Err(UnsupportedStream(
323                    "bitstream serial number changed unexpectedly",
324                ));
325            }
326        }
327        if max_segment > BUFFER_SIZE {
328            return Err(BufferTooSmallError(BUFFER_SIZE, max_segment));
329        }
330        let (next_data, page) = Page::parse(data)?;
331        let (remaining, next_data) = take(next_data.len() - remaining.len())(next_data)?;
332        Ok((
333            remaining,
334            Packets {
335                data: next_data,
336                page,
337                segments: PageHeader::iter_segment_table(data),
338                buffer: [0; BUFFER_SIZE],
339            },
340        ))
341    }
342
343    /// Returns page sequence number for the page being read.
344    pub fn current_page_sequence_number(&self) -> u32 {
345        self.page.page_sequence_number()
346    }
347
348    /// Returns page sequence number of the last page.
349    pub fn last_page_sequence_number(&self) -> u32 {
350        if self.data.is_empty() {
351            self.current_page_sequence_number()
352        } else {
353            // These have been parsed already, we can expect them to succeed
354            let (mut remaining, mut page) = Page::parse(self.data).unwrap();
355            while page.last_packet_continues() {
356                (remaining, page) = Page::parse(remaining).unwrap();
357            }
358            page.page_sequence_number()
359        }
360    }
361
362    /// Returns bitstream serial number for the page being read.
363    pub fn bitstream_serial_number(&self) -> u32 {
364        self.page.bitstream_serial_number()
365    }
366
367    /// Returns whether the current page is the end of the stream.
368    pub fn end_of_stream(&self) -> bool {
369        self.page
370            .header
371            .header_type
372            .contains(HeaderFlags::EndOfStream)
373    }
374
375    /// Iterates to the next packet and returns it, or [`None`] if the last packet has been read.
376    #[allow(clippy::should_implement_trait)]
377    pub fn next(&mut self) -> Option<Packet<'_>> {
378        let mut buf = 0;
379        loop {
380            if let Some(Segment {
381                before,
382                size,
383                complete,
384            }) = self.segments.next()
385            {
386                self.buffer[buf..buf + size]
387                    .copy_from_slice(&self.page.data[before..before + size]);
388                buf += size;
389                if complete {
390                    return Some(Packet {
391                        data: &self.buffer[0..buf],
392                    });
393                }
394            } else if self.page.last_packet_continues() {
395                assert!(!self.data.is_empty());
396                // These have been parsed already, we can expect them to succeed
397                self.segments = PageHeader::iter_segment_table(self.data);
398                (self.data, self.page) = Page::parse(self.data).unwrap();
399                assert!(
400                    (self.page.last_packet_continues() && !self.data.is_empty())
401                        || (!self.page.last_packet_continues() && self.data.is_empty())
402                );
403            } else {
404                assert!(self.data.is_empty());
405                return None;
406            }
407        }
408    }
409}
410
411#[cfg(test)]
412mod test {
413    use super::*;
414    use core::error::Error;
415
416    #[test]
417    fn parse_empty_page() {
418        let data = include_bytes!("test/empty.ogg");
419        let (remaining, page) = Page::parse(data).unwrap();
420        assert_eq!(remaining.len(), 0);
421        assert_eq!(page.data.len(), 0);
422        assert_eq!(page.header.version, 0);
423        assert_eq!(page.header.header_type, HeaderFlags::BeginOfStream);
424        assert_eq!(page.header._granule_position, 0);
425        assert_eq!(page.header.bitstream_serial_number, 2132339074);
426        assert_eq!(page.header.page_sequence_number, 0);
427        assert_eq!(page.header.segment_table, &[0]);
428    }
429
430    #[test]
431    fn parse_single_segment() {
432        let data = include_bytes!("test/single.ogg");
433        let (remaining, page) = Page::parse(data).unwrap();
434        assert_eq!(remaining.len(), 0);
435        assert_eq!(page.data.len(), 0x13);
436        assert_eq!(page.header.version, 0);
437        assert_eq!(page.header.header_type, HeaderFlags::BeginOfStream);
438        assert_eq!(page.header._granule_position, 0);
439        assert_eq!(page.header.bitstream_serial_number, 2132339074);
440        assert_eq!(page.header.page_sequence_number, 0);
441        assert_eq!(page.header.segment_table, &[0x13]);
442        for (a, b) in (1u8..=0x19).zip(page.data) {
443            assert_eq!(a, *b);
444        }
445    }
446
447    #[test]
448    fn parse_packet() -> core::result::Result<(), String> {
449        let data = include_bytes!("test/split.ogg");
450        let (remaining, mut packets) = Packets::<512>::parse(data).unwrap();
451        assert_eq!(remaining.len(), 0);
452        assert_eq!(packets.last_page_sequence_number(), 17);
453        let packet = packets.next().unwrap();
454        assert_eq!(packet.data.len(), 300);
455        for (i, (a, b)) in (0u8..=99)
456            .chain(0u8..=99)
457            .chain(0u8..=99)
458            .zip(packet.data)
459            .enumerate()
460        {
461            if a != *b {
462                return Err(format!("{a} != {b} at {i}"));
463            }
464        }
465        for (i, (a, b)) in (0u8..=99)
466            .chain(0u8..=99)
467            .chain(0u8..=99)
468            .chain(core::iter::repeat(0))
469            .zip(packet.data.iter())
470            .enumerate()
471        {
472            if a != *b {
473                return Err(format!("{a} != {b} at {i}"));
474            }
475        }
476        assert_eq!(packets.last_page_sequence_number(), 17);
477        assert_eq!(packets.end_of_stream(), false);
478        Ok(())
479    }
480
481    #[test]
482    fn incomplete_page() {
483        let data = include_bytes!("test/single.ogg");
484        let result = Page::parse(&data[..40]);
485        assert_eq!(result, Err(OggError::EndOfStreamError(None)));
486        assert_eq!(result.unwrap_err().to_string(), "ogg stream ended abruptly");
487    }
488
489    #[test]
490    fn incomplete_packet() {
491        let data = include_bytes!("test/split.ogg");
492        let result = Packets::<512>::parse(&data[..350]);
493        assert_eq!(result, Err(OggError::EndOfStreamError(None)));
494        let error = result.unwrap_err();
495        assert!(error.source().is_none());
496        assert_eq!(error.to_string(), "ogg stream ended abruptly");
497        let result = Packets::<512>::parse(&data[..300]);
498        assert_eq!(
499            result,
500            Err(OggError::EndOfStreamError(Some(1.try_into().unwrap())))
501        );
502        let error = result.unwrap_err();
503        assert!(error.source().is_none());
504        assert_eq!(
505            error.to_string(),
506            "ogg stream ended abruptly with 1 more bytes needed"
507        );
508    }
509
510    #[test]
511    fn invalid_version() {
512        let mut data = Vec::from(include_bytes!("test/empty.ogg"));
513        data[4] = 1;
514        let result = Page::parse(&data);
515        assert_eq!(result, Err(OggError::UnsupportedVersion(1)));
516        let error = result.unwrap_err();
517        assert!(error.source().is_none());
518        assert_eq!(error.to_string(), "unsupported ogg version: 1");
519    }
520
521    #[test]
522    fn test_skip() {
523        let data = include_bytes!("test/split.ogg");
524        let (remaining, page) = Page::skip(data).unwrap();
525        assert_eq!(remaining.len(), 0);
526        assert_eq!(page.data.len(), 45);
527        assert_eq!(page.header.version, 0);
528        assert_eq!(page.header.header_type, HeaderFlags::Continuation);
529        assert_eq!(page.header._granule_position, 0);
530        assert_eq!(page.header.bitstream_serial_number, 2132339074);
531        assert_eq!(page.header.page_sequence_number, 17);
532        assert_eq!(page.header.segment_table, &[45]);
533    }
534
535    #[test]
536    fn bad_sequence() {
537        let mut data = Vec::from(include_bytes!("test/split.ogg"));
538        data[0x12d] = 9;
539        let result = Page::skip(&data);
540        assert_eq!(
541            result,
542            Err(OggError::InvalidStream(
543                ErrorValues::SequenceNumberMismatch(16, 9)
544            ))
545        );
546        let error = result.unwrap_err();
547        assert!(error.source().is_none());
548        assert_eq!(
549            error.to_string(),
550            "invalid stream: page sequence numbers are not sequential, previous: 16, current: 9"
551        );
552        let result = Packets::<512>::parse(&data);
553        assert_eq!(
554            result,
555            Err(OggError::InvalidStream(
556                ErrorValues::SequenceNumberMismatch(16, 9)
557            ))
558        );
559        let error = result.unwrap_err();
560        assert!(error.source().is_none());
561        assert_eq!(
562            error.to_string(),
563            "invalid stream: page sequence numbers are not sequential, previous: 16, current: 9"
564        );
565    }
566
567    #[test]
568    fn bitstream_changed() {
569        let mut data = Vec::from(include_bytes!("test/split.ogg"));
570        data[0x129] = 0x81;
571        let result = Page::skip(&data);
572        assert_eq!(
573            result,
574            Err(OggError::UnsupportedStream(
575                "bitstream serial number changed unexpectedly"
576            ))
577        );
578        let error = result.unwrap_err();
579        assert!(error.source().is_none());
580        assert_eq!(
581            error.to_string(),
582            "unsupported stream: bitstream serial number changed unexpectedly"
583        );
584        let result = Packets::<512>::parse(&data);
585        assert_eq!(
586            result,
587            Err(OggError::UnsupportedStream(
588                "bitstream serial number changed unexpectedly"
589            ))
590        );
591        let error = result.unwrap_err();
592        assert!(error.source().is_none());
593        assert_eq!(
594            error.to_string(),
595            "unsupported stream: bitstream serial number changed unexpectedly"
596        );
597    }
598
599    #[test]
600    fn too_small_buffer() {
601        let data = include_bytes!("test/split.ogg");
602        let result = Packets::<64>::parse(data);
603        assert_eq!(result, Err(OggError::BufferTooSmallError(64, 300)));
604        let error = result.unwrap_err();
605        assert!(error.source().is_none());
606        assert_eq!(
607            error.to_string(),
608            "buffer is too small: got 64 but needed 300"
609        );
610    }
611}