Files
zstd-rs/rust/src/bitstream.rs
T
ddidderr 089f8e5b4d feat(rust): add common primitive compatibility layer
Introduce the Rust static library and move the shared byte, bitstream, CPU,
error, xxHash, debug, and public-common implementations into it. Thin C shims
preserve the existing header-driven C build while original tests link the Rust
archive.

This establishes the ABI-safe foundation for later codec and CLI ports;
entropy coding, runtime support, codecs, dictionaries, and the CLI remain C.
The top-down migration map documents that boundary and its validation path.

Test Plan:
- cargo fmt --check
- cargo test --all-targets
- cargo clippy --all-targets -- -D warnings
- cargo build --release

Refs: rust/README.md
2026-07-10 20:02:17 +02:00

457 lines
14 KiB
Rust

#![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::<BitContainerType>() {
s.endPtr = s.startPtr;
return ERROR(ZstdErrorCode::DstSizeTooSmall);
}
s.endPtr = s
.startPtr
.add(dstCapacity - std::mem::size_of::<BitContainerType>());
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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>());
if srcSize >= std::mem::size_of::<BitContainerType>() {
s.ptr = s
.start
.add(srcSize)
.offset(-(std::mem::size_of::<BitContainerType>() 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::<BitContainerType>() 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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>() * 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::<BitContainerType>();
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::<c_void>(), 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::<c_void>(),
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::<c_void>(),
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::<BitContainerType>();
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::<c_void>(), 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::<c_void>(), 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::<BitContainerType>();
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::<c_void>(),
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)
);
}
}