diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 26867aa4d..b7782611d 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -774,6 +774,43 @@ size_t ZSTD_compressStream2_c(ZSTD_CCtx* cctx, ZSTD_outBuffer* output, ZSTD_inBuffer* input, ZSTD_EndDirective endOp); +typedef void (*ZSTD_rust_compressStream2SetBufferExpectations_f)( + void* context, const void* output, const void* input); +typedef struct { + void* callbackContext; + ZSTD_rust_compressStream2SetBufferExpectations_f setBufferExpectations; + const void* output; + const void* input; + size_t compressResult; + size_t outBuffContentSize; + size_t outBuffFlushedSize; +} ZSTD_rust_compressStream2ResultPolicyState; +size_t ZSTD_rust_compressStream2ResultPolicy( + const ZSTD_rust_compressStream2ResultPolicyState* state); +typedef char ZSTD_rust_compress_stream2_result_policy_state_layout[ + (offsetof(ZSTD_rust_compressStream2ResultPolicyState, callbackContext) + == 0 + && offsetof(ZSTD_rust_compressStream2ResultPolicyState, + setBufferExpectations) + == sizeof(void*) + && offsetof(ZSTD_rust_compressStream2ResultPolicyState, output) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_compressStream2ResultPolicyState, input) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_compressStream2ResultPolicyState, + compressResult) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_compressStream2ResultPolicyState, + outBuffContentSize) + == 5 * sizeof(void*) + && offsetof(ZSTD_rust_compressStream2ResultPolicyState, + outBuffFlushedSize) + == 6 * sizeof(void*) + && sizeof(ZSTD_rust_compressStream2SetBufferExpectations_f) + == sizeof(void*) + && sizeof(ZSTD_rust_compressStream2ResultPolicyState) + == 7 * sizeof(void*)) + ? 1 : -1]; typedef struct { size_t outputPos; size_t outputSize; @@ -8005,6 +8042,15 @@ ZSTD_setBufferExpectations(ZSTD_CCtx* cctx, const ZSTD_outBuffer* output, const input); } +static void ZSTD_rust_compressStream2_setBufferExpectations( + void* context, const void* output, const void* input) +{ + ZSTD_setBufferExpectations( + (ZSTD_CCtx*)context, + (const ZSTD_outBuffer*)output, + (const ZSTD_inBuffer*)input); +} + /* Validate that the input/output buffers match the expectations set by * ZSTD_setBufferExpectations. */ @@ -8529,10 +8575,27 @@ size_t ZSTD_compressStream2_c( ZSTD_CCtx* cctx, return flushMin; } #endif /* ZSTD_MULTITHREAD */ - FORWARD_IF_ERROR( ZSTD_compressStream_generic(cctx, output, input, endOp) , ""); - DEBUGLOG(5, "completed ZSTD_compressStream2"); - ZSTD_setBufferExpectations(cctx, output, input); - return cctx->outBuffContentSize - cctx->outBuffFlushedSize; /* remaining to flush */ + { size_t const compressResult = + ZSTD_compressStream_generic(cctx, output, input, endOp); + if (!ERR_isError(compressResult)) { + DEBUGLOG(5, "completed ZSTD_compressStream2"); + } + { ZSTD_rust_compressStream2ResultPolicyState const state = { + cctx, + ZSTD_rust_compressStream2_setBufferExpectations, + output, + input, + compressResult, + cctx->outBuffContentSize, + cctx->outBuffFlushedSize + }; + { size_t const policyResult = + ZSTD_rust_compressStream2ResultPolicy(&state); + FORWARD_IF_ERROR(policyResult, ""); + return policyResult; + } + } + } } size_t ZSTD_compressStream2_simpleArgs ( diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 85cd8327c..79bfc94c4 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -257,6 +257,96 @@ pub unsafe extern "C" fn ZSTD_rust_compressStream2Policy( compress_stream2_policy(state.end_op) } +type CompressStream2SetBufferExpectationsFn = + unsafe extern "C" fn(*mut c_void, *const c_void, *const c_void); + +/// Projection for the post-call result policy of `ZSTD_compressStream2_c`. +/// +/// C retains the stream adapter and the private buffer-expectation fields. +/// Rust owns the error short-circuit, publication ordering, and pending-output +/// subtraction after the adapter has returned. +#[repr(C)] +pub struct ZSTD_rust_compressStream2ResultPolicyState { + callback_context: *mut c_void, + set_buffer_expectations: Option, + output: *const c_void, + input: *const c_void, + compress_result: usize, + out_buff_content_size: usize, + out_buff_flushed_size: usize, +} + +const _: () = { + assert!(size_of::>() == size_of::()); + assert!( + offset_of!(ZSTD_rust_compressStream2ResultPolicyState, callback_context) == 0 + ); + assert!( + offset_of!( + ZSTD_rust_compressStream2ResultPolicyState, + set_buffer_expectations + ) == size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2ResultPolicyState, output) + == 2 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2ResultPolicyState, input) + == 3 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStream2ResultPolicyState, compress_result) + == 4 * size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_compressStream2ResultPolicyState, + out_buff_content_size + ) == 5 * size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_compressStream2ResultPolicyState, + out_buff_flushed_size + ) == 6 * size_of::() + ); + assert!( + size_of::() == 7 * size_of::() + ); +}; + +#[inline] +unsafe fn compress_stream2_result_policy( + state: &ZSTD_rust_compressStream2ResultPolicyState, +) -> usize { + if ERR_isError(state.compress_result) { + return state.compress_result; + } + + let Some(set_buffer_expectations) = state.set_buffer_expectations else { + return ERROR(ZstdErrorCode::Generic); + }; + unsafe { + set_buffer_expectations(state.callback_context, state.output, state.input); + } + state + .out_buff_content_size + .wrapping_sub(state.out_buff_flushed_size) +} + +/// Apply the public stream result policy without exposing the private C context +/// layout to Rust. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressStream2ResultPolicy( + state: *const ZSTD_rust_compressStream2ResultPolicyState, +) -> usize { + let Some(state) = (unsafe { state.as_ref() }) else { + return ERROR(ZstdErrorCode::Generic); + }; + unsafe { compress_stream2_result_policy(state) } +} + /// Projection for the stable-input part of `ZSTD_compressStream2_c`'s /// transparent initialization stage. /// @@ -14426,6 +14516,69 @@ mod tests { assert_eq!(unsafe { ZSTD_rust_compressStream2Policy(ptr::null()) }, 0); } + #[derive(Default)] + struct CompressStream2ResultPolicyTestContext { + set_buffer_expectations_calls: usize, + } + + unsafe extern "C" fn compress_stream2_result_policy_test_set_buffer_expectations( + context: *mut c_void, + output: *const c_void, + input: *const c_void, + ) { + assert!(!output.is_null()); + assert!(!input.is_null()); + let context = unsafe { &mut *context.cast::() }; + context.set_buffer_expectations_calls += 1; + } + + #[test] + fn compress_stream2_result_policy_publishes_buffers_before_accounting() { + let mut context = CompressStream2ResultPolicyTestContext::default(); + let context_ptr = (&mut context as *mut CompressStream2ResultPolicyTestContext).cast(); + let state = ZSTD_rust_compressStream2ResultPolicyState { + callback_context: context_ptr, + set_buffer_expectations: Some( + compress_stream2_result_policy_test_set_buffer_expectations, + ), + output: context_ptr, + input: context_ptr, + compress_result: 0, + out_buff_content_size: 17, + out_buff_flushed_size: 5, + }; + + assert_eq!( + unsafe { ZSTD_rust_compressStream2ResultPolicy(&state) }, + 12 + ); + assert_eq!(context.set_buffer_expectations_calls, 1); + } + + #[test] + fn compress_stream2_result_policy_returns_adapter_error_without_publishing() { + let mut context = CompressStream2ResultPolicyTestContext::default(); + let context_ptr = (&mut context as *mut CompressStream2ResultPolicyTestContext).cast(); + let adapter_error = ERROR(ZstdErrorCode::Generic); + let state = ZSTD_rust_compressStream2ResultPolicyState { + callback_context: context_ptr, + set_buffer_expectations: Some( + compress_stream2_result_policy_test_set_buffer_expectations, + ), + output: ptr::null(), + input: ptr::null(), + compress_result: adapter_error, + out_buff_content_size: 17, + out_buff_flushed_size: 5, + }; + + assert_eq!( + unsafe { ZSTD_rust_compressStream2ResultPolicy(&state) }, + adapter_error + ); + assert_eq!(context.set_buffer_expectations_calls, 0); + } + fn compress_stream2_mt_loop_policy_state( end_op: c_int, flush_min: usize,