1use std::sync::OnceLock;
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum SimdAvailability {
19 Avx2,
21 Avx,
23 Sse2,
25 Neon,
27 None,
29}
30
31impl SimdAvailability {
32 pub fn is_available(&self) -> bool {
34 *self != SimdAvailability::None
35 }
36}
37
38static DETECTED: OnceLock<SimdAvailability> = OnceLock::new();
39
40pub fn detect() -> SimdAvailability {
42 *DETECTED.get_or_init(detect_impl)
43}
44
45#[cfg(target_arch = "x86_64")]
46fn detect_impl() -> SimdAvailability {
47 if is_x86_feature_detected!("avx2") {
48 SimdAvailability::Avx2
49 } else if is_x86_feature_detected!("avx") {
50 SimdAvailability::Avx
51 } else if is_x86_feature_detected!("sse2") {
52 SimdAvailability::Sse2
53 } else {
54 SimdAvailability::None
55 }
56}
57
58#[cfg(target_arch = "x86")]
59fn detect_impl() -> SimdAvailability {
60 if is_x86_feature_detected!("avx2") {
61 SimdAvailability::Avx2
62 } else if is_x86_feature_detected!("avx") {
63 SimdAvailability::Avx
64 } else if is_x86_feature_detected!("sse2") {
65 SimdAvailability::Sse2
66 } else {
67 SimdAvailability::None
68 }
69}
70
71#[cfg(target_arch = "aarch64")]
72fn detect_impl() -> SimdAvailability {
73 if std::arch::is_aarch64_feature_detected!("neon") {
74 SimdAvailability::Neon
75 } else {
76 SimdAvailability::None
77 }
78}
79
80#[cfg(not(any(target_arch = "x86_64", target_arch = "x86", target_arch = "aarch64")))]
81fn detect_impl() -> SimdAvailability {
82 SimdAvailability::None
83}
84
85pub const SIMD_THRESHOLD: usize = 1024;
87
88pub fn batch_decode_integers(buf: &[u8], count: usize, avail: SimdAvailability) -> Vec<i64> {
99 if count >= SIMD_THRESHOLD && avail.is_available() {
100 simd_decode_integers(buf, count)
101 } else {
102 scalar_decode_integers(buf, count)
103 }
104}
105
106pub fn scalar_decode_integers(buf: &[u8], count: usize) -> Vec<i64> {
108 let n = count.min(buf.len() / 8);
109 (0..n)
110 .map(|i| {
111 let offset = i * 8;
112 i64::from_le_bytes(buf[offset..offset + 8].try_into().unwrap())
113 })
114 .collect()
115}
116
117#[cfg(feature = "simd")]
118fn simd_decode_integers(buf: &[u8], count: usize) -> Vec<i64> {
119 use wide::i64x4;
120
121 let n = count.min(buf.len() / 8);
122 let mut result = Vec::with_capacity(n);
123
124 let chunk_count = n / 4;
125 let remainder = n % 4;
126
127 for chunk in 0..chunk_count {
128 let base = chunk * 4;
129 let v0 = i64::from_le_bytes(buf[base * 8..base * 8 + 8].try_into().unwrap());
130 let v1 = i64::from_le_bytes(buf[(base + 1) * 8..(base + 1) * 8 + 8].try_into().unwrap());
131 let v2 = i64::from_le_bytes(buf[(base + 2) * 8..(base + 2) * 8 + 8].try_into().unwrap());
132 let v3 = i64::from_le_bytes(buf[(base + 3) * 8..(base + 3) * 8 + 8].try_into().unwrap());
133
134 let vec = i64x4::from([v0, v1, v2, v3]);
135 let arr: [i64; 4] = vec.into();
136 result.extend_from_slice(&arr);
137 }
138
139 for i in 0..remainder {
140 let idx = chunk_count * 4 + i;
141 let offset = idx * 8;
142 result.push(i64::from_le_bytes(
143 buf[offset..offset + 8].try_into().unwrap(),
144 ));
145 }
146
147 result
148}
149
150#[cfg(not(feature = "simd"))]
151fn simd_decode_integers(buf: &[u8], count: usize) -> Vec<i64> {
152 scalar_decode_integers(buf, count)
153}
154
155pub fn batch_compare_eq(values: &[i64], target: i64, avail: SimdAvailability) -> Vec<bool> {
166 if values.len() >= SIMD_THRESHOLD && avail.is_available() {
167 simd_compare_eq(values, target)
168 } else {
169 scalar_compare_eq(values, target)
170 }
171}
172
173pub fn scalar_compare_eq(values: &[i64], target: i64) -> Vec<bool> {
175 values.iter().map(|&v| v == target).collect()
176}
177
178#[cfg(feature = "simd")]
179fn simd_compare_eq(values: &[i64], target: i64) -> Vec<bool> {
180 use wide::{i64x4, CmpEq};
181
182 let n = values.len();
183 let mut result = Vec::with_capacity(n);
184
185 let chunk_count = n / 4;
186 let remainder = n % 4;
187 let target_vec = i64x4::splat(target);
188
189 for chunk in 0..chunk_count {
190 let slice = &values[chunk * 4..chunk * 4 + 4];
191 let vec = i64x4::from([slice[0], slice[1], slice[2], slice[3]]);
192 let cmp = vec.cmp_eq(target_vec);
193 let mask: [i64; 4] = cmp.into();
194 for &m in &mask {
195 result.push(m != 0);
196 }
197 }
198
199 for i in 0..remainder {
200 let idx = chunk_count * 4 + i;
201 result.push(values[idx] == target);
202 }
203
204 result
205}
206
207#[cfg(not(feature = "simd"))]
208fn simd_compare_eq(values: &[i64], target: i64) -> Vec<bool> {
209 scalar_compare_eq(values, target)
210}
211
212pub fn batch_compare_in(values: &[i64], set: &[i64], avail: SimdAvailability) -> Vec<bool> {
219 if values.len() >= SIMD_THRESHOLD && avail.is_available() {
220 simd_compare_in(values, set)
221 } else {
222 scalar_compare_in(values, set)
223 }
224}
225
226pub fn scalar_compare_in(values: &[i64], set: &[i64]) -> Vec<bool> {
228 values.iter().map(|&v| set.contains(&v)).collect()
229}
230
231#[cfg(feature = "simd")]
232fn simd_compare_in(values: &[i64], set: &[i64]) -> Vec<bool> {
233 use wide::{i64x4, CmpEq};
234
235 let n = values.len();
236 let mut result = Vec::with_capacity(n);
237
238 let chunk_count = n / 4;
239 let remainder = n % 4;
240
241 for chunk in 0..chunk_count {
242 let slice = &values[chunk * 4..chunk * 4 + 4];
243 let vec = i64x4::from([slice[0], slice[1], slice[2], slice[3]]);
244
245 let mut any_match = [false; 4];
246 for &s in set {
247 let target_vec = i64x4::splat(s);
248 let cmp = vec.cmp_eq(target_vec);
249 let mask: [i64; 4] = cmp.into();
250 for j in 0..4 {
251 if mask[j] != 0 {
252 any_match[j] = true;
253 }
254 }
255 }
256 result.extend_from_slice(&any_match);
257 }
258
259 for i in 0..remainder {
260 let idx = chunk_count * 4 + i;
261 result.push(set.contains(&values[idx]));
262 }
263
264 result
265}
266
267#[cfg(not(feature = "simd"))]
268fn simd_compare_in(values: &[i64], set: &[i64]) -> Vec<bool> {
269 scalar_compare_in(values, set)
270}
271
272#[cfg(test)]
277mod tests {
278 use super::*;
279
280 #[test]
281 fn test_simd_availability_is_available() {
282 assert!(SimdAvailability::Avx2.is_available());
283 assert!(SimdAvailability::Avx.is_available());
284 assert!(SimdAvailability::Sse2.is_available());
285 assert!(SimdAvailability::Neon.is_available());
286 assert!(!SimdAvailability::None.is_available());
287 }
288
289 #[test]
290 fn test_detect_returns_cached() {
291 let d1 = detect();
292 let d2 = detect();
293 assert_eq!(d1, d2);
294 }
295
296 #[test]
297 fn test_scalar_decode_integers() {
298 let values: Vec<i64> = vec![1, 2, 3, 4, 5];
299 let mut buf = Vec::new();
300 for v in &values {
301 buf.extend_from_slice(&v.to_le_bytes());
302 }
303 let result = scalar_decode_integers(&buf, 5);
304 assert_eq!(result, values);
305 }
306
307 #[test]
308 fn test_batch_decode_integers_small_count() {
309 let values: Vec<i64> = vec![1, 2, 3];
310 let mut buf = Vec::new();
311 for v in &values {
312 buf.extend_from_slice(&v.to_le_bytes());
313 }
314 let result = batch_decode_integers(&buf, 3, SimdAvailability::Avx2);
315 assert_eq!(result, values);
316 }
317
318 #[test]
319 fn test_batch_decode_integers_large_count() {
320 let n: usize = 2000;
321 let values: Vec<i64> = (0..n as i64).map(|i| i * 2 - 1).collect();
322 let mut buf = Vec::new();
323 for v in &values {
324 buf.extend_from_slice(&v.to_le_bytes());
325 }
326 let avail = detect();
327 let result = batch_decode_integers(&buf, n, avail);
328 assert_eq!(result, values);
329 }
330
331 #[test]
332 fn test_batch_decode_integers_none_avail() {
333 let n: usize = 2000;
334 let values: Vec<i64> = (0..n as i64).collect();
335 let mut buf = Vec::new();
336 for v in &values {
337 buf.extend_from_slice(&v.to_le_bytes());
338 }
339 let result = batch_decode_integers(&buf, n, SimdAvailability::None);
340 assert_eq!(result, values);
341 }
342
343 #[test]
344 fn test_scalar_compare_eq() {
345 let values = vec![1, 2, 3, 4, 5, 3, 3];
346 let result = scalar_compare_eq(&values, 3);
347 assert_eq!(result, vec![false, false, true, false, false, true, true]);
348 }
349
350 #[test]
351 fn test_batch_compare_eq_small() {
352 let values = vec![1, 2, 3, 4, 5];
353 let result = batch_compare_eq(&values, 3, SimdAvailability::Avx2);
354 assert_eq!(result, vec![false, false, true, false, false]);
355 }
356
357 #[test]
358 fn test_batch_compare_eq_large() {
359 let n: usize = 2000;
360 let values: Vec<i64> = (0..n as i64).collect();
361 let target = 500_i64;
362 let avail = detect();
363 let result = batch_compare_eq(&values, target, avail);
364 assert_eq!(result.len(), n);
365 assert!(result[500]);
366 assert!(!result[499]);
367 assert!(!result[501]);
368 }
369
370 #[test]
371 fn test_scalar_compare_in() {
372 let values = vec![1, 2, 3, 4, 5];
373 let set = vec![2, 4];
374 let result = scalar_compare_in(&values, &set);
375 assert_eq!(result, vec![false, true, false, true, false]);
376 }
377
378 #[test]
379 fn test_batch_compare_in_small() {
380 let values = vec![1, 2, 3, 4, 5];
381 let set = vec![2, 4];
382 let result = batch_compare_in(&values, &set, SimdAvailability::Avx2);
383 assert_eq!(result, vec![false, true, false, true, false]);
384 }
385
386 #[test]
387 fn test_batch_compare_in_large() {
388 let n: usize = 2000;
389 let values: Vec<i64> = (0..n as i64).collect();
390 let set: Vec<i64> = vec![100, 500, 1500];
391 let avail = detect();
392 let result = batch_compare_in(&values, &set, avail);
393 assert_eq!(result.len(), n);
394 assert!(result[100]);
395 assert!(result[500]);
396 assert!(result[1500]);
397 assert!(!result[200]);
398 }
399
400 #[test]
401 fn test_batch_compare_eq_none_avail() {
402 let n: usize = 2000;
403 let values: Vec<i64> = (0..n as i64).collect();
404 let result = batch_compare_eq(&values, 500, SimdAvailability::None);
405 assert_eq!(result.len(), n);
406 assert!(result[500]);
407 }
408
409 #[test]
410 fn test_batch_compare_in_empty_set() {
411 let values = vec![1, 2, 3];
412 let set: Vec<i64> = vec![];
413 let result = batch_compare_in(&values, &set, SimdAvailability::Avx2);
414 assert_eq!(result, vec![false, false, false]);
415 }
416
417 #[test]
418 fn test_batch_decode_integers_count_exceeds_buf() {
419 let values: Vec<i64> = vec![1, 2, 3];
420 let mut buf = Vec::new();
421 for v in &values {
422 buf.extend_from_slice(&v.to_le_bytes());
423 }
424 let result = batch_decode_integers(&buf, 100, SimdAvailability::None);
425 assert_eq!(result, values);
426 }
427
428 #[test]
429 fn test_batch_decode_integers_empty() {
430 let result = batch_decode_integers(&[], 0, SimdAvailability::Avx2);
431 assert!(result.is_empty());
432 }
433
434 #[test]
435 fn test_simd_threshold_constant() {
436 assert_eq!(SIMD_THRESHOLD, 1024);
437 }
438
439 #[test]
440 fn test_batch_compare_eq_boundary_1023() {
441 let n = 1023;
442 let values: Vec<i64> = vec![42; n];
443 let result = batch_compare_eq(&values, 42, SimdAvailability::Avx2);
444 assert!(result.iter().all(|&b| b));
445 }
446
447 #[test]
448 fn test_batch_compare_eq_boundary_1024() {
449 let n = 1024;
450 let values: Vec<i64> = vec![42; n];
451 let avail = detect();
452 let result = batch_compare_eq(&values, 42, avail);
453 assert!(result.iter().all(|&b| b));
454 }
455
456 #[test]
457 fn test_batch_compare_eq_boundary_1025() {
458 let n = 1025;
459 let values: Vec<i64> = vec![42; n];
460 let avail = detect();
461 let result = batch_compare_eq(&values, 42, avail);
462 assert!(result.iter().all(|&b| b));
463 }
464}