use std::io::{Read, Seek, SeekFrom};
use std::ops::Range;
use super::header::{OpusHead, OpusTags};
use super::page::{
CAPTURE_PATTERN, HEADER_LEN, MAX_PAGE_PAYLOAD, MAX_SEGMENTS, PageHeader, verify_crc,
};
use crate::{Error, Result};
pub(crate) const MAX_OGG_PACKET_BYTES: usize = 16 * 1024 * 1024;
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct OggPacket {
pub data: Vec<u8>,
pub page_granule: i64,
pub end_of_stream: bool,
}
impl Default for OggPacket {
fn default() -> Self {
OggPacket::new(Vec::new(), -1, false)
}
}
impl OggPacket {
pub fn new(data: Vec<u8>, page_granule: i64, end_of_stream: bool) -> Self {
OggPacket {
data,
page_granule,
end_of_stream,
}
}
}
pub struct OggOpusReader<R: Read> {
source: Counted<R>,
head: OpusHead,
tags: OpusTags,
serial: u32,
packets: Vec<u8>,
ends: Vec<usize>,
taken: usize,
page_granule: i64,
last_is_eos: bool,
saw_eos: bool,
exhausted: bool,
audio_start: u64,
audio_start_eos: bool,
}
impl<R: Read> std::fmt::Debug for OggOpusReader<R> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OggOpusReader")
.field("head", &self.head)
.field("serial", &format_args!("{:#010x}", self.serial))
.field("packets_ready", &(self.ends.len() - self.taken))
.field("saw_eos", &self.saw_eos)
.field("exhausted", &self.exhausted)
.finish_non_exhaustive()
}
}
impl<R: Read> OggOpusReader<R> {
pub fn new(source: R) -> Result<Self> {
let mut r = OggOpusReader {
source: Counted {
inner: source,
position: 0,
},
head: OpusHead::new(1, 0)?,
tags: OpusTags::new(),
serial: 0,
packets: Vec::new(),
ends: Vec::with_capacity(MAX_SEGMENTS),
taken: 0,
page_granule: -1,
last_is_eos: false,
saw_eos: false,
exhausted: false,
audio_start: 0,
audio_start_eos: false,
};
let first = r
.next_packet()?
.ok_or(Error::InvalidStream("stream ends before OpusHead"))?;
r.head = OpusHead::parse(&r.packets[first])?;
let second = r
.next_packet()?
.ok_or(Error::InvalidStream("stream ends before OpusTags"))?;
r.tags = OpusTags::parse(&r.packets[second])?;
if r.taken < r.ends.len() || r.packets.len() > r.partial_start() {
return Err(Error::InvalidStream(
"comment header does not finish its page",
));
}
r.audio_start = r.source.position;
r.audio_start_eos = r.saw_eos;
Ok(r)
}
pub fn head(&self) -> &OpusHead {
&self.head
}
pub fn tags(&self) -> &OpusTags {
&self.tags
}
pub fn serial(&self) -> u32 {
self.serial
}
pub fn read_packet(&mut self) -> Result<Option<OggPacket>> {
let Some(range) = self.next_packet()? else {
return Ok(None);
};
Ok(Some(OggPacket {
data: self.packets[range].to_vec(),
page_granule: self.page_granule,
end_of_stream: self.completed_eos(),
}))
}
pub fn read_packet_into(&mut self, packet: &mut OggPacket) -> Result<bool> {
let Some(range) = self.next_packet()? else {
return Ok(false);
};
packet.data.clear();
packet.data.extend_from_slice(&self.packets[range]);
packet.page_granule = self.page_granule;
packet.end_of_stream = self.completed_eos();
Ok(true)
}
pub fn packets(&mut self) -> Packets<'_, R> {
Packets {
reader: self,
done: false,
}
}
pub fn into_inner(self) -> R {
self.source.inner
}
pub fn get_ref(&self) -> &R {
&self.source.inner
}
fn next_packet(&mut self) -> Result<Option<Range<usize>>> {
loop {
if self.taken < self.ends.len() {
let start = if self.taken == 0 {
0
} else {
self.ends[self.taken - 1]
};
let end = self.ends[self.taken];
self.taken += 1;
return Ok(Some(start..end));
}
if self.saw_eos || self.exhausted {
if self.packets.len() > self.partial_start() {
self.packets.clear();
self.ends.clear();
self.taken = 0;
return Err(Error::InvalidStream(
"stream ends in the middle of a packet",
));
}
return Ok(None);
}
self.read_page()?;
}
}
fn partial_start(&self) -> usize {
self.ends.last().copied().unwrap_or(0)
}
fn completed_eos(&self) -> bool {
self.last_is_eos && self.taken == self.ends.len()
}
fn read_page(&mut self) -> Result<()> {
debug_assert_eq!(self.taken, self.ends.len());
self.packets.drain(..self.partial_start());
self.ends.clear();
self.taken = 0;
self.last_is_eos = false;
let Some(raw) = self.read_page_header()? else {
self.exhausted = true;
return Ok(());
};
let header = PageHeader::parse(&raw)?;
if header.segment_count == 0 {
return Err(Error::InvalidStream("page has an empty segment table"));
}
let mut segments_arr = [0u8; MAX_SEGMENTS];
let segments = &mut segments_arr[..header.segment_count as usize];
read_exact(&mut self.source, segments)?;
let payload_len: usize = segments.iter().map(|&s| s as usize).sum();
debug_assert!(payload_len <= MAX_PAGE_PAYLOAD);
let base = self.packets.len();
self.packets.resize(base + payload_len, 0);
if let Err(e) = read_exact(&mut self.source, &mut self.packets[base..]) {
self.packets.truncate(base);
return Err(e);
}
if !verify_crc(&raw, segments, &self.packets[base..], header.crc) {
self.packets.truncate(base);
return Err(Error::InvalidStream("page CRC mismatch"));
}
if header.is_bos() {
self.serial = header.serial;
} else if header.serial != self.serial {
self.packets.truncate(base);
return Err(Error::InvalidStream(
"stream contains more than one logical bitstream",
));
}
if header.is_continued() == (base == 0) {
self.packets.clear();
return Err(Error::InvalidStream(
"page continuation flag does not match the pending packet",
));
}
self.saw_eos = header.is_eos();
self.page_granule = header.granule_position;
let mut start = 0usize;
let mut end = base;
for (i, &lace) in segments.iter().enumerate() {
end += lace as usize;
if end - start > MAX_OGG_PACKET_BYTES {
self.packets.truncate(start);
return Err(Error::InvalidStream("packet exceeds maximum allowed size"));
}
if lace < 255 {
if end > start {
self.ends.push(end);
self.last_is_eos = self.saw_eos && i + 1 == segments.len();
}
start = end;
}
}
Ok(())
}
fn read_page_header(&mut self) -> Result<Option<[u8; HEADER_LEN]>> {
let mut buf = [0u8; HEADER_LEN];
match read_exact_or_eof(&mut self.source, &mut buf)? {
0 => return Ok(None),
n if n < HEADER_LEN => {
return Err(Error::InvalidStream("stream ends inside a page header"));
}
_ => {}
}
if &buf[0..4] == CAPTURE_PATTERN {
return Ok(Some(buf));
}
const RESYNC_LIMIT: usize = 1 << 20;
for _ in 0..RESYNC_LIMIT {
buf.copy_within(1..HEADER_LEN, 0);
let mut b = [0u8; 1];
if read_exact_or_eof(&mut self.source, &mut b)? == 0 {
return Err(Error::InvalidStream("stream ends without a valid page"));
}
buf[HEADER_LEN - 1] = b[0];
if &buf[0..4] == CAPTURE_PATTERN {
return Ok(Some(buf));
}
}
Err(Error::InvalidStream(
"no Ogg page found while resynchronising",
))
}
}
impl<R: Read + Seek> OggOpusReader<R> {
pub fn rewind(&mut self) -> Result<()> {
let back = self.source.position - self.audio_start;
let back = i64::try_from(back).map_err(|_| {
Error::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"stream too long to seek back over",
))
})?;
self.source.inner.seek(SeekFrom::Current(-back))?;
self.source.position = self.audio_start;
self.packets.clear();
self.ends.clear();
self.taken = 0;
self.last_is_eos = false;
self.saw_eos = self.audio_start_eos;
self.exhausted = false;
Ok(())
}
}
struct Counted<R> {
inner: R,
position: u64,
}
impl<R: Read> Read for Counted<R> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
let n = self.inner.read(buf)?;
self.position += n as u64;
Ok(n)
}
}
fn read_exact<R: Read>(source: &mut R, buf: &mut [u8]) -> Result<()> {
if read_exact_or_eof(source, buf)? < buf.len() {
return Err(Error::InvalidStream("stream ends inside a page"));
}
Ok(())
}
fn read_exact_or_eof<R: Read>(source: &mut R, buf: &mut [u8]) -> Result<usize> {
let mut filled = 0;
while filled < buf.len() {
match source.read(&mut buf[filled..]) {
Ok(0) => break,
Ok(n) => filled += n,
Err(e) if e.kind() == std::io::ErrorKind::Interrupted => {}
Err(e) => return Err(Error::Io(e)),
}
}
Ok(filled)
}
#[derive(Debug)]
pub struct Packets<'a, R: Read> {
reader: &'a mut OggOpusReader<R>,
done: bool,
}
impl<R: Read> Iterator for Packets<'_, R> {
type Item = Result<OggPacket>;
fn next(&mut self) -> Option<Self::Item> {
if self.done {
return None;
}
match self.reader.read_packet() {
Ok(Some(p)) => Some(Ok(p)),
Ok(None) => {
self.done = true;
None
}
Err(e) => {
self.done = true;
Some(Err(e))
}
}
}
}