refactor(compress): move sequence generation policy to Rust

ZSTD_generateSequences() previously mixed public parameter validation,
temporary-output allocation, sequence-collector setup, compression, and cleanup
inside the C translation unit. Move that ordering into Rust behind callbacks so
private ZSTD_CCtx and SeqCollector layouts remain C-owned. The adapter retains
C allocation, context access, and the existing ZSTD_compress2 operation, while
Rust guarantees temporary-output cleanup after collector setup failures and
compression errors and preserves the original target-block-size and worker
parameter checks.

Focused callback tests cover validation-before-allocation, cleanup on setup
failure, and successful callback ordering/count publication.

Test Plan:
- `cargo fmt --manifest-path rust/Cargo.toml --all -- --check` -- passed
- `ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo test --manifest-path rust/Cargo.toml generate_sequences --lib` -- passed (3 tests)
- `ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo test --manifest-path rust/Cargo.toml --all-targets` -- passed (762 tests)
- `ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo clippy --manifest-path rust/Cargo.toml --all-targets -- -D warnings` -- passed
- `ulimit -v 41943040; make -j1` -- passed
This commit is contained in:
2026-07-20 04:30:33 +02:00
parent 53218c5537
commit 66e7799f04
2 changed files with 360 additions and 30 deletions
+95 -30
View File
@@ -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. */
+265
View File
@@ -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::<usize>());
assert!(offset_of!(ZSTD_rust_generateSequencesState, allocate) == 2 * size_of::<usize>());
assert!(offset_of!(ZSTD_rust_generateSequencesState, free) == 3 * size_of::<usize>());
assert!(offset_of!(ZSTD_rust_generateSequencesState, set_collector) == 4 * size_of::<usize>());
assert!(offset_of!(ZSTD_rust_generateSequencesState, compress2) == 5 * size_of::<usize>());
assert!(offset_of!(ZSTD_rust_generateSequencesState, get_count) == 6 * size_of::<usize>());
assert!(size_of::<ZSTD_rust_generateSequencesState>() == 7 * size_of::<usize>());
};
/// 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::<GenerateSequencesTestContext>() };
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::<GenerateSequencesTestContext>() };
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::<GenerateSequencesTestContext>() };
context.events.push("free");
if !pointer.is_null() {
unsafe { drop(Box::from_raw(pointer.cast::<u8>())) };
}
}
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::<GenerateSequencesTestContext>() };
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::<GenerateSequencesTestContext>() };
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::<GenerateSequencesTestContext>() };
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,