#![allow(unsafe_code)]
use std::collections::VecDeque;
use crate::{DecodeError, VideoDecoder, VideoDecoderConfig, VideoOutputPreference};
use mediaway_common::{
CodecKind, Packet, PixelFormat, StreamInfo, VideoFrame, VideoFrameStorage, VideoGeometry,
};
use windows::Win32::Media::MediaFoundation::{
IMFSample, IMFTransform, MF_E_NO_MORE_TYPES, MF_MT_SUBTYPE, MFT_MESSAGE_COMMAND_DRAIN,
MFT_MESSAGE_NOTIFY_END_OF_STREAM, MFVideoFormat_NV12,
};
use super::codec::{is_supported_video_codec, video_subtype};
use super::cpu::{nv12_bytes_from_output_sample, open_sw_decoder};
use super::shared::{
Drain, begin_streaming, configure_decode_types, notify_end_streaming, output_buffer_size,
packet_to_sample, process_one_output, read_output_dimensions,
};
fn negotiate_nv12_output_type(transform: &IMFTransform) -> Result<(), DecodeError> {
for i in 0.. {
let media_type = match unsafe { transform.GetOutputAvailableType(0, i) } {
Ok(mt) => mt,
Err(e) if e.code() == MF_E_NO_MORE_TYPES => return Err(DecodeError::Unsupported),
Err(_) => return Err(DecodeError::Backend),
};
let subtype = unsafe { media_type.GetGUID(&MF_MT_SUBTYPE) }.unwrap_or_default();
if subtype == MFVideoFormat_NV12 {
unsafe { transform.SetOutputType(0, &media_type, 0) }
.map_err(|_| DecodeError::Backend)?;
return Ok(());
}
}
Err(DecodeError::Unsupported)
}
pub(crate) struct WmfMultiCodecCpuDecoder {
transform: IMFTransform,
info: StreamInfo,
time_base_num: u64,
time_base_den: u32,
pending: VecDeque<VideoFrame>,
flushed: bool,
output_buf_size: u32,
}
impl WmfMultiCodecCpuDecoder {
pub(crate) fn open(config: &VideoDecoderConfig) -> Result<Self, DecodeError> {
validate(config)?;
super::runtime::ensure_mf()?;
if config.output != VideoOutputPreference::CpuFramesOk {
return Err(DecodeError::Unsupported);
}
let input_subtype = video_subtype(config.codec)?;
let transform = open_sw_decoder(&input_subtype)?;
configure_decode_types(
&transform,
config.width,
config.height,
&config.extra_data,
&input_subtype,
)?;
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,
pending: VecDeque::new(),
flushed: false,
output_buf_size,
})
}
fn push_transform_packet(&mut self, packet: &Packet) -> Result<(), DecodeError> {
let sample = packet_to_sample(packet, self.time_base_num, self.time_base_den, None)?;
unsafe { self.transform.ProcessInput(0, &sample, 0) }.map_err(|_| DecodeError::Backend)?;
self.drain_output()?;
Ok(())
}
fn drain_output(&mut self) -> Result<(), DecodeError> {
loop {
match process_one_output(&self.transform, false, self.output_buf_size)? {
Drain::Sample(sample) => self.adopt_output_sample(&sample)?,
Drain::NeedMore => break,
Drain::StreamChange => self.apply_stream_change()?,
}
}
Ok(())
}
fn apply_stream_change(&mut self) -> Result<(), DecodeError> {
negotiate_nv12_output_type(&self.transform)?;
let (width, height) = read_output_dimensions(&self.transform)?;
if let StreamInfo::Video { geometry, .. } = &mut self.info {
*geometry = VideoGeometry { width, height };
}
self.output_buf_size = output_buffer_size(&self.transform)?;
Ok(())
}
fn adopt_output_sample(&mut self, sample: &IMFSample) -> Result<(), DecodeError> {
let pts_hns = unsafe { sample.GetSampleTime() }.unwrap_or(0);
let dur_hns = unsafe { sample.GetSampleDuration() }.unwrap_or(0);
let pts = super::runtime::from_hns(pts_hns, self.time_base_num, self.time_base_den);
let duration = u64::try_from(
super::runtime::from_hns(dur_hns, self.time_base_num, self.time_base_den).max(0),
)
.unwrap_or(0);
let geometry = self.info.geometry().unwrap_or(VideoGeometry {
width: 0,
height: 0,
});
let width = geometry.width;
let height = geometry.height;
let data = nv12_bytes_from_output_sample(sample, width, height)?;
self.pending.push_back(VideoFrame {
pts,
duration,
width,
height,
format: PixelFormat::Nv12,
storage: VideoFrameStorage::Cpu { data },
});
Ok(())
}
}
impl VideoDecoder for WmfMultiCodecCpuDecoder {
fn stream_info(&self) -> &StreamInfo {
&self.info
}
fn push_packet(&mut self, packet: &Packet) -> Result<(), DecodeError> {
if self.flushed {
return Err(DecodeError::Closed);
}
if packet.is_discard {
return Ok(());
}
self.push_transform_packet(packet)
}
fn poll_frame(&mut self) -> Result<Option<VideoFrame>, DecodeError> {
if self.pending.is_empty() {
self.drain_output()?;
}
Ok(self.pending.pop_front())
}
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 validate(config: &VideoDecoderConfig) -> Result<(), DecodeError> {
if !is_supported_video_codec(config.codec) || config.codec == CodecKind::H264 {
return Err(DecodeError::Unsupported);
}
if config.pixel_format != PixelFormat::Nv12 {
return Err(DecodeError::Unsupported);
}
if config.time_base.den == 0 {
return Err(DecodeError::InvalidInput);
}
Ok(())
}
#[allow(clippy::missing_const_for_fn, reason = "StreamInfo holds Bytes")]
fn stream_info_from(config: &VideoDecoderConfig) -> StreamInfo {
StreamInfo::Video {
id: 0,
codec: config.codec,
time_base: config.time_base,
geometry: VideoGeometry {
width: config.width,
height: config.height,
},
extra_data: config.extra_data.clone(), }
}
#[cfg(test)]
#[path = "video_cpu_tests.rs"]
mod tests;