#![allow(unsafe_code)]
use std::collections::VecDeque;
use crate::{AudioDecoder, DecodeError};
use mediaway_common::{AudioFrame, Bytes, CodecKind, Packet, Rational, SampleFormat, StreamInfo};
use windows::Win32::Media::MediaFoundation::{
CMSAACDecMFT, IMFTransform, MF_MT_AAC_PAYLOAD_TYPE, MF_MT_AUDIO_NUM_CHANNELS,
MF_MT_AUDIO_SAMPLES_PER_SECOND, MF_MT_MAJOR_TYPE, MF_MT_SUBTYPE, MF_MT_USER_DATA,
MFAudioFormat_AAC, MFAudioFormat_Float, MFCreateMediaType, MFMediaType_Audio,
MFT_MESSAGE_COMMAND_DRAIN, MFT_MESSAGE_NOTIFY_END_OF_STREAM,
};
use windows::Win32::System::Com::{CLSCTX_INPROC_SERVER, CoCreateInstance};
use super::audio_mft::{
Drain, OutputPayload, begin_streaming, notify_end_streaming, output_buffer_size,
packet_to_sample, process_one_output,
};
use super::runtime::from_hns;
const HEAAC_WAVEINFO_TAIL_LEN: usize = 12;
const PAYLOAD_TYPE_RAW_AAC: u16 = 0;
const PROFILE_LEVEL_UNSPECIFIED: u16 = 0xFE;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AacDecoderConfig {
pub sample_rate: u32,
pub channels: u16,
pub time_base: Rational,
pub extra_data: Bytes,
}
impl AacDecoderConfig {
#[must_use]
pub const fn new(sample_rate: u32, channels: u16, extra_data: Bytes) -> Self {
Self {
sample_rate,
channels,
time_base: Rational::new(1, sample_rate),
extra_data,
}
}
}
pub struct WmfAacDecoder {
transform: IMFTransform,
info: StreamInfo,
time_base_num: u64,
time_base_den: u32,
channels: u16,
output_buf_size: u32,
pending: VecDeque<AudioFrame>,
flushed: bool,
}
impl WmfAacDecoder {
pub fn open(config: &AacDecoderConfig) -> Result<Self, DecodeError> {
validate(config)?;
super::runtime::ensure_mf()?;
let transform: IMFTransform =
unsafe { CoCreateInstance(&CMSAACDecMFT, None, CLSCTX_INPROC_SERVER) }
.map_err(|_| DecodeError::Backend)?;
configure_types(
&transform,
config.sample_rate,
config.channels,
&config.extra_data,
)?;
begin_streaming(&transform)?;
let output_buf_size = output_buffer_size(&transform)?;
Ok(Self {
transform,
info: stream_info_from(config),
time_base_num: config.time_base.num,
time_base_den: config.time_base.den,
channels: config.channels,
output_buf_size,
pending: VecDeque::new(),
flushed: false,
})
}
pub const fn stream_info(&self) -> &StreamInfo {
&self.info
}
pub fn push_packet(&mut self, packet: &Packet) -> Result<(), DecodeError> {
if self.flushed {
return Err(DecodeError::Closed);
}
if packet.is_discard {
return Ok(());
}
let sample = packet_to_sample(packet, self.time_base_num, self.time_base_den)?;
unsafe { self.transform.ProcessInput(0, &sample, 0) }.map_err(|_| DecodeError::Backend)?;
self.drain_output()
}
pub fn poll_frame(&mut self) -> Result<Option<AudioFrame>, DecodeError> {
if self.pending.is_empty() {
self.drain_output()?;
}
Ok(self.pending.pop_front())
}
pub fn flush(&mut self) -> Result<(), DecodeError> {
if self.flushed {
return Ok(());
}
self.flushed = true;
unsafe {
self.transform
.ProcessMessage(MFT_MESSAGE_NOTIFY_END_OF_STREAM, 0)
.map_err(|_| DecodeError::Backend)?;
self.transform
.ProcessMessage(MFT_MESSAGE_COMMAND_DRAIN, 0)
.map_err(|_| DecodeError::Backend)?;
}
self.drain_output()?;
notify_end_streaming(&self.transform);
Ok(())
}
fn drain_output(&mut self) -> Result<(), DecodeError> {
while let Drain::Frame(payload) = process_one_output(&self.transform, self.output_buf_size)?
{
self.pending.push_back(self.frame_from_payload(payload));
}
Ok(())
}
fn frame_from_payload(&self, payload: OutputPayload) -> AudioFrame {
let channels = usize::from(self.channels).max(1);
let samples_per_channel = payload.data.len() / 4 / channels;
let pts = from_hns(payload.pts_hns, self.time_base_num, self.time_base_den);
AudioFrame {
pts,
duration: u64::try_from(samples_per_channel).unwrap_or(0),
sample_rate: self.info.sample_rate().unwrap_or(0),
channels: self.channels,
format: SampleFormat::F32,
data: payload.data,
}
}
}
impl AudioDecoder for WmfAacDecoder {
fn stream_info(&self) -> &StreamInfo {
self.stream_info()
}
fn push_packet(&mut self, packet: &Packet) -> Result<(), DecodeError> {
self.push_packet(packet)
}
fn poll_frame(&mut self) -> Result<Option<AudioFrame>, DecodeError> {
self.poll_frame()
}
fn flush(&mut self) -> Result<(), DecodeError> {
self.flush()
}
}
fn user_data_blob(asc: &[u8]) -> Vec<u8> {
let mut blob = Vec::with_capacity(HEAAC_WAVEINFO_TAIL_LEN + asc.len());
blob.extend_from_slice(&PAYLOAD_TYPE_RAW_AAC.to_le_bytes());
blob.extend_from_slice(&PROFILE_LEVEL_UNSPECIFIED.to_le_bytes());
blob.extend_from_slice(&0u16.to_le_bytes()); blob.extend_from_slice(&0u16.to_le_bytes()); blob.extend_from_slice(&0u32.to_le_bytes()); debug_assert_eq!(blob.len(), HEAAC_WAVEINFO_TAIL_LEN);
blob.extend_from_slice(asc);
blob
}
fn configure_types(
transform: &IMFTransform,
sample_rate: u32,
channels: u16,
asc: &[u8],
) -> Result<(), DecodeError> {
let in_type = unsafe { MFCreateMediaType() }.map_err(|_| DecodeError::Backend)?;
let user_data = user_data_blob(asc);
unsafe {
in_type
.SetGUID(&MF_MT_MAJOR_TYPE, &MFMediaType_Audio)
.map_err(|_| DecodeError::Backend)?;
in_type
.SetGUID(&MF_MT_SUBTYPE, &MFAudioFormat_AAC)
.map_err(|_| DecodeError::Backend)?;
in_type
.SetUINT32(&MF_MT_AUDIO_NUM_CHANNELS, u32::from(channels))
.map_err(|_| DecodeError::Backend)?;
in_type
.SetUINT32(&MF_MT_AUDIO_SAMPLES_PER_SECOND, sample_rate)
.map_err(|_| DecodeError::Backend)?;
in_type
.SetUINT32(&MF_MT_AAC_PAYLOAD_TYPE, u32::from(PAYLOAD_TYPE_RAW_AAC))
.map_err(|_| DecodeError::Backend)?;
in_type
.SetBlob(&MF_MT_USER_DATA, &user_data)
.map_err(|_| DecodeError::Backend)?;
transform
.SetInputType(0, &in_type, 0)
.map_err(|_| DecodeError::Backend)?;
}
select_float_output(transform)
}
fn select_float_output(transform: &IMFTransform) -> Result<(), DecodeError> {
let mut fallback = None;
for index in 0..16u32 {
let Ok(candidate) = (unsafe { transform.GetOutputAvailableType(0, index) }) else {
break;
};
let subtype = unsafe { candidate.GetGUID(&MF_MT_SUBTYPE) };
if subtype.is_ok_and(|guid| guid == MFAudioFormat_Float) {
unsafe {
transform
.SetOutputType(0, &candidate, 0)
.map_err(|_| DecodeError::Backend)?;
}
return Ok(());
}
if fallback.is_none() {
fallback = Some(candidate);
}
}
let chosen = fallback.ok_or(DecodeError::Backend)?;
unsafe {
transform
.SetOutputType(0, &chosen, 0)
.map_err(|_| DecodeError::Backend)?;
}
Ok(())
}
const fn validate(config: &AacDecoderConfig) -> Result<(), DecodeError> {
if config.sample_rate == 0 || config.channels == 0 || config.time_base.den == 0 {
return Err(DecodeError::InvalidInput);
}
if config.extra_data.is_empty() {
return Err(DecodeError::Unsupported);
}
Ok(())
}
#[allow(clippy::missing_const_for_fn, reason = "StreamInfo holds Bytes")]
fn stream_info_from(config: &AacDecoderConfig) -> StreamInfo {
StreamInfo::Audio {
id: 0,
codec: CodecKind::Aac,
time_base: config.time_base,
extra_data: config.extra_data.clone(),
sample_rate: config.sample_rate,
channels: config.channels,
}
}
#[cfg(test)]
#[path = "aac_tests.rs"]
mod tests;