use std::sync::OnceLock;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SimdAvailability {
Avx2,
Avx,
Sse2,
Neon,
None,
}
impl SimdAvailability {
pub fn is_available(&self) -> bool {
*self != SimdAvailability::None
}
}
static DETECTED: OnceLock<SimdAvailability> = OnceLock::new();
pub fn detect() -> SimdAvailability {
*DETECTED.get_or_init(detect_impl)
}
#[cfg(target_arch = "x86_64")]
fn detect_impl() -> SimdAvailability {
if is_x86_feature_detected!("avx2") {
SimdAvailability::Avx2
} else if is_x86_feature_detected!("avx") {
SimdAvailability::Avx
} else if is_x86_feature_detected!("sse2") {
SimdAvailability::Sse2
} else {
SimdAvailability::None
}
}
#[cfg(target_arch = "x86")]
fn detect_impl() -> SimdAvailability {
if is_x86_feature_detected!("avx2") {
SimdAvailability::Avx2
} else if is_x86_feature_detected!("avx") {
SimdAvailability::Avx
} else if is_x86_feature_detected!("sse2") {
SimdAvailability::Sse2
} else {
SimdAvailability::None
}
}
#[cfg(target_arch = "aarch64")]
fn detect_impl() -> SimdAvailability {
if std::arch::is_aarch64_feature_detected!("neon") {
SimdAvailability::Neon
} else {
SimdAvailability::None
}
}
#[cfg(not(any(target_arch = "x86_64", target_arch = "x86", target_arch = "aarch64")))]
fn detect_impl() -> SimdAvailability {
SimdAvailability::None
}
pub const SIMD_THRESHOLD: usize = 1024;
pub fn batch_decode_integers(buf: &[u8], count: usize, _avail: SimdAvailability) -> Vec<i64> {
scalar_decode_integers(buf, count)
}
pub fn scalar_decode_integers(buf: &[u8], count: usize) -> Vec<i64> {
let n = count.min(buf.len() / 8);
(0..n)
.map(|i| {
let offset = i * 8;
i64::from_le_bytes(buf[offset..offset + 8].try_into().unwrap())
})
.collect()
}
pub fn batch_compare_eq(values: &[i64], target: i64, _avail: SimdAvailability) -> Vec<bool> {
scalar_compare_eq(values, target)
}
pub fn scalar_compare_eq(values: &[i64], target: i64) -> Vec<bool> {
values.iter().map(|&v| v == target).collect()
}
pub fn batch_compare_in(values: &[i64], set: &[i64], _avail: SimdAvailability) -> Vec<bool> {
if set.len() >= 8 {
let hash_set: std::collections::HashSet<i64> = set.iter().copied().collect();
values.iter().map(|&v| hash_set.contains(&v)).collect()
} else {
scalar_compare_in(values, set)
}
}
pub fn scalar_compare_in(values: &[i64], set: &[i64]) -> Vec<bool> {
values.iter().map(|&v| set.contains(&v)).collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_simd_availability_is_available() {
assert!(SimdAvailability::Avx2.is_available());
assert!(SimdAvailability::Avx.is_available());
assert!(SimdAvailability::Sse2.is_available());
assert!(SimdAvailability::Neon.is_available());
assert!(!SimdAvailability::None.is_available());
}
#[test]
fn test_detect_returns_cached() {
let d1 = detect();
let d2 = detect();
assert_eq!(d1, d2);
}
#[test]
fn test_scalar_decode_integers() {
let values: Vec<i64> = vec![1, 2, 3, 4, 5];
let mut buf = Vec::new();
for v in &values {
buf.extend_from_slice(&v.to_le_bytes());
}
let result = scalar_decode_integers(&buf, 5);
assert_eq!(result, values);
}
#[test]
fn test_batch_decode_integers_small_count() {
let values: Vec<i64> = vec![1, 2, 3];
let mut buf = Vec::new();
for v in &values {
buf.extend_from_slice(&v.to_le_bytes());
}
let result = batch_decode_integers(&buf, 3, SimdAvailability::Avx2);
assert_eq!(result, values);
}
#[test]
fn test_batch_decode_integers_large_count() {
let n: usize = 2000;
let values: Vec<i64> = (0..n as i64).map(|i| i * 2 - 1).collect();
let mut buf = Vec::new();
for v in &values {
buf.extend_from_slice(&v.to_le_bytes());
}
let avail = detect();
let result = batch_decode_integers(&buf, n, avail);
assert_eq!(result, values);
}
#[test]
fn test_batch_decode_integers_none_avail() {
let n: usize = 2000;
let values: Vec<i64> = (0..n as i64).collect();
let mut buf = Vec::new();
for v in &values {
buf.extend_from_slice(&v.to_le_bytes());
}
let result = batch_decode_integers(&buf, n, SimdAvailability::None);
assert_eq!(result, values);
}
#[test]
fn test_scalar_compare_eq() {
let values = vec![1, 2, 3, 4, 5, 3, 3];
let result = scalar_compare_eq(&values, 3);
assert_eq!(result, vec![false, false, true, false, false, true, true]);
}
#[test]
fn test_batch_compare_eq_small() {
let values = vec![1, 2, 3, 4, 5];
let result = batch_compare_eq(&values, 3, SimdAvailability::Avx2);
assert_eq!(result, vec![false, false, true, false, false]);
}
#[test]
fn test_batch_compare_eq_large() {
let n: usize = 2000;
let values: Vec<i64> = (0..n as i64).collect();
let target = 500_i64;
let avail = detect();
let result = batch_compare_eq(&values, target, avail);
assert_eq!(result.len(), n);
assert!(result[500]);
assert!(!result[499]);
assert!(!result[501]);
}
#[test]
fn test_scalar_compare_in() {
let values = vec![1, 2, 3, 4, 5];
let set = vec![2, 4];
let result = scalar_compare_in(&values, &set);
assert_eq!(result, vec![false, true, false, true, false]);
}
#[test]
fn test_batch_compare_in_small() {
let values = vec![1, 2, 3, 4, 5];
let set = vec![2, 4];
let result = batch_compare_in(&values, &set, SimdAvailability::Avx2);
assert_eq!(result, vec![false, true, false, true, false]);
}
#[test]
fn test_batch_compare_in_large() {
let n: usize = 2000;
let values: Vec<i64> = (0..n as i64).collect();
let set: Vec<i64> = vec![100, 500, 1500];
let avail = detect();
let result = batch_compare_in(&values, &set, avail);
assert_eq!(result.len(), n);
assert!(result[100]);
assert!(result[500]);
assert!(result[1500]);
assert!(!result[200]);
}
#[test]
fn test_batch_compare_eq_none_avail() {
let n: usize = 2000;
let values: Vec<i64> = (0..n as i64).collect();
let result = batch_compare_eq(&values, 500, SimdAvailability::None);
assert_eq!(result.len(), n);
assert!(result[500]);
}
#[test]
fn test_batch_compare_in_empty_set() {
let values = vec![1, 2, 3];
let set: Vec<i64> = vec![];
let result = batch_compare_in(&values, &set, SimdAvailability::Avx2);
assert_eq!(result, vec![false, false, false]);
}
#[test]
fn test_batch_decode_integers_count_exceeds_buf() {
let values: Vec<i64> = vec![1, 2, 3];
let mut buf = Vec::new();
for v in &values {
buf.extend_from_slice(&v.to_le_bytes());
}
let result = batch_decode_integers(&buf, 100, SimdAvailability::None);
assert_eq!(result, values);
}
#[test]
fn test_batch_decode_integers_empty() {
let result = batch_decode_integers(&[], 0, SimdAvailability::Avx2);
assert!(result.is_empty());
}
#[test]
fn test_simd_threshold_constant() {
assert_eq!(SIMD_THRESHOLD, 1024);
}
#[test]
fn test_batch_compare_eq_boundary_1023() {
let n = 1023;
let values: Vec<i64> = vec![42; n];
let result = batch_compare_eq(&values, 42, SimdAvailability::Avx2);
assert!(result.iter().all(|&b| b));
}
#[test]
fn test_batch_compare_eq_boundary_1024() {
let n = 1024;
let values: Vec<i64> = vec![42; n];
let avail = detect();
let result = batch_compare_eq(&values, 42, avail);
assert!(result.iter().all(|&b| b));
}
#[test]
fn test_batch_compare_eq_boundary_1025() {
let n = 1025;
let values: Vec<i64> = vec![42; n];
let avail = detect();
let result = batch_compare_eq(&values, 42, avail);
assert!(result.iter().all(|&b| b));
}
}