#![allow(dead_code)]
use crate::error::{Error, Result, ScanRead, ScanResult};
#[derive(Debug)]
pub struct BitWriter {
buffer: Vec<u8>,
bit_buffer: u64,
bits_in_buffer: u8,
}
impl BitWriter {
#[must_use]
pub fn new() -> Self {
Self {
buffer: Vec::new(),
bit_buffer: 0,
bits_in_buffer: 0,
}
}
#[must_use]
pub fn with_capacity(capacity: usize) -> Self {
Self {
buffer: Vec::with_capacity(capacity),
bit_buffer: 0,
bits_in_buffer: 0,
}
}
#[inline(always)]
pub fn write_bits(&mut self, bits: u32, count: u8) {
debug_assert!(count <= 24);
debug_assert!(bits < (1 << count) || count == 0);
self.bit_buffer = (self.bit_buffer << count) | (bits as u64);
self.bits_in_buffer += count;
if self.bits_in_buffer >= 32 {
self.flush_bytes();
}
}
#[inline(always)]
pub fn write_code_and_extra(&mut self, code: u32, code_len: u8, extra: u16, extra_len: u8) {
debug_assert!(code_len <= 16);
debug_assert!(extra_len <= 16);
let total_len = code_len + extra_len;
debug_assert!(total_len <= 32);
let combined = ((code as u64) << extra_len) | (extra as u64);
self.bit_buffer = (self.bit_buffer << total_len) | combined;
self.bits_in_buffer += total_len;
if self.bits_in_buffer >= 32 {
self.flush_bytes();
}
}
#[inline(never)]
#[cold]
fn flush_bytes(&mut self) {
while self.bits_in_buffer >= 64 {
self.bits_in_buffer -= 64;
let word = self.bit_buffer; self.emit_8_bytes(word);
}
while self.bits_in_buffer >= 32 {
self.bits_in_buffer -= 32;
let word = (self.bit_buffer >> self.bits_in_buffer) as u32;
self.emit_4_bytes(word);
}
while self.bits_in_buffer >= 8 {
self.bits_in_buffer -= 8;
let byte = (self.bit_buffer >> self.bits_in_buffer) as u8;
self.emit_byte(byte);
}
}
#[inline(always)]
fn emit_byte(&mut self, byte: u8) {
self.buffer.push(byte);
if byte == 0xFF {
self.buffer.push(0x00);
}
}
#[inline(always)]
fn emit_4_bytes(&mut self, word: u32) {
if !has_byte_0xff_u32(word) {
self.buffer.extend_from_slice(&word.to_be_bytes());
} else {
self.emit_byte((word >> 24) as u8);
self.emit_byte((word >> 16) as u8);
self.emit_byte((word >> 8) as u8);
self.emit_byte(word as u8);
}
}
#[inline(always)]
fn emit_8_bytes(&mut self, word: u64) {
if !has_byte_0xff_u64(word) {
self.buffer.extend_from_slice(&word.to_be_bytes());
} else {
self.emit_byte((word >> 56) as u8);
self.emit_byte((word >> 48) as u8);
self.emit_byte((word >> 40) as u8);
self.emit_byte((word >> 32) as u8);
self.emit_byte((word >> 24) as u8);
self.emit_byte((word >> 16) as u8);
self.emit_byte((word >> 8) as u8);
self.emit_byte(word as u8);
}
}
#[inline]
pub fn write_byte_raw(&mut self, byte: u8) {
self.buffer.push(byte);
}
pub fn write_bytes_raw(&mut self, bytes: &[u8]) {
self.buffer.extend_from_slice(bytes);
}
#[inline]
pub fn write_u16_be(&mut self, value: u16) {
self.buffer.push((value >> 8) as u8);
self.buffer.push(value as u8);
}
pub fn flush(&mut self) {
while self.bits_in_buffer >= 8 {
self.bits_in_buffer -= 8;
let byte = (self.bit_buffer >> self.bits_in_buffer) as u8;
self.buffer.push(byte);
if byte == 0xFF {
self.buffer.push(0x00);
}
}
if self.bits_in_buffer > 0 {
let padding = 8 - self.bits_in_buffer;
let padded = (self.bit_buffer << padding) | ((1u64 << padding) - 1);
let byte = padded as u8;
self.buffer.push(byte);
if byte == 0xFF {
self.buffer.push(0x00);
}
self.bit_buffer = 0;
self.bits_in_buffer = 0;
}
}
#[must_use]
pub fn into_bytes(mut self) -> Vec<u8> {
self.flush();
self.buffer
}
#[must_use]
pub fn as_bytes(&self) -> &[u8] {
&self.buffer
}
#[must_use]
pub fn position(&self) -> usize {
self.buffer.len()
}
}
impl Default for BitWriter {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug)]
pub struct BitReader<'a> {
data: &'a [u8],
position: usize,
bit_buffer: u64,
aligned_buffer: u64,
bits_in_buffer: u8,
marker_found: Option<u8>,
overread_by: usize,
}
#[derive(Clone, Copy)]
pub struct BitReaderState {
position: usize,
bit_buffer: u64,
aligned_buffer: u64,
bits_in_buffer: u8,
marker_found: Option<u8>,
overread_by: usize,
}
#[inline(always)]
const fn has_byte_0xff_u32(v: u32) -> bool {
let x = v ^ 0xFFFF_FFFF;
(((x.wrapping_sub(0x0101_0101)) & !x) & 0x8080_8080) != 0
}
#[inline(always)]
const fn has_byte_0xff_u64(v: u64) -> bool {
let x = v ^ 0xFFFF_FFFF_FFFF_FFFF;
(((x.wrapping_sub(0x0101_0101_0101_0101)) & !x) & 0x8080_8080_8080_8080) != 0
}
impl<'a> BitReader<'a> {
#[must_use]
pub fn new(data: &'a [u8]) -> Self {
Self {
data,
position: 0,
bit_buffer: 0,
aligned_buffer: 0,
bits_in_buffer: 0,
marker_found: None,
overread_by: 0,
}
}
#[inline]
fn read_byte_slow(&mut self) -> Option<u8> {
if self.marker_found.is_some() {
return None;
}
if self.position >= self.data.len() {
self.overread_by += 1;
return None;
}
let byte = self.data[self.position];
self.position += 1;
if byte == 0xFF {
while self.position < self.data.len() && self.data[self.position] == 0xFF {
self.position += 1;
}
if self.position >= self.data.len() {
self.overread_by += 1;
return None;
}
let next = self.data[self.position];
if next == 0x00 {
self.position += 1;
} else {
self.position -= 1;
self.marker_found = Some(next);
return None;
}
}
Some(byte)
}
#[inline(always)]
fn sync_aligned(&mut self) {
self.aligned_buffer = if self.bits_in_buffer > 0 && self.bits_in_buffer < 64 {
self.bit_buffer << (64 - self.bits_in_buffer)
} else if self.bits_in_buffer == 64 {
self.bit_buffer
} else {
0
};
}
#[inline(always)]
pub fn refill(&mut self) -> Result<bool> {
if self.bits_in_buffer >= 32 {
return Ok(true);
}
if self.marker_found.is_some() || self.overread_by > 0 {
self.bit_buffer <<= 32;
self.bits_in_buffer = self.bits_in_buffer.saturating_add(32).min(64);
self.sync_aligned();
return Ok(true);
}
if self.position + 4 <= self.data.len() {
let bytes = [
self.data[self.position],
self.data[self.position + 1],
self.data[self.position + 2],
self.data[self.position + 3],
];
let word = u32::from_be_bytes(bytes);
if !has_byte_0xff_u32(word) {
self.position += 4;
self.bit_buffer = (self.bit_buffer << 32) | (word as u64);
self.bits_in_buffer += 32;
self.sync_aligned();
return Ok(true);
}
}
while self.bits_in_buffer <= 56 {
match self.read_byte_slow() {
Some(byte) => {
self.bit_buffer = (self.bit_buffer << 8) | (byte as u64);
self.bits_in_buffer += 8;
}
None => break,
}
if self.bits_in_buffer >= 32 {
break;
}
}
self.sync_aligned();
Ok(self.bits_in_buffer > 0)
}
#[inline]
fn fill_buffer(&mut self, count: u8) -> Result<bool> {
if self.bits_in_buffer < count {
self.refill()?;
}
Ok(self.bits_in_buffer >= count)
}
#[inline]
pub fn peek_bits(&mut self, count: u8) -> ScanResult<u32> {
debug_assert!(count <= 32);
self.fill_buffer(count)?;
if self.bits_in_buffer < count {
return Ok(self.end_state());
}
Ok(ScanRead::Value(
(self.aligned_buffer >> (64 - count)) as u32,
))
}
#[inline(always)]
pub fn peek_bits_refill(&mut self, count: u8) -> Option<u32> {
if self.bits_in_buffer < count {
let _ = self.refill();
if self.bits_in_buffer < count {
return None;
}
}
Some((self.aligned_buffer >> (64 - count)) as u32)
}
#[inline(always)]
pub fn skip_bits_fast(&mut self, count: u8) {
self.bits_in_buffer -= count;
self.aligned_buffer <<= count;
let mask = if self.bits_in_buffer >= 64 {
u64::MAX
} else {
(1u64 << self.bits_in_buffer).wrapping_sub(1)
};
self.bit_buffer &= mask;
}
#[inline(always)]
pub fn read_bits_fast(&mut self, count: u8) -> u32 {
let bits = (self.aligned_buffer >> (64 - count)) as u32;
self.bits_in_buffer -= count;
self.aligned_buffer <<= count;
let mask = if self.bits_in_buffer >= 64 {
u64::MAX
} else {
(1u64 << self.bits_in_buffer).wrapping_sub(1)
};
self.bit_buffer &= mask;
bits
}
#[inline(always)]
pub fn ensure_bits(&mut self) -> bool {
if self.bits_in_buffer < 32 {
let _ = self.refill();
}
self.bits_in_buffer >= 32
}
#[inline(always)]
pub fn peek_top(&self, count: u8) -> u32 {
(self.aligned_buffer >> (64 - count)) as u32
}
#[inline(always)]
pub fn get_bits_rotate(&mut self, n_bits: u8) -> i32 {
let mask = (1_u64 << n_bits) - 1;
self.aligned_buffer = self.aligned_buffer.rotate_left(u32::from(n_bits));
let bits = (self.aligned_buffer & mask) as i32;
self.bits_in_buffer = self.bits_in_buffer.wrapping_sub(n_bits);
bits
}
#[inline]
pub fn read_bits(&mut self, count: u8) -> ScanResult<u32> {
self.fill_buffer(count)?;
if self.bits_in_buffer < count {
return Ok(self.end_state());
}
let bits = (self.aligned_buffer >> (64 - count)) as u32;
self.drop_bits(count);
Ok(ScanRead::Value(bits))
}
#[inline(always)]
fn drop_bits(&mut self, count: u8) {
self.bits_in_buffer = self.bits_in_buffer.saturating_sub(count);
self.aligned_buffer <<= count;
let mask = if self.bits_in_buffer >= 64 {
u64::MAX
} else {
(1u64 << self.bits_in_buffer).wrapping_sub(1)
};
self.bit_buffer &= mask;
}
#[inline]
pub fn skip_bits(&mut self, count: u8) {
self.drop_bits(count);
}
#[inline]
pub fn read_bit(&mut self) -> ScanResult<bool> {
match self.read_bits(1)? {
ScanRead::Value(v) => Ok(ScanRead::Value(v != 0)),
ScanRead::EndOfScan => Ok(ScanRead::EndOfScan),
ScanRead::Truncated => Ok(ScanRead::Truncated),
}
}
pub fn read_signed(&mut self, bits: u8) -> ScanResult<i16> {
if bits == 0 {
return Ok(ScanRead::Value(0));
}
let value = match self.read_bits(bits)? {
ScanRead::Value(v) => v as i16,
ScanRead::EndOfScan => return Ok(ScanRead::EndOfScan),
ScanRead::Truncated => return Ok(ScanRead::Truncated),
};
let half = 1i16 << (bits - 1);
if value < half {
Ok(ScanRead::Value(value - (2 * half - 1)))
} else {
Ok(ScanRead::Value(value))
}
}
pub fn align_to_byte(&mut self) {
self.bits_in_buffer = 0;
self.aligned_buffer = 0;
}
#[must_use]
pub fn save_state(&self) -> BitReaderState {
BitReaderState {
position: self.position,
bit_buffer: self.bit_buffer,
aligned_buffer: self.aligned_buffer,
bits_in_buffer: self.bits_in_buffer,
marker_found: self.marker_found,
overread_by: self.overread_by,
}
}
pub fn restore_state(&mut self, state: BitReaderState) {
self.position = state.position;
self.bit_buffer = state.bit_buffer;
self.aligned_buffer = state.aligned_buffer;
self.bits_in_buffer = state.bits_in_buffer;
self.marker_found = state.marker_found;
self.overread_by = state.overread_by;
}
pub fn read_restart_marker(&mut self, expected_num: u8) -> Result<()> {
self.marker_found = None;
if self.position >= self.data.len() {
return Err(Error::invalid_jpeg_data(
"unexpected end of data before restart marker",
));
}
let first = self.data[self.position];
if first != 0xFF {
return Err(Error::invalid_jpeg_data("expected 0xFF for restart marker"));
}
self.position += 1;
if self.position >= self.data.len() {
return Err(Error::invalid_jpeg_data(
"unexpected end of data in restart marker",
));
}
let second = self.data[self.position];
let expected_marker = 0xD0 + (expected_num & 7);
if second != expected_marker {
if (0xD0..=0xD7).contains(&second) {
return Err(Error::invalid_jpeg_data("restart marker sequence mismatch"));
}
return Err(Error::invalid_jpeg_data(
"expected restart marker not found",
));
}
self.position += 1;
Ok(())
}
pub fn read_byte_raw(&mut self) -> Result<u8> {
if self.position >= self.data.len() {
return Err(Error::truncated_data("reading raw byte"));
}
let byte = self.data[self.position];
self.position += 1;
Ok(byte)
}
pub fn read_u16_be(&mut self) -> Result<u16> {
let high = self.read_byte_raw()? as u16;
let low = self.read_byte_raw()? as u16;
Ok((high << 8) | low)
}
#[must_use]
pub fn marker_found(&self) -> Option<u8> {
self.marker_found
}
#[must_use]
pub fn position(&self) -> usize {
self.position
}
#[must_use]
pub fn remaining(&self) -> usize {
self.data.len().saturating_sub(self.position)
}
#[must_use]
pub fn is_exhausted(&self) -> bool {
self.marker_found.is_some() || self.position >= self.data.len()
}
#[must_use]
pub fn bits_available(&self) -> u8 {
self.bits_in_buffer
}
#[inline]
fn end_state<T>(&self) -> ScanRead<T> {
if self.marker_found.is_some() {
ScanRead::EndOfScan
} else {
ScanRead::Truncated
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_write_read_bits() {
let mut writer = BitWriter::new();
writer.write_bits(0b101, 3);
writer.write_bits(0b1100, 4);
writer.write_bits(0b1, 1);
let bytes = writer.into_bytes();
let mut reader = BitReader::new(&bytes);
assert_eq!(reader.read_bits(3).unwrap(), ScanRead::Value(0b101));
assert_eq!(reader.read_bits(4).unwrap(), ScanRead::Value(0b1100));
assert_eq!(reader.read_bits(1).unwrap(), ScanRead::Value(0b1));
}
#[test]
fn test_byte_stuffing() {
let mut writer = BitWriter::new();
writer.write_bits(0xFF, 8);
let bytes = writer.into_bytes();
assert_eq!(bytes[0], 0xFF);
assert_eq!(bytes[1], 0x00);
}
#[test]
fn test_byte_unstuffing() {
let data = [0xFF, 0x00, 0xAB];
let mut reader = BitReader::new(&data);
assert_eq!(reader.read_bits(8).unwrap(), ScanRead::Value(0xFF));
assert_eq!(reader.read_bits(8).unwrap(), ScanRead::Value(0xAB));
}
#[test]
fn test_signed_values() {
let data = [0b0100_0000]; let mut reader = BitReader::new(&data);
assert_eq!(reader.read_signed(1).unwrap(), ScanRead::Value(-1));
assert_eq!(reader.read_signed(1).unwrap(), ScanRead::Value(1));
}
}