#![allow(unsafe_code)]
use crate::DecodeError;
use mediaway_common::{Bytes, Packet};
use windows::Win32::Media::MediaFoundation::{
IMFMediaBuffer, IMFSample, IMFTransform, MF_E_TRANSFORM_NEED_MORE_INPUT, MFCreateMemoryBuffer,
MFCreateSample, MFT_MESSAGE_NOTIFY_BEGIN_STREAMING, MFT_MESSAGE_NOTIFY_END_STREAMING,
MFT_MESSAGE_NOTIFY_START_OF_STREAM, MFT_OUTPUT_DATA_BUFFER,
};
use super::runtime::to_hns;
pub(super) enum Drain {
Frame(OutputPayload),
NeedMore,
}
pub(super) struct OutputPayload {
pub(super) data: Bytes,
pub(super) pts_hns: i64,
}
pub(super) fn begin_streaming(transform: &IMFTransform) -> Result<(), DecodeError> {
unsafe {
transform
.ProcessMessage(MFT_MESSAGE_NOTIFY_BEGIN_STREAMING, 0)
.map_err(|_| DecodeError::Backend)?;
transform
.ProcessMessage(MFT_MESSAGE_NOTIFY_START_OF_STREAM, 0)
.map_err(|_| DecodeError::Backend)?;
}
Ok(())
}
pub(super) fn output_buffer_size(transform: &IMFTransform) -> Result<u32, DecodeError> {
let out_info = unsafe { transform.GetOutputStreamInfo(0) }.map_err(|_| DecodeError::Backend)?;
Ok(out_info.cbSize.max(1))
}
pub(super) fn process_one_output(
transform: &IMFTransform,
output_buf_size: u32,
) -> Result<Drain, DecodeError> {
let mut status = 0u32;
let out_sample: IMFSample = unsafe { MFCreateSample() }.map_err(|_| DecodeError::Backend)?;
let out_buffer =
unsafe { MFCreateMemoryBuffer(output_buf_size) }.map_err(|_| DecodeError::Backend)?;
unsafe { out_sample.AddBuffer(&out_buffer) }.map_err(|_| DecodeError::Backend)?;
let mut buffers = [MFT_OUTPUT_DATA_BUFFER {
dwStreamID: 0,
pSample: std::mem::ManuallyDrop::new(Some(out_sample)),
dwStatus: 0,
pEvents: std::mem::ManuallyDrop::new(None),
}];
let hr = unsafe { transform.ProcessOutput(0, &mut buffers, &raw mut status) };
let sample = unsafe { std::mem::ManuallyDrop::take(&mut buffers[0].pSample) };
let _ = unsafe { std::mem::ManuallyDrop::take(&mut buffers[0].pEvents) };
if let Err(e) = hr {
if e.code() == MF_E_TRANSFORM_NEED_MORE_INPUT {
return Ok(Drain::NeedMore);
}
return Err(DecodeError::Backend);
}
let Some(sample) = sample else {
return Ok(Drain::NeedMore);
};
Ok(Drain::Frame(payload_from_sample(&sample)?))
}
pub(super) fn payload_from_sample(sample: &IMFSample) -> Result<OutputPayload, DecodeError> {
let buffer = unsafe { sample.ConvertToContiguousBuffer() }.map_err(|_| DecodeError::Backend)?;
let mut ptr = std::ptr::null_mut();
let mut cur_len = 0u32;
unsafe {
buffer
.Lock(&raw mut ptr, None, Some(std::ptr::from_mut(&mut cur_len)))
.map_err(|_| DecodeError::Backend)?;
}
if ptr.is_null() {
unsafe {
let _: windows::core::Result<()> = buffer.Unlock();
}
return Err(DecodeError::Backend);
}
let mut data = vec![0u8; cur_len as usize];
unsafe {
std::ptr::copy_nonoverlapping(ptr, data.as_mut_ptr(), cur_len as usize);
buffer.Unlock().map_err(|_| DecodeError::Backend)?;
}
let pts_hns = unsafe { sample.GetSampleTime() }.unwrap_or(0);
Ok(OutputPayload {
data: Bytes::from(data),
pts_hns,
})
}
pub(super) fn packet_to_sample(
packet: &Packet,
time_base_num: u64,
time_base_den: u32,
) -> Result<IMFSample, DecodeError> {
if packet.payload.is_empty() {
return Err(DecodeError::InvalidInput);
}
let len = u32::try_from(packet.payload.len()).map_err(|_| DecodeError::InvalidInput)?;
let sample: IMFSample = unsafe { MFCreateSample() }.map_err(|_| DecodeError::Backend)?;
let buffer: IMFMediaBuffer =
unsafe { MFCreateMemoryBuffer(len) }.map_err(|_| DecodeError::Backend)?;
unsafe {
let mut ptr = std::ptr::null_mut();
let mut max_len = 0u32;
buffer
.Lock(&raw mut ptr, Some(std::ptr::from_mut(&mut max_len)), None)
.map_err(|_| DecodeError::Backend)?;
if ptr.is_null() || max_len < len {
let _: windows::core::Result<()> = buffer.Unlock();
return Err(DecodeError::Backend);
}
std::ptr::copy_nonoverlapping(packet.payload.as_ref().as_ptr(), ptr, packet.payload.len());
buffer
.SetCurrentLength(len)
.map_err(|_| DecodeError::Backend)?;
buffer.Unlock().map_err(|_| DecodeError::Backend)?;
}
unsafe { sample.AddBuffer(&buffer) }.map_err(|_| DecodeError::Backend)?;
let hns = to_hns(packet.pts, time_base_num, time_base_den);
unsafe {
sample
.SetSampleTime(hns)
.map_err(|_| DecodeError::Backend)?;
}
Ok(sample)
}
pub(super) fn notify_end_streaming(transform: &IMFTransform) {
unsafe {
let _: windows::core::Result<()> =
transform.ProcessMessage(MFT_MESSAGE_NOTIFY_END_STREAMING, 0);
}
}