use anyhow::{Context, Result};
use bytes::Bytes;
use symphonia::core::codecs::audio::well_known::CODEC_ID_OPUS;
use symphonia::core::codecs::audio::{AudioDecoder, AudioDecoderOptions};
use symphonia::core::formats::probe::Hint;
use symphonia::core::formats::{FormatOptions, FormatReader, TrackType};
use symphonia::core::io::MediaSourceStream;
use symphonia::core::meta::MetadataOptions;
use super::super::decode::BytesMediaSource;
use super::super::opus::{OPUS_DECODE_RATE, OpusStream, next_demux_packet};
use super::super::resample::{RESAMPLE_STAGING_FRAMES, ResampleTo16k, SampleRate};
use super::super::wave::{WaveRecv, WaveSource, check_budget, try_open_bytes, try_open_path};
use super::super::{MAX_SAMPLE_RATE, audio_too_long_err, decode_error, resolve_budget};
use super::{ChannelSelect, PcmWindow, PcmWindows, WindowCursor, WindowSpec};
use crate::error::GigasttError;
#[cfg(feature = "file-decode")]
enum Source {
Streaming {
format: Box<dyn FormatReader>,
decoder: Box<dyn AudioDecoder>,
track_id: u32,
sample_rate: u32,
source_frames: usize,
max_samples: usize,
limit_secs: f64,
resampler: Box<ResampleTo16k>,
interleaved: Vec<f32>,
},
Opus {
format: Box<dyn FormatReader>,
track_id: u32,
stream: Box<OpusStream>,
decoded_48k: usize,
max_samples: usize,
limit_secs: f64,
resampler: Box<ResampleTo16k>,
pending: Vec<f32>,
},
Wave(Box<WaveSource>),
}
#[cfg(feature = "file-decode")]
pub(crate) struct FileWindows {
src: Source,
eof: bool,
finished: bool,
buf: Vec<f32>,
buf_start_abs: usize,
decoded_16k_total: usize,
spec: WindowSpec,
channel: ChannelSelect,
cursor: WindowCursor,
}
#[cfg(feature = "file-decode")]
impl FileWindows {
pub(crate) fn open(path: &str, spec: WindowSpec, max_audio_secs: Option<f64>) -> Result<Self> {
if let Some(wave) = try_open_path(path, ChannelSelect::Mono, max_audio_secs)? {
return Ok(Self::from_wave(wave, spec, ChannelSelect::Mono));
}
let file = std::fs::File::open(path)
.with_context(|| format!("Failed to open audio file: {path}"))?;
let mss = MediaSourceStream::new(Box::new(file), Default::default());
let mut hint = Hint::new();
if let Some(ext) = std::path::Path::new(path)
.extension()
.and_then(|e| e.to_str())
{
hint.with_extension(ext);
}
Self::from_mss(mss, hint, spec, max_audio_secs, ChannelSelect::Mono)
}
pub(crate) fn from_bytes(
data: Bytes,
spec: WindowSpec,
max_audio_secs: Option<f64>,
) -> Result<Self> {
if let Some(wave) = try_open_bytes(data.clone(), ChannelSelect::Mono, max_audio_secs)? {
return Ok(Self::from_wave(wave, spec, ChannelSelect::Mono));
}
let source = BytesMediaSource::new(data);
let mss = MediaSourceStream::new(Box::new(source), Default::default());
Self::from_mss(mss, Hint::new(), spec, max_audio_secs, ChannelSelect::Mono)
}
fn from_mss(
mss: MediaSourceStream<'static>,
hint: Hint,
spec: WindowSpec,
max_audio_secs: Option<f64>,
channel: ChannelSelect,
) -> Result<Self> {
let format = symphonia::default::get_probe()
.probe(
&hint,
mss,
FormatOptions::default(),
MetadataOptions::default(),
)
.context("Unsupported audio format")?;
let (track_id, sample_rate, channels, decoder_opt) = {
let track = format
.default_track(TrackType::Audio)
.context("No audio track found")?;
let track_id = track.id;
let audio_params = track
.codec_params
.as_ref()
.and_then(|p| p.audio())
.context("No audio codec parameters")?;
let sample_rate = audio_params.sample_rate.context("Unknown sample rate")?;
if sample_rate == 0 || sample_rate > MAX_SAMPLE_RATE {
anyhow::bail!("Unsupported sample rate: {sample_rate}Hz");
}
let channels = audio_params
.channels
.as_ref()
.map(|c| c.count())
.unwrap_or(1);
let decoder_opt = if audio_params.codec == CODEC_ID_OPUS {
None
} else {
Some(
symphonia::default::get_codecs()
.make_audio_decoder(audio_params, &AudioDecoderOptions::default())
.context("Unsupported audio codec")?,
)
};
(track_id, sample_rate, channels, decoder_opt)
};
let (max_samples, limit_secs) = resolve_budget(max_audio_secs, sample_rate);
match decoder_opt {
None => {
let (max_samples, limit_secs) = resolve_budget(max_audio_secs, OPUS_DECODE_RATE);
tracing::info!("Audio (opus): {sample_rate}Hz, {channels}ch (streaming windows)");
Ok(Self {
src: Source::Opus {
format,
track_id,
stream: Box::new(OpusStream::new(channels, channel)?),
decoded_48k: 0,
max_samples,
limit_secs,
resampler: Box::new(ResampleTo16k::new(SampleRate(sample_rate), None)),
pending: Vec::with_capacity(RESAMPLE_STAGING_FRAMES),
},
eof: false,
finished: false,
buf: Vec::new(),
buf_start_abs: 0,
decoded_16k_total: 0,
spec,
channel,
cursor: WindowCursor::new(spec),
})
}
Some(decoder) => {
tracing::info!("Audio: {sample_rate}Hz, {channels}ch (streaming windows)");
Ok(Self {
src: Source::Streaming {
format,
decoder,
track_id,
sample_rate,
source_frames: 0,
max_samples,
limit_secs,
resampler: Box::new(ResampleTo16k::new(SampleRate(sample_rate), None)),
interleaved: Vec::new(),
},
eof: false,
finished: false,
buf: Vec::new(),
buf_start_abs: 0,
decoded_16k_total: 0,
spec,
channel,
cursor: WindowCursor::new(spec),
})
}
}
}
pub(crate) fn from_bytes_channel(
data: Bytes,
spec: WindowSpec,
max_audio_secs: Option<f64>,
channel: usize,
) -> Result<Self> {
if let Some(wave) =
try_open_bytes(data.clone(), ChannelSelect::One(channel), max_audio_secs)?
{
return Ok(Self::from_wave(wave, spec, ChannelSelect::One(channel)));
}
let source = BytesMediaSource::new(data);
let mss = MediaSourceStream::new(Box::new(source), Default::default());
Self::from_mss(
mss,
Hint::new(),
spec,
max_audio_secs,
ChannelSelect::One(channel),
)
}
fn from_wave(wave: WaveSource, spec: WindowSpec, channel: ChannelSelect) -> Self {
Self {
src: Source::Wave(Box::new(wave)),
eof: false,
finished: false,
buf: Vec::new(),
buf_start_abs: 0,
decoded_16k_total: 0,
spec,
channel,
cursor: WindowCursor::new(spec),
}
}
pub(crate) fn total_16k_samples(&self) -> usize {
self.decoded_16k_total
}
pub(crate) fn drain_to_vec(mut self) -> Result<Vec<f32>> {
self.fill_to(usize::MAX)?;
Ok(std::mem::take(&mut self.buf))
}
pub(crate) fn decode_file(path: &str, max_audio_secs: Option<f64>) -> Result<Vec<f32>> {
Self::open(path, WindowSpec::flat(), max_audio_secs)?.drain_to_vec()
}
pub(crate) fn decode_bytes(data: Bytes, max_audio_secs: Option<f64>) -> Result<Vec<f32>> {
Self::from_bytes(data, WindowSpec::flat(), max_audio_secs)?.drain_to_vec()
}
fn fill_to(&mut self, target: usize) -> Result<()> {
match self.src {
Source::Streaming { .. } => self.fill_streaming(target),
Source::Opus { .. } => self.fill_opus(target),
Source::Wave(_) => self.fill_wave(target),
}
}
fn fill_streaming(&mut self, target: usize) -> Result<()> {
let channel = self.channel;
let Source::Streaming {
format,
decoder,
track_id,
sample_rate,
source_frames,
max_samples,
limit_secs,
resampler,
interleaved,
} = &mut self.src
else {
return Ok(());
};
while !self.eof && self.decoded_16k_total < target {
let have_pcm = *source_frames > 0;
let packet = match next_demux_packet(&mut **format, have_pcm)? {
Some(p) => p,
None => {
self.eof = true;
break;
}
};
if packet.track_id != *track_id {
continue;
}
let decoded = decoder.decode(&packet).context("Decode error")?;
let num_frames = decoded.frames();
let ch = decoded.spec().channels().count();
if ch > 1 {
interleaved.clear();
decoded.copy_to_vec_interleaved(interleaved);
let stage = resampler.stage();
match channel {
ChannelSelect::Mono => {
for frame in 0..num_frames {
let mut sum = 0.0_f32;
for c in 0..ch {
sum += interleaved[frame * ch + c];
}
stage.push(sum / ch as f32);
}
}
ChannelSelect::One(k) if k < ch => {
for frame in 0..num_frames {
stage.push(interleaved[frame * ch + k]);
}
}
ChannelSelect::One(_) => {}
}
} else if matches!(channel, ChannelSelect::Mono | ChannelSelect::One(0)) {
let stage = resampler.stage();
let offset = stage.len();
stage.resize(offset + num_frames, 0.0);
decoded.copy_to_slice_interleaved(&mut stage[offset..]);
}
*source_frames += num_frames;
if *source_frames > *max_samples {
return Err(audio_too_long_err(
*source_frames,
*sample_rate,
*limit_secs,
));
}
resampler.flush_full()?;
let before = self.buf.len();
resampler.drain_ready_into(&mut self.buf);
self.decoded_16k_total += self.buf.len() - before;
}
if self.eof && !self.finished {
let before = self.buf.len();
resampler.finish_into(&mut self.buf)?;
self.decoded_16k_total += self.buf.len() - before;
self.finished = true;
}
Ok(())
}
fn fill_opus(&mut self, target: usize) -> Result<()> {
let Source::Opus {
format,
track_id,
stream,
decoded_48k,
max_samples,
limit_secs,
resampler,
pending,
} = &mut self.src
else {
return Ok(());
};
while !self.eof && self.decoded_16k_total < target {
let have_pcm = *decoded_48k > 0;
let packet = match next_demux_packet(&mut **format, have_pcm)? {
Some(p) => p,
None => {
self.eof = true;
break;
}
};
if packet.track_id != *track_id {
continue;
}
*decoded_48k += stream.decode_packet(&packet.data, pending)?;
if *decoded_48k > *max_samples {
return Err(audio_too_long_err(
*decoded_48k,
OPUS_DECODE_RATE,
*limit_secs,
));
}
while pending.len() >= RESAMPLE_STAGING_FRAMES {
resampler
.stage()
.extend_from_slice(&pending[..RESAMPLE_STAGING_FRAMES]);
pending.drain(..RESAMPLE_STAGING_FRAMES);
resampler.flush_full()?;
}
let before = self.buf.len();
resampler.drain_ready_into(&mut self.buf);
self.decoded_16k_total += self.buf.len() - before;
}
if self.eof && !self.finished {
resampler.stage().extend_from_slice(pending);
pending.clear();
let before = self.buf.len();
resampler.finish_into(&mut self.buf)?;
self.decoded_16k_total += self.buf.len() - before;
self.finished = true;
}
Ok(())
}
fn fill_wave(&mut self, target: usize) -> Result<()> {
let Source::Wave(wave) = &mut self.src else {
return Ok(());
};
while !self.eof && self.decoded_16k_total < target {
match wave.recv_block()? {
WaveRecv::Block(block) => {
wave.source_frames += block.frames;
check_budget(
wave.source_frames,
wave.sample_rate,
wave.max_samples,
wave.limit_secs,
)?;
if !block.samples.is_empty() {
wave.resampler.stage().extend_from_slice(&block.samples);
wave.resampler.flush_full()?;
}
let before = self.buf.len();
wave.resampler.drain_ready_into(&mut self.buf);
self.decoded_16k_total += self.buf.len() - before;
}
WaveRecv::Eof => {
self.eof = true;
break;
}
}
}
if self.eof && !self.finished {
let before = self.buf.len();
wave.resampler.finish_into(&mut self.buf)?;
self.decoded_16k_total += self.buf.len() - before;
self.finished = true;
}
Ok(())
}
}
#[cfg(feature = "file-decode")]
impl PcmWindows for FileWindows {
fn spec(&self) -> WindowSpec {
self.spec
}
fn next_window(&mut self) -> Result<Option<PcmWindow<'_>>, GigasttError> {
if self.cursor.is_done() {
return Ok(None);
}
let drop = self
.cursor
.next_start()
.saturating_sub(self.buf_start_abs)
.min(self.buf.len());
if drop > 0 {
self.buf.drain(0..drop);
self.buf_start_abs += drop;
}
self.fill_to(self.cursor.fill_target())
.map_err(decode_error)?;
let Some((start, end)) = self.cursor.take(self.decoded_16k_total, self.eof) else {
return Ok(None);
};
let s = start - self.buf_start_abs;
let e = end - self.buf_start_abs;
Ok(Some(PcmWindow {
start_sample: start,
samples: &self.buf[s..e],
}))
}
}