use std::io::Write;
use oxideav_core::{Error, MediaType, Packet, Result, StreamInfo};
use oxideav_core::{Muxer, WriteSeek};
use crate::codec_id;
use crate::ebml::{write_element_id, write_vint, VINT_UNKNOWN_SIZE};
use crate::ids;
const CLUSTER_DURATION_MS: i64 = 5_000;
pub fn open(output: Box<dyn WriteSeek>, streams: &[StreamInfo]) -> Result<Box<dyn Muxer>> {
MkvMuxer::new(output, streams, DocType::Matroska).map(|m| Box::new(m) as Box<dyn Muxer>)
}
pub fn open_webm(output: Box<dyn WriteSeek>, streams: &[StreamInfo]) -> Result<Box<dyn Muxer>> {
MkvMuxer::new(output, streams, DocType::Webm).map(|m| Box::new(m) as Box<dyn Muxer>)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum DocType {
Matroska,
Webm,
}
impl DocType {
fn as_str(self) -> &'static str {
match self {
DocType::Matroska => "matroska",
DocType::Webm => "webm",
}
}
}
pub struct MkvMuxer {
output: Box<dyn WriteSeek>,
streams: Vec<StreamInfo>,
track_numbers: Vec<u64>,
stream_pts: Vec<i64>,
cluster_open: bool,
cluster_timecode_ms: i64,
cluster_offset_rel: u64,
segment_data_start: u64,
cues: Vec<CueRecord>,
cue_seen_in_cluster: Vec<bool>,
seek_cues_entry_offset: u64,
seek_head_written: bool,
header_written: bool,
trailer_written: bool,
doc_type: DocType,
}
#[derive(Clone, Copy, Debug)]
struct CueRecord {
track: u64,
time_ms: u64,
cluster_offset: u64,
}
impl MkvMuxer {
fn new(output: Box<dyn WriteSeek>, streams: &[StreamInfo], doc_type: DocType) -> Result<Self> {
if streams.is_empty() {
return Err(Error::invalid("MKV muxer: need at least one stream"));
}
if doc_type == DocType::Webm {
for (i, s) in streams.iter().enumerate() {
if !codec_id::is_webm_codec(&s.params.codec_id) {
return Err(Error::unsupported(format!(
"WebM muxer: stream {i} uses codec '{}' which is not in the WebM whitelist (allowed: vp8, vp9, av1, vorbis, opus)",
s.params.codec_id.as_str()
)));
}
}
}
let stream_track_numbers: Vec<u64> = (0..streams.len() as u64).map(|i| i + 1).collect();
let n = streams.len();
Ok(MkvMuxer {
output,
streams: streams.to_vec(),
track_numbers: stream_track_numbers,
stream_pts: vec![0i64; n],
cluster_open: false,
cluster_timecode_ms: 0,
cluster_offset_rel: 0,
segment_data_start: 0,
cues: Vec::new(),
cue_seen_in_cluster: vec![false; n],
seek_cues_entry_offset: 0,
seek_head_written: false,
header_written: false,
trailer_written: false,
doc_type,
})
}
pub fn new_matroska(output: Box<dyn WriteSeek>, streams: &[StreamInfo]) -> Result<Self> {
Self::new(output, streams, DocType::Matroska)
}
pub fn new_webm(output: Box<dyn WriteSeek>, streams: &[StreamInfo]) -> Result<Self> {
Self::new(output, streams, DocType::Webm)
}
}
impl Muxer for MkvMuxer {
fn format_name(&self) -> &str {
self.doc_type.as_str()
}
fn write_header(&mut self) -> Result<()> {
if self.header_written {
return Err(Error::other("MKV muxer: write_header called twice"));
}
let base_pos = self.output.stream_position().unwrap_or(0);
let mut ebml_body = Vec::new();
write_uint_element(&mut ebml_body, ids::EBML_VERSION, 1);
write_uint_element(&mut ebml_body, ids::EBML_READ_VERSION, 1);
write_uint_element(&mut ebml_body, ids::EBML_MAX_ID_LENGTH, 4);
write_uint_element(&mut ebml_body, ids::EBML_MAX_SIZE_LENGTH, 8);
write_string_element(&mut ebml_body, ids::EBML_DOC_TYPE, self.doc_type.as_str());
write_uint_element(&mut ebml_body, ids::EBML_DOC_TYPE_VERSION, 4);
write_uint_element(&mut ebml_body, ids::EBML_DOC_TYPE_READ_VERSION, 2);
let mut all = Vec::new();
write_master_element(&mut all, ids::EBML_HEADER, &ebml_body);
all.extend_from_slice(&write_element_id(ids::SEGMENT));
all.extend_from_slice(&write_vint(VINT_UNKNOWN_SIZE, 0));
let segment_data_start_in_buf = all.len() as u64;
let seek_head_offset_in_buf = all.len() as u64 - segment_data_start_in_buf;
let seek_head_bytes = build_initial_seek_head();
let seek_head_start_in_buf = all.len();
all.extend_from_slice(&seek_head_bytes);
let info_seek_entry_in_buf = seek_head_start_in_buf + SEEK_HEAD_HEADER_LEN;
let tracks_seek_entry_in_buf = info_seek_entry_in_buf + SEEK_ENTRY_LEN;
let cues_seek_entry_in_buf = tracks_seek_entry_in_buf + SEEK_ENTRY_LEN;
debug_assert_eq!(seek_head_bytes.len(), SEEK_HEAD_TOTAL_LEN);
let _ = seek_head_offset_in_buf;
let info_offset_in_buf = all.len() as u64 - segment_data_start_in_buf;
let mut info_body = Vec::new();
write_uint_element(&mut info_body, ids::TIMECODE_SCALE, 1_000_000); write_string_element(&mut info_body, ids::MUXING_APP, "oxideav");
write_string_element(&mut info_body, ids::WRITING_APP, "oxideav");
write_master_element(&mut all, ids::INFO, &info_body);
let tracks_offset_in_buf = all.len() as u64 - segment_data_start_in_buf;
let mut tracks_body = Vec::new();
for (i, s) in self.streams.iter().enumerate() {
let track_number = self.track_numbers[i];
let mut t = Vec::new();
write_uint_element(&mut t, ids::TRACK_NUMBER, track_number);
write_uint_element(&mut t, ids::TRACK_UID, track_number);
let track_type = match s.params.media_type {
MediaType::Audio => ids::TRACK_TYPE_AUDIO,
MediaType::Video => ids::TRACK_TYPE_VIDEO,
MediaType::Subtitle => ids::TRACK_TYPE_SUBTITLE,
_ => 17, };
write_uint_element(&mut t, ids::TRACK_TYPE, track_type);
write_uint_element(&mut t, ids::FLAG_LACING, 0);
if let Some(name) = codec_id::to_matroska(&s.params.codec_id) {
write_string_element(&mut t, ids::CODEC_ID, name);
} else {
let raw = format!("X_{}", s.params.codec_id);
write_string_element(&mut t, ids::CODEC_ID, &raw);
}
let cp = encode_codec_private(&s.params.codec_id, &s.params.extradata);
if !cp.is_empty() {
write_bytes_element(&mut t, ids::CODEC_PRIVATE, &cp);
}
if s.params.codec_id.as_str() == "opus" {
let pre_skip_samples = parse_opus_pre_skip(&s.params.extradata);
let codec_delay_ns = pre_skip_samples as u64 * 1_000_000_000 / 48_000;
write_uint_element(&mut t, ids::CODEC_DELAY, codec_delay_ns);
write_uint_element(&mut t, ids::SEEK_PRE_ROLL, 80_000_000);
}
if s.params.media_type == MediaType::Audio {
let mut audio = Vec::new();
if let Some(sr) = s.params.sample_rate {
write_float_element(&mut audio, ids::SAMPLING_FREQUENCY, sr as f64);
}
if let Some(ch) = s.params.channels {
write_uint_element(&mut audio, ids::CHANNELS, ch as u64);
}
if let Some(fmt) = s.params.sample_format {
let bd = (fmt.bytes_per_sample() * 8) as u64;
write_uint_element(&mut audio, ids::BIT_DEPTH, bd);
}
write_master_element(&mut t, ids::AUDIO, &audio);
}
if s.params.media_type == MediaType::Video {
let mut video = Vec::new();
if let Some(w) = s.params.width {
write_uint_element(&mut video, ids::PIXEL_WIDTH, w as u64);
}
if let Some(h) = s.params.height {
write_uint_element(&mut video, ids::PIXEL_HEIGHT, h as u64);
}
write_master_element(&mut t, ids::VIDEO, &video);
}
write_master_element(&mut tracks_body, ids::TRACK_ENTRY, &t);
}
write_master_element(&mut all, ids::TRACKS, &tracks_body);
write_u64_be_at(
&mut all,
info_seek_entry_in_buf + SEEK_POS_PAYLOAD_OFFSET,
info_offset_in_buf,
);
write_u64_be_at(
&mut all,
tracks_seek_entry_in_buf + SEEK_POS_PAYLOAD_OFFSET,
tracks_offset_in_buf,
);
self.segment_data_start = base_pos + segment_data_start_in_buf;
self.seek_cues_entry_offset = base_pos + cues_seek_entry_in_buf as u64;
self.seek_head_written = true;
self.output.write_all(&all)?;
self.header_written = true;
Ok(())
}
fn write_packet(&mut self, packet: &Packet) -> Result<()> {
if !self.header_written {
return Err(Error::other("MKV muxer: write_header not called"));
}
let stream_idx = packet.stream_index as usize;
if stream_idx >= self.streams.len() {
return Err(Error::invalid(format!(
"MKV muxer: unknown stream index {}",
stream_idx
)));
}
let track_number = self.track_numbers[stream_idx];
let stream_time_base = self.streams[stream_idx].time_base;
let media_type = self.streams[stream_idx].params.media_type;
let codec = self.streams[stream_idx].params.codec_id.as_str().to_owned();
let derived_duration: Option<i64> = match codec.as_str() {
"opus" => opus_packet_duration_samples(&packet.data).map(|s| s as i64),
_ => packet.duration,
};
let effective_pts = match packet.pts {
Some(v) => v,
None => self.stream_pts[stream_idx],
};
if let Some(d) = derived_duration {
self.stream_pts[stream_idx] = effective_pts + d;
} else if packet.pts.is_some() {
self.stream_pts[stream_idx] = effective_pts;
}
let pts_ms = pts_to_ms(effective_pts, stream_time_base);
if !self.cluster_open
|| pts_ms - self.cluster_timecode_ms > CLUSTER_DURATION_MS
|| pts_ms - self.cluster_timecode_ms > i16::MAX as i64
|| pts_ms - self.cluster_timecode_ms < 0
{
self.start_cluster(pts_ms)?;
}
let timecode_offset = pts_ms - self.cluster_timecode_ms;
if timecode_offset < i16::MIN as i64 || timecode_offset > i16::MAX as i64 {
return Err(Error::other(
"MKV muxer: packet timecode delta exceeds i16 range",
));
}
if !self.cue_seen_in_cluster[stream_idx] {
let indexable = match media_type {
MediaType::Video => packet.flags.keyframe,
_ => true,
};
if indexable {
self.cues.push(CueRecord {
track: track_number,
time_ms: pts_ms.max(0) as u64,
cluster_offset: self.cluster_offset_rel,
});
self.cue_seen_in_cluster[stream_idx] = true;
}
}
let block_bytes =
build_simple_block(track_number, timecode_offset as i16, packet, &packet.data);
self.output.write_all(&block_bytes)?;
Ok(())
}
fn write_trailer(&mut self) -> Result<()> {
if self.trailer_written {
return Ok(());
}
let cues_offset_rel = self.write_cues()?;
if self.seek_head_written {
self.patch_cues_seek_entry(cues_offset_rel)?;
}
self.output.flush()?;
self.trailer_written = true;
Ok(())
}
}
impl MkvMuxer {
fn start_cluster(&mut self, timecode_ms: i64) -> Result<()> {
let cluster_abs = self.output.stream_position().unwrap_or(0);
self.cluster_offset_rel = cluster_abs.saturating_sub(self.segment_data_start);
self.output.write_all(&write_element_id(ids::CLUSTER))?;
self.output.write_all(&write_vint(VINT_UNKNOWN_SIZE, 0))?;
let mut tc = Vec::new();
write_uint_element(&mut tc, ids::TIMECODE, timecode_ms.max(0) as u64);
self.output.write_all(&tc)?;
self.cluster_timecode_ms = timecode_ms.max(0);
self.cluster_open = true;
for s in self.cue_seen_in_cluster.iter_mut() {
*s = false;
}
Ok(())
}
fn write_cues(&mut self) -> Result<Option<u64>> {
if self.cues.is_empty() {
return Ok(None);
}
let mut by_time: std::collections::BTreeMap<u64, Vec<CueRecord>> =
std::collections::BTreeMap::new();
for c in &self.cues {
by_time.entry(c.time_ms).or_default().push(*c);
}
let mut body = Vec::new();
for (time, entries) in by_time {
let mut cp = Vec::new();
write_uint_element(&mut cp, ids::CUE_TIME, time);
for e in entries {
let mut ctp = Vec::new();
write_uint_element(&mut ctp, ids::CUE_TRACK, e.track);
write_uint_element(&mut ctp, ids::CUE_CLUSTER_POSITION, e.cluster_offset);
write_master_element(&mut cp, ids::CUE_TRACK_POSITIONS, &ctp);
}
write_master_element(&mut body, ids::CUE_POINT, &cp);
}
let mut out = Vec::with_capacity(body.len() + 8);
write_master_element(&mut out, ids::CUES, &body);
let cues_abs = self.output.stream_position().unwrap_or(0);
self.output.write_all(&out)?;
Ok(Some(cues_abs.saturating_sub(self.segment_data_start)))
}
fn patch_cues_seek_entry(&mut self, cues_offset_rel: Option<u64>) -> Result<()> {
use std::io::SeekFrom;
let resume_pos = self.output.stream_position().unwrap_or(0);
match cues_offset_rel {
Some(off) => {
let payload_pos = self.seek_cues_entry_offset + SEEK_POS_PAYLOAD_OFFSET as u64;
self.output.seek(SeekFrom::Start(payload_pos))?;
self.output.write_all(&off.to_be_bytes())?;
}
None => {
self.output
.seek(SeekFrom::Start(self.seek_cues_entry_offset))?;
self.output.write_all(&void_seek_entry())?;
}
}
self.output.seek(SeekFrom::Start(resume_pos))?;
Ok(())
}
}
fn build_simple_block(track: u64, tc_offset: i16, packet: &Packet, data: &[u8]) -> Vec<u8> {
let mut body = Vec::with_capacity(4 + data.len());
body.extend_from_slice(&write_vint(track, 0));
body.extend_from_slice(&tc_offset.to_be_bytes());
let mut flags: u8 = 0;
if packet.flags.keyframe {
flags |= 0x80;
}
body.push(flags);
body.extend_from_slice(data);
let mut out = Vec::with_capacity(8 + body.len());
out.extend_from_slice(&write_element_id(ids::SIMPLE_BLOCK));
out.extend_from_slice(&write_vint(body.len() as u64, 0));
out.extend_from_slice(&body);
out
}
fn pts_to_ms(value: i64, tb: oxideav_core::TimeBase) -> i64 {
let r = tb.as_rational();
if r.den == 0 {
return value;
}
let v = value as i128 * r.num as i128 * 1000;
(v / r.den as i128) as i64
}
fn opus_packet_duration_samples(packet: &[u8]) -> Option<u32> {
if packet.is_empty() {
return None;
}
let toc = packet[0];
let config = toc >> 3;
let frame_size_48k: u32 = match config {
0 | 4 | 8 => 480,
1 | 5 | 9 => 960,
2 | 6 | 10 => 1920,
3 | 7 | 11 => 2880,
12 | 14 => 480,
13 | 15 => 960,
16 | 20 | 24 | 28 => 120,
17 | 21 | 25 | 29 => 240,
18 | 22 | 26 | 30 => 480,
19 | 23 | 27 | 31 => 960,
_ => return None,
};
let n_frames: u32 = match toc & 0x03 {
0 => 1,
1 | 2 => 2,
3 => {
if packet.len() < 2 {
return None;
}
(packet[1] & 0x3F) as u32
}
_ => unreachable!(),
};
Some(frame_size_48k * n_frames)
}
fn parse_opus_pre_skip(extradata: &[u8]) -> u16 {
if extradata.len() < 12 || &extradata[0..8] != b"OpusHead" {
return 0;
}
u16::from_le_bytes([extradata[10], extradata[11]])
}
fn encode_codec_private(codec_id: &oxideav_core::CodecId, extradata: &[u8]) -> Vec<u8> {
match codec_id.as_str() {
"flac" => {
let mut out = Vec::with_capacity(4 + extradata.len());
out.extend_from_slice(b"fLaC");
out.extend_from_slice(extradata);
out
}
_ => extradata.to_vec(),
}
}
fn write_uint_element(buf: &mut Vec<u8>, id: u32, value: u64) {
let n = if value == 0 {
1
} else {
(64 - value.leading_zeros()).div_ceil(8) as usize
};
buf.extend_from_slice(&write_element_id(id));
buf.extend_from_slice(&write_vint(n as u64, 0));
for i in (0..n).rev() {
buf.push(((value >> (i * 8)) & 0xFF) as u8);
}
}
fn write_string_element(buf: &mut Vec<u8>, id: u32, value: &str) {
buf.extend_from_slice(&write_element_id(id));
buf.extend_from_slice(&write_vint(value.len() as u64, 0));
buf.extend_from_slice(value.as_bytes());
}
fn write_bytes_element(buf: &mut Vec<u8>, id: u32, value: &[u8]) {
buf.extend_from_slice(&write_element_id(id));
buf.extend_from_slice(&write_vint(value.len() as u64, 0));
buf.extend_from_slice(value);
}
fn write_float_element(buf: &mut Vec<u8>, id: u32, value: f64) {
buf.extend_from_slice(&write_element_id(id));
buf.extend_from_slice(&write_vint(8, 0));
buf.extend_from_slice(&value.to_be_bytes());
}
fn write_master_element(buf: &mut Vec<u8>, id: u32, body: &[u8]) {
buf.extend_from_slice(&write_element_id(id));
buf.extend_from_slice(&write_vint(body.len() as u64, 0));
buf.extend_from_slice(body);
}
const SEEK_HEAD_HEADER_LEN: usize = 5;
const SEEK_HEAD_TOTAL_LEN: usize = SEEK_HEAD_HEADER_LEN + 3 * SEEK_ENTRY_LEN;
const SEEK_ENTRY_LEN: usize = 21;
const SEEK_POS_PAYLOAD_OFFSET: usize = 13;
fn build_initial_seek_head() -> Vec<u8> {
let mut body = Vec::with_capacity(3 * SEEK_ENTRY_LEN);
body.extend_from_slice(&seek_entry(ids::INFO, 0));
body.extend_from_slice(&seek_entry(ids::TRACKS, 0));
body.extend_from_slice(&seek_entry(ids::CUES, 0));
debug_assert_eq!(body.len(), 3 * SEEK_ENTRY_LEN);
let mut out = Vec::with_capacity(SEEK_HEAD_TOTAL_LEN);
write_master_element(&mut out, ids::SEEK_HEAD, &body);
debug_assert_eq!(out.len(), SEEK_HEAD_TOTAL_LEN);
out
}
fn seek_entry(target_id: u32, position: u64) -> Vec<u8> {
let mut body = Vec::with_capacity(SEEK_ENTRY_LEN - 3);
body.extend_from_slice(&write_element_id(ids::SEEK_ID));
body.extend_from_slice(&write_vint(4, 0));
body.extend_from_slice(&target_id.to_be_bytes());
body.extend_from_slice(&write_element_id(ids::SEEK_POSITION));
body.extend_from_slice(&write_vint(8, 0));
body.extend_from_slice(&position.to_be_bytes());
debug_assert_eq!(body.len(), SEEK_ENTRY_LEN - 3);
let mut entry = Vec::with_capacity(SEEK_ENTRY_LEN);
write_master_element(&mut entry, ids::SEEK, &body);
debug_assert_eq!(entry.len(), SEEK_ENTRY_LEN);
entry
}
fn void_seek_entry() -> Vec<u8> {
let mut out = Vec::with_capacity(SEEK_ENTRY_LEN);
out.push(ids::VOID as u8); out.push(0x93); out.resize(SEEK_ENTRY_LEN, 0u8);
debug_assert_eq!(out.len(), SEEK_ENTRY_LEN);
out
}
fn write_u64_be_at(buf: &mut [u8], pos: usize, value: u64) {
buf[pos..pos + 8].copy_from_slice(&value.to_be_bytes());
}