Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
171 changes: 151 additions & 20 deletions crates/jxr-native/src/entropy/bit_reader.rs
Original file line number Diff line number Diff line change
Expand Up @@ -50,41 +50,115 @@ impl<'a> PacketBitReader<'a> {
}

/// Reads one bit.
#[inline]
pub fn read_bit(&mut self) -> Result<bool, EntropyError> {
Ok(self.read_bits(1)? != 0)
}

/// Reads up to 64 bits, most-significant bit first.
#[inline]
pub fn read_bits(&mut self, count: u8) -> Result<u64, EntropyError> {
// A byte-aligned 64-bit window always holds 57 bits after the in-byte
// offset, so every entropy-syntax read needs exactly one load.
if count > 57 {
return self.read_wide_bits(count);
}
let end = self.checked_end(count)?;
let value = match count {
0 => 0,
_ => self.window() >> (64 - count),
};
self.bit_position = end;
Ok(value)
}

/// Returns the next eight bits without consuming them, zero-filled past the packet end.
#[inline]
pub(crate) fn peek_byte(&self) -> u8 {
let remaining = self.bits_remaining();
let mut window = self.window();
if remaining < 8 {
window &= !(u64::MAX >> remaining);
}
window.to_be_bytes()[0]
}

/// Advances past bits already inspected with [`Self::peek_byte`].
#[inline]
pub(crate) fn skip_bits(&mut self, count: u8) -> Result<(), EntropyError> {
self.bit_position = self.checked_end(count)?;
Ok(())
}

/// Builds the error a bit-serial reader reports after consuming `consumed` more bits.
#[cold]
pub(crate) fn unexpected_end_after(&self, consumed: usize) -> EntropyError {
EntropyError::UnexpectedEnd {
bit_position: self.bit_position.saturating_add(consumed),
requested_bits: 1,
bit_length: self.bit_length,
}
}

#[inline(never)]
fn read_wide_bits(&mut self, count: u8) -> Result<u64, EntropyError> {
if count > 64 {
return Err(EntropyError::InvalidParameter {
parameter: "bit count",
value: i64::from(count),
});
}
let Some(end) = self.bit_position.checked_add(usize::from(count)) else {
return Err(EntropyError::UnexpectedEnd {
bit_position: self.bit_position,
requested_bits: count,
bit_length: self.bit_length,
});
};
if end > self.bit_length {
return Err(EntropyError::UnexpectedEnd {
bit_position: self.bit_position,
requested_bits: count,
bit_length: self.bit_length,
});
let end = self.checked_end(count)?;
let high_bits = count - 32;
let high = self.window() >> (64 - high_bits);
self.bit_position += usize::from(high_bits);
let low = self.window() >> 32;
self.bit_position = end;
Ok((high << 32) | low)
}

#[inline]
fn checked_end(&self, count: u8) -> Result<usize, EntropyError> {
match self.bit_position.checked_add(usize::from(count)) {
Some(end) if end <= self.bit_length => Ok(end),
_ => Err(self.end_error(count)),
}
}

let mut value = 0_u64;
while self.bit_position < end {
let byte = self.bytes[self.bit_position / 8];
let shift = 7 - (self.bit_position % 8);
value = (value << 1) | u64::from((byte >> shift) & 1);
self.bit_position += 1;
#[cold]
#[inline(never)]
fn end_error(&self, count: u8) -> EntropyError {
EntropyError::UnexpectedEnd {
bit_position: self.bit_position,
requested_bits: count,
bit_length: self.bit_length,
}
Ok(value)
}

/// Returns 64 bits starting at `bit_position`, most-significant-bit aligned.
///
/// Bytes beyond the backing slice read as zero. Bits between the packet
/// bit length and the end of its final byte are returned unchanged, so
/// callers must bound their use by `bits_remaining`.
#[inline]
fn window(&self) -> u64 {
let byte = self.bit_position / 8;
let word = match self.bytes.get(byte..byte + 8) {
Some(chunk) => {
u64::from_be_bytes(chunk.try_into().expect("slice is exactly eight bytes"))
}
None => Self::padded_word(self.bytes, byte),
};
word << (self.bit_position % 8)
}

#[cold]
#[inline(never)]
fn padded_word(bytes: &[u8], byte: usize) -> u64 {
let mut padded = [0_u8; 8];
let tail = bytes.get(byte..).unwrap_or_default();
padded[..tail.len()].copy_from_slice(tail);
u64::from_be_bytes(padded)
}
}

Expand Down Expand Up @@ -125,4 +199,61 @@ mod tests {
assert_eq!(reader.read_bits(0).unwrap(), 0);
assert_eq!(reader.bit_position(), 0);
}

/// Bit-serial reference for the word-at-a-time reader.
fn reference_bits(bytes: &[u8], start: usize, count: usize) -> u64 {
(start..start + count).fold(0, |value, position| {
(value << 1) | u64::from((bytes[position / 8] >> (7 - position % 8)) & 1)
})
}

#[test]
fn every_width_and_offset_matches_bit_serial_reads() {
let bytes: Vec<u8> = (0_u8..24)
.map(|index| index.wrapping_mul(157) ^ 0xa5)
.collect();
let bit_length = bytes.len() * 8 - 3;
for start in 0..bit_length {
for count in 0..=64_u8 {
let mut reader = PacketBitReader::with_bit_length(&bytes, bit_length).unwrap();
if start > 0 {
reader.read_bits(u8::try_from(start % 64).unwrap()).unwrap();
for _ in 0..start / 64 {
reader.read_bits(64).unwrap();
}
}
assert_eq!(reader.bit_position(), start);
let result = reader.read_bits(count);
if start + usize::from(count) > bit_length {
assert_eq!(
result,
Err(EntropyError::UnexpectedEnd {
bit_position: start,
requested_bits: count,
bit_length,
})
);
assert_eq!(reader.bit_position(), start);
} else {
assert_eq!(
result.unwrap(),
reference_bits(&bytes, start, usize::from(count)),
"start {start}, count {count}"
);
assert_eq!(reader.bit_position(), start + usize::from(count));
}
}
}
}

#[test]
fn peek_zero_fills_past_the_bounded_packet() {
let mut reader = PacketBitReader::with_bit_length(&[0xff, 0xff], 11).unwrap();
assert_eq!(reader.peek_byte(), 0xff);
reader.read_bits(5).unwrap();
assert_eq!(reader.peek_byte(), 0b1111_1100);
reader.read_bits(6).unwrap();
assert_eq!(reader.peek_byte(), 0);
assert_eq!(reader.bit_position(), 11);
}
}
20 changes: 12 additions & 8 deletions crates/jxr-native/src/entropy/coefficients.rs
Original file line number Diff line number Diff line change
Expand Up @@ -245,7 +245,7 @@ fn decode_abs_level(
let index = vlc::decode(
reader,
"ABS_LEVEL_INDEX",
vlc::ABS_LEVEL[state.table_index()],
&vlc::ABS_LEVEL[state.table_index()],
)?;
observe_abs_level(state, index);
if index < 6 {
Expand Down Expand Up @@ -279,7 +279,11 @@ fn decode_first_index(
reader: &mut PacketBitReader<'_>,
state: &mut AdaptiveVlc,
) -> Result<u8, EntropyError> {
let symbol = vlc::decode(reader, "FIRST_INDEX", vlc::FIRST_INDEX[state.table_index()])?;
let symbol = vlc::decode(
reader,
"FIRST_INDEX",
&vlc::FIRST_INDEX[state.table_index()],
)?;
observe_first_index(state, symbol);
Ok(symbol)
}
Expand All @@ -291,11 +295,11 @@ fn decode_index(
) -> Result<u8, EntropyError> {
match location.cmp(&15) {
core::cmp::Ordering::Less => {
let symbol = vlc::decode(reader, "INDEX_A", vlc::INDEX_A[state.table_index()])?;
let symbol = vlc::decode(reader, "INDEX_A", &vlc::INDEX_A[state.table_index()])?;
observe_index(state, symbol);
Ok(symbol)
}
core::cmp::Ordering::Equal => vlc::decode(reader, "INDEX_B", vlc::INDEX_B),
core::cmp::Ordering::Equal => vlc::decode(reader, "INDEX_B", &vlc::INDEX_B),
core::cmp::Ordering::Greater => Ok(u8::from(reader.read_bit()?)),
}
}
Expand All @@ -310,14 +314,14 @@ fn decode_run(reader: &mut PacketBitReader<'_>, max_run: u8) -> Result<u8, Entro
if max_run < 5 {
return match max_run {
1 => Ok(1),
2 => vlc::decode(reader, "RUN_VALUE", vlc::RUN_VALUE_2),
3 => vlc::decode(reader, "RUN_VALUE", vlc::RUN_VALUE_3),
4 => vlc::decode(reader, "RUN_VALUE", vlc::RUN_VALUE_4),
2 => vlc::decode(reader, "RUN_VALUE", &vlc::RUN_VALUE_2),
3 => vlc::decode(reader, "RUN_VALUE", &vlc::RUN_VALUE_3),
4 => vlc::decode(reader, "RUN_VALUE", &vlc::RUN_VALUE_4),
_ => unreachable!(),
};
}

let run_index = vlc::decode(reader, "RUN_INDEX", vlc::RUN_INDEX)?;
let run_index = vlc::decode(reader, "RUN_INDEX", &vlc::RUN_INDEX)?;
let table_index = i16::from(run_index) + 5 * i16::from(RUN_BIN[usize::from(max_run)]);
let table_index = usize::try_from(table_index).map_err(|_| EntropyError::InvalidParameter {
parameter: "run table index",
Expand Down
3 changes: 3 additions & 0 deletions crates/jxr-native/src/entropy/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,3 +20,6 @@ pub use refinement::{
decode_flex, decode_flex_block, decode_lp_refinement, decode_lp_refinement_at,
};
pub use scan::{AdaptiveHpScan, AdaptiveLpScan, HpScanDirection};
#[cfg(test)]
pub(crate) use vlc::tests::assert_matches_serial;
pub(crate) use vlc::{PrefixTable, decode as decode_prefix};
Loading
Loading