diff --git a/lib/decompress/zstd_decompress.c b/lib/decompress/zstd_decompress.c index 294b11e00..c8a68052c 100644 --- a/lib/decompress/zstd_decompress.c +++ b/lib/decompress/zstd_decompress.c @@ -133,7 +133,8 @@ void ZSTD_rust_dctx_view(ZSTD_DCtx* dctx, ZSTD_rustDctxView* out); void ZSTD_rust_dctx_trace_view(ZSTD_DCtx* dctx, ZSTD_rustDctxTraceView* out); size_t ZSTD_rust_dctx_sizeof(void); -void ZSTD_rust_dctx_init_platform(ZSTD_DCtx* dctx); +void ZSTD_dctx_init_platform(ZSTD_DCtx* dctx); +void ZSTD_rust_dctx_init_platform(void* bmi2, int dynamic_bmi2); size_t ZSTD_rust_dctx_default_max_window_size(void); int ZSTD_rust_no_forward_progress_max(void); int ZSTD_rust_heapmode(void); @@ -245,12 +246,13 @@ size_t ZSTD_rust_dctx_sizeof(void) return sizeof(ZSTD_DCtx); } -void ZSTD_rust_dctx_init_platform(ZSTD_DCtx* dctx) +void ZSTD_dctx_init_platform(ZSTD_DCtx* dctx) { #if DYNAMIC_BMI2 - dctx->bmi2 = ZSTD_cpuSupportsBmi2(); + ZSTD_rust_dctx_init_platform(&dctx->bmi2, DYNAMIC_BMI2); #else (void)dctx; + ZSTD_rust_dctx_init_platform(NULL, DYNAMIC_BMI2); #endif } diff --git a/rust/src/zstd_decompress.rs b/rust/src/zstd_decompress.rs index b1cb8fcf1..853cd85ec 100644 --- a/rust/src/zstd_decompress.rs +++ b/rust/src/zstd_decompress.rs @@ -344,7 +344,7 @@ unsafe extern "C" { fn ZSTD_rust_dctx_view(dctx: *mut ZSTD_DCtx, out: *mut ZSTD_rustDctxView); fn ZSTD_rust_dctx_trace_view(dctx: *mut ZSTD_DCtx, out: *mut ZSTD_rustDctxTraceView); fn ZSTD_rust_dctx_sizeof() -> usize; - fn ZSTD_rust_dctx_init_platform(dctx: *mut ZSTD_DCtx); + fn ZSTD_dctx_init_platform(dctx: *mut ZSTD_DCtx); fn ZSTD_rust_dctx_default_max_window_size() -> usize; fn ZSTD_rust_no_forward_progress_max() -> c_int; fn ZSTD_rust_heapmode() -> c_int; @@ -378,6 +378,37 @@ unsafe extern "C" { fn ZSTD_checkContinuity(dctx: *mut ZSTD_DCtx, dst: *const c_void, dst_size: usize); } +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_dctx_init_platform(bmi2: *mut c_void, dynamic_bmi2: c_int) { + if dynamic_bmi2 != 0 && !bmi2.is_null() { + unsafe { + bmi2.cast::() + .write(c_int::from(crate::cpu::ZSTD_cpuSupportsBmi2())); + } + } +} + +#[cfg(test)] +mod dctx_platform_tests { + use super::*; + + #[test] + fn initializes_bmi2_only_when_dynamic_support_is_enabled() { + let mut bmi2 = c_int::MAX; + unsafe { + ZSTD_rust_dctx_init_platform((&mut bmi2 as *mut c_int).cast(), 0); + } + assert_eq!(bmi2, c_int::MAX); + + let expected = c_int::from(crate::cpu::ZSTD_cpuSupportsBmi2()); + bmi2 = c_int::MAX; + unsafe { + ZSTD_rust_dctx_init_platform((&mut bmi2 as *mut c_int).cast(), 1); + } + assert_eq!(bmi2, expected); + } +} + type ZstdDecompressDCtxFn = unsafe extern "C" fn(*mut ZSTD_DCtx, *mut c_void, usize, *const c_void, usize) -> usize; @@ -2576,7 +2607,7 @@ unsafe fn init_dctx_internal(view: &ZSTD_rustDctxView) { if !view.fuzz_end.is_null() { set_pointer(view.fuzz_end, ptr::null()); } - ZSTD_rust_dctx_init_platform(view.dctx.cast()); + ZSTD_dctx_init_platform(view.dctx.cast()); } }