#[cfg(target_arch = "aarch64")]
use core::arch::aarch64;
#[cfg(target_arch = "x86_64")]
use core::arch::x86_64;
#[cfg(opthash_x86_16_group)]
use core::arch::x86_64::__m128i;
use super::bitmask::BitMask;
#[cfg(opthash_scalar_group)]
use super::config::GROUP_SIZE;
#[cfg(any(opthash_neon_group, opthash_x86_16_group))]
use super::control::FINGERPRINT_MASK;
#[cfg(opthash_scalar_group)]
const SWAR_LO7: u64 = 0x7f7f_7f7f_7f7f_7f7f;
#[cfg(opthash_scalar_group)]
const SWAR_HI: u64 = 0x8080_8080_8080_8080;
#[cfg(opthash_scalar_group)]
const SWAR_ONES: u64 = 0x0101_0101_0101_0101;
#[cfg(opthash_scalar_group)]
#[inline]
unsafe fn swar_word(ptr: *const u8) -> u64 {
#[allow(clippy::cast_ptr_alignment)]
let raw = unsafe { ptr.cast::<u64>().read_unaligned() };
raw.to_le()
}
#[cfg(opthash_scalar_group)]
#[inline]
fn swar_eq_mask(word: u64, target: u8) -> u64 {
let cmp = word ^ (u64::from(target).wrapping_mul(SWAR_ONES));
let ne = (((cmp & SWAR_LO7).wrapping_add(SWAR_LO7)) | cmp) & SWAR_HI; ne ^ SWAR_HI
}
#[cfg(opthash_scalar_group)]
#[inline]
fn swar_occupied_mask(word: u64) -> u64 {
((word & SWAR_LO7).wrapping_add(SWAR_LO7)) & SWAR_HI
}
#[cfg(opthash_scalar_group)]
#[inline]
fn swar_free_mask(word: u64) -> u64 {
swar_occupied_mask(word) ^ SWAR_HI
}
#[inline]
#[must_use]
pub(crate) unsafe fn eq_mask_group(ptr: *const u8, target: u8) -> BitMask {
#[cfg(opthash_neon_group)]
let mask = unsafe { eq_mask_8_neon(ptr, target) };
#[cfg(opthash_x86_16_group)]
let mask = unsafe { eq_mask_16_sse2(ptr, target) };
#[cfg(opthash_scalar_group)]
let mask = unsafe { BitMask(swar_eq_mask(swar_word(ptr), target)) };
mask
}
#[inline]
#[must_use]
pub(crate) unsafe fn free_mask_group(ptr: *const u8) -> BitMask {
#[cfg(opthash_neon_group)]
let mask = unsafe { free_mask_8_neon(ptr) };
#[cfg(opthash_x86_16_group)]
let mask = unsafe { free_mask_16_sse2(ptr) };
#[cfg(opthash_scalar_group)]
let mask = unsafe { BitMask(swar_free_mask(swar_word(ptr))) };
mask
}
#[inline]
#[must_use]
pub(crate) unsafe fn occupied_mask_group(ptr: *const u8) -> BitMask {
#[cfg(opthash_neon_group)]
let mask = unsafe { occupied_mask_8_neon(ptr) };
#[cfg(opthash_x86_16_group)]
let mask = unsafe { occupied_mask_16_sse2(ptr) };
#[cfg(opthash_scalar_group)]
let mask = unsafe { BitMask(swar_occupied_mask(swar_word(ptr))) };
mask
}
#[cfg(opthash_neon_group)]
#[inline]
unsafe fn eq_mask_8_neon(ptr: *const u8, target: u8) -> BitMask {
unsafe {
let bytes = aarch64::vld1_u8(ptr);
let cmp = aarch64::vceq_u8(bytes, aarch64::vdup_n_u8(target));
BitMask(aarch64::vget_lane_u64(aarch64::vreinterpret_u64_u8(cmp), 0))
}
}
#[cfg(opthash_neon_group)]
#[inline]
unsafe fn free_mask_8_neon(ptr: *const u8) -> BitMask {
unsafe {
let bytes = aarch64::vld1_u8(ptr);
let masked = aarch64::vand_u8(bytes, aarch64::vdup_n_u8(FINGERPRINT_MASK));
let free_cmp = aarch64::vceq_u8(masked, aarch64::vdup_n_u8(0));
BitMask(aarch64::vget_lane_u64(
aarch64::vreinterpret_u64_u8(free_cmp),
0,
))
}
}
#[cfg(opthash_neon_group)]
#[inline]
unsafe fn occupied_mask_8_neon(ptr: *const u8) -> BitMask {
unsafe {
let bytes = aarch64::vld1_u8(ptr);
let occ_cmp = aarch64::vtst_u8(bytes, aarch64::vdup_n_u8(FINGERPRINT_MASK));
BitMask(aarch64::vget_lane_u64(
aarch64::vreinterpret_u64_u8(occ_cmp),
0,
))
}
}
#[allow(clippy::cast_ptr_alignment)]
#[cfg(opthash_x86_16_group)]
#[inline]
unsafe fn eq_mask_16_sse2(ptr: *const u8, target: u8) -> BitMask {
unsafe {
let data = x86_64::_mm_loadu_si128(ptr.cast::<__m128i>());
let target_vec = x86_64::_mm_set1_epi8(target.cast_signed());
let cmp = x86_64::_mm_cmpeq_epi8(data, target_vec);
let bits = x86_64::_mm_movemask_epi8(cmp).cast_unsigned() & 0xFFFF;
BitMask(u64::from(bits))
}
}
#[allow(clippy::cast_ptr_alignment)]
#[cfg(opthash_x86_16_group)]
#[inline]
unsafe fn free_mask_16_sse2(ptr: *const u8) -> BitMask {
unsafe {
let data = x86_64::_mm_loadu_si128(ptr.cast::<__m128i>());
let masked =
x86_64::_mm_and_si128(data, x86_64::_mm_set1_epi8(FINGERPRINT_MASK.cast_signed()));
let free = x86_64::_mm_cmpeq_epi8(masked, x86_64::_mm_setzero_si128());
let bits = x86_64::_mm_movemask_epi8(free).cast_unsigned() & 0xFFFF;
BitMask(u64::from(bits))
}
}
#[allow(clippy::cast_ptr_alignment)]
#[cfg(opthash_x86_16_group)]
#[inline]
unsafe fn occupied_mask_16_sse2(ptr: *const u8) -> BitMask {
unsafe {
let data = x86_64::_mm_loadu_si128(ptr.cast::<__m128i>());
let masked =
x86_64::_mm_and_si128(data, x86_64::_mm_set1_epi8(FINGERPRINT_MASK.cast_signed()));
let occ = x86_64::_mm_cmpgt_epi8(masked, x86_64::_mm_setzero_si128());
let bits = x86_64::_mm_movemask_epi8(occ).cast_unsigned() & 0xFFFF;
BitMask(u64::from(bits))
}
}
#[cfg(all(test, opthash_neon_group))]
mod neon_tests {
use super::*;
use crate::common::control::{CTRL_EMPTY, CTRL_TOMBSTONE};
#[test]
fn eight_lane_masks_preserve_control_order() {
let controls = [
CTRL_EMPTY,
CTRL_TOMBSTONE,
1,
7,
CTRL_EMPTY,
7,
CTRL_TOMBSTONE,
2,
];
let matches: alloc::vec::Vec<_> = unsafe { eq_mask_8_neon(controls.as_ptr(), 7) }.collect();
let free: alloc::vec::Vec<_> = unsafe { free_mask_8_neon(controls.as_ptr()) }.collect();
let occupied: alloc::vec::Vec<_> =
unsafe { occupied_mask_8_neon(controls.as_ptr()) }.collect();
assert_eq!(matches, [3, 5]);
assert_eq!(free, [0, 1, 4, 6]);
assert_eq!(occupied, [2, 3, 5, 7]);
}
}
#[cfg(all(test, opthash_scalar_group))]
mod swar_tests {
use super::*;
use crate::common::control::{CTRL_EMPTY, CTRL_TOMBSTONE, FINGERPRINT_MASK};
fn word(bytes: [u8; 8]) -> u64 {
u64::from_le_bytes(bytes)
}
fn ref_eq(bytes: [u8; 8], target: u8) -> u64 {
let mut m = 0u64;
for (i, &b) in bytes.iter().enumerate() {
if b == target {
m |= 0x80u64 << (8 * i);
}
}
m
}
fn ref_occupied(bytes: [u8; 8]) -> u64 {
let mut m = 0u64;
for (i, &b) in bytes.iter().enumerate() {
if b & FINGERPRINT_MASK != 0 {
m |= 0x80u64 << (8 * i);
}
}
m
}
fn sample_words() -> [[u8; 8]; 10] {
[
[CTRL_EMPTY; 8],
[CTRL_TOMBSTONE; 8],
[0x2a; 8],
[1, 2, 3, 4, 5, 6, 7, 8],
[0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07],
[0x80, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00],
[0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00],
[0x7f, 0x80, 0x00, 0x01, 0x7f, 0x80, 0x00, 0x01],
[0xff, 0xfe, 0xfd, 0xfc, 0x80, 0x7f, 0x01, 0x00],
[0x00, 0x80, 0x00, 0x80, 0x00, 0x80, 0x00, 0x80],
]
}
#[test]
fn eq_matches_reference_for_every_target() {
for bytes in sample_words() {
for target in 0..=u8::MAX {
assert_eq!(
swar_eq_mask(word(bytes), target),
ref_eq(bytes, target),
"word={bytes:02x?} target={target:#04x}"
);
}
}
}
#[test]
fn eq_empty_flags_only_exact_zero() {
for bytes in sample_words() {
assert_eq!(
swar_eq_mask(word(bytes), CTRL_EMPTY),
ref_eq(bytes, CTRL_EMPTY)
);
}
assert_eq!(swar_eq_mask(word([CTRL_TOMBSTONE; 8]), CTRL_EMPTY), 0);
assert_eq!(swar_eq_mask(word([0x2a; 8]), CTRL_EMPTY), 0);
}
#[test]
fn free_and_occupied_partition_all_lanes() {
for bytes in sample_words() {
let f = swar_free_mask(word(bytes));
let o = swar_occupied_mask(word(bytes));
assert_eq!(o, ref_occupied(bytes), "occupied mismatch: {bytes:02x?}");
assert_eq!(f & o, 0, "free/occupied overlap: {bytes:02x?}");
assert_eq!(f | o, SWAR_HI, "free|occupied != all lanes: {bytes:02x?}");
}
}
#[test]
fn free_is_empty_or_tombstone() {
assert_eq!(swar_free_mask(word([CTRL_EMPTY; 8])), SWAR_HI);
assert_eq!(swar_free_mask(word([CTRL_TOMBSTONE; 8])), SWAR_HI);
assert_eq!(swar_occupied_mask(word([CTRL_EMPTY; 8])), 0);
assert_eq!(swar_occupied_mask(word([CTRL_TOMBSTONE; 8])), 0);
for fp in 1..=FINGERPRINT_MASK {
assert_eq!(swar_occupied_mask(word([fp; 8])), SWAR_HI, "fp {fp:#04x}");
assert_eq!(swar_free_mask(word([fp; 8])), 0, "fp {fp:#04x}");
}
}
#[test]
fn no_cross_byte_borrow() {
let bytes = [0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00];
assert_eq!(swar_occupied_mask(word(bytes)), 0x80u64 << 8);
assert_eq!(swar_eq_mask(word(bytes), 0x01), 0x80u64 << 8);
let asc = [0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07];
for (i, &b) in asc.iter().enumerate() {
assert_eq!(
swar_eq_mask(word(asc), b),
0x80u64 << (8 * i),
"byte {b:#04x}"
);
}
}
#[test]
fn match_lands_in_lane_high_bit() {
assert_eq!(swar_eq_mask(word([9; 8]), 9), SWAR_HI);
assert_eq!(
swar_eq_mask(word([0, 0, 0, 9, 0, 0, 0, 0]), 9),
0x80u64 << (8 * 3)
);
}
}