use std::collections::VecDeque;
use crate::{EncodeError, VideoEncoder, VideoEncoderConfig, VideoInputPreference};
use mediaway_common::{
Bytes, CodecKind, Packet, PixelFormat, Rational, StreamInfo, VideoFrame, VideoFrameStorage,
VideoGeometry,
};
use nvenc::bitstream::BitStream;
use nvenc::encoder::{Encoder, RegisteredResource};
use nvenc::session::{InitParams, Session};
use nvenc::sys::enums::{
NVencBufferFormat, NVencParamsRcMode, NVencPicStruct, NVencPicType, NVencTuningInfo,
};
use nvenc::sys::guids::{
NV_ENC_CODEC_AV1_GUID, NV_ENC_CODEC_H264_GUID, NV_ENC_CODEC_HEVC_GUID, NV_ENC_PRESET_P3_GUID,
};
use nvenc::sys::structs::Guid;
use windows::Win32::Graphics::Direct3D11::ID3D11Device;
use super::device::{self, Dx11Upload};
const DEFAULT_BITRATE_BPS: u32 = 2_000_000;
pub(crate) struct NvencSession {
registered: RegisteredResource,
bitstream: BitStream,
upload: Dx11Upload,
encoder: Encoder,
_device: ID3D11Device,
info: StreamInfo,
width: u32,
height: u32,
codec: CodecKind,
frame_idx: u32,
pending: VecDeque<Packet>,
flushed: bool,
}
impl NvencSession {
pub(crate) fn open(config: &VideoEncoderConfig) -> Result<Self, EncodeError> {
validate(config)?;
let codec_guid = codec_guid(config.codec).ok_or(EncodeError::Unsupported)?;
let (device, ctx) = device::open_device()?;
let upload = Dx11Upload::new(&device, ctx, config.width, config.height)?;
if !nvenc_runtime_available() {
return Err(EncodeError::Backend);
}
let session = Session::open_dx(&device).map_err(|_| EncodeError::Backend)?;
let codecs = session
.get_encode_codecs()
.map_err(|_| EncodeError::Backend)?;
if !codecs.contains(&codec_guid) {
return Err(EncodeError::Unsupported);
}
let (session, mut preset_config) = session
.get_encode_preset_config_ex(
codec_guid.clone(),
NV_ENC_PRESET_P3_GUID,
NVencTuningInfo::HighQuality,
)
.map_err(|_| EncodeError::Backend)?;
let bitrate = if config.bitrate_bps == 0 {
DEFAULT_BITRATE_BPS
} else {
config.bitrate_bps
};
preset_config.preset_cfg.rc_params.rate_control_mode = NVencParamsRcMode::VBR;
preset_config.preset_cfg.rc_params.average_bit_rate = bitrate;
preset_config.preset_cfg.gop_len = 30;
preset_config.preset_cfg.frame_interval_p = 1;
let init_params = InitParams {
encode_guid: codec_guid,
preset_guid: NV_ENC_PRESET_P3_GUID,
resolution: [config.width, config.height],
aspect_ratio: [config.width, config.height],
frame_rate: frame_rate(config.time_base),
tuning_info: NVencTuningInfo::HighQuality,
buffer_format: NVencBufferFormat::NV12,
encode_config: &mut preset_config.preset_cfg,
enable_ptd: true,
max_encoder_resolution: [0, 0],
};
let encoder = session
.init_encoder(init_params)
.map_err(|_| EncodeError::Backend)?;
let registered = encoder
.register_resource_dx11(upload.gpu_texture(), NVencBufferFormat::NV12, 0)
.map_err(|_| EncodeError::Backend)?;
let bitstream = encoder
.create_bitstream_buffer()
.map_err(|_| EncodeError::Backend)?;
Ok(Self {
registered,
bitstream,
upload,
encoder,
_device: device,
info: stream_info_from(config),
width: config.width,
height: config.height,
codec: config.codec,
frame_idx: 0,
pending: VecDeque::new(),
flushed: false,
})
}
}
const fn codec_guid(codec: CodecKind) -> Option<Guid> {
match codec {
CodecKind::H264 => Some(NV_ENC_CODEC_H264_GUID),
CodecKind::Hevc => Some(NV_ENC_CODEC_HEVC_GUID),
CodecKind::Av1 => Some(NV_ENC_CODEC_AV1_GUID),
_ => None,
}
}
impl VideoEncoder for NvencSession {
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);
}
self.upload.upload_cpu_nv12(data, self.width, self.height)?;
let timestamp = u64::try_from(frame.pts).unwrap_or(0);
let frame_idx = self.frame_idx;
self.frame_idx = self.frame_idx.wrapping_add(1);
self.encoder
.encode_picture(
&self.registered,
&self.bitstream,
frame_idx as usize,
timestamp,
NVencBufferFormat::NV12,
NVencPicStruct::Frame,
NVencPicType::UNKNOWN,
None,
)
.map_err(|_| EncodeError::Backend)?;
let lock = self
.bitstream
.try_lock(true)
.map_err(|_| EncodeError::Backend)?;
let payload = lock.as_slice();
let is_keyframe = is_keyframe_packet(self.codec, payload);
let packet = Packet {
stream_id: 0,
pts: frame.pts,
dts: frame.pts,
duration: frame.duration,
is_keyframe,
is_discard: false,
payload: Bytes::from(payload.to_vec()),
};
drop(lock);
self.pending.push_back(packet);
Ok(())
}
fn poll_packet(&mut self) -> Result<Option<Packet>, EncodeError> {
Ok(self.pending.pop_front())
}
fn flush(&mut self) -> Result<(), EncodeError> {
if self.flushed {
return Ok(());
}
self.flushed = true;
self.encoder
.end_encode()
.map_err(|_| EncodeError::Backend)?;
Ok(())
}
}
fn frame_rate(time_base: Rational) -> [u32; 2] {
let fps_num = time_base.den;
let fps_den = u32::try_from(time_base.num.max(1)).unwrap_or(1);
[fps_num, fps_den]
}
fn is_keyframe_packet(codec: CodecKind, data: &[u8]) -> bool {
match codec {
CodecKind::H264 => contains_h264_idr_nal(data),
CodecKind::Hevc => contains_hevc_idr_nal(data),
CodecKind::Av1 => contains_av1_sequence_header_obu(data),
_ => false,
}
}
fn contains_h264_idr_nal(data: &[u8]) -> bool {
data.windows(5)
.any(|w| w[..4] == [0, 0, 0, 1] && (w[4] & 0x1F) == 5)
}
fn contains_hevc_idr_nal(data: &[u8]) -> bool {
data.windows(5)
.any(|w| w[..4] == [0, 0, 0, 1] && matches!((w[4] >> 1) & 0x3F, 19 | 20))
}
fn contains_av1_sequence_header_obu(data: &[u8]) -> bool {
let mut i = 0usize;
while i < data.len() {
let header = data[i];
if header & 0x80 != 0 {
break; }
let obu_type = (header >> 3) & 0x0F;
let has_extension = header & 0x04 != 0;
let has_size_field = header & 0x02 != 0;
let mut pos = i + 1;
if has_extension {
pos += 1;
}
if !has_size_field {
break;
}
let Some((obu_size, leb_len)) = read_leb128(data.get(pos..).unwrap_or_default()) else {
break;
};
if obu_type == 1 {
return true;
}
pos += leb_len;
let Some(next) = pos.checked_add(obu_size) else {
break;
};
i = next;
}
false
}
fn read_leb128(data: &[u8]) -> Option<(usize, usize)> {
let mut value: u64 = 0;
for (i, &byte) in data.iter().take(8).enumerate() {
value |= u64::from(byte & 0x7F) << (i * 7);
if byte & 0x80 == 0 {
return Some((usize::try_from(value).ok()?, i + 1));
}
}
None
}
fn validate(config: &VideoEncoderConfig) -> Result<(), EncodeError> {
if codec_guid(config.codec).is_none() {
return Err(EncodeError::Unsupported);
}
if config.input != VideoInputPreference::CpuUploadOk {
return Err(EncodeError::Unsupported);
}
if config.width == 0 || config.height == 0 || config.width % 2 != 0 || config.height % 2 != 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);
}
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(),
}
}
fn nvenc_runtime_available() -> bool {
unsafe {
windows::Win32::System::LibraryLoader::LoadLibraryExW(
windows::core::w!("nvEncodeAPI64.dll"),
None,
windows::Win32::System::LibraryLoader::LOAD_LIBRARY_SEARCH_DEFAULT_DIRS,
)
.is_ok()
}
}
#[cfg(test)]
#[path = "video_tests.rs"]
mod tests;