use bytes::Bytes;
use super::decoder::{Config, Decoder};
use crate::resample::{Resampler, remix, validate_channels};
use crate::{Error, Frame};
pub struct Consumer {
decoder: Decoder,
track: moq_mux::container::Consumer<moq_mux::catalog::hang::Container>,
resampler: Option<Resampler>,
config: Config,
resolved_sample_rate: u32,
resolved_channels: u32,
tail: Option<moq_net::Timestamp>,
epoch: Option<moq_net::Timestamp>,
frames_decoded: usize,
end: Option<moq_net::Timestamp>,
terminal_start: Option<moq_net::Timestamp>,
discontinuity: u64,
}
impl Consumer {
pub async fn new(
broadcast: &moq_net::broadcast::Consumer,
catalog: &hang::catalog::AudioConfig,
name: impl Into<String>,
config: Config,
) -> Result<Self, Error> {
let decoder = Decoder::new(catalog)?;
let sample_rate = config.sample_rate.unwrap_or_else(|| decoder.sample_rate());
let channels = config.channels.unwrap_or_else(|| decoder.channel_count());
validate_channels(channels)?;
let resampler = if sample_rate == decoder.sample_rate() {
None
} else {
let chunk_frames = (decoder.sample_rate() as usize * 20) / 1000;
Some(Resampler::new(
decoder.sample_rate(),
sample_rate,
decoder.channel_count(),
chunk_frames,
)?)
};
let name = name.into();
let track = broadcast
.track(&name)?
.subscribe(moq_net::track::Subscription::default().with_priority(hang::catalog::PRIORITY.audio))
.await?;
let container = moq_mux::catalog::hang::Container::try_from(&catalog.container)?;
let mut track = moq_mux::container::Consumer::new(track, container);
if let Some(latency) = config.latency_max {
track = track.with_latency(latency);
}
Ok(Self {
decoder,
track,
resampler,
config,
resolved_sample_rate: sample_rate,
resolved_channels: channels,
tail: None,
epoch: None,
frames_decoded: 0,
end: None,
terminal_start: None,
discontinuity: 0,
})
}
pub fn config(&self) -> &Config {
&self.config
}
pub fn sample_rate(&self) -> u32 {
self.resolved_sample_rate
}
pub fn channels(&self) -> u32 {
self.resolved_channels
}
pub async fn read(&mut self) -> Result<Option<Frame>, Error> {
loop {
let mux_frame = self.track.read().await?;
self.apply_discontinuity()?;
let Some(mux_frame) = mux_frame else {
return self.flush();
};
if let Some(end) = self.track.end()
&& self.end != Some(end)
{
self.end = Some(end);
self.frames_decoded = 0;
self.terminal_start = None;
}
let rate = self.decoder.sample_rate();
let epoch = *self.epoch.get_or_insert(mux_frame.timestamp);
let mut decoded = self.decoder.decode(&mux_frame.payload)?;
if let Some(end) = self.end {
let terminal_start = *self
.terminal_start
.get_or_insert(rewind(mux_frame.timestamp, self.decoder.delay(), rate)?.max(epoch));
let total = frames_between(terminal_start, end, rate)?;
let remaining = total.saturating_sub(self.frames_decoded);
decoded.truncate(remaining.saturating_mul(self.decoder.channel_count() as usize));
}
let frames = decoded.len() / self.decoder.channel_count().max(1) as usize;
let decoded_at = if let Some(terminal_start) = self.terminal_start {
advance(terminal_start, self.frames_decoded, rate)?
} else {
mux_frame.timestamp
};
if self.end.is_some() {
self.frames_decoded += frames;
}
if decoded.is_empty() {
continue;
}
let (pcm, timestamp) = match self.resampler.as_mut() {
Some(r) => {
let pending = r.pending_frames();
let skipped = r.skipped();
let pcm = r.process(&decoded)?;
(pcm, self.starts_at(decoded_at, pending, skipped, rate)?)
}
None => (decoded, decoded_at),
};
self.tail = Some(advance(decoded_at, frames, rate)?);
return Ok(Some(self.frame(pcm, timestamp)?));
}
}
fn apply_discontinuity(&mut self) -> Result<(), Error> {
let discontinuity = self.track.discontinuity();
if discontinuity == self.discontinuity {
return Ok(());
}
self.discontinuity = discontinuity;
self.decoder.reset()?;
if let Some(resampler) = self.resampler.as_mut() {
resampler.reset();
}
self.tail = None;
self.epoch = None;
self.frames_decoded = 0;
self.end = None;
self.terminal_start = None;
Ok(())
}
fn flush(&mut self) -> Result<Option<Frame>, Error> {
let (Some(resampler), Some(tail)) = (self.resampler.take(), self.tail) else {
return Ok(None);
};
let pending = resampler.pending_frames();
let skipped = resampler.skipped();
let pcm = resampler.flush()?;
if pcm.is_empty() {
return Ok(None);
}
let timestamp = self.starts_at(tail, pending, skipped, self.decoder.sample_rate())?;
Ok(Some(self.frame(pcm, timestamp)?))
}
fn starts_at(
&self,
timestamp: moq_net::Timestamp,
pending: usize,
skipped: usize,
rate: u32,
) -> Result<moq_net::Timestamp, Error> {
let timestamp = rewind(timestamp, pending, rate)?;
rewind(timestamp, skipped, self.resolved_sample_rate)
}
fn frame(&self, pcm: Vec<f32>, timestamp: moq_net::Timestamp) -> Result<Frame, Error> {
let pcm = if self.decoder.channel_count() == self.resolved_channels {
pcm
} else {
remix(&pcm, self.decoder.channel_count(), self.resolved_channels)?
};
let bytes = self.config.format.from_interleaved_f32(&pcm, self.resolved_channels)?;
Ok(Frame {
timestamp,
data: Bytes::from(bytes),
})
}
}
fn advance(timestamp: moq_net::Timestamp, frames: usize, sample_rate: u32) -> Result<moq_net::Timestamp, Error> {
if frames == 0 {
return Ok(timestamp);
}
let offset = moq_net::Timestamp::from_scale(frames as u64, sample_rate as u64)?.convert(timestamp.scale())?;
Ok(timestamp.checked_add(offset)?)
}
fn frames_between(start: moq_net::Timestamp, end: moq_net::Timestamp, sample_rate: u32) -> Result<usize, Error> {
let duration = end.checked_sub(start)?;
let frames = (std::time::Duration::from(duration).as_nanos() * sample_rate as u128 + 500_000_000) / 1_000_000_000;
usize::try_from(frames).map_err(|_| Error::Unsupported("audio duration does not fit in memory".into()))
}
fn rewind(timestamp: moq_net::Timestamp, frames: usize, sample_rate: u32) -> Result<moq_net::Timestamp, Error> {
if frames == 0 {
return Ok(timestamp);
}
let offset = moq_net::Timestamp::from_scale(frames as u64, sample_rate as u64)?.convert(timestamp.scale())?;
Ok(timestamp
.checked_sub(offset)
.unwrap_or(moq_net::Timestamp::new(0, timestamp.scale())?))
}
#[cfg(test)]
mod tests {
use moq_net::Timestamp;
use super::*;
use crate::Format;
use crate::encode::{Encoder, Input, Options, Producer};
#[tokio::test]
async fn remixes_mono_stream_to_stereo_output() {
let mut broadcast = moq_net::broadcast::Info::new().produce();
let catalog = moq_mux::catalog::Producer::new(&mut broadcast).unwrap();
let subscriber = broadcast.consume();
let input = Input {
format: Format::F32,
sample_rate: 48_000,
channels: 1,
};
let options = Options {
track: Some("audio".to_string()),
..Options::default()
};
let mut producer = Producer::new(&mut broadcast, catalog, input.clone(), &options).unwrap();
let catalog = Encoder::new(&crate::encode::Config::new(input)).unwrap().catalog();
let mut consumer = Consumer::new(
&subscriber,
&catalog,
"audio",
Config {
channels: Some(2),
..Config::new()
},
)
.await
.unwrap();
let samples = vec![0.1f32; 960];
let mut data = Vec::with_capacity(samples.len() * size_of::<f32>());
for sample in samples {
data.extend_from_slice(&sample.to_le_bytes());
}
producer
.write(&Frame {
timestamp: Timestamp::ZERO,
data: data.into(),
})
.unwrap();
let frame = consumer.read().await.unwrap().expect("decoded frame");
let samples = Format::F32.as_interleaved_f32(&frame.data, 2).unwrap();
assert_eq!(samples.len(), (960 - 312) * 2);
for pair in samples.chunks_exact(2) {
assert_eq!(pair[0], pair[1]);
}
}
#[tokio::test]
async fn resampled_timestamps_follow_the_samples() {
let mut broadcast = moq_net::broadcast::Info::new().produce();
let track = broadcast.create_track("audio", hang::container::track_info()).unwrap();
let subscriber = broadcast.consume();
let catalog = hang::catalog::AudioConfig::new(hang::catalog::AudioCodec::Pcm, 44_100, 1);
let mut producer = moq_mux::container::Producer::new(track, moq_mux::catalog::hang::Container::Legacy);
let mut consumer = Consumer::new(
&subscriber,
&catalog,
"audio",
Config {
sample_rate: Some(48_000),
..Config::new()
},
)
.await
.unwrap();
const FRAMES: u64 = 1024;
let payload: Bytes = vec![0u8; FRAMES as usize * size_of::<f32>()].into();
for packet in 0..2 {
producer
.write(moq_mux::container::Frame {
timestamp: moq_net::Timestamp::from_scale(packet * FRAMES, 44_100).unwrap(),
duration: None,
payload: payload.clone(),
keyframe: true,
})
.unwrap();
}
let first = consumer.read().await.unwrap().expect("decoded frame");
assert_eq!(first.timestamp.as_micros(), 0);
let second = consumer.read().await.unwrap().expect("decoded frame");
let first_frames = (first.data.len() / size_of::<f32>()) as u128;
let ends_at = first_frames * 1_000_000 / 48_000;
let gap = second.timestamp.as_micros().abs_diff(ends_at);
assert!(gap < 100, "expected the frames to meet, got a {gap} us gap");
}
#[tokio::test]
async fn resampled_tail_survives_the_end_of_the_track() {
let mut broadcast = moq_net::broadcast::Info::new().produce();
let track = broadcast.create_track("audio", hang::container::track_info()).unwrap();
let subscriber = broadcast.consume();
let catalog = hang::catalog::AudioConfig::new(hang::catalog::AudioCodec::Pcm, 44_100, 1);
let mut producer = moq_mux::container::Producer::new(track, moq_mux::catalog::hang::Container::Legacy);
let mut consumer = Consumer::new(
&subscriber,
&catalog,
"audio",
Config {
sample_rate: Some(48_000),
..Config::new()
},
)
.await
.unwrap();
const FRAMES: usize = 1024;
let payload: Bytes = vec![0u8; FRAMES * size_of::<f32>()].into();
producer
.write(moq_mux::container::Frame {
timestamp: moq_net::Timestamp::ZERO,
duration: None,
payload,
keyframe: true,
})
.unwrap();
producer.finish().unwrap();
let first = consumer.read().await.unwrap().expect("decoded frame");
let first_frames = first.data.len() / size_of::<f32>();
let tail = consumer.read().await.unwrap().expect("flushed tail");
let tail_frames = tail.data.len() / size_of::<f32>();
assert!((215..=230).contains(&tail_frames), "unexpected tail: {tail_frames}");
let ends_at = (first_frames as u128) * 1_000_000 / 48_000;
let gap = tail.timestamp.as_micros().abs_diff(ends_at);
assert!(gap < 100, "expected the tail to meet the body, got a {gap} us gap");
let total = first_frames + tail_frames;
assert!((1105..=1120).contains(&total), "unexpected total: {total}");
assert!(consumer.read().await.unwrap().is_none());
}
#[tokio::test]
async fn reads_the_container_the_catalog_declares() {
let mut broadcast = moq_net::broadcast::Info::new().produce();
let track = broadcast.create_track("audio", hang::container::track_info()).unwrap();
let subscriber = broadcast.consume();
let mut catalog = hang::catalog::AudioConfig::new(hang::catalog::AudioCodec::Pcm, 48_000, 1);
catalog.container = hang::catalog::Container::Loc;
let mut producer = moq_mux::container::Producer::new(track, moq_mux::catalog::hang::Container::Loc);
let mut consumer = Consumer::new(
&subscriber,
&catalog,
"audio",
Config {
format: Format::F32,
..Config::new()
},
)
.await
.unwrap();
let samples = [0.25f32, -0.5, 0.75, -1.0];
let payload: Vec<u8> = samples.iter().flat_map(|sample| sample.to_le_bytes()).collect();
producer
.write(moq_mux::container::Frame {
timestamp: Timestamp::ZERO,
duration: None,
payload: payload.into(),
keyframe: true,
})
.unwrap();
let frame = consumer.read().await.unwrap().expect("decoded frame");
assert_eq!(
Format::F32.as_interleaved_f32(&frame.data, 1).unwrap().as_ref(),
samples
);
}
#[tokio::test]
async fn decodes_a_cmaf_framed_track() {
let input = Input {
format: Format::F32,
sample_rate: 48_000,
channels: 2,
};
let mut encoder = Encoder::new(&crate::encode::Config::new(input.clone())).unwrap();
let mut catalog = encoder.catalog();
let pcm = vec![0.0f32; encoder.frame_size() * encoder.codec_channels() as usize];
let packet = encoder.encode(&pcm).unwrap();
let muxer = moq_mux::container::fmp4::Muxer::audio(&catalog).unwrap();
let init = muxer.init().unwrap().expect("an out-of-band codec has an init segment");
catalog.container = hang::catalog::Container::Cmaf { init };
let mut broadcast = moq_net::broadcast::Info::new().produce();
let subscriber = broadcast.consume();
let track = broadcast.create_track("audio", hang::container::track_info()).unwrap();
let container = moq_mux::catalog::hang::Container::try_from(&catalog.container).unwrap();
let mut producer = moq_mux::container::Producer::new(track, container);
let mut consumer = Consumer::new(&subscriber, &catalog, "audio", Config::new())
.await
.unwrap();
producer
.write(moq_mux::container::Frame {
timestamp: Timestamp::ZERO,
payload: packet,
keyframe: true,
duration: None,
})
.unwrap();
producer.cut(None).unwrap();
let frame = consumer.read().await.unwrap().expect("decoded frame");
assert_eq!(frame.timestamp.as_micros(), 0);
let samples = Format::F32.as_interleaved_f32(&frame.data, 2).unwrap();
assert_eq!(samples.len(), (960 - 312) * 2);
}
}