use std::sync::OnceLock;
use std::sync::atomic::{AtomicUsize, Ordering};
use crate::threadpool::ProgressGate;
use crate::cabac::{ContextSet, IntraModeContexts, PaletteContexts};
use crate::decode::FullDecoder;
use crate::error::DecodeError;
use crate::palette::PalettePredictor;
use crate::threadpool::ThreadPool;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct RowSubstream {
pub start: usize,
pub end: usize,
}
pub(crate) fn row_substreams(
src_of: &[usize],
cabac_rbsp_off: usize,
entry_points: &[u32],
rbsp_len: usize,
ctb_rows: usize,
) -> Option<Vec<RowSubstream>> {
if ctb_rows == 0 {
return None;
}
if entry_points.len() + 1 != ctb_rows {
return None;
}
if cabac_rbsp_off > rbsp_len {
return None;
}
let nal_data_start = if cabac_rbsp_off < src_of.len() {
src_of[cabac_rbsp_off]
} else {
return None;
};
let mut rows = Vec::with_capacity(ctb_rows);
let mut nal_cursor = nal_data_start;
let mut rbsp_start = cabac_rbsp_off;
for (i, _) in (0..ctb_rows).enumerate() {
let rbsp_end = if i + 1 < ctb_rows {
nal_cursor += entry_points[i] as usize;
let e = crate::bitreader::nal_to_rbsp_offset(src_of, nal_cursor);
e.min(rbsp_len)
} else {
rbsp_len
};
if rbsp_end < rbsp_start {
return None;
}
rows.push(RowSubstream {
start: rbsp_start,
end: rbsp_end,
});
rbsp_start = rbsp_end;
}
Some(rows)
}
pub(crate) fn run_wavefront(
template: &mut FullDecoder<'_>,
rbsp: &[u8],
rows: &[RowSubstream],
pool: &ThreadPool,
) -> Result<(), DecodeError> {
let ctb_rows = rows.len();
debug_assert_eq!(ctb_rows, template.ctb_rows_pub());
let progress: Vec<ProgressGate> = (0..ctb_rows).map(|_| ProgressGate::new()).collect();
let snapshots: Vec<
OnceLock<(
ContextSet,
IntraModeContexts,
PaletteContexts,
PalettePredictor,
)>,
> = (0..ctb_rows).map(|_| OnceLock::new()).collect();
let (init_ctx, init_ictx, init_pctx, init_palette) = template.init_contexts_pub();
let factory = template.row_factory();
let first_err: OnceLock<DecodeError> = OnceLock::new();
let next_row = AtomicUsize::new(0);
let n_runners = pool.threads().max(1).min(ctb_rows);
pool.scope(|scope| {
for _ in 0..n_runners {
let progress_ref = &progress;
let snapshots_ref = &snapshots;
let first_err_ref = &first_err;
let factory_ref = &factory;
let next_row_ref = &next_row;
let init_ctx = init_ctx.clone();
let init_palette = init_palette.clone();
scope.spawn(move || {
loop {
let ry = next_row_ref.fetch_add(1, Ordering::Relaxed);
if ry >= ctb_rows {
break;
}
if first_err_ref.get().is_some() {
let _ = snapshots_ref[ry].set(default_contexts());
progress_ref[ry].publish(usize::MAX);
continue;
}
let sub = rows[ry];
if sub.start > sub.end || sub.end > rbsp.len() {
let _ =
first_err_ref.set(DecodeError::Bitstream("wpp substream range".into()));
let _ = snapshots_ref[ry].set(default_contexts());
progress_ref[ry].publish(usize::MAX);
continue;
}
let row_cabac = &rbsp[sub.start..sub.end];
let (ctx, ictx, pctx, palette_predictor) = if ry == 0 {
(init_ctx.clone(), init_ictx, init_pctx, init_palette.clone())
} else {
progress_ref[ry - 1].wait_at_least(2);
snapshots_ref[ry - 1]
.get()
.cloned()
.unwrap_or_else(default_contexts)
};
let mut row = match unsafe {
factory_ref.make(row_cabac, ctx, ictx, pctx, palette_predictor)
} {
Ok(r) => r,
Err(e) => {
let _ = first_err_ref.set(e);
let _ = snapshots_ref[ry].set(default_contexts());
progress_ref[ry].publish(usize::MAX);
continue;
}
};
let above = if ry == 0 {
None
} else {
Some(&progress_ref[ry - 1])
};
if let Err(e) =
row.decode_wavefront_row(ry, &progress_ref[ry], above, &snapshots_ref[ry])
{
let _ = first_err_ref.set(e);
}
}
});
}
});
if let Some(e) = first_err.into_inner() {
return Err(e);
}
Ok(())
}
fn default_contexts() -> (
ContextSet,
IntraModeContexts,
PaletteContexts,
PalettePredictor,
) {
(
ContextSet::init_islice(26),
IntraModeContexts::init_islice(26),
PaletteContexts::init(26),
PalettePredictor::default(),
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::bitreader::unescape_rbsp_with_map;
#[test]
fn substreams_without_emulation_bytes() {
let nal = vec![0xAAu8; 20];
let (_rbsp, src_of) = unescape_rbsp_with_map(&nal);
let rows = row_substreams(&src_of, 2, &[4, 5], 20, 3).unwrap();
assert_eq!(rows[0], RowSubstream { start: 2, end: 6 });
assert_eq!(rows[1], RowSubstream { start: 6, end: 11 });
assert_eq!(rows[2], RowSubstream { start: 11, end: 20 });
}
#[test]
fn wrong_entry_point_count_rejected() {
let nal = vec![0xAAu8; 20];
let (_r, src_of) = unescape_rbsp_with_map(&nal);
assert!(row_substreams(&src_of, 2, &[4], 20, 3).is_none());
}
#[test]
fn emulation_bytes_shift_rbsp_offsets() {
let mut nal = vec![0x00, 0x00, 0x03, 0x00, 0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF];
nal.extend_from_slice(&[0x11, 0x22, 0x33, 0x44]);
let (rbsp, src_of) = unescape_rbsp_with_map(&nal);
assert_eq!(rbsp.len(), nal.len() - 1);
let rows = row_substreams(&src_of, 0, &[], rbsp.len(), 1).unwrap();
assert_eq!(
rows[0],
RowSubstream {
start: 0,
end: rbsp.len()
}
);
let rows2 = row_substreams(&src_of, 0, &[5], rbsp.len(), 2).unwrap();
assert_eq!(rows2[0].start, 0);
assert_eq!(rows2[0].end, 4);
assert_eq!(rows2[1].start, 4);
assert_eq!(rows2[1].end, rbsp.len());
}
}
pub(crate) fn substream_starts_rbsp_rel(
src_of: &[usize],
cabac_rbsp_off: usize,
entry_points: &[u32],
rbsp_len: usize,
) -> Vec<usize> {
let mut starts = Vec::with_capacity(entry_points.len() + 1);
starts.push(0usize);
if cabac_rbsp_off >= src_of.len() {
let mut acc = 0usize;
for &len in entry_points {
acc = acc.saturating_add(len as usize);
starts.push(acc.min(rbsp_len.saturating_sub(cabac_rbsp_off)));
}
return starts;
}
let nal_data_start = src_of[cabac_rbsp_off];
let mut nal_cursor = nal_data_start;
for &len in entry_points {
nal_cursor += len as usize;
let rbsp_abs = crate::bitreader::nal_to_rbsp_offset(src_of, nal_cursor).min(rbsp_len);
starts.push(rbsp_abs.saturating_sub(cabac_rbsp_off));
}
starts
}