feat(dict): port the public zdict builder to Rust

Move the public ZDICT helpers, legacy trainer, dictionary finalization,
entropy-table construction, and supporting dictionary logic into Rust. Keep
the C translation unit as an ABI anchor and register the Rust module only
when compression and dictionary-builder features are enabled.

Test Plan:
- RUSTC_WRAPPER= CARGO_BUILD_RUSTC_WRAPPER= cargo test --manifest-path rust/Cargo.toml dict_builder_zdict
- RUSTC_WRAPPER= CARGO_BUILD_RUSTC_WRAPPER= cargo clippy --manifest-path rust/Cargo.toml --all-targets -- -D warnings
- rustfmt +nightly --check --edition 2021 rust/src/dict_builder_zdict.rs rust/src/lib.rs
- git diff --cached --check
This commit is contained in:
2026-07-12 10:38:41 +02:00
parent 63fdbc3fc2
commit 6a9b2d30b5
3 changed files with 1319 additions and 1120 deletions
+1310
View File
@@ -0,0 +1,1310 @@
#![allow(non_camel_case_types)]
#![allow(non_snake_case)]
#![allow(clippy::missing_safety_doc)]
#![allow(clippy::too_many_arguments)]
//! Public dictionary-builder wrappers.
//!
//! This is the Rust translation of `lib/dictBuilder/zdict.c`. The COVER and
//! fastCOVER implementations remain separate translation units, but their
//! finalization calls, the legacy trainer, entropy-header construction, and
//! public helper functions all live here. Compression contexts remain opaque:
//! the statistics pass uses the existing C context API and only projects the
//! stable `SeqStore_t` leaf into Rust.
use crate::bits::ZSTD_highbit32;
use crate::common::{LL_FSE_LOG, MAX_LL, MAX_ML, ML_FSE_LOG, OFF_FSE_LOG, ZSTD_REP_NUM};
use crate::divsufsort::divsufsort;
use crate::errors::{ERR_getErrorName, ERR_isError, ZstdErrorCode, ERROR};
use crate::fse_compress::{FSE_normalizeCount, FSE_writeNCount};
use crate::huf_compress::{HUF_buildCTable_wksp, HUF_writeCTable_wksp};
use crate::mem::{MEM_readLE32, MEM_writeLE32};
use crate::xxhash::XXH64;
use crate::zstd_compress_params::{ZSTD_compressionParameters, ZSTD_parameters};
use crate::zstd_compress_sequences::SeqDef;
use crate::zstd_compress_stats::{SeqStore_t, ZSTD_compressedBlockState_t, ZSTD_seqToCodes};
use std::ffi::{c_char, c_void};
use std::mem::{size_of, MaybeUninit};
use std::os::raw::{c_int, c_uint};
use std::ptr;
const ZSTD_MAGIC_DICTIONARY: u32 = 0xEC30_A437;
const ZSTD_CLEVEL_DEFAULT: c_int = 3;
const ZSTD_BLOCKSIZE_MAX: usize = 1 << 17;
const HUF_WORKSPACE_SIZE: usize = (8 << 10) + 512;
const HUF_CTABLE_WORKSPACE_SIZE_U32: usize = 4 * (255 + 1) + 192;
const ZDICT_DICTSIZE_MIN: usize = 256;
const ZDICT_CONTENTSIZE_MIN: usize = 128;
const ZDICT_MAX_SAMPLES_SIZE: usize = 2000 << 20;
const ZDICT_MIN_SAMPLES_SIZE: usize = ZDICT_CONTENTSIZE_MIN * 4;
const DICTLISTSIZE_DEFAULT: usize = 10_000;
const NOISELENGTH: usize = 32;
const MINRATIO: usize = 4;
const LLIMIT: usize = 64;
const MINMATCHLENGTH: usize = 7;
const MAXREPOFFSET: usize = 1024;
const OFFCODE_MAX: usize = 30;
const ZSTD_DLM_BY_REF: c_int = 1;
const ZSTD_DCT_RAW_CONTENT: c_int = 1;
type ZstdAllocFunction = unsafe extern "C" fn(*mut c_void, usize) -> *mut c_void;
type ZstdFreeFunction = unsafe extern "C" fn(*mut c_void, *mut c_void);
/// ABI-compatible `ZDICT_params_t` from `zdict.h`.
#[repr(C)]
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct ZDICT_params_t {
pub compressionLevel: c_int,
pub notificationLevel: c_uint,
pub dictID: c_uint,
}
/// ABI-compatible `ZDICT_legacy_params_t` from `zdict.h`.
#[repr(C)]
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct ZDICT_legacy_params_t {
pub selectivityLevel: c_uint,
pub zParams: ZDICT_params_t,
}
/// ABI-compatible fastCOVER parameters used by `ZDICT_trainFromBuffer()`.
///
/// The type is local to this module because the public fastCOVER declaration
/// is still provided by the C header and implementation.
#[repr(C)]
#[derive(Clone, Copy, Debug, Default)]
struct ZDICT_fastCover_params_t {
k: c_uint,
d: c_uint,
f: c_uint,
steps: c_uint,
nbThreads: c_uint,
splitPoint: f64,
accel: c_uint,
shrinkDict: c_uint,
shrinkDictMaxRegression: c_uint,
zParams: ZDICT_params_t,
}
type ZSTD_CCtx = c_void;
type ZSTD_CDict = c_void;
#[repr(C)]
#[derive(Clone, Copy)]
struct ZSTD_customMem {
customAlloc: Option<ZstdAllocFunction>,
customFree: Option<ZstdFreeFunction>,
opaque: *mut c_void,
}
const ZSTD_DEFAULT_CMEM: ZSTD_customMem = ZSTD_customMem {
customAlloc: None,
customFree: None,
opaque: ptr::null_mut(),
};
unsafe extern "C" {
fn ZSTD_loadCEntropy(
bs: *mut ZSTD_compressedBlockState_t,
workspace: *mut c_void,
dict: *const c_void,
dict_size: usize,
) -> usize;
fn ZSTD_reset_compressedBlockState(bs: *mut ZSTD_compressedBlockState_t);
fn ZSTD_getParams(
compression_level: c_int,
estimated_src_size: u64,
dict_size: usize,
) -> ZSTD_parameters;
fn ZSTD_createCDict_advanced(
dict: *const c_void,
dict_size: usize,
dict_load_method: c_int,
dict_content_type: c_int,
c_params: ZSTD_compressionParameters,
custom_mem: ZSTD_customMem,
) -> *mut ZSTD_CDict;
fn ZSTD_freeCDict(cdict: *mut ZSTD_CDict) -> usize;
fn ZSTD_createCCtx() -> *mut ZSTD_CCtx;
fn ZSTD_freeCCtx(cctx: *mut ZSTD_CCtx) -> usize;
fn ZSTD_compressBegin_usingCDict_deprecated(
cctx: *mut ZSTD_CCtx,
cdict: *const ZSTD_CDict,
) -> usize;
fn ZSTD_compressBlock_deprecated(
cctx: *mut ZSTD_CCtx,
dst: *mut c_void,
dst_capacity: usize,
src: *const c_void,
src_size: usize,
) -> usize;
fn ZSTD_getSeqStore(cctx: *const ZSTD_CCtx) -> *const SeqStore_t;
fn ZDICT_optimizeTrainFromBuffer_fastCover(
dict_buffer: *mut c_void,
dict_buffer_capacity: usize,
samples_buffer: *const c_void,
samples_sizes: *const usize,
nb_samples: c_uint,
parameters: *mut ZDICT_fastCover_params_t,
) -> usize;
}
#[inline]
fn dictionary_error(code: ZstdErrorCode) -> usize {
ERROR(code)
}
#[inline]
fn read_u16(data: &[u8], offset: usize) -> u16 {
u16::from_le_bytes([data[offset], data[offset + 1]])
}
#[inline]
fn count_common(data: &[u8], input: usize, matched: usize) -> usize {
let start = input;
let mut input = input;
let mut matched = matched;
while input < data.len() && matched < data.len() && data[input] == data[matched] {
input += 1;
matched += 1;
}
input - start
}
#[inline]
fn suffix_at(suffix0: &[i32], c_index: isize) -> usize {
suffix0[(c_index + 1) as usize] as usize
}
#[derive(Clone, Copy, Debug, Default)]
struct DictItem {
pos: u32,
length: u32,
savings: u32,
}
#[inline]
fn init_dict_item(item: &mut DictItem) {
item.pos = 1;
item.length = 0;
item.savings = u32::MAX;
}
fn analyze_pos(
done_marks: &mut [u8],
suffix0: &[i32],
start: usize,
data: &[u8],
min_ratio: usize,
_notification_level: c_uint,
) -> DictItem {
let mut length_list = [0u32; LLIMIT];
let mut cumulative_length = [0u32; LLIMIT];
let mut savings = [0u32; LLIMIT];
let mut pos = suffix_at(suffix0, start as isize);
let mut end = start;
let mut solution = DictItem::default();
done_marks[pos] = 1;
if read_u16(data, pos) == read_u16(data, pos + 2)
|| read_u16(data, pos + 1) == read_u16(data, pos + 3)
|| read_u16(data, pos + 2) == read_u16(data, pos + 4)
{
let pattern = read_u16(data, pos + 4);
let mut pattern_end = 6;
while read_u16(data, pos + pattern_end) == pattern {
pattern_end += 2;
}
if data[pos + pattern_end] == data[pos + pattern_end - 1] {
pattern_end += 1;
}
for offset in 1..pattern_end {
done_marks[pos + offset] = 1;
}
return solution;
}
loop {
end += 1;
let length = count_common(data, pos, suffix_at(suffix0, end as isize));
if length < MINMATCHLENGTH {
break;
}
}
let mut start = start;
loop {
let length = count_common(data, pos, suffix_at(suffix0, start as isize - 1));
if length < MINMATCHLENGTH {
break;
}
start -= 1;
}
if end - start < min_ratio {
for index in start..end {
done_marks[suffix_at(suffix0, index as isize)] = 1;
}
return solution;
}
let mut refined_start = start;
let mut refined_end = end;
let mut mml = MINMATCHLENGTH;
loop {
let mut current_char = 0u8;
let mut current_count = 0usize;
let mut current_id = refined_start;
let mut selected_count = 0usize;
let mut selected_id = current_id;
for id in refined_start..refined_end {
let byte = data[suffix_at(suffix0, id as isize) + mml];
if byte != current_char {
if current_count > selected_count {
selected_count = current_count;
selected_id = current_id;
}
current_id = id;
current_char = byte;
current_count = 0;
}
current_count += 1;
}
if current_count > selected_count {
selected_count = current_count;
selected_id = current_id;
}
if selected_count < min_ratio {
break;
}
refined_start = selected_id;
refined_end = refined_start + selected_count;
mml += 1;
}
start = refined_start;
pos = suffix_at(suffix0, refined_start as isize);
end = start;
loop {
end += 1;
let original_length = count_common(data, pos, suffix_at(suffix0, end as isize));
let length = original_length.min(LLIMIT - 1);
length_list[length] += 1;
if original_length < MINMATCHLENGTH {
break;
}
}
let mut length = MINMATCHLENGTH;
while length >= MINMATCHLENGTH && start > 0 {
let original_length = count_common(data, pos, suffix_at(suffix0, start as isize - 1));
length = original_length.min(LLIMIT - 1);
length_list[length] += 1;
if original_length >= MINMATCHLENGTH {
start -= 1;
}
}
cumulative_length[LLIMIT - 1] = length_list[LLIMIT - 1];
for index in (0..LLIMIT - 1).rev() {
cumulative_length[index] = cumulative_length[index + 1] + length_list[index];
}
let mut useful_length = MINMATCHLENGTH - 1;
for index in (MINMATCHLENGTH..LLIMIT).rev() {
if cumulative_length[index] >= min_ratio as u32 {
useful_length = index;
break;
}
}
let mut max_length = useful_length;
let repeated = data[pos + max_length - 1];
let mut reduced = max_length as u32;
while data[pos + reduced as usize - 2] == repeated {
reduced -= 1;
}
max_length = reduced as usize;
if max_length < MINMATCHLENGTH {
return solution;
}
savings[5] = 0;
for index in MINMATCHLENGTH..=max_length {
savings[index] =
savings[index - 1].wrapping_add(length_list[index].wrapping_mul((index - 3) as u32));
}
solution.pos = pos as u32;
solution.length = max_length as u32;
solution.savings = savings[max_length];
for index in start..end {
let tested_pos = suffix_at(suffix0, index as isize);
let length = if tested_pos == pos {
solution.length as usize
} else {
count_common(data, pos, tested_pos).min(solution.length as usize)
};
for mark in done_marks.iter_mut().skip(tested_pos).take(length) {
*mark = 1;
}
}
solution
}
#[inline]
fn is_included(data: &[u8], input: usize, container: usize, length: usize) -> bool {
data[input..input + length] == data[container..container + length]
}
fn resort_item(table: &mut [DictItem], mut index: usize) {
let item = table[index];
while index > 1 && table[index - 1].savings < item.savings {
table[index] = table[index - 1];
index -= 1;
}
table[index] = item;
}
fn try_merge(table: &mut [DictItem], elt: DictItem, skip: usize, data: &[u8]) -> usize {
let table_size = table[0].pos as usize;
let elt_end = elt.pos as usize + elt.length as usize;
for index in 1..table_size {
if index == skip {
continue;
}
let item_pos = table[index].pos as usize;
if item_pos > elt.pos as usize && item_pos <= elt_end {
let added = item_pos - elt.pos as usize;
table[index].length = table[index].length.wrapping_add(added as u32);
table[index].pos = elt.pos;
table[index].savings = table[index]
.savings
.wrapping_add(elt.savings.wrapping_mul(added as u32) / elt.length)
.wrapping_add(elt.length / 8);
resort_item(table, index);
return index;
}
}
for index in 1..table_size {
if index == skip {
continue;
}
let item_pos = table[index].pos as usize;
let item_end = item_pos + table[index].length as usize;
if item_end >= elt.pos as usize && item_pos < elt.pos as usize {
let added = elt_end as isize - item_end as isize;
table[index].savings = table[index].savings.wrapping_add(elt.length / 8);
if added > 0 {
table[index].length = table[index].length.wrapping_add(added as u32);
table[index].savings = table[index]
.savings
.wrapping_add(elt.savings.wrapping_mul(added as u32) / elt.length);
}
resort_item(table, index);
return index;
}
let left = item_pos;
let right = elt.pos as usize + 1;
if left + 8 <= data.len()
&& right + 8 <= data.len()
&& data[left..left + 8] == data[right..right + 8]
&& is_included(data, left, right, table[index].length as usize)
{
let added = ((elt.length as isize - table[index].length as isize).max(1)) as usize;
table[index].pos = elt.pos;
table[index].savings = table[index]
.savings
.wrapping_add(elt.savings.wrapping_mul(added as u32) / elt.length);
table[index].length =
(elt.length as usize).min(table[index].length as usize + 1) as u32;
return index;
}
}
0
}
fn remove_dict_item(table: &mut [DictItem], id: usize) {
if id == 0 {
return;
}
let max = table[0].pos as usize;
for index in id..max - 1 {
table[index] = table[index + 1];
}
table[0].pos -= 1;
}
fn insert_dict_item(table: &mut [DictItem], max_size: usize, elt: DictItem, data: &[u8]) {
let merge_id = try_merge(table, elt, 0, data);
if merge_id != 0 {
let mut merge = merge_id;
while merge != 0 {
let next = try_merge(table, table[merge], merge, data);
if next != 0 {
remove_dict_item(table, merge);
}
merge = next;
}
return;
}
let mut next_elt = table[0].pos as usize;
if next_elt >= max_size {
next_elt = max_size - 1;
}
let mut current = next_elt - 1;
while table[current].savings < elt.savings {
table[current + 1] = table[current];
current -= 1;
}
table[current + 1] = elt;
table[0].pos = (next_elt + 1) as u32;
}
fn dict_size(table: &[DictItem]) -> usize {
(1..table[0].pos as usize)
.map(|index| table[index].length as usize)
.sum()
}
fn fill_noise(buffer: &mut [u8]) {
let prime2 = 2_246_822_519u32;
let mut accumulator = 2_654_435_761u32;
for byte in buffer {
accumulator = accumulator.wrapping_mul(prime2);
*byte = (accumulator >> 21) as u8;
}
}
fn total_sample_size(file_sizes: &[usize]) -> usize {
file_sizes
.iter()
.fold(0usize, |total, size| total.wrapping_add(*size))
}
fn train_buffer_legacy(
dict_list: &mut [DictItem],
buffer: &[u8],
file_sizes: &[usize],
min_ratio: usize,
notification_level: c_uint,
) -> usize {
let buffer_size = buffer.len() - NOISELENGTH;
let mut suffix0 = vec![0i32; buffer_size.saturating_add(2)];
let mut reverse_suffix = vec![0u32; buffer_size];
let mut done_marks = vec![0u8; buffer_size.saturating_add(16)];
let mut file_pos = vec![0u32; file_sizes.len()];
let mut effective_buffer_size = buffer_size;
let mut effective_files = file_sizes.len();
while effective_buffer_size > ZDICT_MAX_SAMPLES_SIZE {
if effective_files == 0 {
break;
}
effective_files -= 1;
effective_buffer_size -= file_sizes[effective_files];
}
if effective_buffer_size > ZDICT_MAX_SAMPLES_SIZE {
eprintln!(
"sample set too large : reduced to {} MB ...",
ZDICT_MAX_SAMPLES_SIZE >> 20
);
}
if effective_files > 0 {
for index in 1..effective_files {
file_pos[index] = file_pos[index - 1].wrapping_add(file_sizes[index - 1] as u32);
}
}
let result = unsafe {
divsufsort(
buffer.as_ptr(),
suffix0.as_mut_ptr().add(1),
effective_buffer_size as c_int,
0,
)
};
if result != 0 {
return dictionary_error(ZstdErrorCode::Generic);
}
suffix0[0] = effective_buffer_size as i32;
suffix0[effective_buffer_size + 1] = effective_buffer_size as i32;
for position in 0..effective_buffer_size {
reverse_suffix[suffix_at(&suffix0, position as isize)] = position as u32;
}
done_marks.fill(0);
let mut cursor = 0usize;
while cursor < effective_buffer_size {
if done_marks[cursor] != 0 {
cursor += 1;
continue;
}
let solution = analyze_pos(
&mut done_marks,
&suffix0,
reverse_suffix[cursor] as usize,
buffer,
min_ratio.max(MINRATIO),
notification_level,
);
if solution.length == 0 {
cursor += 1;
continue;
}
insert_dict_item(dict_list, dict_list.len(), solution, buffer);
cursor += solution.length as usize;
}
0
}
fn count_entropy_stats(
cdict: *mut ZSTD_CDict,
cctx: *mut ZSTD_CCtx,
workplace: &mut [u8],
params: &ZSTD_parameters,
count_lit: &mut [u32; 256],
offset_counts: &mut [u32; OFFCODE_MAX + 1],
match_counts: &mut [u32; MAX_ML + 1],
lit_counts: &mut [u32; MAX_LL + 1],
rep_offsets: &mut [u32; MAXREPOFFSET],
src: *const c_void,
mut src_size: usize,
) {
let block_size_max = ZSTD_BLOCKSIZE_MAX.min(1usize << params.cParams.windowLog);
src_size = src_size.min(block_size_max);
let begin = unsafe { ZSTD_compressBegin_usingCDict_deprecated(cctx, cdict) };
if ERR_isError(begin) {
return;
}
let compressed = unsafe {
ZSTD_compressBlock_deprecated(
cctx,
workplace.as_mut_ptr().cast(),
workplace.len(),
src,
src_size,
)
};
if ERR_isError(compressed) || compressed == 0 {
return;
}
let store = unsafe { &*ZSTD_getSeqStore(cctx) };
let literal_count = unsafe { store.lit.offset_from(store.litStart) as usize };
for byte in unsafe { std::slice::from_raw_parts(store.litStart, literal_count) } {
count_lit[*byte as usize] += 1;
}
let nb_sequences = unsafe { store.sequences.offset_from(store.sequencesStart) as usize };
unsafe { ZSTD_seqToCodes(store) };
for index in 0..nb_sequences {
let of_code = unsafe { *store.ofCode.add(index) as usize };
let ml_code = unsafe { *store.mlCode.add(index) as usize };
let ll_code = unsafe { *store.llCode.add(index) as usize };
if of_code <= OFFCODE_MAX {
offset_counts[of_code] += 1;
}
if ml_code <= MAX_ML {
match_counts[ml_code] += 1;
}
if ll_code <= MAX_LL {
lit_counts[ll_code] += 1;
}
}
if nb_sequences >= 2 {
let first = unsafe { &*store.sequencesStart.cast::<SeqDef>() };
let second = unsafe { &*store.sequencesStart.add(1).cast::<SeqDef>() };
let offset1 = first.offBase.wrapping_sub(ZSTD_REP_NUM as u32);
let offset2 = second.offBase.wrapping_sub(ZSTD_REP_NUM as u32);
rep_offsets[if offset1 < MAXREPOFFSET as u32 {
offset1 as usize
} else {
0
}] += 3;
rep_offsets[if offset2 < MAXREPOFFSET as u32 {
offset2 as usize
} else {
0
}] += 1;
}
}
unsafe fn analyze_entropy(
dst_buffer: *mut u8,
max_dst_size: usize,
compression_level: c_int,
src_buffer: *const u8,
file_sizes: *const usize,
nb_files: c_uint,
dict_buffer: *const u8,
dict_buffer_size: usize,
_notification_level: c_uint,
) -> usize {
let offcode_max = ZSTD_highbit32((dict_buffer_size + (128 << 10)) as u32) as usize;
if offcode_max > OFFCODE_MAX {
return dictionary_error(ZstdErrorCode::DictionaryCreationFailed);
}
let file_sizes = if nb_files == 0 {
&[][..]
} else {
unsafe { std::slice::from_raw_parts(file_sizes, nb_files as usize) }
};
let total_src_size = total_sample_size(file_sizes);
let average_sample_size = total_src_size / (nb_files as usize + usize::from(nb_files == 0));
let mut count_lit = [1u32; 256];
let mut offset_counts = [0u32; OFFCODE_MAX + 1];
offset_counts[..=offcode_max].fill(1);
let mut match_counts = [1u32; MAX_ML + 1];
let mut lit_counts = [1u32; MAX_LL + 1];
let mut rep_offsets = [0u32; MAXREPOFFSET];
rep_offsets[1] = 1;
rep_offsets[4] = 1;
rep_offsets[8] = 1;
let level = if compression_level == 0 {
ZSTD_CLEVEL_DEFAULT
} else {
compression_level
};
let params = unsafe { ZSTD_getParams(level, average_sample_size as u64, dict_buffer_size) };
let cdict = unsafe {
ZSTD_createCDict_advanced(
dict_buffer.cast(),
dict_buffer_size,
ZSTD_DLM_BY_REF,
ZSTD_DCT_RAW_CONTENT,
params.cParams,
ZSTD_DEFAULT_CMEM,
)
};
let cctx = unsafe { ZSTD_createCCtx() };
let mut workplace = vec![0u8; ZSTD_BLOCKSIZE_MAX];
if cdict.is_null() || cctx.is_null() {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return dictionary_error(ZstdErrorCode::MemoryAllocation);
}
let mut source_offset = 0usize;
for &sample_size in file_sizes {
count_entropy_stats(
cdict,
cctx,
&mut workplace,
&params,
&mut count_lit,
&mut offset_counts,
&mut match_counts,
&mut lit_counts,
&mut rep_offsets,
unsafe { src_buffer.add(source_offset).cast() },
sample_size,
);
source_offset = source_offset.wrapping_add(sample_size);
}
let mut huf_table = [0usize; 257];
let mut huf_workspace = [0u32; HUF_CTABLE_WORKSPACE_SIZE_U32];
let mut huff_log = 11u32;
let mut written = unsafe {
HUF_buildCTable_wksp(
huf_table.as_mut_ptr(),
count_lit.as_ptr(),
255,
huff_log,
huf_workspace.as_mut_ptr().cast(),
size_of_val(&huf_workspace),
)
};
if ERR_isError(written) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return written;
}
if written == 8 {
for count in count_lit.iter_mut().skip(1) {
*count = 2;
}
count_lit[0] = 4;
count_lit[253] = 1;
count_lit[254] = 1;
written = unsafe {
HUF_buildCTable_wksp(
huf_table.as_mut_ptr(),
count_lit.as_ptr(),
255,
huff_log,
huf_workspace.as_mut_ptr().cast(),
size_of_val(&huf_workspace),
)
};
if ERR_isError(written) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return written;
}
}
huff_log = written as u32;
let mut offcode_ncount = [0i16; OFFCODE_MAX + 1];
let mut match_ncount = [0i16; MAX_ML + 1];
let mut lit_ncount = [0i16; MAX_LL + 1];
let total = offset_counts[..=offcode_max]
.iter()
.fold(0usize, |sum, count| sum + *count as usize);
let mut off_log = OFF_FSE_LOG as u32;
let mut match_log = ML_FSE_LOG as u32;
let mut lit_log = LL_FSE_LOG as u32;
let normalized = unsafe {
FSE_normalizeCount(
offcode_ncount.as_mut_ptr(),
off_log,
offset_counts.as_ptr(),
total,
offcode_max as u32,
1,
)
};
if ERR_isError(normalized) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return normalized;
}
off_log = normalized as u32;
let total = match_counts
.iter()
.fold(0usize, |sum, count| sum + *count as usize);
let normalized = unsafe {
FSE_normalizeCount(
match_ncount.as_mut_ptr(),
match_log,
match_counts.as_ptr(),
total,
MAX_ML as u32,
1,
)
};
if ERR_isError(normalized) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return normalized;
}
match_log = normalized as u32;
let total = lit_counts
.iter()
.fold(0usize, |sum, count| sum + *count as usize);
let normalized = unsafe {
FSE_normalizeCount(
lit_ncount.as_mut_ptr(),
lit_log,
lit_counts.as_ptr(),
total,
MAX_LL as u32,
1,
)
};
if ERR_isError(normalized) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return normalized;
}
lit_log = normalized as u32;
let mut dst = dst_buffer;
let mut remaining = max_dst_size;
let mut entropy_size = unsafe {
HUF_writeCTable_wksp(
dst.cast(),
remaining,
huf_table.as_ptr(),
255,
huff_log,
huf_workspace.as_mut_ptr().cast(),
size_of_val(&huf_workspace),
)
};
if ERR_isError(entropy_size) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return entropy_size;
}
dst = unsafe { dst.add(entropy_size) };
remaining -= entropy_size;
let header_size = unsafe {
FSE_writeNCount(
dst.cast(),
remaining,
offcode_ncount.as_ptr(),
OFFCODE_MAX as u32,
off_log,
)
};
if ERR_isError(header_size) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return header_size;
}
entropy_size += header_size;
dst = unsafe { dst.add(header_size) };
remaining -= header_size;
let header_size = unsafe {
FSE_writeNCount(
dst.cast(),
remaining,
match_ncount.as_ptr(),
MAX_ML as u32,
match_log,
)
};
if ERR_isError(header_size) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return header_size;
}
entropy_size += header_size;
dst = unsafe { dst.add(header_size) };
remaining -= header_size;
let header_size = unsafe {
FSE_writeNCount(
dst.cast(),
remaining,
lit_ncount.as_ptr(),
MAX_LL as u32,
lit_log,
)
};
if ERR_isError(header_size) {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return header_size;
}
entropy_size += header_size;
dst = unsafe { dst.add(header_size) };
remaining -= header_size;
if remaining < 12 {
unsafe {
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
return dictionary_error(ZstdErrorCode::DstSizeTooSmall);
}
unsafe {
MEM_writeLE32(dst.cast(), 1);
MEM_writeLE32(dst.add(4).cast(), 4);
MEM_writeLE32(dst.add(8).cast(), 8);
ZSTD_freeCDict(cdict);
ZSTD_freeCCtx(cctx);
}
entropy_size + 12
}
#[no_mangle]
pub unsafe extern "C" fn ZDICT_isError(error_code: usize) -> c_uint {
c_uint::from(ERR_isError(error_code))
}
#[no_mangle]
pub unsafe extern "C" fn ZDICT_getErrorName(error_code: usize) -> *const c_char {
ERR_getErrorName(error_code)
}
#[no_mangle]
pub unsafe extern "C" fn ZDICT_getDictID(dict_buffer: *const c_void, dict_size: usize) -> c_uint {
if dict_size < 8 {
return 0;
}
if unsafe { MEM_readLE32(dict_buffer) } != ZSTD_MAGIC_DICTIONARY {
return 0;
}
unsafe { MEM_readLE32(dict_buffer.cast::<u8>().add(4).cast()) }
}
#[no_mangle]
pub unsafe extern "C" fn ZDICT_getDictHeaderSize(
dict_buffer: *const c_void,
dict_size: usize,
) -> usize {
if dict_size <= 8 || unsafe { MEM_readLE32(dict_buffer) } != ZSTD_MAGIC_DICTIONARY {
return dictionary_error(ZstdErrorCode::DictionaryCorrupted);
}
let mut state = unsafe { MaybeUninit::<ZSTD_compressedBlockState_t>::zeroed().assume_init() };
let mut workspace = vec![0u32; HUF_WORKSPACE_SIZE / size_of::<u32>()];
unsafe {
ZSTD_reset_compressedBlockState(&mut state);
ZSTD_loadCEntropy(
&mut state,
workspace.as_mut_ptr().cast(),
dict_buffer,
dict_size,
)
}
}
#[no_mangle]
pub unsafe extern "C" fn ZDICT_finalizeDictionary(
dict_buffer: *mut c_void,
dict_buffer_capacity: usize,
custom_dict_content: *const c_void,
mut dict_content_size: usize,
samples_buffer: *const c_void,
samples_sizes: *const usize,
nb_samples: c_uint,
params: ZDICT_params_t,
) -> usize {
if dict_buffer_capacity < dict_content_size || dict_buffer_capacity < ZDICT_DICTSIZE_MIN {
return dictionary_error(ZstdErrorCode::DstSizeTooSmall);
}
let mut header = [0u8; 256];
unsafe {
MEM_writeLE32(header.as_mut_ptr().cast(), ZSTD_MAGIC_DICTIONARY);
let hash = XXH64(custom_dict_content, dict_content_size, 0);
let compliant_id = (hash % ((1u64 << 31) - 32_768)) as u32 + 32_768;
let dict_id = if params.dictID != 0 {
params.dictID
} else {
compliant_id
};
MEM_writeLE32(header.as_mut_ptr().add(4).cast(), dict_id);
}
let entropy_size = analyze_entropy(
header.as_mut_ptr().add(8),
header.len() - 8,
if params.compressionLevel == 0 {
ZSTD_CLEVEL_DEFAULT
} else {
params.compressionLevel
},
samples_buffer.cast(),
samples_sizes,
nb_samples,
custom_dict_content.cast(),
dict_content_size,
params.notificationLevel,
);
if ERR_isError(entropy_size) {
return entropy_size;
}
let header_size = 8 + entropy_size;
if header_size + dict_content_size > dict_buffer_capacity {
dict_content_size = dict_buffer_capacity - header_size;
}
let min_content_size = 8usize;
let padding_size = if dict_content_size < min_content_size {
if header_size + min_content_size > dict_buffer_capacity {
return dictionary_error(ZstdErrorCode::DstSizeTooSmall);
}
min_content_size - dict_content_size
} else {
0
};
let dictionary_size = header_size + padding_size + dict_content_size;
unsafe {
let output = dict_buffer.cast::<u8>();
let content = output.add(header_size + padding_size);
ptr::copy(custom_dict_content.cast::<u8>(), content, dict_content_size);
ptr::copy_nonoverlapping(header.as_ptr(), output, header_size);
ptr::write_bytes(output.add(header_size), 0, padding_size);
}
dictionary_size
}
unsafe fn add_entropy_tables_advanced(
dict_buffer: *mut c_void,
dict_content_size: usize,
dict_buffer_capacity: usize,
samples_buffer: *const c_void,
samples_sizes: *const usize,
nb_samples: c_uint,
params: ZDICT_params_t,
) -> usize {
if dict_buffer_capacity < 8 || dict_content_size > dict_buffer_capacity {
return dictionary_error(ZstdErrorCode::DstSizeTooSmall);
}
let content = dict_buffer
.cast::<u8>()
.add(dict_buffer_capacity - dict_content_size);
let entropy_size = analyze_entropy(
dict_buffer.cast::<u8>().add(8),
dict_buffer_capacity - 8,
if params.compressionLevel == 0 {
ZSTD_CLEVEL_DEFAULT
} else {
params.compressionLevel
},
samples_buffer.cast(),
samples_sizes,
nb_samples,
content,
dict_content_size,
params.notificationLevel,
);
if ERR_isError(entropy_size) {
return entropy_size;
}
let header_size = 8 + entropy_size;
unsafe {
MEM_writeLE32(dict_buffer, ZSTD_MAGIC_DICTIONARY);
let hash = XXH64(content.cast(), dict_content_size, 0);
let compliant_id = (hash % ((1u64 << 31) - 32_768)) as u32 + 32_768;
MEM_writeLE32(
dict_buffer.cast::<u8>().add(4).cast(),
if params.dictID != 0 {
params.dictID
} else {
compliant_id
},
);
if header_size + dict_content_size < dict_buffer_capacity {
ptr::copy(
content,
dict_buffer.cast::<u8>().add(header_size),
dict_content_size,
);
}
}
dict_buffer_capacity.min(header_size + dict_content_size)
}
unsafe fn train_from_buffer_unsafe_legacy(
dict_buffer: *mut c_void,
max_dict_size: usize,
samples_buffer: &[u8],
samples_sizes: &[usize],
params: ZDICT_legacy_params_t,
) -> usize {
let dict_list_size = DICTLISTSIZE_DEFAULT
.max(samples_sizes.len())
.max(max_dict_size / 16);
let mut dict_list = Vec::<DictItem>::new();
if dict_list.try_reserve_exact(dict_list_size).is_err() {
return dictionary_error(ZstdErrorCode::MemoryAllocation);
}
dict_list.resize(dict_list_size, DictItem::default());
init_dict_item(&mut dict_list[0]);
let selectivity = if params.selectivityLevel == 0 {
9usize
} else {
params.selectivityLevel as usize
};
let min_rep = if selectivity > 30 {
MINRATIO
} else {
samples_sizes.len() >> selectivity
};
let sample_size = total_sample_size(samples_sizes);
if max_dict_size < ZDICT_DICTSIZE_MIN {
return dictionary_error(ZstdErrorCode::DstSizeTooSmall);
}
if sample_size < ZDICT_MIN_SAMPLES_SIZE {
return dictionary_error(ZstdErrorCode::DictionaryCreationFailed);
}
let _ = train_buffer_legacy(
&mut dict_list,
samples_buffer,
samples_sizes,
min_rep,
params.zParams.notificationLevel,
);
let mut content_size = dict_size(&dict_list);
if content_size < ZDICT_CONTENTSIZE_MIN {
return dictionary_error(ZstdErrorCode::DictionaryCreationFailed);
}
let output = unsafe { std::slice::from_raw_parts_mut(dict_buffer.cast::<u8>(), max_dict_size) };
let mut write_at = max_dict_size;
for item in dict_list.iter().take(dict_list[0].pos as usize).skip(1) {
let length = item.length as usize;
if length > write_at {
return dictionary_error(ZstdErrorCode::Generic);
}
write_at -= length;
if write_at + length > output.len()
|| item.pos as usize + length > samples_buffer.len().saturating_sub(NOISELENGTH)
{
return dictionary_error(ZstdErrorCode::Generic);
}
output[write_at..write_at + length]
.copy_from_slice(&samples_buffer[item.pos as usize..item.pos as usize + length]);
}
let max = dict_list[0].pos as usize;
let mut current_size = 0usize;
let mut count = 1usize;
while count < max {
current_size += dict_list[count].length as usize;
if current_size > max_dict_size {
current_size -= dict_list[count].length as usize;
break;
}
count += 1;
}
dict_list[0].pos = count as u32;
content_size = current_size;
add_entropy_tables_advanced(
dict_buffer,
content_size,
max_dict_size,
samples_buffer.as_ptr().cast(),
samples_sizes.as_ptr(),
samples_sizes.len() as c_uint,
params.zParams,
)
}
#[no_mangle]
pub unsafe extern "C" fn ZDICT_trainFromBuffer_legacy(
dict_buffer: *mut c_void,
dict_buffer_capacity: usize,
samples_buffer: *const c_void,
samples_sizes: *const usize,
nb_samples: c_uint,
params: ZDICT_legacy_params_t,
) -> usize {
let sizes = if nb_samples == 0 {
&[][..]
} else {
unsafe { std::slice::from_raw_parts(samples_sizes, nb_samples as usize) }
};
let sample_size = total_sample_size(sizes);
if sample_size < ZDICT_MIN_SAMPLES_SIZE {
return 0;
}
let mut guarded = Vec::<u8>::new();
if guarded
.try_reserve_exact(sample_size.saturating_add(NOISELENGTH))
.is_err()
{
return dictionary_error(ZstdErrorCode::MemoryAllocation);
}
guarded.resize(sample_size + NOISELENGTH, 0);
if sample_size != 0 {
unsafe {
ptr::copy_nonoverlapping(
samples_buffer.cast::<u8>(),
guarded.as_mut_ptr(),
sample_size,
);
}
}
fill_noise(&mut guarded[sample_size..]);
unsafe {
train_from_buffer_unsafe_legacy(dict_buffer, dict_buffer_capacity, &guarded, sizes, params)
}
}
#[no_mangle]
pub unsafe extern "C" fn ZDICT_trainFromBuffer(
dict_buffer: *mut c_void,
dict_buffer_capacity: usize,
samples_buffer: *const c_void,
samples_sizes: *const usize,
nb_samples: c_uint,
) -> usize {
let mut params = ZDICT_fastCover_params_t {
d: 8,
steps: 4,
zParams: ZDICT_params_t {
compressionLevel: ZSTD_CLEVEL_DEFAULT,
..ZDICT_params_t::default()
},
..ZDICT_fastCover_params_t::default()
};
unsafe {
ZDICT_optimizeTrainFromBuffer_fastCover(
dict_buffer,
dict_buffer_capacity,
samples_buffer,
samples_sizes,
nb_samples,
&mut params,
)
}
}
#[no_mangle]
pub unsafe extern "C" fn ZDICT_addEntropyTablesFromBuffer(
dict_buffer: *mut c_void,
dict_content_size: usize,
dict_buffer_capacity: usize,
samples_buffer: *const c_void,
samples_sizes: *const usize,
nb_samples: c_uint,
) -> usize {
unsafe {
add_entropy_tables_advanced(
dict_buffer,
dict_content_size,
dict_buffer_capacity,
samples_buffer,
samples_sizes,
nb_samples,
ZDICT_params_t::default(),
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::ffi::CStr;
#[test]
fn helper_abi_and_error_names_match() {
assert_eq!(unsafe { ZDICT_getDictID(ptr::null(), 0) }, 0);
let error = dictionary_error(ZstdErrorCode::DstSizeTooSmall);
assert_eq!(unsafe { ZDICT_isError(error) }, 1);
let name = unsafe { CStr::from_ptr(ZDICT_getErrorName(error)) };
assert!(name
.to_bytes()
.windows(11)
.any(|part| part == b"Destination"));
}
#[test]
fn dictionary_id_reads_only_valid_headers() {
let mut dict = [0u8; 8];
unsafe {
MEM_writeLE32(dict.as_mut_ptr().cast(), ZSTD_MAGIC_DICTIONARY);
MEM_writeLE32(dict.as_mut_ptr().add(4).cast(), 1234);
}
assert_eq!(
unsafe { ZDICT_getDictID(dict.as_ptr().cast(), dict.len()) },
1234
);
dict[0] ^= 1;
assert_eq!(
unsafe { ZDICT_getDictID(dict.as_ptr().cast(), dict.len()) },
0
);
}
}
+4 -2
View File
@@ -6,6 +6,10 @@ pub mod common;
pub mod cpu;
pub mod debug;
#[cfg(feature = "dict-builder")]
pub mod dict_builder_cover;
#[cfg(all(feature = "compression", feature = "dict-builder"))]
pub mod dict_builder_zdict;
#[cfg(feature = "dict-builder")]
pub mod divsufsort;
pub mod entropy_common;
pub mod errors;
@@ -18,8 +22,6 @@ pub mod hist;
pub mod huf_compress;
#[cfg(feature = "decompression")]
pub mod huf_decompress;
#[cfg(feature = "dict-builder")]
pub mod dict_builder_cover;
pub mod legacy;
pub mod mem;
pub mod pool;