diff --git a/lib/common/zstd_internal.h b/lib/common/zstd_internal.h index 2789a3591..5e612e1df 100644 --- a/lib/common/zstd_internal.h +++ b/lib/common/zstd_internal.h @@ -167,20 +167,11 @@ static UNUSED_ATTR const U32 OF_defaultNormLog = OF_DEFAULTNORMLOG; /*-******************************************* * Shared functions to include for inlining *********************************************/ -static void ZSTD_copy8(void* dst, const void* src) { -#if defined(ZSTD_ARCH_ARM_NEON) - vst1_u8((uint8_t*)dst, vld1_u8((const uint8_t*)src)); -#else - ZSTD_memcpy(dst, src, 8); -#endif -} -#define COPY8(d,s) do { ZSTD_copy8(d,s); d+=8; s+=8; } while (0) - /* Need to use memmove here since the literal buffer can now be located within the dst buffer. In circumstances where the op "catches up" to where the literal buffer is, there can be partial overlaps in this call on the final copy if the literal is being shifted by less than 16 bytes. */ -static void ZSTD_copy16(void* dst, const void* src) { +static UNUSED_ATTR void ZSTD_copy16(void* dst, const void* src) { #if defined(ZSTD_ARCH_ARM_NEON) vst1q_u8((uint8_t*)dst, vld1q_u8((const uint8_t*)src)); #elif defined(ZSTD_ARCH_X86_SSE2) @@ -205,6 +196,8 @@ typedef enum { /* ZSTD_overlap_dst_before_src, */ } ZSTD_overlap_e; +void ZSTD_rust_wildcopy(void* dst, const void* src, ptrdiff_t length, int ovtype); + /*! ZSTD_wildcopy() : * Custom version of ZSTD_memcpy(), can over read/write up to WILDCOPY_OVERLENGTH bytes (if length==0) * @param ovtype controls the overlap detection @@ -215,34 +208,7 @@ typedef enum { MEM_STATIC FORCE_INLINE_ATTR void ZSTD_wildcopy(void* dst, const void* src, ptrdiff_t length, ZSTD_overlap_e const ovtype) { - ptrdiff_t diff = (BYTE*)dst - (const BYTE*)src; - const BYTE* ip = (const BYTE*)src; - BYTE* op = (BYTE*)dst; - BYTE* const oend = op + length; - - if (ovtype == ZSTD_overlap_src_before_dst && diff < WILDCOPY_VECLEN) { - /* Handle short offset copies. */ - do { - COPY8(op, ip); - } while (op < oend); - } else { - assert(diff >= WILDCOPY_VECLEN || diff <= -WILDCOPY_VECLEN); - /* Separate out the first COPY16() call because the copy length is - * almost certain to be short, so the branches have different - * probabilities. Since it is almost certain to be short, only do - * one COPY16() in the first call. Then, do two calls per loop since - * at that point it is more likely to have a high trip count. - */ - ZSTD_copy16(op, ip); - if (16 >= length) return; - op += 16; - ip += 16; - do { - COPY16(op, ip); - COPY16(op, ip); - } - while (op < oend); - } + ZSTD_rust_wildcopy(dst, src, length, (int)ovtype); } MEM_STATIC size_t ZSTD_limitCopy(void* dst, size_t dstCapacity, const void* src, size_t srcSize) diff --git a/rust/src/mem.rs b/rust/src/mem.rs index 8ff14630b..0c7109c69 100644 --- a/rust/src/mem.rs +++ b/rust/src/mem.rs @@ -1,5 +1,6 @@ #![allow(non_snake_case)] use std::mem; +use std::os::raw::c_int; pub type U8 = u8; pub type S8 = i8; @@ -202,6 +203,7 @@ pub unsafe fn ZSTD_wildcopy( let diff = (op as isize).wrapping_sub(ip as isize); if ovtype == ZstdOverlap::OverlapSrcBeforeDst && diff < WILDCOPY_VECLEN { + // Handle short offset copies. loop { std::ptr::copy_nonoverlapping(ip, op, 8); op = op.add(8); @@ -211,23 +213,48 @@ pub unsafe fn ZSTD_wildcopy( } } } else { + assert!(diff >= WILDCOPY_VECLEN || diff <= -WILDCOPY_VECLEN); + // Keep the first COPY16 separate because short copies are common. std::ptr::copy_nonoverlapping(ip, op, 16); - if length <= 16 { + if WILDCOPY_VECLEN >= length { return; } op = op.add(16); ip = ip.add(16); - while op < oend { + loop { std::ptr::copy_nonoverlapping(ip, op, 16); op = op.add(16); ip = ip.add(16); std::ptr::copy_nonoverlapping(ip, op, 16); op = op.add(16); ip = ip.add(16); + if op >= oend { + break; + } } } } +/// Rust implementation of the C `ZSTD_wildcopy()` ABI wrapper. +/// +/// # Safety +/// `dst`/`src` must be valid for the wildcopy contract (may over-read/write up to 32 bytes), +/// and `ovtype` must be one of the two `ZSTD_overlap_e` values. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_wildcopy( + dst: *mut std::ffi::c_void, + src: *const std::ffi::c_void, + length: isize, + ovtype: c_int, +) { + let ovtype = match ovtype { + 0 => ZstdOverlap::NoOverlap, + 1 => ZstdOverlap::OverlapSrcBeforeDst, + _ => unreachable!("invalid ZSTD_overlap_e value"), + }; + ZSTD_wildcopy(dst, src, length, ovtype); +} + /// # Safety /// `dst` must be valid for `dst_capacity` bytes; `src` for `src_size` bytes. #[inline] @@ -271,40 +298,101 @@ mod tests { } } - #[test] - fn wildcopy_always_performs_its_first_copy() { - let mut bytes = [0u8; 96]; - for (index, byte) in bytes[..32].iter_mut().enumerate() { - *byte = index as u8; - } + const WILDCOPY_REDZONE: usize = 32; + const WILDCOPY_LENGTHS: &[usize] = &[0, 1, 16, 17, 31, 32, 33, 63, 64]; - unsafe { - ZSTD_wildcopy( - bytes.as_mut_ptr().add(64).cast(), - bytes.as_ptr().cast(), - 0, - ZstdOverlap::NoOverlap, - ); + fn no_overlap_write_len(length: usize) -> usize { + if length <= 16 { + 16 + } else { + 16 + (length - 16).div_ceil(32) * 32 } - assert_eq!(&bytes[64..80], &bytes[..16]); + } + + fn short_overlap_write_len(length: usize) -> usize { + length.max(8).div_ceil(8) * 8 + } + + fn assert_filled(bytes: &[u8], value: u8) { + assert!(bytes.iter().all(|byte| *byte == value)); } #[test] - fn short_offset_wildcopy_is_do_while_like() { - let mut bytes = [0u8; 32]; - for (index, byte) in bytes[..8].iter_mut().enumerate() { - *byte = (index + 1) as u8; - } + fn wildcopy_abi_no_overlap_respects_contract_boundaries() { + for &length in WILDCOPY_LENGTHS { + let written = no_overlap_write_len(length); + let source_start = WILDCOPY_REDZONE; + let destination_start = source_start + written + 2 * WILDCOPY_REDZONE; + let total_size = destination_start + written + WILDCOPY_REDZONE; + let mut bytes = vec![0xa5; total_size]; - unsafe { - ZSTD_wildcopy( - bytes.as_mut_ptr().add(8).cast(), - bytes.as_ptr().cast(), - 0, - ZstdOverlap::OverlapSrcBeforeDst, + for (index, byte) in bytes[source_start..source_start + written] + .iter_mut() + .enumerate() + { + *byte = (index as u8).wrapping_mul(17).wrapping_add(3); + } + bytes[destination_start..destination_start + written].fill(0xcc); + let source = bytes[source_start..source_start + written].to_vec(); + + unsafe { + ZSTD_rust_wildcopy( + bytes.as_mut_ptr().add(destination_start).cast(), + bytes.as_ptr().add(source_start).cast(), + length as isize, + 0, + ); + } + + assert_filled(&bytes[..source_start], 0xa5); + assert_eq!(&bytes[source_start..source_start + written], &source); + assert_filled(&bytes[source_start + written..destination_start], 0xa5); + assert_eq!( + &bytes[destination_start..destination_start + written], + &source, + "length={length}" ); + assert_filled(&bytes[destination_start + written..], 0xa5); + } + } + + #[test] + fn wildcopy_abi_short_offsets_respect_contract_boundaries() { + for &offset in &[8usize, 15] { + for &length in WILDCOPY_LENGTHS { + let written = short_overlap_write_len(length); + let source_start = WILDCOPY_REDZONE; + let destination_start = source_start + offset; + let data_end = destination_start + written; + let total_size = data_end + WILDCOPY_REDZONE; + let mut bytes = vec![0xa5; total_size]; + + for (index, byte) in bytes[source_start..data_end].iter_mut().enumerate() { + *byte = (index % offset) as u8 + 1; + } + + unsafe { + ZSTD_rust_wildcopy( + bytes.as_mut_ptr().add(destination_start).cast(), + bytes.as_ptr().add(source_start).cast(), + length as isize, + 1, + ); + } + + assert_filled(&bytes[..source_start], 0xa5); + assert_eq!( + &bytes[destination_start..data_end], + &(1..=offset) + .cycle() + .take(written) + .map(|value| value as u8) + .collect::>(), + "offset={offset}, length={length}" + ); + assert_filled(&bytes[data_end..], 0xa5); + } } - assert_eq!(&bytes[8..16], &[1, 2, 3, 4, 5, 6, 7, 8]); } #[test]