#![allow(non_snake_case)] #![allow(non_upper_case_globals)] use crate::bits::ZSTD_highbit32; use crate::errors::{ZstdErrorCode, ERROR}; use crate::mem::{MEM_readLEST, MEM_writeLEST}; use std::os::raw::c_void; pub type BitContainerType = usize; pub const STREAM_ACCUMULATOR_MIN_32: u32 = 25; pub const STREAM_ACCUMULATOR_MIN_64: u32 = 57; #[repr(C)] pub struct BIT_CStream_t { pub bitContainer: BitContainerType, pub bitPos: u32, pub startPtr: *mut i8, pub ptr: *mut i8, pub endPtr: *mut i8, } #[repr(C)] pub struct BIT_DStream_t { pub bitContainer: BitContainerType, pub bitsConsumed: u32, pub ptr: *const i8, pub start: *const i8, pub limitPtr: *const i8, } #[repr(C)] #[derive(Debug, Copy, Clone, PartialEq, Eq)] pub enum BIT_DStream_status { Unfinished = 0, EndOfBuffer = 1, Completed = 2, Overflow = 3, } const BIT_MASK: [usize; 32] = [ 0, 1, 3, 7, 0xF, 0x1F, 0x3F, 0x7F, 0xFF, 0x1FF, 0x3FF, 0x7FF, 0xFFF, 0x1FFF, 0x3FFF, 0x7FFF, 0xFFFF, 0x1FFFF, 0x3FFFF, 0x7FFFF, 0xFFFFF, 0x1FFFFF, 0x3FFFFF, 0x7FFFFF, 0xFFFFFF, 0x1FFFFFF, 0x3FFFFFF, 0x7FFFFFF, 0xFFFFFFF, 0x1FFFFFFF, 0x3FFFFFFF, 0x7FFFFFFF, ]; #[inline] pub fn BIT_getLowerBits(bitContainer: BitContainerType, nbBits: u32) -> BitContainerType { assert!(nbBits < 32); bitContainer & BIT_MASK[nbBits as usize] } pub unsafe fn BIT_initCStream( bitC: *mut BIT_CStream_t, startPtr: *mut c_void, dstCapacity: usize, ) -> usize { let s = &mut *bitC; s.bitContainer = 0; s.bitPos = 0; s.startPtr = startPtr as *mut i8; s.ptr = s.startPtr; if dstCapacity <= std::mem::size_of::() { s.endPtr = s.startPtr; return ERROR(ZstdErrorCode::DstSizeTooSmall); } s.endPtr = s .startPtr .add(dstCapacity - std::mem::size_of::()); 0 } #[inline] pub unsafe fn BIT_addBits(bitC: *mut BIT_CStream_t, value: BitContainerType, nbBits: u32) { let s = &mut *bitC; assert!(nbBits < 32); assert!(nbBits + s.bitPos < (std::mem::size_of::() * 8) as u32); s.bitContainer |= BIT_getLowerBits(value, nbBits) << s.bitPos; s.bitPos += nbBits; } #[inline] pub unsafe fn BIT_addBitsFast(bitC: *mut BIT_CStream_t, value: BitContainerType, nbBits: u32) { let s = &mut *bitC; assert!((value >> nbBits) == 0); assert!(nbBits + s.bitPos < (std::mem::size_of::() * 8) as u32); s.bitContainer |= value << s.bitPos; s.bitPos += nbBits; } #[inline] pub unsafe fn BIT_flushBitsFast(bitC: *mut BIT_CStream_t) { let s = &mut *bitC; let nbBytes = (s.bitPos >> 3) as usize; assert!(s.bitPos < (std::mem::size_of::() * 8) as u32); assert!(s.ptr <= s.endPtr); MEM_writeLEST(s.ptr as *mut c_void, s.bitContainer); s.ptr = s.ptr.add(nbBytes); s.bitPos &= 7; s.bitContainer >>= nbBytes * 8; } #[inline] pub unsafe fn BIT_flushBits(bitC: *mut BIT_CStream_t) { let s = &mut *bitC; let nbBytes = (s.bitPos >> 3) as usize; assert!(s.bitPos < (std::mem::size_of::() * 8) as u32); assert!(s.ptr <= s.endPtr); MEM_writeLEST(s.ptr as *mut c_void, s.bitContainer); let next_ptr = s.ptr.add(nbBytes); if next_ptr > s.endPtr { s.ptr = s.endPtr; } else { s.ptr = next_ptr; } s.bitPos &= 7; s.bitContainer >>= nbBytes * 8; } pub unsafe fn BIT_closeCStream(bitC: *mut BIT_CStream_t) -> usize { BIT_addBitsFast(bitC, 1, 1); BIT_flushBits(bitC); let s = &*bitC; if s.ptr >= s.endPtr { return 0; } (s.ptr as isize - s.startPtr as isize) as usize + (if s.bitPos > 0 { 1 } else { 0 }) } pub unsafe fn BIT_initDStream( bitD: *mut BIT_DStream_t, srcBuffer: *const c_void, srcSize: usize, ) -> usize { let s = &mut *bitD; if srcSize < 1 { std::ptr::write_bytes(bitD, 0, 1); return ERROR(ZstdErrorCode::SrcSizeWrong); } s.start = srcBuffer as *const i8; s.limitPtr = s.start.add(std::mem::size_of::()); if srcSize >= std::mem::size_of::() { s.ptr = s .start .add(srcSize) .offset(-(std::mem::size_of::() as isize)); s.bitContainer = MEM_readLEST(s.ptr as *const c_void); let lastByte = *s.start.offset(srcSize as isize - 1) as u8; s.bitsConsumed = if lastByte != 0 { 8 - ZSTD_highbit32(lastByte as u32) } else { 0 }; if lastByte == 0 { return ERROR(ZstdErrorCode::Generic); } } else { s.ptr = s.start; s.bitContainer = *(s.start as *const u8) as BitContainerType; let mut container = s.bitContainer; let src_bytes = std::slice::from_raw_parts(s.start as *const u8, srcSize); for (i, &byte) in src_bytes.iter().enumerate().skip(1) { container += (byte as BitContainerType) << (i * 8); } s.bitContainer = container; let lastByte = src_bytes[srcSize - 1]; s.bitsConsumed = if lastByte != 0 { 8 - ZSTD_highbit32(lastByte as u32) } else { 0 }; if lastByte == 0 { return ERROR(ZstdErrorCode::CorruptionDetected); } s.bitsConsumed += (std::mem::size_of::() as u32 - srcSize as u32) * 8; } srcSize } #[inline] pub fn BIT_getUpperBits(bitContainer: BitContainerType, start: u32) -> BitContainerType { bitContainer >> start } #[inline] pub fn BIT_getMiddleBits( bitContainer: BitContainerType, start: u32, nbBits: u32, ) -> BitContainerType { let regMask = (std::mem::size_of::() * 8 - 1) as u32; assert!(nbBits < 32); (bitContainer >> (start & regMask)) & ((1usize << nbBits) - 1) } #[inline] pub unsafe fn BIT_lookBits(bitD: *const BIT_DStream_t, nbBits: u32) -> BitContainerType { let s = &*bitD; let register_bits = (std::mem::size_of::() * 8) as u32; BIT_getMiddleBits( s.bitContainer, register_bits .wrapping_sub(s.bitsConsumed) .wrapping_sub(nbBits), nbBits, ) } #[inline] pub unsafe fn BIT_lookBitsFast(bitD: *const BIT_DStream_t, nbBits: u32) -> BitContainerType { let s = &*bitD; let regMask = (std::mem::size_of::() * 8 - 1) as u32; assert!(nbBits >= 1); (s.bitContainer << (s.bitsConsumed & regMask)) >> ((regMask + 1 - nbBits) & regMask) } #[inline] pub unsafe fn BIT_skipBits(bitD: *mut BIT_DStream_t, nbBits: u32) { let s = &mut *bitD; s.bitsConsumed += nbBits; } #[inline] pub unsafe fn BIT_readBits(bitD: *mut BIT_DStream_t, nbBits: u32) -> BitContainerType { let value = BIT_lookBits(bitD, nbBits); BIT_skipBits(bitD, nbBits); value } #[inline] pub unsafe fn BIT_readBitsFast(bitD: *mut BIT_DStream_t, nbBits: u32) -> BitContainerType { let value = BIT_lookBitsFast(bitD, nbBits); assert!(nbBits >= 1); BIT_skipBits(bitD, nbBits); value } #[inline] unsafe fn BIT_reloadDStream_internal(bitD: *mut BIT_DStream_t) -> BIT_DStream_status { let s = &mut *bitD; assert!(s.bitsConsumed <= (std::mem::size_of::() * 8) as u32); s.ptr = s.ptr.offset(-(s.bitsConsumed as isize / 8)); assert!(s.ptr >= s.start); s.bitsConsumed &= 7; s.bitContainer = MEM_readLEST(s.ptr as *const c_void); BIT_DStream_status::Unfinished } #[inline] pub unsafe fn BIT_reloadDStreamFast(bitD: *mut BIT_DStream_t) -> BIT_DStream_status { let s = &*bitD; if s.ptr < s.limitPtr { return BIT_DStream_status::Overflow; } BIT_reloadDStream_internal(bitD) } #[inline] pub unsafe fn BIT_reloadDStream(bitD: *mut BIT_DStream_t) -> BIT_DStream_status { let s = &mut *bitD; if s.bitsConsumed > (std::mem::size_of::() * 8) as u32 { static ZERO_FILLED: BitContainerType = 0; s.ptr = &ZERO_FILLED as *const _ as *const i8; return BIT_DStream_status::Overflow; } assert!(s.ptr >= s.start); if s.ptr >= s.limitPtr { return BIT_reloadDStream_internal(bitD); } if s.ptr == s.start { if s.bitsConsumed < (std::mem::size_of::() * 8) as u32 { return BIT_DStream_status::EndOfBuffer; } return BIT_DStream_status::Completed; } let nbBytes = s.bitsConsumed / 8; let mut result = BIT_DStream_status::Unfinished; let available = s.ptr.offset_from(s.start) as u32; let ptr_offset = if nbBytes > available { result = BIT_DStream_status::EndOfBuffer; available } else { nbBytes }; s.ptr = s.ptr.sub(ptr_offset as usize); s.bitsConsumed -= ptr_offset * 8; s.bitContainer = MEM_readLEST(s.ptr as *const c_void); result } #[inline] pub unsafe fn BIT_endOfDStream(DStream: *const BIT_DStream_t) -> u32 { let s = &*DStream; if s.ptr == s.start && s.bitsConsumed == (std::mem::size_of::() * 8) as u32 { 1 } else { 0 } } #[cfg(test)] mod tests { use super::*; use std::mem::{size_of, zeroed}; #[test] fn short_sources_are_loaded_little_endian_and_counted_in_bits() { let width = size_of::(); for src_size in 1..width { let mut source = vec![0u8; src_size]; for (index, byte) in source.iter_mut().enumerate() { *byte = (index as u8).wrapping_mul(37).wrapping_add(1); } source[src_size - 1] |= 0x80; let mut stream: BIT_DStream_t = unsafe { zeroed() }; let result = unsafe { BIT_initDStream(&mut stream, source.as_ptr().cast::(), source.len()) }; assert_eq!(result, src_size); let expected = source .iter() .enumerate() .fold(0usize, |value, (index, byte)| { value | (usize::from(*byte) << (index * 8)) }); assert_eq!(stream.bitContainer, expected); let last = u32::from(source[src_size - 1]); let expected_consumed = (width as u32 - src_size as u32) * 8 + 8 - (31 - last.leading_zeros()); assert_eq!(stream.bitsConsumed, expected_consumed); } } #[test] fn short_stream_round_trip_is_lifo() { let mut encoded = [0u8; 64]; let mut encoder: BIT_CStream_t = unsafe { zeroed() }; assert_eq!( unsafe { BIT_initCStream( &mut encoder, encoded.as_mut_ptr().cast::(), encoded.len(), ) }, 0 ); unsafe { BIT_addBits(&mut encoder, 0b101, 3); BIT_addBits(&mut encoder, 0b1_0010, 5); BIT_addBits(&mut encoder, 0b10_1010_1010, 10); } let encoded_size = unsafe { BIT_closeCStream(&mut encoder) }; assert_eq!(encoded_size, 3); let mut decoder: BIT_DStream_t = unsafe { zeroed() }; assert_eq!( unsafe { BIT_initDStream( &mut decoder, encoded.as_ptr().cast::(), encoded_size, ) }, encoded_size ); assert_eq!(unsafe { BIT_readBits(&mut decoder, 10) }, 0b10_1010_1010); assert_eq!(unsafe { BIT_readBits(&mut decoder, 5) }, 0b1_0010); assert_eq!(unsafe { BIT_readBits(&mut decoder, 3) }, 0b101); assert_eq!( unsafe { BIT_reloadDStream(&mut decoder) }, BIT_DStream_status::Completed ); assert_eq!(unsafe { BIT_endOfDStream(&decoder) }, 1); } #[test] fn reload_covers_fast_and_cautious_paths() { let width = size_of::(); let mut source = vec![0u8; width * 2]; for (index, byte) in source.iter_mut().enumerate() { *byte = index as u8 + 1; } *source.last_mut().unwrap() |= 0x80; let mut stream: BIT_DStream_t = unsafe { zeroed() }; assert_eq!( unsafe { BIT_initDStream(&mut stream, source.as_ptr().cast::(), source.len()) }, source.len() ); stream.bitsConsumed = (width * 8) as u32; assert_eq!( unsafe { BIT_reloadDStream(&mut stream) }, BIT_DStream_status::Unfinished ); assert_eq!(stream.ptr, stream.start); assert_eq!(stream.bitsConsumed, 0); source.truncate(width + 3); *source.last_mut().unwrap() |= 0x80; assert_eq!( unsafe { BIT_initDStream(&mut stream, source.as_ptr().cast::(), source.len()) }, source.len() ); stream.bitsConsumed = (width * 8) as u32; assert_eq!( unsafe { BIT_reloadDStream(&mut stream) }, BIT_DStream_status::EndOfBuffer ); assert_eq!(stream.ptr, stream.start); assert_eq!(stream.bitsConsumed, (width * 8 - 24) as u32); } #[test] fn invalid_sizes_and_missing_end_mark_report_c_errors() { let width = size_of::(); let mut encoder: BIT_CStream_t = unsafe { zeroed() }; let mut destination = vec![0u8; width]; assert_eq!( unsafe { BIT_initCStream( &mut encoder, destination.as_mut_ptr().cast::(), destination.len(), ) }, ERROR(ZstdErrorCode::DstSizeTooSmall) ); let mut decoder: BIT_DStream_t = unsafe { zeroed() }; assert_eq!( unsafe { BIT_initDStream(&mut decoder, std::ptr::null(), 0) }, ERROR(ZstdErrorCode::SrcSizeWrong) ); let zero = [0u8; 1]; assert_eq!( unsafe { BIT_initDStream(&mut decoder, zero.as_ptr().cast(), zero.len()) }, ERROR(ZstdErrorCode::CorruptionDetected) ); } }