#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
use fearless_simd::{Level, dispatch, prelude::*, u8x16};
#[inline]
pub fn fast_key_eq(a: &[u8], b: &[u8]) -> bool {
let len = a.len();
if len != b.len() {
return false;
}
if a.as_ptr() == b.as_ptr() || len == 0 {
return true;
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
if len >= 16 {
return dispatch!(Level::new(), simd => simd_key_eq(simd, a, b, len));
}
unsafe { scalar_fallback_key_eq(a.as_ptr(), b.as_ptr(), len) }
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn simd_key_eq<S: Simd>(simd: S, a: &[u8], b: &[u8], len: usize) -> bool {
let a_ptr = a.as_ptr();
let b_ptr = b.as_ptr();
let mut offset = 0;
while offset + 16 <= len {
let sa = unsafe { &*(a_ptr.add(offset) as *const [u8; 16]) };
let sb = unsafe { &*(b_ptr.add(offset) as *const [u8; 16]) };
if !eq16(simd, sa, sb) {
return false;
}
offset += 16;
}
if offset < len {
let tail_offset = len - 16;
let sa = unsafe { &*(a_ptr.add(tail_offset) as *const [u8; 16]) };
let sb = unsafe { &*(b_ptr.add(tail_offset) as *const [u8; 16]) };
return eq16(simd, sa, sb);
}
true
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn eq16<S: Simd>(simd: S, a: &[u8; 16], b: &[u8; 16]) -> bool {
u8x16::load_array_ref(simd, a)
.simd_eq(u8x16::load_array_ref(simd, b))
.all_true()
}
#[inline]
unsafe fn scalar_fallback_key_eq(a: *const u8, b: *const u8, len: usize) -> bool {
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
if len >= 8 {
let a_head = unsafe { (a as *const u64).read_unaligned() };
let b_head = unsafe { (b as *const u64).read_unaligned() };
let a_tail = unsafe { (a.add(len - 8) as *const u64).read_unaligned() };
let b_tail = unsafe { (b.add(len - 8) as *const u64).read_unaligned() };
return ((a_head ^ b_head) | (a_tail ^ b_tail)) == 0;
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
if len >= 8 {
let mut offset = 0;
while offset + 8 <= len {
if unsafe { (a.add(offset) as *const u64).read_unaligned() }
!= unsafe { (b.add(offset) as *const u64).read_unaligned() }
{
return false;
}
offset += 8;
}
let tail = len - 8;
return unsafe { (a.add(tail) as *const u64).read_unaligned() }
== unsafe { (b.add(tail) as *const u64).read_unaligned() };
}
if len >= 4 {
let a_head = unsafe { (a as *const u32).read_unaligned() };
let b_head = unsafe { (b as *const u32).read_unaligned() };
let a_tail = unsafe { (a.add(len - 4) as *const u32).read_unaligned() };
let b_tail = unsafe { (b.add(len - 4) as *const u32).read_unaligned() };
return ((a_head ^ b_head) | (a_tail ^ b_tail)) == 0;
}
if len == 3 {
let head_diff =
unsafe { ((a as *const u16).read_unaligned() ^ (b as *const u16).read_unaligned()) as u32 };
let tail_diff = unsafe { (*a.add(2) ^ *b.add(2)) as u32 };
return (head_diff | tail_diff) == 0;
}
if len == 2 {
return unsafe { (a as *const u16).read_unaligned() == (b as *const u16).read_unaligned() };
}
unsafe { *a == *b }
}