use std::io::{Read, Seek, SeekFrom};
use oxideav_core::{Error, Result};
pub const VINT_UNKNOWN_SIZE: u64 = u64::MAX;
pub fn read_vint(r: &mut dyn Read, keep_marker: bool) -> Result<(u64, usize)> {
let mut first = [0u8; 1];
r.read_exact(&mut first)?;
let b0 = first[0];
if b0 == 0 {
return Err(Error::invalid("EBML VINT: invalid leading byte 0x00"));
}
let len = (b0.leading_zeros() + 1) as usize;
if len > 8 {
return Err(Error::invalid("EBML VINT: width > 8 bytes"));
}
let mut value: u64 = if keep_marker {
b0 as u64
} else {
(b0 & ((1u8 << (8 - len)) - 1)) as u64
};
let mut buf = [0u8; 8];
let extra = len - 1;
if extra > 0 {
r.read_exact(&mut buf[..extra])?;
for &b in &buf[..extra] {
value = (value << 8) | (b as u64);
}
}
if !keep_marker && len <= 8 {
let payload_bits = (8 - len) as u32 + 8 * extra as u32;
let all_ones = if payload_bits >= 64 {
u64::MAX
} else {
(1u64 << payload_bits) - 1
};
if value == all_ones {
return Ok((VINT_UNKNOWN_SIZE, len));
}
}
Ok((value, len))
}
pub fn write_vint(value: u64, min_width: u8) -> Vec<u8> {
if value == VINT_UNKNOWN_SIZE {
return vec![0xFF];
}
let mut width = min_width.max(1);
loop {
let payload_bits = (8 - width as u32) + 8 * (width as u32 - 1);
let all_ones = if payload_bits >= 64 {
u64::MAX
} else {
(1u64 << payload_bits) - 1
};
if value < all_ones {
break;
}
width += 1;
if width > 8 {
panic!("EBML VINT value too large to encode");
}
}
let mut out = vec![0u8; width as usize];
out[0] = 1u8 << (8 - width);
let mut v = value;
for i in (0..width as usize).rev() {
out[i] |= (v & 0xFF) as u8;
v >>= 8;
}
out
}
pub fn write_element_id(id: u32) -> Vec<u8> {
let bytes = if id < 0x100 {
1
} else if id < 0x10000 {
2
} else if id < 0x1000000 {
3
} else {
4
};
let mut out = Vec::with_capacity(bytes);
for i in (0..bytes).rev() {
out.push(((id >> (i * 8)) & 0xFF) as u8);
}
out
}
#[derive(Clone, Debug)]
pub struct ElementHeader {
pub id: u32,
pub size: u64,
pub header_len: usize,
}
pub fn read_element_header(r: &mut dyn Read) -> Result<ElementHeader> {
let (id, id_len) = read_vint(r, true)?;
if id > u32::MAX as u64 {
return Err(Error::invalid("EBML: element id exceeds 32 bits"));
}
let (size, size_len) = read_vint(r, false)?;
Ok(ElementHeader {
id: id as u32,
size,
header_len: id_len + size_len,
})
}
pub fn read_uint(r: &mut dyn Read, n: usize) -> Result<u64> {
if n > 8 {
return Err(Error::invalid("EBML uint > 8 bytes"));
}
if n == 0 {
return Ok(0);
}
let mut buf = [0u8; 8];
r.read_exact(&mut buf[..n])?;
let mut v = 0u64;
for &b in &buf[..n] {
v = (v << 8) | (b as u64);
}
Ok(v)
}
pub fn read_int(r: &mut dyn Read, n: usize) -> Result<i64> {
if n == 0 {
return Ok(0);
}
let raw = read_uint(r, n)?;
let shift = 64 - 8 * n as u32;
Ok(((raw << shift) as i64) >> shift)
}
pub fn read_float(r: &mut dyn Read, n: usize) -> Result<f64> {
match n {
0 => Ok(0.0),
4 => {
let mut buf = [0u8; 4];
r.read_exact(&mut buf)?;
Ok(f32::from_be_bytes(buf) as f64)
}
8 => {
let mut buf = [0u8; 8];
r.read_exact(&mut buf)?;
Ok(f64::from_be_bytes(buf))
}
_ => Err(Error::invalid(format!(
"EBML float must be 4 or 8 bytes (got {n})"
))),
}
}
pub fn read_string(r: &mut dyn Read, n: usize) -> Result<String> {
let mut buf = read_bytes(r, n)?;
while buf.last() == Some(&0) {
buf.pop();
}
String::from_utf8(buf).map_err(|e| Error::invalid(format!("EBML string not UTF-8: {e}")))
}
pub fn read_bytes(r: &mut dyn Read, n: usize) -> Result<Vec<u8>> {
let mut buf = Vec::new();
let read = r.take(n as u64).read_to_end(&mut buf)?;
if read != n {
return Err(Error::from(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
format!("EBML: short read ({read} of {n} bytes)"),
)));
}
Ok(buf)
}
pub fn skip<R: Seek + ?Sized>(r: &mut R, n: u64) -> Result<()> {
if n > 0 {
let cur = r.stream_position()?;
let target = cur.saturating_add(n);
r.seek(SeekFrom::Start(target))?;
}
Ok(())
}
pub fn crc32_ieee(data: &[u8]) -> u32 {
use std::sync::OnceLock;
static TABLE: OnceLock<[u32; 256]> = OnceLock::new();
let table = TABLE.get_or_init(|| {
let mut t = [0u32; 256];
let mut n = 0usize;
while n < 256 {
let mut c = n as u32;
let mut k = 0;
while k < 8 {
c = if c & 1 != 0 {
0xEDB8_8320 ^ (c >> 1)
} else {
c >> 1
};
k += 1;
}
t[n] = c;
n += 1;
}
t
});
let mut crc = 0xFFFF_FFFFu32;
for &b in data {
crc = table[((crc ^ b as u32) & 0xFF) as usize] ^ (crc >> 8);
}
crc ^ 0xFFFF_FFFF
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn vint_round_trip_small() {
for v in [
0u64,
1,
126,
127,
128,
255,
16_000,
1_000_000,
1_234_567_890,
]
.iter()
{
let bytes = write_vint(*v, 0);
let mut c = Cursor::new(&bytes);
let (got, len) = read_vint(&mut c, false).unwrap();
assert_eq!(got, *v, "v={v}");
assert_eq!(len, bytes.len());
}
}
#[test]
fn vint_known_widths() {
assert_eq!(write_vint(0, 0), vec![0x80]);
assert_eq!(write_vint(126, 0), vec![0xFE]);
assert_eq!(write_vint(127, 0), vec![0x40, 0x7F]);
}
#[test]
fn id_round_trip() {
let bytes = write_element_id(0x1A45DFA3);
assert_eq!(bytes, vec![0x1A, 0x45, 0xDF, 0xA3]);
let mut c = Cursor::new(&bytes);
let (got, len) = read_vint(&mut c, true).unwrap();
assert_eq!(got as u32, 0x1A45DFA3);
assert_eq!(len, 4);
}
#[test]
fn unknown_size_sentinel() {
let mut c = Cursor::new(&[0xFFu8]);
let (v, _) = read_vint(&mut c, false).unwrap();
assert_eq!(v, VINT_UNKNOWN_SIZE);
}
#[test]
fn crc32_check_value() {
assert_eq!(crc32_ieee(b"123456789"), 0xCBF4_3926);
}
#[test]
fn crc32_empty_is_zero() {
assert_eq!(crc32_ieee(b""), 0);
}
#[test]
fn crc32_little_endian_storage() {
let crc = crc32_ieee(b"123456789");
assert_eq!(crc.to_le_bytes(), [0x26, 0x39, 0xF4, 0xCB]);
}
}