diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 6b1fbb4e5..26867aa4d 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -852,6 +852,44 @@ typedef char ZSTD_rust_compress_stream2_init_policy_state_layout[ && sizeof(ZSTD_rust_compressStream2InitPolicyState) == 2 * sizeof(int) + 7 * sizeof(size_t)) ? 1 : -1]; +typedef struct { + int endOp; + size_t flushMin; + size_t inputPos; + size_t inputSize; + size_t inputPosBefore; + size_t outputPos; + size_t outputSize; + size_t outputPosBefore; +} ZSTD_rust_compressStream2MTLoopPolicyState; +enum { + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_CONTINUE = 0, + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK = 1, + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR = 2, + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE = 3 +}; +int ZSTD_rust_compressStream2MTLoopPolicy( + const ZSTD_rust_compressStream2MTLoopPolicyState* state); +typedef char ZSTD_rust_compress_stream2_mt_loop_policy_state_layout[ + (offsetof(ZSTD_rust_compressStream2MTLoopPolicyState, endOp) == 0 + && offsetof(ZSTD_rust_compressStream2MTLoopPolicyState, flushMin) + == sizeof(size_t) + && offsetof(ZSTD_rust_compressStream2MTLoopPolicyState, inputPos) + == 2 * sizeof(size_t) + && offsetof(ZSTD_rust_compressStream2MTLoopPolicyState, inputSize) + == 3 * sizeof(size_t) + && offsetof(ZSTD_rust_compressStream2MTLoopPolicyState, inputPosBefore) + == 4 * sizeof(size_t) + && offsetof(ZSTD_rust_compressStream2MTLoopPolicyState, outputPos) + == 5 * sizeof(size_t) + && offsetof(ZSTD_rust_compressStream2MTLoopPolicyState, outputSize) + == 6 * sizeof(size_t) + && offsetof(ZSTD_rust_compressStream2MTLoopPolicyState, + outputPosBefore) + == 7 * sizeof(size_t) + && sizeof(ZSTD_rust_compressStream2MTLoopPolicyState) + == 8 * sizeof(size_t)) + ? 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 { @@ -8458,27 +8496,27 @@ size_t ZSTD_compressStream2_c( ZSTD_CCtx* cctx, flushMin = ZSTDMT_compressStream_generic(cctx->mtctx, output, input, endOp); cctx->consumedSrcSize += (U64)(input->pos - ipos); cctx->producedCSize += (U64)(output->pos - opos); - if ( ZSTD_isError(flushMin) - || (endOp == ZSTD_e_end && flushMin == 0) ) { /* compression completed */ - if (flushMin == 0) - ZSTD_CCtx_trace(cctx, 0); - ZSTD_CCtx_reset(cctx, ZSTD_reset_session_only); - } - FORWARD_IF_ERROR(flushMin, "ZSTDMT_compressStream_generic failed"); - - if (endOp == ZSTD_e_continue) { - /* We only require some progress with ZSTD_e_continue, not maximal progress. - * We're done if we've consumed or produced any bytes, or either buffer is - * full. - */ - if (input->pos != ipos || output->pos != opos || input->pos == input->size || output->pos == output->size) - break; - } else { - assert(endOp == ZSTD_e_flush || endOp == ZSTD_e_end); - /* We require maximal progress. We're done when the flush is complete or the - * output buffer is full. - */ - if (flushMin == 0 || output->pos == output->size) + { 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); + if (loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR + || loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE) { + if (loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE) + ZSTD_CCtx_trace(cctx, 0); + ZSTD_CCtx_reset(cctx, ZSTD_reset_session_only); + } + 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; } } diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 82e83b61c..39026cd4b 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -361,6 +361,100 @@ pub unsafe extern "C" fn ZSTD_rust_compressStream2InitPolicy( compress_stream2_init_policy(state) } +/// Scalar projection for the MT loop after a C-owned compression call. +/// +/// Rust owns the result classification and progress/termination decision. C +/// retains the MT call, private progress counters, trace callback, reset, and +/// error encoding. The ordering matches `ZSTD_compressStream2_c`: errors win, +/// then a completed end directive, then the directive-specific progress rule. +#[repr(C)] +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub struct ZSTD_rust_compressStream2MTLoopPolicyState { + pub end_op: c_int, + pub flush_min: usize, + pub input_pos: usize, + pub input_size: usize, + pub input_pos_before: usize, + pub output_pos: usize, + pub output_size: usize, + pub output_pos_before: usize, +} + +const ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_CONTINUE: c_int = 0; +const ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK: c_int = 1; +const ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR: c_int = 2; +const ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE: c_int = 3; + +const _: () = { + assert!(offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, end_op) == 0); + assert!( + offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, flush_min) == size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, input_pos) == 2 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, input_size) + == 3 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, input_pos_before) + == 4 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, output_pos) + == 5 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, output_size) + == 6 * size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_compressStream2MTLoopPolicyState, + output_pos_before + ) == 7 * size_of::() + ); + assert!(size_of::() == 8 * size_of::()); +}; + +#[inline] +fn compress_stream2_mt_loop_policy(state: &ZSTD_rust_compressStream2MTLoopPolicyState) -> c_int { + if ERR_isError(state.flush_min) { + return ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR; + } + if state.end_op == ZSTD_E_END && state.flush_min == 0 { + return ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE; + } + + let stop = if state.end_op == ZSTD_E_CONTINUE { + state.input_pos != state.input_pos_before + || state.output_pos != state.output_pos_before + || state.input_pos == state.input_size + || state.output_pos == state.output_size + } else { + debug_assert!(state.end_op == ZSTD_E_FLUSH || state.end_op == ZSTD_E_END); + state.flush_min == 0 || state.output_pos == state.output_size + }; + if stop { + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK + } else { + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_CONTINUE + } +} + +/// Classify one MT compression result while leaving all private state changes +/// and callbacks in the C caller. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressStream2MTLoopPolicy( + state: *const ZSTD_rust_compressStream2MTLoopPolicyState, +) -> c_int { + let Some(state) = (unsafe { state.as_ref() }) else { + return ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR; + }; + compress_stream2_mt_loop_policy(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; @@ -14326,6 +14420,92 @@ mod tests { assert_eq!(unsafe { ZSTD_rust_compressStream2Policy(ptr::null()) }, 0); } + fn compress_stream2_mt_loop_policy_state( + end_op: c_int, + flush_min: usize, + input_pos: usize, + input_size: usize, + input_pos_before: usize, + output_pos: usize, + output_size: usize, + output_pos_before: usize, + ) -> ZSTD_rust_compressStream2MTLoopPolicyState { + ZSTD_rust_compressStream2MTLoopPolicyState { + end_op, + flush_min, + input_pos, + input_size, + input_pos_before, + output_pos, + output_size, + output_pos_before, + } + } + + #[test] + fn compress_stream2_mt_loop_policy_classifies_errors_before_completion() { + let state = compress_stream2_mt_loop_policy_state( + ZSTD_E_END, + ERROR(ZstdErrorCode::Generic), + 3, + 3, + 3, + 4, + 4, + 4, + ); + assert_eq!( + compress_stream2_mt_loop_policy(&state), + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR + ); + } + + #[test] + fn compress_stream2_mt_loop_policy_classifies_completed_end() { + let state = compress_stream2_mt_loop_policy_state(ZSTD_E_END, 0, 3, 8, 3, 4, 8, 4); + assert_eq!( + compress_stream2_mt_loop_policy(&state), + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE + ); + } + + #[test] + fn compress_stream2_mt_loop_policy_breaks_continue_after_input_or_output_progress() { + for (input_pos, output_pos) in [(4, 3), (3, 4)] { + let state = compress_stream2_mt_loop_policy_state( + ZSTD_E_CONTINUE, + 7, + input_pos, + 8, + 3, + output_pos, + 8, + 3, + ); + assert_eq!( + compress_stream2_mt_loop_policy(&state), + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK + ); + } + } + + #[test] + fn compress_stream2_mt_loop_policy_handles_pending_flush_and_end_output() { + for end_op in [ZSTD_E_FLUSH, ZSTD_E_END] { + let pending_output = compress_stream2_mt_loop_policy_state(end_op, 7, 3, 8, 3, 4, 8, 4); + assert_eq!( + compress_stream2_mt_loop_policy(&pending_output), + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_CONTINUE + ); + + let full_output = compress_stream2_mt_loop_policy_state(end_op, 7, 3, 8, 3, 8, 8, 8); + assert_eq!( + compress_stream2_mt_loop_policy(&full_output), + ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK + ); + } + } + fn compress_stream2_init_policy_state( in_buffer_mode: c_int, end_op: c_int,