refactor(compress): move MT stream loop policy to Rust

The multithreaded branch of `ZSTD_compressStream2_c` performed the complete
post-call result classification in C: it prioritized errors, recognized a
completed end directive, and applied different progress rules for continue,
flush, and end operations. Those decisions were scalar policy around a C-owned
MT call, but remained embedded beside private counters, reset, and tracing.

Rust now classifies one projected iteration and returns an explicit action for
continue, break, error, or completed end. The C shim still owns the MT call,
input/output accounting, error forwarding, private context mutation, reset,
and trace callback; the existing condition ordering and progress comparisons
are preserved. ABI layout assertions and focused tests cover error precedence,
completed end, progress-based continue termination, and pending/full output.

Test Plan:
- `git diff --cached --check` -- passed
- `rustfmt --edition 2021 --check rust/src/zstd_compress.rs` -- passed
- Full capped Rust/native verification is the next serial step.
This commit is contained in:
2026-07-21 06:53:09 +02:00
parent 35150c3cff
commit 0faf3e4cf4
2 changed files with 239 additions and 21 deletions
+59 -21
View File
@@ -852,6 +852,44 @@ typedef char ZSTD_rust_compress_stream2_init_policy_state_layout[
&& sizeof(ZSTD_rust_compressStream2InitPolicyState) && sizeof(ZSTD_rust_compressStream2InitPolicyState)
== 2 * sizeof(int) + 7 * sizeof(size_t)) == 2 * sizeof(int) + 7 * sizeof(size_t))
? 1 : -1]; ? 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); int ZSTD_rust_simpleCompress2Level(const void* cctx, size_t srcSize);
typedef int (*ZSTD_rust_simpleCompress2Level_f)(const void* cctx, size_t srcSize); typedef int (*ZSTD_rust_simpleCompress2Level_f)(const void* cctx, size_t srcSize);
typedef struct { typedef struct {
@@ -8458,27 +8496,27 @@ size_t ZSTD_compressStream2_c( ZSTD_CCtx* cctx,
flushMin = ZSTDMT_compressStream_generic(cctx->mtctx, output, input, endOp); flushMin = ZSTDMT_compressStream_generic(cctx->mtctx, output, input, endOp);
cctx->consumedSrcSize += (U64)(input->pos - ipos); cctx->consumedSrcSize += (U64)(input->pos - ipos);
cctx->producedCSize += (U64)(output->pos - opos); cctx->producedCSize += (U64)(output->pos - opos);
if ( ZSTD_isError(flushMin) { ZSTD_rust_compressStream2MTLoopPolicyState const state = {
|| (endOp == ZSTD_e_end && flushMin == 0) ) { /* compression completed */ (int)endOp,
if (flushMin == 0) flushMin,
ZSTD_CCtx_trace(cctx, 0); input->pos,
ZSTD_CCtx_reset(cctx, ZSTD_reset_session_only); input->size,
} ipos,
FORWARD_IF_ERROR(flushMin, "ZSTDMT_compressStream_generic failed"); output->pos,
output->size,
if (endOp == ZSTD_e_continue) { opos
/* 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 int const loopPolicy =
* full. ZSTD_rust_compressStream2MTLoopPolicy(&state);
*/ if (loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_ERROR
if (input->pos != ipos || output->pos != opos || input->pos == input->size || output->pos == output->size) || loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE) {
break; if (loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE)
} else { ZSTD_CCtx_trace(cctx, 0);
assert(endOp == ZSTD_e_flush || endOp == ZSTD_e_end); ZSTD_CCtx_reset(cctx, ZSTD_reset_session_only);
/* We require maximal progress. We're done when the flush is complete or the }
* output buffer is full. FORWARD_IF_ERROR(flushMin, "ZSTDMT_compressStream_generic failed");
*/ if (loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_BREAK
if (flushMin == 0 || output->pos == output->size) || loopPolicy == ZSTD_RUST_COMPRESS_STREAM2_MT_LOOP_POLICY_END_COMPLETE)
break; break;
} }
} }
+180
View File
@@ -361,6 +361,100 @@ pub unsafe extern "C" fn ZSTD_rust_compressStream2InitPolicy(
compress_stream2_init_policy(state) 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::<usize>()
);
assert!(
offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, input_pos) == 2 * size_of::<usize>()
);
assert!(
offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, input_size)
== 3 * size_of::<usize>()
);
assert!(
offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, input_pos_before)
== 4 * size_of::<usize>()
);
assert!(
offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, output_pos)
== 5 * size_of::<usize>()
);
assert!(
offset_of!(ZSTD_rust_compressStream2MTLoopPolicyState, output_size)
== 6 * size_of::<usize>()
);
assert!(
offset_of!(
ZSTD_rust_compressStream2MTLoopPolicyState,
output_pos_before
) == 7 * size_of::<usize>()
);
assert!(size_of::<ZSTD_rust_compressStream2MTLoopPolicyState>() == 8 * size_of::<usize>());
};
#[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_WINDOW_LOG: c_int = 101;
const ZSTD_C_HASH_LOG: c_int = 102; const ZSTD_C_HASH_LOG: c_int = 102;
const ZSTD_C_CHAIN_LOG: c_int = 103; const ZSTD_C_CHAIN_LOG: c_int = 103;
@@ -14326,6 +14420,92 @@ mod tests {
assert_eq!(unsafe { ZSTD_rust_compressStream2Policy(ptr::null()) }, 0); 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( fn compress_stream2_init_policy_state(
in_buffer_mode: c_int, in_buffer_mode: c_int,
end_op: c_int, end_op: c_int,