use core::slice::from_raw_parts;
#[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(target_arch = "aarch64")]
{
if len >= 16 {
return unsafe { neon_key_eq(a.as_ptr(), b.as_ptr(), len) };
}
}
#[cfg(target_arch = "x86_64")]
{
if len >= 16 {
return unsafe { sse2_key_eq(a.as_ptr(), b.as_ptr(), len) };
}
}
unsafe { scalar_fallback_key_eq(a.as_ptr(), b.as_ptr(), len) }
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn neon_key_eq(a: *const u8, b: *const u8, len: usize) -> bool {
use core::arch::aarch64::{vceqq_u8, vld1q_u8, vminvq_u8};
unsafe {
let mut offset = 0;
while offset + 16 <= len {
let va = vld1q_u8(a.add(offset));
let vb = vld1q_u8(b.add(offset));
let vcmp = vceqq_u8(va, vb);
if vminvq_u8(vcmp) != 0xFF {
return false;
}
offset += 16;
}
if offset < len {
let va = vld1q_u8(a.add(len - 16));
let vb = vld1q_u8(b.add(len - 16));
let vcmp = vceqq_u8(va, vb);
return vminvq_u8(vcmp) == 0xFF;
}
true
}
}
#[cfg(target_arch = "x86_64")]
#[inline]
unsafe fn sse2_key_eq(a: *const u8, b: *const u8, len: usize) -> bool {
use core::arch::x86_64::{__m128i, _mm_cmpeq_epi8, _mm_loadu_si128, _mm_movemask_epi8};
unsafe {
let mut offset = 0;
while offset + 16 <= len {
let va = _mm_loadu_si128(a.add(offset) as *const __m128i);
let vb = _mm_loadu_si128(b.add(offset) as *const __m128i);
let vcmp = _mm_cmpeq_epi8(va, vb);
let mask = _mm_movemask_epi8(vcmp);
if mask != 0xFFFF {
return false;
}
offset += 16;
}
if offset < len {
let va = _mm_loadu_si128(a.add(len - 16) as *const __m128i);
let vb = _mm_loadu_si128(b.add(len - 16) as *const __m128i);
let vcmp = _mm_cmpeq_epi8(va, vb);
return _mm_movemask_epi8(vcmp) == 0xFFFF;
}
true
}
}
#[inline]
unsafe fn scalar_fallback_key_eq(a: *const u8, b: *const u8, len: usize) -> bool {
unsafe {
if len >= 8 {
let mut offset = 0;
while offset + 8 <= len {
if (a.add(offset) as *const u64).read_unaligned()
!= (b.add(offset) as *const u64).read_unaligned()
{
return false;
}
offset += 8;
}
let tail = len - 8;
return (a.add(tail) as *const u64).read_unaligned()
== (b.add(tail) as *const u64).read_unaligned();
}
if len >= 4 {
let a_head = (a as *const u32).read_unaligned();
let b_head = (b as *const u32).read_unaligned();
let a_tail = (a.add(len - 4) as *const u32).read_unaligned();
let b_tail = (b.add(len - 4) as *const u32).read_unaligned();
return a_head == b_head && a_tail == b_tail;
}
from_raw_parts(a, len) == from_raw_parts(b, len)
}
}
#[cfg(test)]
mod tests {
use super::fast_key_eq;
#[test]
fn test_fast_key_eq() {
assert!(fast_key_eq(b"", b""));
let empty_a: &[u8] = &[];
let empty_b: &[u8] = &[];
assert!(fast_key_eq(empty_a, empty_b));
assert!(fast_key_eq(b"hello", b"hello"));
assert!(!fast_key_eq(b"hello", b"world"));
assert!(!fast_key_eq(b"short", b"shorter"));
let overlap_buf = b"0123456789abcdef0123456789abcdef";
assert!(fast_key_eq(&overlap_buf[0..16], &overlap_buf[16..32]));
assert!(!fast_key_eq(&overlap_buf[0..16], &overlap_buf[1..17]));
assert!(fast_key_eq(&overlap_buf[0..8], &overlap_buf[16..24]));
assert!(!fast_key_eq(&overlap_buf[0..8], &overlap_buf[1..9]));
for len in 1..=64 {
let v1 = vec![0x5Au8; len];
let v2 = vec![0x5Au8; len];
assert!(fast_key_eq(&v1, &v2));
for pos in 0..len {
let mut v3 = v1.clone();
v3[pos] ^= 0xFF;
assert!(
!fast_key_eq(&v1, &v3),
"fast_key_eq 应检测出 len={len} 在 pos={pos} 处的差异"
);
}
}
for &long_len in &[128, 256, 1024] {
let l1 = vec![0x33u8; long_len];
let l2 = vec![0x33u8; long_len];
assert!(fast_key_eq(&l1, &l2));
let mut l_head = l1.clone();
l_head[0] = 0x44;
assert!(!fast_key_eq(&l1, &l_head));
let mut l_tail = l1.clone();
l_tail[long_len - 1] = 0x44;
assert!(!fast_key_eq(&l1, &l_tail));
let mut l_mid = l1.clone();
l_mid[long_len / 2] = 0x44;
assert!(!fast_key_eq(&l1, &l_mid));
for boundary in [15, 16, 31, 32, long_len - 17, long_len - 16] {
let mut l_bound = l1.clone();
l_bound[boundary] = 0x44;
assert!(!fast_key_eq(&l1, &l_bound));
}
}
}
}