diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 8e9bd1e49..89dce1148 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -688,6 +688,19 @@ U32 ZSTD_rust_resolveRepcodeToRawOffset(const U32 rep[ZSTD_REP_NUM], U32 offBase, U32 ll0); size_t ZSTD_rust_loadCEntropy(ZSTD_compressedBlockState_t* bs, void* workspace, const void* dict, size_t dictSize); +/* Rust owns dictionary-ingestion policy. The content loader remains a narrow + * callback because it needs configuration-sensitive C-private layouts. */ +typedef size_t (*ZSTD_rust_loadDictionaryContent_f)( + void* matchState, void* ldmState, void* workspaceState, + const void* params, const void* src, size_t srcSize, + int dtlm, int tfp); +size_t ZSTD_rust_compressInsertDictionary( + ZSTD_compressedBlockState_t* bs, + void* matchState, void* ldmState, void* workspaceState, + const void* params, const void* dict, size_t dictSize, + int dictContentType, int dtlm, int tfp, + void* workspace, int noDictIDFlag, + ZSTD_rust_loadDictionaryContent_f loadDictionaryContent); size_t ZSTD_rust_transferSequencesWBlockDelim( SeqStore_t* seqStore, ZSTD_SequencePosition* seqPos, const ZSTD_Sequence* inSeqs, size_t inSeqsSize, @@ -3566,43 +3579,21 @@ size_t ZSTD_loadCEntropy(ZSTD_compressedBlockState_t* bs, void* workspace, return ZSTD_rust_loadCEntropy(bs, workspace, dict, dictSize); } -/* Dictionary format : - * See : - * https://github.com/facebook/zstd/blob/release/doc/zstd_compression_format.md#dictionary-format - */ -/*! ZSTD_loadZstdDictionary() : - * @return : dictID, or an error code - * assumptions : magic number supposed already checked - * dictSize supposed >= 8 - */ -static size_t ZSTD_loadZstdDictionary(ZSTD_compressedBlockState_t* bs, - ZSTD_MatchState_t* ms, - ZSTD_cwksp* ws, - ZSTD_CCtx_params const* params, - const void* dict, size_t dictSize, - ZSTD_dictTableLoadMethod_e dtlm, - ZSTD_tableFillPurpose_e tfp, - void* workspace) +/* Keep the match-state/content operation private to C. Rust passes only + * opaque pointers here after it has selected the dictionary path. */ +static size_t ZSTD_loadDictionaryContent_callback( + void* matchState, void* ldmState, void* workspaceState, + const void* params, const void* src, size_t srcSize, + int dtlm, int tfp) { - const BYTE* dictPtr = (const BYTE*)dict; - const BYTE* const dictEnd = dictPtr + dictSize; - size_t dictID; - size_t eSize; - ZSTD_STATIC_ASSERT(HUF_WORKSPACE_SIZE >= (1<= 8); - assert(MEM_readLE32(dictPtr) == ZSTD_MAGIC_DICTIONARY); - - dictID = params->fParams.noDictIDFlag ? 0 : MEM_readLE32(dictPtr + 4 /* skip magic number */ ); - eSize = ZSTD_loadCEntropy(bs, workspace, dict, dictSize); - FORWARD_IF_ERROR(eSize, "ZSTD_loadCEntropy failed"); - dictPtr += eSize; - - { - size_t const dictContentSize = (size_t)(dictEnd - dictPtr); - FORWARD_IF_ERROR(ZSTD_loadDictionaryContent( - ms, NULL, ws, params, dictPtr, dictContentSize, dtlm, tfp), ""); - } - return dictID; + return ZSTD_loadDictionaryContent( + (ZSTD_MatchState_t*)matchState, + (ldmState_t*)ldmState, + (ZSTD_cwksp*)workspaceState, + (const ZSTD_CCtx_params*)params, + src, srcSize, + (ZSTD_dictTableLoadMethod_e)dtlm, + (ZSTD_tableFillPurpose_e)tfp); } /** ZSTD_compress_insertDictionary() : @@ -3619,31 +3610,15 @@ ZSTD_compress_insertDictionary(ZSTD_compressedBlockState_t* bs, ZSTD_tableFillPurpose_e tfp, void* workspace) { + int const noDictIDFlag = (params != NULL && dict != NULL && dictSize >= 8) + ? params->fParams.noDictIDFlag + : 0; DEBUGLOG(4, "ZSTD_compress_insertDictionary (dictSize=%u)", (U32)dictSize); - if ((dict==NULL) || (dictSize<8)) { - RETURN_ERROR_IF(dictContentType == ZSTD_dct_fullDict, dictionary_wrong, ""); - return 0; - } - - ZSTD_reset_compressedBlockState(bs); - - /* dict restricted modes */ - if (dictContentType == ZSTD_dct_rawContent) - return ZSTD_loadDictionaryContent(ms, ls, ws, params, dict, dictSize, dtlm, tfp); - - if (MEM_readLE32(dict) != ZSTD_MAGIC_DICTIONARY) { - if (dictContentType == ZSTD_dct_auto) { - DEBUGLOG(4, "raw content dictionary detected"); - return ZSTD_loadDictionaryContent( - ms, ls, ws, params, dict, dictSize, dtlm, tfp); - } - RETURN_ERROR_IF(dictContentType == ZSTD_dct_fullDict, dictionary_wrong, ""); - assert(0); /* impossible */ - } - - /* dict as full zstd dictionary */ - return ZSTD_loadZstdDictionary( - bs, ms, ws, params, dict, dictSize, dtlm, tfp, workspace); + return ZSTD_rust_compressInsertDictionary( + bs, ms, ls, ws, params, dict, dictSize, + (int)dictContentType, (int)dtlm, (int)tfp, + workspace, noDictIDFlag, + ZSTD_loadDictionaryContent_callback); } #define ZSTD_USE_CDICT_PARAMS_SRCSIZE_CUTOFF (128 KB) diff --git a/rust/src/zstd_compress_dictionary.rs b/rust/src/zstd_compress_dictionary.rs index 176ced20b..f04670f4f 100644 --- a/rust/src/zstd_compress_dictionary.rs +++ b/rust/src/zstd_compress_dictionary.rs @@ -15,9 +15,12 @@ use crate::entropy_common::FSE_readNCount; use crate::errors::{ERR_isError, ZstdErrorCode, ERROR}; use crate::fse_compress::FSE_buildCTable_wksp; use crate::huf_compress::HUF_readCTable; -use crate::zstd_compress_stats::ZSTD_compressedBlockState_t; +use crate::zstd_compress_stats::{ + ZSTD_compressedBlockState_t, ZSTD_rust_resetCompressedBlockState, +}; use std::ffi::c_void; use std::os::raw::{c_int, c_short, c_uint}; +use std::ptr; use std::slice; const HUF_REPEAT_CHECK: c_int = 1; @@ -27,6 +30,25 @@ const FSE_REPEAT_VALID: c_int = 2; const HUF_WORKSPACE_SIZE: usize = (8 << 10) + 512; const DICTIONARY_ID_AND_MAGIC_SIZE: usize = 8; const REPCODE_SECTION_SIZE: usize = 12; +const ZSTD_MAGIC_DICTIONARY: u32 = 0xEC30_A437; +const ZSTD_DCT_AUTO: c_int = 0; +const ZSTD_DCT_RAW_CONTENT: c_int = 1; +const ZSTD_DCT_FULL_DICT: c_int = 2; + +/// C keeps match-state/content insertion private because it depends on the +/// configuration-sensitive `ZSTD_MatchState_t`, `ldmState_t`, workspace, and +/// parameter layouts. Rust owns the dictionary dispatch and calls this narrow +/// operation only after selecting the raw or full-dictionary path. +pub type LoadDictionaryContentFn = unsafe extern "C" fn( + match_state: *mut c_void, + ldm_state: *mut c_void, + workspace_state: *mut c_void, + params: *const c_void, + src: *const c_void, + src_size: usize, + dtlm: c_int, + tfp: c_int, +) -> usize; #[inline] fn dictionary_corrupted() -> usize { @@ -257,6 +279,107 @@ pub unsafe extern "C" fn ZSTD_rust_loadCEntropy( offset } +/// Rust-owned dispatch for C's `ZSTD_compress_insertDictionary()`. +/// +/// The callback is the only operation that still crosses back into C. It +/// receives opaque pointers because the content loader needs private C +/// layouts, while Rust owns short-dictionary handling, mode selection, block +/// state reset, dictionary magic/ID handling, and entropy setup. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressInsertDictionary( + bs: *mut ZSTD_compressedBlockState_t, + match_state: *mut c_void, + ldm_state: *mut c_void, + workspace_state: *mut c_void, + params: *const c_void, + dict: *const c_void, + dict_size: usize, + dict_content_type: c_int, + dtlm: c_int, + tfp: c_int, + workspace: *mut c_void, + no_dict_id_flag: c_int, + load_dictionary_content: LoadDictionaryContentFn, +) -> usize { + if dict.is_null() || dict_size < DICTIONARY_ID_AND_MAGIC_SIZE { + if dict_content_type == ZSTD_DCT_FULL_DICT { + return ERROR(ZstdErrorCode::DictionaryWrong); + } + return 0; + } + + unsafe { ZSTD_rust_resetCompressedBlockState(bs) }; + + if dict_content_type == ZSTD_DCT_RAW_CONTENT { + return unsafe { + load_dictionary_content( + match_state, + ldm_state, + workspace_state, + params, + dict, + dict_size, + dtlm, + tfp, + ) + }; + } + + let dict_magic = unsafe { u32::from_le(ptr::read_unaligned(dict.cast::())) }; + if dict_magic != ZSTD_MAGIC_DICTIONARY { + if dict_content_type == ZSTD_DCT_AUTO { + return unsafe { + load_dictionary_content( + match_state, + ldm_state, + workspace_state, + params, + dict, + dict_size, + dtlm, + tfp, + ) + }; + } + if dict_content_type == ZSTD_DCT_FULL_DICT { + return ERROR(ZstdErrorCode::DictionaryWrong); + } + debug_assert!(false, "invalid dictionary content type"); + } + + let dict_id = if no_dict_id_flag != 0 { + 0 + } else { + unsafe { + u32::from_le(ptr::read_unaligned(dict.cast::().add(4).cast::())) as usize + } + }; + let entropy_size = unsafe { ZSTD_rust_loadCEntropy(bs, workspace, dict, dict_size) }; + if ERR_isError(entropy_size) { + return entropy_size; + } + if entropy_size > dict_size { + return dictionary_corrupted(); + } + + let content_result = unsafe { + load_dictionary_content( + match_state, + ptr::null_mut(), + workspace_state, + params, + dict.cast::().add(entropy_size).cast(), + dict_size - entropy_size, + dtlm, + tfp, + ) + }; + if ERR_isError(content_result) { + return content_result; + } + dict_id +} + #[cfg(test)] mod tests { use super::*; @@ -277,6 +400,77 @@ mod tests { Offset, } + struct DictionaryLoadProbe { + calls: usize, + ldm_state: *mut c_void, + dict: *const c_void, + dict_size: usize, + dtlm: c_int, + tfp: c_int, + } + + impl Default for DictionaryLoadProbe { + fn default() -> Self { + Self { + calls: 0, + ldm_state: ptr::null_mut(), + dict: ptr::null(), + dict_size: 0, + dtlm: 0, + tfp: 0, + } + } + } + + unsafe extern "C" fn record_dictionary_load( + match_state: *mut c_void, + ldm_state: *mut c_void, + _workspace_state: *mut c_void, + _params: *const c_void, + dict: *const c_void, + dict_size: usize, + dtlm: c_int, + tfp: c_int, + ) -> usize { + let probe = unsafe { &mut *match_state.cast::() }; + probe.calls += 1; + probe.ldm_state = ldm_state; + probe.dict = dict; + probe.dict_size = dict_size; + probe.dtlm = dtlm; + probe.tfp = tfp; + 37 + } + + unsafe fn dispatch_for_test( + state: &mut ZSTD_compressedBlockState_t, + probe: &mut DictionaryLoadProbe, + dict: *const c_void, + dict_size: usize, + dict_content_type: c_int, + workspace: *mut c_void, + no_dict_id_flag: c_int, + ldm_state: *mut c_void, + ) -> usize { + unsafe { + ZSTD_rust_compressInsertDictionary( + state, + (probe as *mut DictionaryLoadProbe).cast(), + ldm_state, + ptr::null_mut(), + ptr::null(), + dict, + dict_size, + dict_content_type, + 1, + 2, + workspace, + no_dict_id_flag, + record_dictionary_load, + ) + } + } + fn assert_dictionary_corrupted(result: usize) { assert!(ERR_isError(result)); assert_eq!( @@ -467,6 +661,154 @@ mod tests { assert_dictionary_corrupted(result); } + #[test] + fn short_or_missing_dictionary_preserves_state_and_full_mode_error() { + let mut state = + unsafe { MaybeUninit::::zeroed().assume_init() }; + state.rep = [9, 10, 11]; + let mut probe = DictionaryLoadProbe::default(); + + let result = unsafe { + dispatch_for_test( + &mut state, + &mut probe, + ptr::null(), + 0, + ZSTD_DCT_AUTO, + ptr::null_mut(), + 0, + ptr::null_mut(), + ) + }; + assert_eq!(result, 0); + assert_eq!(state.rep, [9, 10, 11]); + + let short_dictionary = [0u8; DICTIONARY_ID_AND_MAGIC_SIZE - 1]; + let result = unsafe { + dispatch_for_test( + &mut state, + &mut probe, + short_dictionary.as_ptr().cast(), + short_dictionary.len(), + ZSTD_DCT_FULL_DICT, + ptr::null_mut(), + 0, + ptr::null_mut(), + ) + }; + assert!(ERR_isError(result)); + assert_eq!( + ERR_getErrorCode(result), + ZstdErrorCode::DictionaryWrong as i32 + ); + assert_eq!(state.rep, [9, 10, 11]); + + state.rep = [12, 13, 14]; + let wrong_magic_dictionary = [0u8; DICTIONARY_ID_AND_MAGIC_SIZE]; + let result = unsafe { + dispatch_for_test( + &mut state, + &mut probe, + wrong_magic_dictionary.as_ptr().cast(), + wrong_magic_dictionary.len(), + ZSTD_DCT_FULL_DICT, + ptr::null_mut(), + 0, + ptr::null_mut(), + ) + }; + assert!(ERR_isError(result)); + assert_eq!( + ERR_getErrorCode(result), + ZstdErrorCode::DictionaryWrong as i32 + ); + assert_eq!(state.rep, [1, 4, 8]); + assert_eq!(probe.calls, 0); + } + + #[test] + fn raw_and_auto_dictionary_modes_use_the_content_callback() { + let mut state = + unsafe { MaybeUninit::::zeroed().assume_init() }; + let mut probe = DictionaryLoadProbe::default(); + let raw_dictionary = [0xEC, 0x30, 0xA4, 0x37, 0, 0, 0, 0]; + let result = unsafe { + dispatch_for_test( + &mut state, + &mut probe, + raw_dictionary.as_ptr().cast(), + raw_dictionary.len(), + ZSTD_DCT_RAW_CONTENT, + ptr::null_mut(), + 0, + ptr::dangling_mut::(), + ) + }; + assert_eq!(result, 37); + assert_eq!(probe.calls, 1); + assert_eq!(probe.dict, raw_dictionary.as_ptr().cast()); + assert_eq!(probe.dict_size, raw_dictionary.len()); + assert_eq!(probe.ldm_state, ptr::dangling_mut::()); + + let auto_dictionary = [0u8; DICTIONARY_ID_AND_MAGIC_SIZE]; + let result = unsafe { + dispatch_for_test( + &mut state, + &mut probe, + auto_dictionary.as_ptr().cast(), + auto_dictionary.len(), + ZSTD_DCT_AUTO, + ptr::null_mut(), + 0, + ptr::dangling_mut::(), + ) + }; + assert_eq!(result, 37); + assert_eq!(probe.calls, 2); + assert_eq!(probe.dict, auto_dictionary.as_ptr().cast()); + assert_eq!(probe.dict_size, auto_dictionary.len()); + } + + #[test] + fn full_dictionary_loads_entropy_and_returns_or_suppresses_id() { + let dictionary = make_dictionary(ZeroWeight::None); + let mut workspace = [0u32; HUF_WORKSPACE_SIZE / size_of::()]; + let mut state = + unsafe { MaybeUninit::::zeroed().assume_init() }; + let mut probe = DictionaryLoadProbe::default(); + let result = unsafe { + dispatch_for_test( + &mut state, + &mut probe, + dictionary.as_ptr().cast(), + dictionary.len(), + ZSTD_DCT_FULL_DICT, + workspace.as_mut_ptr().cast(), + 0, + ptr::dangling_mut::(), + ) + }; + assert_eq!(result, 1234); + assert_eq!(probe.calls, 1); + assert_eq!(probe.ldm_state, ptr::null_mut()); + assert_eq!(probe.dict_size, DICT_CONTENT_SIZE); + + let result = unsafe { + dispatch_for_test( + &mut state, + &mut probe, + dictionary.as_ptr().cast(), + dictionary.len(), + ZSTD_DCT_FULL_DICT, + workspace.as_mut_ptr().cast(), + 1, + ptr::dangling_mut::(), + ) + }; + assert_eq!(result, 0); + assert_eq!(probe.calls, 2); + } + #[test] fn dict_n_count_repeat_requires_coverage_and_nonzero_weights() { let mut normalized = [1i16; 4];