From 5bc03f4e041e2d3efb1d040734308aeecec664d1 Mon Sep 17 00:00:00 2001 From: ddidderr Date: Tue, 21 Jul 2026 12:14:10 +0200 Subject: [PATCH] refactor(compress): move MT stream loop to Rust Move the multithreaded compressStream2 outer coordinator into the Rust compression module. The C bridge now owns only the private ZSTDMT call and consumed/produced counters, while Rust preserves the progress loop, completion ordering, and error handoff. Keep the MT context mutations, diagnostics, assertions, and buffer-expectation publication in C because they depend on the private context layout. Add callback-driven tests for continue/break, error completion, end trace/reset ordering, and malformed states. Test Plan: - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo +nightly fmt --manifest-path rust/Cargo.toml --all -- --check - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo clippy --manifest-path rust/Cargo.toml --all-targets -- -D warnings - ulimit -v 41943040; CARGO_BUILD_JOBS=1 make -j1 - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo test --manifest-path rust/cli/Cargo.toml --all-targets - ulimit -v 41943040; CARGO_BUILD_JOBS=1 make -j1 -C tests test - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo clippy --manifest-path rust/cli/Cargo.toml --all-targets -- -D warnings --- lib/compress/zstd_compress.c | 92 ++++++--- rust/src/zstd_compress.rs | 386 +++++++++++++++++++++++++++++++++-- 2 files changed, 431 insertions(+), 47 deletions(-) diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 3d321131d..65f32fd98 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -1012,6 +1012,36 @@ typedef char ZSTD_rust_compress_stream2_mt_completion_state_layout[ && sizeof(ZSTD_rust_compressStream2MTCompletionState) == 2 * sizeof(void*)) ? 1 : -1]; +typedef size_t (*ZSTD_rust_compressStream2MTStep_f)( + void* context, ZSTD_outBuffer* output, ZSTD_inBuffer* input, int endOp); +typedef struct { + void* callbackContext; + ZSTD_rust_compressStream2MTStep_f step; + ZSTD_outBuffer* output; + ZSTD_inBuffer* input; + const ZSTD_rust_compressStream2MTCompletionCallbacks* completionCallbacks; + int endOp; +} ZSTD_rust_compressStream2MTCoordinatorState; +size_t ZSTD_rust_compressStream2MTCoordinator( + const ZSTD_rust_compressStream2MTCoordinatorState* state); +typedef char ZSTD_rust_compress_stream2_mt_coordinator_state_layout[ + (offsetof(ZSTD_rust_compressStream2MTCoordinatorState, callbackContext) + == 0 + && offsetof(ZSTD_rust_compressStream2MTCoordinatorState, step) + == sizeof(void*) + && offsetof(ZSTD_rust_compressStream2MTCoordinatorState, output) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_compressStream2MTCoordinatorState, input) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_compressStream2MTCoordinatorState, + completionCallbacks) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_compressStream2MTCoordinatorState, endOp) + == 5 * sizeof(void*) + && sizeof(ZSTD_rust_compressStream2MTStep_f) == sizeof(void*) + && sizeof(ZSTD_rust_compressStream2MTCoordinatorState) + == 6 * sizeof(void*)) + ? 1 : -1]; int ZSTD_rust_simpleCompress2Level(const void* cctx, size_t srcSize); typedef int (*ZSTD_rust_simpleCompress2Level_f)(const void* cctx, size_t srcSize); typedef struct { @@ -8674,6 +8704,19 @@ static size_t ZSTD_rust_compressStream2MT_reset(void* context) { return ZSTD_CCtx_reset((ZSTD_CCtx*)context, ZSTD_reset_session_only); } + +static size_t ZSTD_rust_compressStream2MT_step( + void* context, ZSTD_outBuffer* output, ZSTD_inBuffer* input, int endOp) +{ + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + size_t const ipos = input->pos; + size_t const opos = output->pos; + size_t const flushMin = ZSTDMT_compressStream_generic( + cctx->mtctx, output, input, (ZSTD_EndDirective)endOp); + cctx->consumedSrcSize += (U64)(input->pos - ipos); + cctx->producedCSize += (U64)(output->pos - opos); + return flushMin; +} #endif /* ZSTD_MULTITHREAD */ size_t ZSTD_compressStream2_c( ZSTD_CCtx* cctx, @@ -8781,41 +8824,22 @@ size_t ZSTD_compressStream2_c( ZSTD_CCtx* cctx, input->pos -= cctx->stableIn_notConsumed; cctx->stableIn_notConsumed = 0; } - for (;;) { - size_t const ipos = input->pos; - size_t const opos = output->pos; - flushMin = ZSTDMT_compressStream_generic(cctx->mtctx, output, input, endOp); - cctx->consumedSrcSize += (U64)(input->pos - ipos); - cctx->producedCSize += (U64)(output->pos - opos); - { ZSTD_rust_compressStream2MTLoopPolicyState const state = { - (int)endOp, - flushMin, - input->pos, - input->size, - ipos, - output->pos, - output->size, - opos - }; - int const loopPolicy = - ZSTD_rust_compressStream2MTLoopPolicy(&state); - { ZSTD_rust_compressStream2MTCompletionCallbacks const callbacks = { - cctx, - ZSTD_rust_compressStream2MT_trace, - ZSTD_rust_compressStream2MT_reset - }; - ZSTD_rust_compressStream2MTCompletionState const completionState = { - loopPolicy, - &callbacks - }; - ZSTD_rust_compressStream2MTCompletion(&completionState); - } - FORWARD_IF_ERROR(flushMin, "ZSTDMT_compressStream_generic failed"); - if (loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK - || loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE) - break; - } + { ZSTD_rust_compressStream2MTCompletionCallbacks const callbacks = { + cctx, + ZSTD_rust_compressStream2MT_trace, + ZSTD_rust_compressStream2MT_reset + }; + ZSTD_rust_compressStream2MTCoordinatorState const state = { + cctx, + ZSTD_rust_compressStream2MT_step, + output, + input, + &callbacks, + (int)endOp + }; + flushMin = ZSTD_rust_compressStream2MTCoordinator(&state); } + FORWARD_IF_ERROR(flushMin, "ZSTDMT_compressStream_generic failed"); DEBUGLOG(5, "completed ZSTD_compressStream2 delegating to ZSTDMT_compressStream_generic"); /* Either we don't require maximum forward progress, we've finished the * flush, or we are out of output space. diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 646d131b4..1e066e8ab 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -556,8 +556,8 @@ type CompressStream2MTResetFn = unsafe extern "C" fn(*mut c_void) -> usize; #[repr(C)] pub struct ZSTD_rust_compressStream2MTCompletionCallbacks { callback_context: *mut c_void, - trace: CompressStream2MTTraceFn, - reset: CompressStream2MTResetFn, + trace: Option, + reset: Option, } /// Projected MT loop result and its opaque completion callbacks. @@ -570,6 +570,8 @@ pub struct ZSTD_rust_compressStream2MTCompletionState { const _: () = { assert!(size_of::() == size_of::()); assert!(size_of::() == size_of::()); + assert!(size_of::>() == size_of::()); + assert!(size_of::>() == size_of::()); assert!( offset_of!( ZSTD_rust_compressStream2MTCompletionCallbacks, @@ -598,14 +600,23 @@ unsafe fn compress_stream2_mt_completion(state: &ZSTD_rust_compressStream2MTComp let Some(callbacks) = (unsafe { state.callbacks.as_ref() }) else { return; }; + if callbacks.callback_context.is_null() { + return; + } match state.loop_policy { ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR => { - let _ = unsafe { (callbacks.reset)(callbacks.callback_context) }; + if let Some(reset) = callbacks.reset { + let _ = unsafe { reset(callbacks.callback_context) }; + } } ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE => { - unsafe { (callbacks.trace)(callbacks.callback_context) }; - let _ = unsafe { (callbacks.reset)(callbacks.callback_context) }; + if let Some(trace) = callbacks.trace { + unsafe { trace(callbacks.callback_context) }; + } + if let Some(reset) = callbacks.reset { + let _ = unsafe { reset(callbacks.callback_context) }; + } } ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_CONTINUE | ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK => {} @@ -624,6 +635,129 @@ pub unsafe extern "C" fn ZSTD_rust_compressStream2MTCompletion( unsafe { compress_stream2_mt_completion(state) }; } +type CompressStream2MTStepFn = + unsafe extern "C" fn(*mut c_void, *mut ZSTD_outBuffer, *mut ZSTD_inBuffer, c_int) -> usize; + +/// Rust owns only the MT outer-loop coordination. The step callback retains +/// the private C context and performs one `ZSTDMT_compressStream_generic` call. +#[repr(C)] +pub struct ZSTD_rust_compressStream2MTCoordinatorState { + callback_context: *mut c_void, + step: Option, + output: *mut ZSTD_outBuffer, + input: *mut ZSTD_inBuffer, + completion_callbacks: *const ZSTD_rust_compressStream2MTCompletionCallbacks, + end_op: c_int, +} + +const _: () = { + assert!(size_of::() == size_of::()); + assert!(size_of::>() == size_of::()); + assert!( + offset_of!( + ZSTD_rust_compressStream2MTCoordinatorState, + callback_context + ) == 0 + ); + assert!(offset_of!(ZSTD_rust_compressStream2MTCoordinatorState, step) == size_of::()); + assert!( + offset_of!(ZSTD_rust_compressStream2MTCoordinatorState, output) == 2 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2MTCoordinatorState, input) == 3 * size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_compressStream2MTCoordinatorState, + completion_callbacks + ) == 4 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2MTCoordinatorState, end_op) == 5 * size_of::() + ); + assert!(size_of::() == 6 * size_of::()); +}; + +/// Run the MT outer loop, returning the final C error/result for the caller's +/// existing `FORWARD_IF_ERROR` handling. Completion callbacks run before an +/// error result is returned, matching the former C loop ordering. +#[inline] +unsafe fn compress_stream2_mt_coordinator( + state: &ZSTD_rust_compressStream2MTCoordinatorState, +) -> usize { + let Some(step) = state.step else { + return ERROR(ZstdErrorCode::Generic); + }; + let Some(output) = (unsafe { state.output.as_mut() }) else { + return ERROR(ZstdErrorCode::Generic); + }; + let Some(input) = (unsafe { state.input.as_mut() }) else { + return ERROR(ZstdErrorCode::Generic); + }; + let Some(completion_callbacks) = (unsafe { state.completion_callbacks.as_ref() }) else { + return ERROR(ZstdErrorCode::Generic); + }; + if state.callback_context.is_null() + || completion_callbacks.callback_context.is_null() + || completion_callbacks.trace.is_none() + || completion_callbacks.reset.is_none() + || compress_stream2_policy(state.end_op) == 0 + || output.pos > output.size + || input.pos > input.size + { + return ERROR(ZstdErrorCode::Generic); + } + + loop { + let input_pos_before = input.pos; + let output_pos_before = output.pos; + let flush_min = unsafe { step(state.callback_context, output, input, state.end_op) }; + let loop_policy = if input.pos > input.size || output.pos > output.size { + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR + } else { + let policy_state = ZSTD_rust_compressStream2MTLoopPolicyState { + end_op: state.end_op, + flush_min, + input_pos: input.pos, + input_size: input.size, + input_pos_before, + output_pos: output.pos, + output_size: output.size, + output_pos_before, + }; + compress_stream2_mt_loop_policy(&policy_state) + }; + let completion_state = ZSTD_rust_compressStream2MTCompletionState { + loop_policy, + callbacks: completion_callbacks, + }; + unsafe { compress_stream2_mt_completion(&completion_state) }; + + if ERR_isError(flush_min) { + return flush_min; + } + if loop_policy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR { + return ERROR(ZstdErrorCode::Generic); + } + if loop_policy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK + || loop_policy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE + { + return flush_min; + } + } +} + +/// Coordinate the C-owned MT step callback and the existing loop policy. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressStream2MTCoordinator( + state: *const ZSTD_rust_compressStream2MTCoordinatorState, +) -> usize { + let Some(state) = (unsafe { state.as_ref() }) else { + return ERROR(ZstdErrorCode::Generic); + }; + unsafe { compress_stream2_mt_coordinator(state) } +} + const ZSTD_C_WINDOW_LOG: c_int = 101; const ZSTD_C_HASH_LOG: c_int = 102; const ZSTD_C_CHAIN_LOG: c_int = 103; @@ -15013,8 +15147,8 @@ mod tests { let mut context = CompressStream2MTCompletionTestContext::default(); let mut callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { callback_context: ptr::null_mut(), - trace: compress_stream2_mt_completion_test_trace, - reset: compress_stream2_mt_completion_test_reset, + trace: Some(compress_stream2_mt_completion_test_trace), + reset: Some(compress_stream2_mt_completion_test_reset), }; let state = compress_stream2_mt_completion_test_state( &mut context, @@ -15031,8 +15165,8 @@ mod tests { let mut context = CompressStream2MTCompletionTestContext::default(); let mut callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { callback_context: ptr::null_mut(), - trace: compress_stream2_mt_completion_test_trace, - reset: compress_stream2_mt_completion_test_reset, + trace: Some(compress_stream2_mt_completion_test_trace), + reset: Some(compress_stream2_mt_completion_test_reset), }; let state = compress_stream2_mt_completion_test_state( &mut context, @@ -15049,8 +15183,8 @@ mod tests { let mut context = CompressStream2MTCompletionTestContext::default(); let mut callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { callback_context: ptr::null_mut(), - trace: compress_stream2_mt_completion_test_trace, - reset: compress_stream2_mt_completion_test_reset, + trace: Some(compress_stream2_mt_completion_test_trace), + reset: Some(compress_stream2_mt_completion_test_reset), }; let state = compress_stream2_mt_completion_test_state( &mut context, @@ -15067,8 +15201,8 @@ mod tests { let mut context = CompressStream2MTCompletionTestContext::default(); let mut callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { callback_context: ptr::null_mut(), - trace: compress_stream2_mt_completion_test_trace, - reset: compress_stream2_mt_completion_test_reset, + trace: Some(compress_stream2_mt_completion_test_trace), + reset: Some(compress_stream2_mt_completion_test_reset), }; let state = compress_stream2_mt_completion_test_state( &mut context, @@ -15080,6 +15214,232 @@ mod tests { assert_eq!(context.events, ["trace", "reset"]); } + struct CompressStream2MTCoordinatorTestContext { + events: Vec<&'static str>, + step_results: Vec, + step_positions: Vec<(usize, usize)>, + step_calls: usize, + } + + unsafe extern "C" fn compress_stream2_mt_coordinator_test_step( + context: *mut c_void, + output: *mut ZSTD_outBuffer, + input: *mut ZSTD_inBuffer, + _end_op: c_int, + ) -> usize { + let context = unsafe { &mut *context.cast::() }; + let call = context.step_calls; + context.step_calls += 1; + context.events.push("step"); + let (input_pos, output_pos) = context.step_positions[call]; + unsafe { + (*input).pos = input_pos; + (*output).pos = output_pos; + } + context.step_results[call] + } + + unsafe extern "C" fn compress_stream2_mt_coordinator_test_trace(context: *mut c_void) { + let context = unsafe { &mut *context.cast::() }; + context.events.push("trace"); + } + + unsafe extern "C" fn compress_stream2_mt_coordinator_test_reset(context: *mut c_void) -> usize { + let context = unsafe { &mut *context.cast::() }; + context.events.push("reset"); + ERROR(ZstdErrorCode::Generic) + } + + fn compress_stream2_mt_coordinator_test_state( + context: &mut CompressStream2MTCoordinatorTestContext, + output: &mut ZSTD_outBuffer, + input: &mut ZSTD_inBuffer, + callbacks: &mut ZSTD_rust_compressStream2MTCompletionCallbacks, + end_op: c_int, + ) -> ZSTD_rust_compressStream2MTCoordinatorState { + let context_ptr = (context as *mut CompressStream2MTCoordinatorTestContext).cast(); + callbacks.callback_context = context_ptr; + ZSTD_rust_compressStream2MTCoordinatorState { + callback_context: context_ptr, + step: Some(compress_stream2_mt_coordinator_test_step), + output, + input, + completion_callbacks: callbacks, + end_op, + } + } + + #[test] + fn compress_stream2_mt_coordinator_continues_then_breaks() { + let mut context = CompressStream2MTCoordinatorTestContext { + events: Vec::new(), + step_results: vec![7, 7], + step_positions: vec![(1, 0), (1, 8)], + step_calls: 0, + }; + let mut output = ZSTD_outBuffer { + dst: ptr::null_mut(), + size: 8, + pos: 0, + }; + let mut input = ZSTD_inBuffer { + src: ptr::null(), + size: 8, + pos: 0, + }; + let mut callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { + callback_context: ptr::null_mut(), + trace: Some(compress_stream2_mt_coordinator_test_trace), + reset: Some(compress_stream2_mt_coordinator_test_reset), + }; + let state = compress_stream2_mt_coordinator_test_state( + &mut context, + &mut output, + &mut input, + &mut callbacks, + ZSTD_E_FLUSH, + ); + + let result = unsafe { ZSTD_rust_compressStream2MTCoordinator(&state) }; + assert_eq!(result, 7); + assert_eq!(context.step_calls, 2); + assert_eq!(context.events, ["step", "step"]); + assert_eq!(input.pos, 1); + assert_eq!(output.pos, 8); + } + + #[test] + fn compress_stream2_mt_coordinator_runs_completion_before_returning_error() { + let error = ERROR(ZstdErrorCode::Generic); + let mut context = CompressStream2MTCoordinatorTestContext { + events: Vec::new(), + step_results: vec![error], + step_positions: vec![(0, 0)], + step_calls: 0, + }; + let mut output = ZSTD_outBuffer { + dst: ptr::null_mut(), + size: 8, + pos: 0, + }; + let mut input = ZSTD_inBuffer { + src: ptr::null(), + size: 8, + pos: 0, + }; + let mut callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { + callback_context: ptr::null_mut(), + trace: Some(compress_stream2_mt_coordinator_test_trace), + reset: Some(compress_stream2_mt_coordinator_test_reset), + }; + let state = compress_stream2_mt_coordinator_test_state( + &mut context, + &mut output, + &mut input, + &mut callbacks, + ZSTD_E_FLUSH, + ); + + let result = unsafe { ZSTD_rust_compressStream2MTCoordinator(&state) }; + assert_eq!(result, error); + assert_eq!(context.step_calls, 1); + assert_eq!(context.events, ["step", "reset"]); + } + + #[test] + fn compress_stream2_mt_coordinator_traces_before_resetting_completed_end() { + let mut context = CompressStream2MTCoordinatorTestContext { + events: Vec::new(), + step_results: vec![0], + step_positions: vec![(0, 0)], + step_calls: 0, + }; + let mut output = ZSTD_outBuffer { + dst: ptr::null_mut(), + size: 8, + pos: 0, + }; + let mut input = ZSTD_inBuffer { + src: ptr::null(), + size: 8, + pos: 0, + }; + let mut callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { + callback_context: ptr::null_mut(), + trace: Some(compress_stream2_mt_coordinator_test_trace), + reset: Some(compress_stream2_mt_coordinator_test_reset), + }; + let state = compress_stream2_mt_coordinator_test_state( + &mut context, + &mut output, + &mut input, + &mut callbacks, + ZSTD_E_END, + ); + + let result = unsafe { ZSTD_rust_compressStream2MTCoordinator(&state) }; + assert_eq!(result, 0); + assert_eq!(context.step_calls, 1); + assert_eq!(context.events, ["step", "trace", "reset"]); + } + + #[test] + fn compress_stream2_mt_coordinator_rejects_null_or_malformed_state() { + assert!(ERR_isError(unsafe { + ZSTD_rust_compressStream2MTCoordinator(ptr::null()) + })); + + let mut context = CompressStream2MTCoordinatorTestContext { + events: Vec::new(), + step_results: vec![0], + step_positions: vec![(0, 0)], + step_calls: 0, + }; + let mut output = ZSTD_outBuffer { + dst: ptr::null_mut(), + size: 8, + pos: 0, + }; + let mut input = ZSTD_inBuffer { + src: ptr::null(), + size: 8, + pos: 0, + }; + let mut callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { + callback_context: ptr::null_mut(), + trace: Some(compress_stream2_mt_coordinator_test_trace), + reset: Some(compress_stream2_mt_coordinator_test_reset), + }; + let mut state = compress_stream2_mt_coordinator_test_state( + &mut context, + &mut output, + &mut input, + &mut callbacks, + ZSTD_E_FLUSH, + ); + + state.step = None; + assert!(ERR_isError(unsafe { + ZSTD_rust_compressStream2MTCoordinator(&state) + })); + state.step = Some(compress_stream2_mt_coordinator_test_step); + state.completion_callbacks = ptr::null(); + assert!(ERR_isError(unsafe { + ZSTD_rust_compressStream2MTCoordinator(&state) + })); + let malformed_callbacks = ZSTD_rust_compressStream2MTCompletionCallbacks { + callback_context: callbacks.callback_context, + trace: None, + reset: callbacks.reset, + }; + state.completion_callbacks = &malformed_callbacks; + assert!(ERR_isError(unsafe { + ZSTD_rust_compressStream2MTCoordinator(&state) + })); + assert!(context.events.is_empty()); + assert_eq!(context.step_calls, 0); + } + fn compress_stream2_init_policy_state( in_buffer_mode: c_int, end_op: c_int,