use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use crate::{EncodeError, VideoEncoder, VideoEncoderConfig, VideoInputPreference};
use mediaway_common::{
Bytes, CodecKind, Packet, PixelFormat, StreamInfo, VideoFrame, VideoFrameStorage, VideoGeometry,
};
use shiguredo_amf::amf::{Plane, Surface};
use shiguredo_amf::ffi::AMF_PLANE_TYPE;
use shiguredo_amf::{
Av1EncoderConfig, CodecConfig, EncodeHandler, EncodeOptions, EncodedFrame, Encoder,
EncoderConfig, FrameFormat, H264EncoderConfig, HevcEncoderConfig, PictureType, RateControlMode,
ReconfigureParams,
};
use super::codec;
#[derive(Debug, Clone, Copy)]
struct FrameMeta {
pts: i64,
duration: u64,
}
#[derive(Clone)]
struct PacketSink {
queue: Arc<Mutex<VecDeque<Result<Packet, EncodeError>>>>,
}
impl EncodeHandler for PacketSink {
type UserData = FrameMeta;
type Error = EncodeError;
fn on_encoded(&mut self, result: Result<EncodedFrame<Self::UserData>, Self::Error>) {
let outcome = match result {
Ok(frame) => Ok(packet_from_encoded_frame(&frame)),
Err(e) => Err(e),
};
if let Ok(mut queue) = self.queue.lock() {
queue.push_back(outcome);
}
}
}
impl From<shiguredo_amf::Error> for EncodeError {
fn from(_: shiguredo_amf::Error) -> Self {
Self::Backend
}
}
fn packet_from_encoded_frame(frame: &EncodedFrame<FrameMeta>) -> Packet {
let meta = *frame.user_data();
let is_keyframe = matches!(frame.picture_type(), PictureType::Idr);
let buffer = frame.buffer();
let len = buffer.get_size();
let ptr = buffer.get_native().cast::<u8>();
let bytes: &[u8] = if len == 0 || ptr.is_null() {
&[]
} else {
unsafe { std::slice::from_raw_parts(ptr, len) }
};
Packet {
stream_id: 0,
pts: meta.pts,
dts: meta.pts,
duration: meta.duration,
is_keyframe,
is_discard: false,
payload: Bytes::copy_from_slice(bytes),
}
}
pub(crate) struct AmfSession {
encoder: Encoder<PacketSink>,
queue: Arc<Mutex<VecDeque<Result<Packet, EncodeError>>>>,
info: StreamInfo,
width: u32,
height: u32,
nv12_bytes: usize,
is_cbr: bool,
flushed: bool,
}
impl AmfSession {
pub(crate) fn open(config: &VideoEncoderConfig) -> Result<Self, EncodeError> {
validate(config)?;
match config.input {
VideoInputPreference::CpuUploadOk => Self::open_cpu(config),
_ => Err(EncodeError::Unsupported),
}
}
fn open_cpu(config: &VideoEncoderConfig) -> Result<Self, EncodeError> {
let (framerate_num, framerate_den) = codec::framerate_from_time_base(config.time_base);
let codec_config = codec_config_for(config.codec)?;
let is_cbr = config.rate_control.is_some();
let rate_control_mode = if is_cbr {
RateControlMode::Cbr
} else {
RateControlMode::Cqp
};
let mut encoder_config = EncoderConfig::new(
codec_config,
config.width,
config.height,
FrameFormat::Nv12,
framerate_num,
framerate_den,
rate_control_mode,
);
encoder_config.gop_pic_size = u16::try_from(config.gop_size).ok();
if let Some(rc) = config.rate_control {
encoder_config.target_kbps = Some(codec::bps_to_kbps(rc.target_bitrate_bps));
encoder_config.max_kbps = rc.vbv_buffer_size_bytes.map(codec::vbv_bytes_to_max_kbps);
}
let queue: Arc<Mutex<VecDeque<Result<Packet, EncodeError>>>> =
Arc::new(Mutex::new(VecDeque::new()));
let handler = PacketSink {
queue: Arc::clone(&queue),
};
let encoder = Encoder::new(encoder_config, handler).map_err(|_| EncodeError::Backend)?;
let nv12_bytes = codec::nv12_size(config.width, config.height)?;
Ok(Self {
encoder,
queue,
info: stream_info_from(config),
width: config.width,
height: config.height,
nv12_bytes,
is_cbr,
flushed: false,
})
}
}
impl VideoEncoder for AmfSession {
fn stream_info(&self) -> &StreamInfo {
&self.info
}
fn push_frame(&mut self, frame: &VideoFrame) -> Result<(), EncodeError> {
if self.flushed {
return Err(EncodeError::Closed);
}
let VideoFrameStorage::Cpu { data } = &frame.storage else {
return Err(EncodeError::Unsupported);
};
if frame.width != self.width || frame.height != self.height {
return Err(EncodeError::InvalidInput);
}
if data.len() < self.nv12_bytes {
return Err(EncodeError::InvalidInput);
}
let surface = self
.encoder
.alloc_surface()
.map_err(|_| EncodeError::Backend)?;
upload_cpu_nv12(&surface, data, self.width, self.height)?;
surface.set_pts(frame.pts);
let duration = i64::try_from(frame.duration).unwrap_or(i64::MAX);
surface.set_duration(duration);
let meta = FrameMeta {
pts: frame.pts,
duration: frame.duration,
};
let options = EncodeOptions::default();
self.encoder
.encode(surface, &options, meta)
.map_err(|_| EncodeError::Backend)
}
fn poll_packet(&mut self) -> Result<Option<Packet>, EncodeError> {
let mut queue = self.queue.lock().map_err(|_| EncodeError::Backend)?;
queue.pop_front().transpose()
}
fn flush(&mut self) -> Result<(), EncodeError> {
if self.flushed {
return Ok(());
}
self.flushed = true;
self.encoder.finish().map_err(|_| EncodeError::Backend)
}
fn set_bitrate(&mut self, bitrate_bps: u32) -> Result<(), EncodeError> {
if !self.is_cbr {
return Err(EncodeError::Unsupported);
}
let params = ReconfigureParams {
target_kbps: Some(codec::bps_to_kbps(bitrate_bps)),
..Default::default()
};
self.encoder
.reconfigure(params)
.map_err(|_| EncodeError::Backend)
}
}
fn upload_cpu_nv12(
surface: &Surface,
data: &[u8],
width: u32,
height: u32,
) -> Result<(), EncodeError> {
let w = width as usize;
let h = height as usize;
let y_plane_bytes = w.checked_mul(h).ok_or(EncodeError::InvalidInput)?;
let uv_rows = h / 2;
let y_plane = surface
.get_plane(AMF_PLANE_TYPE::AMF_PLANE_Y)
.map_err(|_| EncodeError::Backend)?;
write_plane_rows(&y_plane, data, 0, w, h)?;
let uv_plane = surface
.get_plane(AMF_PLANE_TYPE::AMF_PLANE_UV)
.map_err(|_| EncodeError::Backend)?;
write_plane_rows(&uv_plane, data, y_plane_bytes, w, uv_rows)?;
Ok(())
}
fn write_plane_rows(
plane: &Plane,
data: &[u8],
src_offset: usize,
row_bytes: usize,
rows: usize,
) -> Result<(), EncodeError> {
let hpitch = usize::try_from(plane.get_hpitch()).map_err(|_| EncodeError::Backend)?;
let plane_width = usize::try_from(plane.get_width()).map_err(|_| EncodeError::Backend)?;
let plane_height = usize::try_from(plane.get_height()).map_err(|_| EncodeError::Backend)?;
if plane_width < row_bytes || plane_height < rows || hpitch < row_bytes {
return Err(EncodeError::Backend);
}
let row_span = rows
.checked_mul(row_bytes)
.ok_or(EncodeError::InvalidInput)?;
let src_end = src_offset
.checked_add(row_span)
.ok_or(EncodeError::InvalidInput)?;
if data.len() < src_end {
return Err(EncodeError::InvalidInput);
}
let base = plane.get_native().cast::<u8>();
if base.is_null() {
return Err(EncodeError::Backend);
}
for row in 0..rows {
let src = src_offset + row * row_bytes;
let dst_row_offset = row * hpitch;
unsafe {
std::ptr::copy_nonoverlapping(
data.as_ptr().add(src),
base.add(dst_row_offset),
row_bytes,
);
}
}
Ok(())
}
const fn codec_config_for(codec: CodecKind) -> Result<CodecConfig, EncodeError> {
match codec {
CodecKind::H264 => Ok(CodecConfig::H264(H264EncoderConfig { profile: None })),
CodecKind::Hevc => Ok(CodecConfig::Hevc(HevcEncoderConfig { profile: None })),
CodecKind::Av1 => Ok(CodecConfig::Av1(Av1EncoderConfig { profile: None })),
_ => Err(EncodeError::Unsupported),
}
}
fn validate(config: &VideoEncoderConfig) -> Result<(), EncodeError> {
if !codec::is_supported_video_codec(config.codec) {
return Err(EncodeError::Unsupported);
}
if config.width == 0 || config.height == 0 {
return Err(EncodeError::InvalidInput);
}
if config.pixel_format != PixelFormat::Nv12 {
return Err(EncodeError::Unsupported);
}
if config.time_base.den == 0 {
return Err(EncodeError::InvalidInput);
}
if config.gop_size == 0 {
return Err(EncodeError::InvalidInput);
}
Ok(())
}
#[allow(clippy::missing_const_for_fn, reason = "StreamInfo holds Bytes")]
fn stream_info_from(config: &VideoEncoderConfig) -> StreamInfo {
StreamInfo::Video {
id: 0,
codec: config.codec,
time_base: config.time_base,
geometry: VideoGeometry {
width: config.width,
height: config.height,
},
extra_data: Bytes::new(),
}
}
#[cfg(test)]
#[path = "session_tests.rs"]
mod tests;