diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index d7f679a4f..be68a18d1 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -5454,39 +5454,104 @@ static void ZSTD_initBuildSeqStoreState( /* ZSTD_sequenceBound() lives in rust/src/zstd_compress_api.rs. */ +typedef size_t (*ZSTD_rust_generateSequencesGetParameter_f)( + void* context, int param, int* value); +typedef void* (*ZSTD_rust_generateSequencesAllocate_f)( + void* context, size_t size); +typedef void (*ZSTD_rust_generateSequencesFree_f)( + void* context, void* pointer); +typedef size_t (*ZSTD_rust_generateSequencesSetCollector_f)( + void* context, ZSTD_Sequence* outSeqs, size_t outSeqsSize); +typedef size_t (*ZSTD_rust_generateSequencesCompress2_f)( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize); +typedef size_t (*ZSTD_rust_generateSequencesGetCount_f)(void* context); +typedef struct { + void* callbackContext; + ZSTD_rust_generateSequencesGetParameter_f getParameter; + ZSTD_rust_generateSequencesAllocate_f allocate; + ZSTD_rust_generateSequencesFree_f free; + ZSTD_rust_generateSequencesSetCollector_f setCollector; + ZSTD_rust_generateSequencesCompress2_f compress2; + ZSTD_rust_generateSequencesGetCount_f getCount; +} ZSTD_rust_generateSequencesState; +typedef char ZSTD_rust_generate_sequences_state_layout[ + (offsetof(ZSTD_rust_generateSequencesState, callbackContext) == 0 + && offsetof(ZSTD_rust_generateSequencesState, getParameter) + == sizeof(void*) + && offsetof(ZSTD_rust_generateSequencesState, allocate) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_generateSequencesState, free) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_generateSequencesState, setCollector) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_generateSequencesState, compress2) + == 5 * sizeof(void*) + && offsetof(ZSTD_rust_generateSequencesState, getCount) + == 6 * sizeof(void*) + && sizeof(ZSTD_rust_generateSequencesState) == 7 * sizeof(void*)) + ? 1 : -1]; +size_t ZSTD_rust_generateSequences( + const ZSTD_rust_generateSequencesState* state, + ZSTD_Sequence* outSeqs, size_t outSeqsSize, + const void* src, size_t srcSize); + +static size_t ZSTD_rust_generateSequences_getParameter( + void* context, int param, int* value) +{ + return ZSTD_CCtx_getParameter( + (const ZSTD_CCtx*)context, (ZSTD_cParameter)param, value); +} + +static void* ZSTD_rust_generateSequences_allocate(void* context, size_t size) +{ + (void)context; + return ZSTD_customMalloc(size, ZSTD_defaultCMem); +} + +static void ZSTD_rust_generateSequences_free(void* context, void* pointer) +{ + (void)context; + ZSTD_customFree(pointer, ZSTD_defaultCMem); +} + +static size_t ZSTD_rust_generateSequences_setCollector( + void* context, ZSTD_Sequence* outSeqs, size_t outSeqsSize) +{ + ZSTD_CCtx* const zc = (ZSTD_CCtx*)context; + zc->seqCollector.collectSequences = 1; + zc->seqCollector.seqStart = outSeqs; + zc->seqCollector.seqIndex = 0; + zc->seqCollector.maxSequences = outSeqsSize; + return 0; +} + +static size_t ZSTD_rust_generateSequences_compress2( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize) +{ + return ZSTD_compress2((ZSTD_CCtx*)context, dst, dstCapacity, src, srcSize); +} + +static size_t ZSTD_rust_generateSequences_getCount(void* context) +{ + return ((const ZSTD_CCtx*)context)->seqCollector.seqIndex; +} + size_t ZSTD_generateSequences(ZSTD_CCtx* zc, ZSTD_Sequence* outSeqs, size_t outSeqsSize, const void* src, size_t srcSize) { - const size_t dstCapacity = ZSTD_compressBound(srcSize); - void* dst; /* Make C90 happy. */ - SeqCollector seqCollector; - { - int targetCBlockSize; - FORWARD_IF_ERROR(ZSTD_CCtx_getParameter(zc, ZSTD_c_targetCBlockSize, &targetCBlockSize), ""); - RETURN_ERROR_IF(targetCBlockSize != 0, parameter_unsupported, "targetCBlockSize != 0"); - } - { - int nbWorkers; - FORWARD_IF_ERROR(ZSTD_CCtx_getParameter(zc, ZSTD_c_nbWorkers, &nbWorkers), ""); - RETURN_ERROR_IF(nbWorkers != 0, parameter_unsupported, "nbWorkers != 0"); - } - - dst = ZSTD_customMalloc(dstCapacity, ZSTD_defaultCMem); - RETURN_ERROR_IF(dst == NULL, memory_allocation, "NULL pointer!"); - - seqCollector.collectSequences = 1; - seqCollector.seqStart = outSeqs; - seqCollector.seqIndex = 0; - seqCollector.maxSequences = outSeqsSize; - zc->seqCollector = seqCollector; - - { - const size_t ret = ZSTD_compress2(zc, dst, dstCapacity, src, srcSize); - ZSTD_customFree(dst, ZSTD_defaultCMem); - FORWARD_IF_ERROR(ret, "ZSTD_compress2 failed"); - } - assert(zc->seqCollector.seqIndex <= ZSTD_sequenceBound(srcSize)); - return zc->seqCollector.seqIndex; + ZSTD_rust_generateSequencesState const state = { + zc, + ZSTD_rust_generateSequences_getParameter, + ZSTD_rust_generateSequences_allocate, + ZSTD_rust_generateSequences_free, + ZSTD_rust_generateSequences_setCollector, + ZSTD_rust_generateSequences_compress2, + ZSTD_rust_generateSequences_getCount + }; + return ZSTD_rust_generateSequences( + &state, outSeqs, outSeqsSize, src, srcSize); } /* ZSTD_mergeBlockDelimiters() lives in rust/src/zstd_compress_api.rs. */ diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index edda9dbdc..71d365206 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -1275,6 +1275,105 @@ const ZSTD_COMPRESSION_STAGE_ONGOING: c_int = 2; #[cfg(test)] const ZSTD_COMPRESSION_STAGE_ENDING: c_int = 3; +type GenerateSequencesGetParameterFn = + unsafe extern "C" fn(*mut c_void, c_int, *mut c_int) -> usize; +type GenerateSequencesAllocateFn = unsafe extern "C" fn(*mut c_void, usize) -> *mut c_void; +type GenerateSequencesFreeFn = unsafe extern "C" fn(*mut c_void, *mut c_void); +type GenerateSequencesSetCollectorFn = + unsafe extern "C" fn(*mut c_void, *mut ZSTD_Sequence, usize) -> usize; +type GenerateSequencesCompress2Fn = + unsafe extern "C" fn(*mut c_void, *mut c_void, usize, *const c_void, usize) -> usize; +type GenerateSequencesGetCountFn = unsafe extern "C" fn(*mut c_void) -> usize; + +/// C-owned context projection for `ZSTD_generateSequences()`. +/// +/// Rust owns the public policy and temporary-buffer lifetime. C callbacks keep +/// the private `ZSTD_CCtx` and `SeqCollector` layouts out of this module. +#[repr(C)] +pub struct ZSTD_rust_generateSequencesState { + callback_context: *mut c_void, + get_parameter: GenerateSequencesGetParameterFn, + allocate: GenerateSequencesAllocateFn, + free: GenerateSequencesFreeFn, + set_collector: GenerateSequencesSetCollectorFn, + compress2: GenerateSequencesCompress2Fn, + get_count: GenerateSequencesGetCountFn, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_generateSequencesState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_generateSequencesState, get_parameter) == size_of::()); + assert!(offset_of!(ZSTD_rust_generateSequencesState, allocate) == 2 * size_of::()); + assert!(offset_of!(ZSTD_rust_generateSequencesState, free) == 3 * size_of::()); + assert!(offset_of!(ZSTD_rust_generateSequencesState, set_collector) == 4 * size_of::()); + assert!(offset_of!(ZSTD_rust_generateSequencesState, compress2) == 5 * size_of::()); + assert!(offset_of!(ZSTD_rust_generateSequencesState, get_count) == 6 * size_of::()); + assert!(size_of::() == 7 * size_of::()); +}; + +/// Own the validation, allocation, collection setup, compression call, and +/// cleanup ordering for the public sequence-generation API. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_generateSequences( + state: *const ZSTD_rust_generateSequencesState, + out_seqs: *mut ZSTD_Sequence, + out_seqs_size: usize, + src: *const c_void, + src_size: usize, +) -> usize { + if state.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let state = unsafe { &*state }; + let dst_capacity = crate::zstd_compress_api::ZSTD_compressBound(src_size); + let mut parameter = 0; + + let result = unsafe { + (state.get_parameter)( + state.callback_context, + ZSTD_C_TARGET_C_BLOCK_SIZE, + &mut parameter, + ) + }; + if ERR_isError(result) { + return result; + } + if parameter != 0 { + return ERROR(ZstdErrorCode::ParameterUnsupported); + } + + let result = + unsafe { (state.get_parameter)(state.callback_context, ZSTD_C_NB_WORKERS, &mut parameter) }; + if ERR_isError(result) { + return result; + } + if parameter != 0 { + return ERROR(ZstdErrorCode::ParameterUnsupported); + } + + let dst = unsafe { (state.allocate)(state.callback_context, dst_capacity) }; + if dst.is_null() { + return ERROR(ZstdErrorCode::MemoryAllocation); + } + + let result = unsafe { (state.set_collector)(state.callback_context, out_seqs, out_seqs_size) }; + if ERR_isError(result) { + unsafe { (state.free)(state.callback_context, dst) }; + return result; + } + + let result = + unsafe { (state.compress2)(state.callback_context, dst, dst_capacity, src, src_size) }; + unsafe { (state.free)(state.callback_context, dst) }; + if ERR_isError(result) { + return result; + } + + let sequence_count = unsafe { (state.get_count)(state.callback_context) }; + debug_assert!(sequence_count <= crate::zstd_compress_api::ZSTD_sequenceBound(src_size)); + sequence_count +} + #[inline] unsafe fn compress_continue_block_body_with( state: &ZSTD_rust_compressContinueBlockState, @@ -9663,6 +9762,172 @@ mod tests { const ZSTD_BTOPT: c_int = 7; const ZSTD_BTULTRA2: c_int = 9; + struct GenerateSequencesTestContext { + events: Vec<&'static str>, + target_c_block_size: c_int, + nb_workers: c_int, + allocation_fails: bool, + collector_result: usize, + compress_result: usize, + sequence_count: usize, + } + + unsafe extern "C" fn generate_sequences_get_parameter_test( + context: *mut c_void, + parameter: c_int, + value: *mut c_int, + ) -> usize { + let context = unsafe { &mut *context.cast::() }; + let (event, result) = match parameter { + ZSTD_C_TARGET_C_BLOCK_SIZE => ("target", context.target_c_block_size), + ZSTD_C_NB_WORKERS => ("workers", context.nb_workers), + _ => ("other", 0), + }; + context.events.push(event); + unsafe { *value = result }; + 0 + } + + unsafe extern "C" fn generate_sequences_allocate_test( + context: *mut c_void, + _size: usize, + ) -> *mut c_void { + let context = unsafe { &mut *context.cast::() }; + context.events.push("allocate"); + if context.allocation_fails { + ptr::null_mut() + } else { + Box::into_raw(Box::new(0_u8)).cast() + } + } + + unsafe extern "C" fn generate_sequences_free_test(context: *mut c_void, pointer: *mut c_void) { + let context = unsafe { &mut *context.cast::() }; + context.events.push("free"); + if !pointer.is_null() { + unsafe { drop(Box::from_raw(pointer.cast::())) }; + } + } + + unsafe extern "C" fn generate_sequences_set_collector_test( + context: *mut c_void, + _out_seqs: *mut ZSTD_Sequence, + _out_seqs_size: usize, + ) -> usize { + let context = unsafe { &mut *context.cast::() }; + context.events.push("collector"); + context.collector_result + } + + unsafe extern "C" fn generate_sequences_compress2_test( + context: *mut c_void, + _dst: *mut c_void, + _dst_capacity: usize, + _src: *const c_void, + _src_size: usize, + ) -> usize { + let context = unsafe { &mut *context.cast::() }; + context.events.push("compress2"); + context.compress_result + } + + unsafe extern "C" fn generate_sequences_get_count_test(context: *mut c_void) -> usize { + let context = unsafe { &mut *context.cast::() }; + context.events.push("count"); + context.sequence_count + } + + fn generate_sequences_test_state( + context: &mut GenerateSequencesTestContext, + ) -> ZSTD_rust_generateSequencesState { + ZSTD_rust_generateSequencesState { + callback_context: context as *mut GenerateSequencesTestContext as *mut c_void, + get_parameter: generate_sequences_get_parameter_test, + allocate: generate_sequences_allocate_test, + free: generate_sequences_free_test, + set_collector: generate_sequences_set_collector_test, + compress2: generate_sequences_compress2_test, + get_count: generate_sequences_get_count_test, + } + } + + #[test] + fn generate_sequences_rejects_unsupported_parameters_before_allocation() { + let mut context = GenerateSequencesTestContext { + events: Vec::new(), + target_c_block_size: 1, + nb_workers: 0, + allocation_fails: false, + collector_result: 0, + compress_result: 0, + sequence_count: 0, + }; + let state = generate_sequences_test_state(&mut context); + + let result = + unsafe { ZSTD_rust_generateSequences(&state, ptr::null_mut(), 0, ptr::null(), 0) }; + + assert_eq!( + ERR_getErrorCode(result), + ZstdErrorCode::ParameterUnsupported as i32 + ); + assert_eq!(context.events, ["target"]); + } + + #[test] + fn generate_sequences_frees_temporary_output_on_setup_failure() { + let mut context = GenerateSequencesTestContext { + events: Vec::new(), + target_c_block_size: 0, + nb_workers: 0, + allocation_fails: false, + collector_result: ERROR(ZstdErrorCode::Generic), + compress_result: 0, + sequence_count: 0, + }; + let state = generate_sequences_test_state(&mut context); + + let result = + unsafe { ZSTD_rust_generateSequences(&state, ptr::null_mut(), 0, ptr::null(), 0) }; + + assert_eq!(ERR_getErrorCode(result), ZstdErrorCode::Generic as i32); + assert_eq!( + context.events, + ["target", "workers", "allocate", "collector", "free"] + ); + } + + #[test] + fn generate_sequences_runs_callbacks_in_order_and_returns_count() { + let mut context = GenerateSequencesTestContext { + events: Vec::new(), + target_c_block_size: 0, + nb_workers: 0, + allocation_fails: false, + collector_result: 0, + compress_result: 0, + sequence_count: 2, + }; + let state = generate_sequences_test_state(&mut context); + + let result = + unsafe { ZSTD_rust_generateSequences(&state, ptr::null_mut(), 0, ptr::null(), 0) }; + + assert_eq!(result, 2); + assert_eq!( + context.events, + [ + "target", + "workers", + "allocate", + "collector", + "compress2", + "free", + "count" + ] + ); + } + unsafe extern "C" fn select_block_compressor_test_default( _match_state: *mut c_void, _seq_store: *mut SeqStore_t,