From 429694a250f9c6e3a82e5c878ceaaea9d58816cb Mon Sep 17 00:00:00 2001 From: ddidderr Date: Mon, 20 Jul 2026 07:31:23 +0200 Subject: [PATCH] refactor(mt): move stream reset ordering to Rust Make Rust own the MT stream reset sequence through a narrow callback projection. C retains the private buffer, job-ID, frame-flag, progress, and sentinel mutations, preserving the original order without exposing private layouts across the ABI. Test Plan: - rustfmt --edition 2021 --check rust/src/zstdmt_compress.rs - git diff --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 cargo test --manifest-path rust/Cargo.toml --all-targets (786 passed) - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo clippy --manifest-path rust/cli/Cargo.toml --all-targets -- -D warnings - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo test --manifest-path rust/cli/Cargo.toml --all-targets (179 passed) - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo +nightly fmt --manifest-path rust/Cargo.toml -- --check - ulimit -v 41943040; make -j1 - ulimit -v 41943040; make -j1 -C tests test (all tests completed successfully) --- lib/compress/zstdmt_compress.c | 73 +++++++++++++--- rust/src/zstdmt_compress.rs | 151 +++++++++++++++++++++++++++++++++ 2 files changed, 214 insertions(+), 10 deletions(-) diff --git a/lib/compress/zstdmt_compress.c b/lib/compress/zstdmt_compress.c index 6d95d8f52..e1c26ade9 100644 --- a/lib/compress/zstdmt_compress.c +++ b/lib/compress/zstdmt_compress.c @@ -695,6 +695,25 @@ size_t ZSTDMT_rust_initCStream( ZSTDMT_initResetStreamFn resetStream, ZSTDMT_initDictionaryFn updateDictionary, ZSTDMT_initSerialResetFn serialReset); +typedef void (*ZSTDMT_resetStreamCallbackFn)(void* opaque); +typedef struct { + void* callbackContext; + ZSTDMT_resetStreamCallbackFn resetRoundBuffer; + ZSTDMT_resetStreamCallbackFn resetInputBuffer; + ZSTDMT_resetStreamCallbackFn resetJobIDs; + ZSTDMT_resetStreamCallbackFn resetFrameFlags; + ZSTDMT_resetStreamCallbackFn resetProgress; +} ZSTDMT_RustResetStreamState; +typedef char ZSTDMT_rust_reset_stream_state_layout[ + (sizeof(ZSTDMT_resetStreamCallbackFn) == sizeof(void*) + && offsetof(ZSTDMT_RustResetStreamState, callbackContext) == 0 + && offsetof(ZSTDMT_RustResetStreamState, resetRoundBuffer) == sizeof(void*) + && offsetof(ZSTDMT_RustResetStreamState, resetInputBuffer) == 2 * sizeof(void*) + && offsetof(ZSTDMT_RustResetStreamState, resetJobIDs) == 3 * sizeof(void*) + && offsetof(ZSTDMT_RustResetStreamState, resetFrameFlags) == 4 * sizeof(void*) + && offsetof(ZSTDMT_RustResetStreamState, resetProgress) == 5 * sizeof(void*) + && sizeof(ZSTDMT_RustResetStreamState) == 6 * sizeof(void*)) ? 1 : -1]; +void ZSTDMT_rust_resetStream(const ZSTDMT_RustResetStreamState* state); typedef struct { void* callbackContext; unsigned nbWorkers; @@ -2377,21 +2396,55 @@ static size_t ZSTDMT_initCStreamResizeRoundBuffer(void* opaque, size_t capacity) return 0; } +static void ZSTDMT_initCStreamResetRoundBuffer(void* opaque) +{ + ZSTDMT_initCStreamState* const state = (ZSTDMT_initCStreamState*)opaque; + state->mtctx->roundBuff.pos = 0; +} + +static void ZSTDMT_initCStreamResetInputBuffer(void* opaque) +{ + ZSTDMT_initCStreamState* const state = (ZSTDMT_initCStreamState*)opaque; + state->mtctx->inBuff.buffer = g_nullBuffer; + state->mtctx->inBuff.filled = 0; + state->mtctx->inBuff.prefix = kNullRange; +} + +static void ZSTDMT_initCStreamResetJobIDs(void* opaque) +{ + ZSTDMT_initCStreamState* const state = (ZSTDMT_initCStreamState*)opaque; + state->mtctx->doneJobID = 0; + state->mtctx->nextJobID = 0; +} + +static void ZSTDMT_initCStreamResetFrameFlags(void* opaque) +{ + ZSTDMT_initCStreamState* const state = (ZSTDMT_initCStreamState*)opaque; + state->mtctx->frameEnded = 0; + state->mtctx->allJobsCompleted = 0; +} + +static void ZSTDMT_initCStreamResetProgress(void* opaque) +{ + ZSTDMT_initCStreamState* const state = (ZSTDMT_initCStreamState*)opaque; + state->mtctx->consumed = 0; + state->mtctx->produced = 0; +} + static void ZSTDMT_initCStreamResetStream(void* opaque) { ZSTDMT_initCStreamState* const state = (ZSTDMT_initCStreamState*)opaque; ZSTDMT_CCtx* const mtctx = state->mtctx; + ZSTDMT_RustResetStreamState const resetState = { + state, + ZSTDMT_initCStreamResetRoundBuffer, + ZSTDMT_initCStreamResetInputBuffer, + ZSTDMT_initCStreamResetJobIDs, + ZSTDMT_initCStreamResetFrameFlags, + ZSTDMT_initCStreamResetProgress + }; DEBUGLOG(4, "roundBuff capacity : %u KB", (U32)(mtctx->roundBuff.capacity >> 10)); - mtctx->roundBuff.pos = 0; - mtctx->inBuff.buffer = g_nullBuffer; - mtctx->inBuff.filled = 0; - mtctx->inBuff.prefix = kNullRange; - mtctx->doneJobID = 0; - mtctx->nextJobID = 0; - mtctx->frameEnded = 0; - mtctx->allJobsCompleted = 0; - mtctx->consumed = 0; - mtctx->produced = 0; + ZSTDMT_rust_resetStream(&resetState); } static size_t ZSTDMT_initCStreamUpdateDictionary(void* opaque) diff --git a/rust/src/zstdmt_compress.rs b/rust/src/zstdmt_compress.rs index 302abdeea..68db326ca 100644 --- a/rust/src/zstdmt_compress.rs +++ b/rust/src/zstdmt_compress.rs @@ -791,12 +791,74 @@ pub type ZSTDMT_initSetBufferSizeFn = unsafe extern "C" fn(*mut c_void, usize); pub type ZSTDMT_initResizeRoundBufferFn = unsafe extern "C" fn(*mut c_void, usize) -> usize; pub type ZSTDMT_initResetStreamFn = unsafe extern "C" fn(*mut c_void); pub type ZSTDMT_initSerialResetFn = unsafe extern "C" fn(*mut c_void, usize) -> usize; +type ZSTDMT_resetStreamCallbackFn = unsafe extern "C" fn(*mut c_void); pub type ZSTDMT_serialResetVoidFn = unsafe extern "C" fn(*mut c_void); pub type ZSTDMT_serialResetSetNbSeqFn = unsafe extern "C" fn(*mut c_void, usize); pub type ZSTDMT_serialResetResizeFn = unsafe extern "C" fn(*mut c_void) -> c_int; pub type ZSTDMT_serialStateVoidFn = unsafe extern "C" fn(*mut c_void); pub type ZSTDMT_serialStateInitFn = unsafe extern "C" fn(*mut c_void) -> c_int; +/// Callback projection for the MT stream reset. C retains the private buffer, +/// job, flag, and progress fields; Rust owns the order in which each group is +/// reset. +#[repr(C)] +pub struct ZSTDMT_RustResetStreamState { + callback_context: *mut c_void, + reset_round_buffer: Option, + reset_input_buffer: Option, + reset_job_ids: Option, + reset_frame_flags: Option, + reset_progress: Option, +} + +const _: () = { + assert!(size_of::() == size_of::()); + assert!(offset_of!(ZSTDMT_RustResetStreamState, callback_context) == 0); + assert!(offset_of!(ZSTDMT_RustResetStreamState, reset_round_buffer) == size_of::()); + assert!(offset_of!(ZSTDMT_RustResetStreamState, reset_input_buffer) == 2 * size_of::()); + assert!(offset_of!(ZSTDMT_RustResetStreamState, reset_job_ids) == 3 * size_of::()); + assert!(offset_of!(ZSTDMT_RustResetStreamState, reset_frame_flags) == 4 * size_of::()); + assert!(offset_of!(ZSTDMT_RustResetStreamState, reset_progress) == 5 * size_of::()); + assert!(size_of::() == 6 * size_of::()); +}; + +/// Reset the MT stream state in the same order as the original C function. +/// Each callback keeps the corresponding private C layout and sentinel values +/// out of the Rust ABI. +#[no_mangle] +pub unsafe extern "C" fn ZSTDMT_rust_resetStream(state: *const ZSTDMT_RustResetStreamState) { + let Some(state) = (unsafe { state.as_ref() }) else { + return; + }; + if state.callback_context.is_null() { + return; + } + let ( + Some(reset_round_buffer), + Some(reset_input_buffer), + Some(reset_job_ids), + Some(reset_frame_flags), + Some(reset_progress), + ) = ( + state.reset_round_buffer, + state.reset_input_buffer, + state.reset_job_ids, + state.reset_frame_flags, + state.reset_progress, + ) + else { + return; + }; + + unsafe { + reset_round_buffer(state.callback_context); + reset_input_buffer(state.callback_context); + reset_job_ids(state.callback_context); + reset_frame_flags(state.callback_context); + reset_progress(state.callback_context); + } +} + #[repr(C)] #[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] pub struct ZSTDMT_flushPublicationResult { @@ -5345,6 +5407,95 @@ mod tests { } } + #[derive(Default)] + struct ResetStreamTestContext { + events: Vec<&'static str>, + } + + unsafe extern "C" fn reset_stream_test_round_buffer(context: *mut c_void) { + unsafe { + (*context.cast::()) + .events + .push("round"); + } + } + + unsafe extern "C" fn reset_stream_test_input_buffer(context: *mut c_void) { + unsafe { + (*context.cast::()) + .events + .push("input"); + } + } + + unsafe extern "C" fn reset_stream_test_job_ids(context: *mut c_void) { + unsafe { + (*context.cast::()) + .events + .push("jobs"); + } + } + + unsafe extern "C" fn reset_stream_test_frame_flags(context: *mut c_void) { + unsafe { + (*context.cast::()) + .events + .push("flags"); + } + } + + unsafe extern "C" fn reset_stream_test_progress(context: *mut c_void) { + unsafe { + (*context.cast::()) + .events + .push("progress"); + } + } + + fn reset_stream_test_state(context: *mut c_void) -> ZSTDMT_RustResetStreamState { + ZSTDMT_RustResetStreamState { + callback_context: context, + reset_round_buffer: Some(reset_stream_test_round_buffer), + reset_input_buffer: Some(reset_stream_test_input_buffer), + reset_job_ids: Some(reset_stream_test_job_ids), + reset_frame_flags: Some(reset_stream_test_frame_flags), + reset_progress: Some(reset_stream_test_progress), + } + } + + #[test] + fn reset_stream_preserves_callback_order() { + let mut context = ResetStreamTestContext::default(); + let state = reset_stream_test_state((&mut context as *mut ResetStreamTestContext).cast()); + + unsafe { ZSTDMT_rust_resetStream(&state) }; + + assert_eq!( + context.events, + vec!["round", "input", "jobs", "flags", "progress"] + ); + } + + #[test] + fn reset_stream_ignores_null_or_incomplete_state() { + let mut context = ResetStreamTestContext::default(); + let state = ZSTDMT_RustResetStreamState { + callback_context: (&mut context as *mut ResetStreamTestContext).cast(), + reset_round_buffer: None, + reset_input_buffer: Some(reset_stream_test_input_buffer), + reset_job_ids: Some(reset_stream_test_job_ids), + reset_frame_flags: Some(reset_stream_test_frame_flags), + reset_progress: Some(reset_stream_test_progress), + }; + + unsafe { + ZSTDMT_rust_resetStream(&state); + ZSTDMT_rust_resetStream(ptr::null()); + } + + assert!(context.events.is_empty()); + } + #[derive(Default)] struct SerialResetTestContext { events: Vec<&'static str>,