#![forbid(unsafe_code)]
use std::collections::VecDeque;
use crate::error::PipelineError;
use crate::filter::{FilterError, FrameFilter};
use mediaway_common::{AudioFrame, VideoFrame, VideoFrameStorage};
use mediaway_container::mp4;
use mediaway_encoder::{AudioEncoder, VideoEncoder};
use mediaway_sw::apm::{AudioProcessor, VoiceActivityDetector};
use smallvec::SmallVec;
struct AudioTrack {
encoder: Box<dyn AudioEncoder>,
track_id: u32,
processor: Option<AudioProcessor>,
vad: Option<VoiceActivityDetector>,
vad_scores: VecDeque<f32>,
}
pub struct EncodeSession<E: VideoEncoder> {
encoder: E,
muxer: mp4::Muxer<mp4::Live>,
track_id: u32,
filters: SmallVec<[Box<dyn FrameFilter>; 4]>,
audio: Option<AudioTrack>,
}
impl<E: VideoEncoder> EncodeSession<E> {
pub fn open(encoder: E) -> Result<Self, PipelineError> {
let mut open = mp4::Muxer::new();
let track_id = open.add_track(encoder.stream_info().clone())?;
Ok(Self {
encoder,
muxer: open.begin(),
track_id,
filters: SmallVec::new(),
audio: None,
})
}
pub fn open_with_audio(
encoder: E,
audio_encoder: impl AudioEncoder + 'static,
) -> Result<Self, PipelineError> {
let mut open = mp4::Muxer::new();
let track_id = open.add_track(encoder.stream_info().clone().with_id(0))?;
let audio_track_id = open.add_track(audio_encoder.stream_info().clone().with_id(1))?;
Ok(Self {
encoder,
muxer: open.begin(),
track_id,
filters: SmallVec::new(),
audio: Some(AudioTrack {
encoder: Box::new(audio_encoder),
track_id: audio_track_id,
processor: None,
vad: None,
vad_scores: VecDeque::new(),
}),
})
}
pub fn attach_audio_processor(
&mut self,
processor: AudioProcessor,
) -> Result<&mut Self, PipelineError> {
let audio = self.audio.as_mut().ok_or(PipelineError::NoAudioTrack)?;
audio.processor = Some(processor);
Ok(self)
}
pub fn attach_vad(&mut self, vad: VoiceActivityDetector) -> Result<&mut Self, PipelineError> {
let audio = self.audio.as_mut().ok_or(PipelineError::NoAudioTrack)?;
audio.vad = Some(vad);
Ok(self)
}
pub fn push_filter<F: FrameFilter>(&mut self, filter: F) -> &mut Self {
self.filters.push(Box::new(filter));
self
}
pub fn write_frame(&mut self, frame: &VideoFrame) -> Result<(), PipelineError> {
if self.filters.is_empty() {
self.encoder.push_frame(frame)?; } else {
if matches!(frame.storage, VideoFrameStorage::Gpu(_)) {
return Err(PipelineError::Filter(FilterError::GpuFrameUnsupported));
}
let mut current = frame.clone();
for filter in &mut self.filters {
current = filter.process(current)?;
}
self.encoder.push_frame(¤t)?;
}
self.drain()
}
pub fn write_audio_frame(&mut self, frame: &AudioFrame) -> Result<(), PipelineError> {
let Self { audio, muxer, .. } = self;
let Some(audio) = audio.as_mut() else {
return Err(PipelineError::NoAudioTrack);
};
if let Some(processor) = audio.processor.as_mut() {
processor.push_capture_frame(frame)?;
while let Some(block) = processor.poll_processed_frame()? {
if let Some(vad) = audio.vad.as_mut() {
if let Ok(score) = vad.analyze(&block) {
audio.vad_scores.push_back(score);
}
}
audio.encoder.push_frame(&block)?;
}
} else {
audio.encoder.push_frame(frame)?;
}
Self::drain_audio(audio, muxer)
}
pub fn write_audio_render_frame(&mut self, frame: &AudioFrame) -> Result<(), PipelineError> {
let Some(audio) = self.audio.as_mut() else {
return Err(PipelineError::NoAudioTrack);
};
if let Some(processor) = audio.processor.as_mut() {
processor.push_render_frame(frame)?;
}
Ok(())
}
pub fn poll_vad_score(&mut self) -> Option<f32> {
self.audio.as_mut()?.vad_scores.pop_front()
}
pub fn finish(mut self) -> Result<Vec<u8>, PipelineError> {
self.encoder.flush()?;
self.drain()?;
if let Some(mut audio) = self.audio.take() {
audio.encoder.flush()?;
Self::drain_audio(&mut audio, &mut self.muxer)?;
}
self.muxer.flush();
let mut bytes = Vec::new();
self.muxer.poll_bytes(&mut bytes);
Ok(bytes)
}
fn drain(&mut self) -> Result<(), PipelineError> {
while let Some(mut pkt) = self.encoder.poll_packet()? {
pkt.stream_id = self.track_id;
self.muxer.push_packet(&pkt)?;
}
Ok(())
}
fn drain_audio(
audio: &mut AudioTrack,
muxer: &mut mp4::Muxer<mp4::Live>,
) -> Result<(), PipelineError> {
while let Some(mut pkt) = audio.encoder.poll_packet()? {
pkt.stream_id = audio.track_id;
muxer.push_packet(&pkt)?;
}
Ok(())
}
}
#[cfg(test)]
#[path = "session_tests.rs"]
mod tests;