1use 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#[derive(Debug, PartialEq)]
27pub enum OggError {
28 UnsupportedVersion(u8),
30 ParsingError(ErrorKind),
32 EndOfStreamError(Option<NonZeroUsize>),
34 InvalidStream(ErrorValues),
36 UnsupportedStream(&'static str),
38 NotOggStream,
40 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 pub fn bitstream_serial_number(&self) -> u32 {
245 self.header.bitstream_serial_number
246 }
247
248 pub fn page_sequence_number(&self) -> u32 {
250 self.header.page_sequence_number
251 }
252
253 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#[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
297pub struct Packet<'buffer> {
299 pub data: &'buffer [u8],
301}
302
303impl<const BUFFER_SIZE: usize> Packets<'_, BUFFER_SIZE> {
304 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 pub fn current_page_sequence_number(&self) -> u32 {
345 self.page.page_sequence_number()
346 }
347
348 pub fn last_page_sequence_number(&self) -> u32 {
350 if self.data.is_empty() {
351 self.current_page_sequence_number()
352 } else {
353 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 pub fn bitstream_serial_number(&self) -> u32 {
364 self.page.bitstream_serial_number()
365 }
366
367 pub fn end_of_stream(&self) -> bool {
369 self.page
370 .header
371 .header_type
372 .contains(HeaderFlags::EndOfStream)
373 }
374
375 #[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 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}