#![allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::print_stderr,
clippy::cast_possible_truncation,
reason = "test file: unwrap/print are fine; picture dimensions are small test constants"
)]
use mediaway_common::{Bytes, Packet, PixelFormat, Rational};
use mediaway_decoder::{VideoDecoder, VideoDecoderConfig, VideoOutputPreference};
use mediaway_decoder_vulkan::VulkanVideoDecoder;
const WIDTH: u32 = 64;
const HEIGHT: u32 = 16;
const MB_SIZE: usize = 16;
struct BitWriter {
bytes: Vec<u8>,
cur: u8,
nbits: u8,
}
impl BitWriter {
const fn new() -> Self {
Self {
bytes: Vec::new(),
cur: 0,
nbits: 0,
}
}
fn push_bit(&mut self, bit: u32) {
let bit_u8 = u8::from(bit & 1 == 1);
self.cur = (self.cur << 1) | bit_u8;
self.nbits += 1;
if self.nbits == 8 {
self.bytes.push(self.cur);
self.cur = 0;
self.nbits = 0;
}
}
fn write_bits(&mut self, value: u32, count: u32) {
for i in (0..count).rev() {
self.push_bit(value >> i);
}
}
fn write_ue(&mut self, value: u32) {
let code = value + 1;
let len = u32::BITS - code.leading_zeros();
for _ in 0..(len - 1) {
self.push_bit(0);
}
self.write_bits(code, len);
}
fn write_se(&mut self, value: i32) {
let magnitude = value.unsigned_abs();
let code = if value <= 0 {
magnitude * 2
} else {
magnitude * 2 - 1
};
self.write_ue(code);
}
fn byte_align(&mut self) {
while self.nbits != 0 {
self.push_bit(0);
}
}
fn write_raw_bytes(&mut self, data: &[u8]) {
debug_assert_eq!(self.nbits, 0, "write_raw_bytes requires byte alignment");
self.bytes.extend_from_slice(data);
}
fn rbsp_trailing_bits(&mut self) {
self.push_bit(1); self.byte_align(); }
fn finish(self) -> Vec<u8> {
self.bytes
}
}
fn write_i_pcm_macroblock(writer: &mut BitWriter, mb_type: u32, luma: u8, cb: u8, cr: u8) {
writer.write_ue(mb_type);
writer.byte_align();
writer.write_raw_bytes(&[luma; MB_SIZE * MB_SIZE]);
writer.write_raw_bytes(&[cb; (MB_SIZE / 2) * (MB_SIZE / 2)]);
writer.write_raw_bytes(&[cr; (MB_SIZE / 2) * (MB_SIZE / 2)]);
}
fn insert_emulation_prevention(rbsp: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(rbsp.len());
let mut zero_run = 0u32;
for &byte in rbsp {
if zero_run >= 2 && byte <= 3 {
out.push(0x03);
zero_run = 0;
}
out.push(byte);
zero_run = if byte == 0 { zero_run + 1 } else { 0 };
}
out
}
fn annex_b_nal(nal_ref_idc: u8, nal_unit_type: u8, rbsp: &[u8]) -> Vec<u8> {
let mut out = vec![0x00, 0x00, 0x00, 0x01];
out.push((nal_ref_idc << 5) | nal_unit_type);
out.extend_from_slice(&insert_emulation_prevention(rbsp));
out
}
fn build_sps() -> Vec<u8> {
let mut writer = BitWriter::new();
writer.write_ue(0); writer.write_ue(0); writer.write_ue(0); writer.write_ue(0); writer.write_ue(1); writer.push_bit(0); writer.write_ue(3); writer.write_ue(0); writer.push_bit(1); writer.push_bit(1); writer.push_bit(0); writer.rbsp_trailing_bits();
let mut rbsp = vec![66u8, 0, 30]; rbsp.extend(writer.finish());
rbsp
}
fn build_pps() -> Vec<u8> {
let mut writer = BitWriter::new();
writer.write_ue(0); writer.write_ue(0); writer.push_bit(0); writer.push_bit(0); writer.write_ue(0); writer.write_ue(0); writer.write_ue(0); writer.push_bit(0); writer.write_bits(0, 2); writer.write_se(0); writer.write_se(0); writer.write_se(0); writer.push_bit(1); writer.push_bit(0); writer.push_bit(0); writer.rbsp_trailing_bits();
writer.finish()
}
fn build_idr_slice() -> Vec<u8> {
let mut writer = BitWriter::new();
writer.write_ue(0); writer.write_ue(2); writer.write_ue(0); writer.write_bits(0, 4); writer.write_ue(0); writer.write_bits(0, 4); writer.push_bit(0); writer.push_bit(0); writer.write_se(0); writer.write_ue(1);
write_i_pcm_macroblock(&mut writer, 25, 200, 100, 150); write_i_pcm_macroblock(&mut writer, 25, 50, 80, 180); write_i_pcm_macroblock(&mut writer, 25, 90, 128, 128); write_i_pcm_macroblock(&mut writer, 25, 90, 128, 128);
writer.rbsp_trailing_bits();
writer.finish()
}
fn build_p_slice() -> Vec<u8> {
let mut writer = BitWriter::new();
writer.write_ue(0); writer.write_ue(0); writer.write_ue(0); writer.write_bits(1, 4); writer.write_bits(2, 4); writer.push_bit(0); writer.push_bit(0); writer.push_bit(0); writer.write_se(0); writer.write_ue(1);
writer.write_ue(1); write_i_pcm_macroblock(&mut writer, 30, 220, 90, 160); writer.write_ue(0); write_i_pcm_macroblock(&mut writer, 30, 90, 128, 128); writer.write_ue(0); write_i_pcm_macroblock(&mut writer, 30, 90, 128, 128);
writer.rbsp_trailing_bits();
writer.finish()
}
fn luma_at(nv12: &[u8], x: usize, y: usize) -> u8 {
nv12[y * WIDTH as usize + x]
}
#[test]
fn decode_idr_then_p_frame_or_skip() {
let mut config = VideoDecoderConfig::h264(WIDTH, HEIGHT, Rational::new(1, 30));
config.output = VideoOutputPreference::CpuFramesOk;
let mut decoder = match VulkanVideoDecoder::open(&config) {
Ok(decoder) => decoder,
Err(error) => {
eprintln!(
"skip: VulkanVideoDecoder::open failed ({error:?}) — no decode-capable Vulkan device?"
);
return;
}
};
let mut first_packet = annex_b_nal(3, 7, &build_sps()); first_packet.extend(annex_b_nal(3, 8, &build_pps())); first_packet.extend(annex_b_nal(1, 5, &build_idr_slice()));
if let Err(error) = decoder.push_packet(&Packet {
stream_id: 0,
pts: 0,
dts: 0,
duration: 1,
is_keyframe: true,
is_discard: false,
payload: Bytes::from(first_packet),
}) {
eprintln!("skip: push_packet(IDR) failed ({error:?})");
return;
}
let idr_frame = match decoder.poll_frame() {
Ok(frame) => frame.expect("expected a decoded IDR frame, got none"),
Err(error) => {
eprintln!("skip: poll_frame(IDR) failed ({error:?})");
return;
}
};
let mediaway_common::VideoFrameStorage::Cpu { data: idr_nv12 } = idr_frame.storage else {
unreachable!(
"expected CPU NV12 storage for the IDR frame (VideoOutputPreference::CpuFramesOk was requested)"
);
};
assert_eq!(idr_frame.format, PixelFormat::Nv12);
assert_eq!(luma_at(&idr_nv12, 4, 4), 200, "IDR MB0 luma");
assert_eq!(luma_at(&idr_nv12, 20, 4), 50, "IDR MB1 luma");
let p_packet = annex_b_nal(1, 1, &build_p_slice()); if let Err(error) = decoder.push_packet(&Packet {
stream_id: 0,
pts: 1,
dts: 1,
duration: 1,
is_keyframe: false,
is_discard: false,
payload: Bytes::from(p_packet),
}) {
eprintln!("skip: push_packet(P) failed ({error:?})");
return;
}
let p_frame = match decoder.poll_frame() {
Ok(frame) => frame.expect("expected a decoded P frame, got none"),
Err(error) => {
eprintln!("skip: poll_frame(P) failed ({error:?})");
return;
}
};
let mediaway_common::VideoFrameStorage::Cpu { data: p_nv12 } = p_frame.storage else {
unreachable!(
"expected CPU NV12 storage for the P frame (VideoOutputPreference::CpuFramesOk was requested)"
);
};
assert_eq!(
luma_at(&p_nv12, 4, 4),
200,
"P MB0 (P_Skip) must match IDR MB0 luma"
);
assert_eq!(luma_at(&p_nv12, 20, 4), 220, "P MB1 (I_PCM) new content");
assert_ne!(
luma_at(&p_nv12, 20, 4),
luma_at(&idr_nv12, 20, 4),
"P-frame output must genuinely differ from the IDR, not just re-emit it"
);
let _ = decoder.flush();
}