#![allow(unsafe_code)]
#[must_use]
pub fn find(haystack: &[u8], needle: &[u8]) -> Option<usize> {
if needle.is_empty() {
return Some(0);
}
if needle.len() > haystack.len() {
return None;
}
#[cfg(target_arch = "x86_64")]
{
match crate::isa::tier() {
crate::isa::Tier::Avx512 => return unsafe { find_avx512(haystack, needle) },
crate::isa::Tier::Avx2 => return unsafe { find_avx2(haystack, needle) },
crate::isa::Tier::Sse2 => return unsafe { find_sse2(haystack, needle) },
crate::isa::Tier::Scalar => {}
}
}
find_scalar(haystack, needle)
}
#[must_use]
pub fn contains(haystack: &[u8], needle: &[u8]) -> bool {
find(haystack, needle).is_some()
}
#[must_use]
pub fn find_all(haystack: &[u8], needle: &[u8]) -> Vec<usize> {
if haystack.len() < FIND_ALL_PARALLEL_THRESHOLD {
crate::trace::rung("find all", "one pass, under the split's size", haystack.len());
return find_all_one_pass(haystack, needle);
}
crate::trace::rung("find all", "split across the cores", haystack.len());
find_all_across(haystack, needle)
}
const FIND_ALL_PARALLEL_THRESHOLD: usize = 2 * 1024 * 1024;
const FIND_ALL_SPAN: usize = 8 * 1024 * 1024;
#[must_use]
pub fn occurrences<'a>(haystack: &'a [u8], needle: &'a [u8]) -> Occurrences<'a> {
occurrences_by_span(haystack, needle, FIND_ALL_SPAN)
}
fn occurrences_by_span<'a>(haystack: &'a [u8], needle: &'a [u8], span: usize) -> Occurrences<'a> {
Occurrences { haystack, needle, span: span.max(1), from: 0, buf: Vec::new(), taken: 0 }
}
#[must_use]
pub fn occurrences_unsplit<'a>(haystack: &'a [u8], needle: &'a [u8]) -> Unsplit<'a> {
Unsplit { haystack, needle, from: 0 }
}
pub struct Unsplit<'a> {
haystack: &'a [u8],
needle: &'a [u8],
from: usize,
}
impl Iterator for Unsplit<'_> {
type Item = usize;
fn next(&mut self) -> Option<usize> {
if self.needle.is_empty() || self.from >= self.haystack.len() {
return None;
}
let rel = find(&self.haystack[self.from..], self.needle)?;
let at = self.from + rel;
self.from = at + 1;
Some(at)
}
}
pub struct Occurrences<'a> {
haystack: &'a [u8],
needle: &'a [u8],
span: usize,
from: usize,
buf: Vec<usize>,
taken: usize,
}
impl Iterator for Occurrences<'_> {
type Item = usize;
fn next(&mut self) -> Option<usize> {
let n = self.haystack.len();
while self.taken == self.buf.len() {
if self.needle.is_empty() || self.needle.len() > n || self.from >= n {
return None;
}
let end = (self.from + self.span).min(n);
let edge = (end + self.needle.len() - 1).min(n);
let base = self.from;
self.buf = find_all(&self.haystack[base..edge], self.needle);
self.buf.retain(|&r| base + r < end);
for at in &mut self.buf {
*at += base;
}
self.taken = 0;
self.from = end;
}
let at = self.buf[self.taken];
self.taken += 1;
Some(at)
}
}
#[must_use]
pub fn find_all_one_pass(haystack: &[u8], needle: &[u8]) -> Vec<usize> {
let mut out = Vec::new();
if needle.is_empty() {
return out;
}
let mut from = 0usize;
while let Some(rel) = find(&haystack[from..], needle) {
let at = from + rel;
out.push(at);
from = at + 1;
}
out
}
#[must_use]
pub fn find_all_across(haystack: &[u8], needle: &[u8]) -> Vec<usize> {
let n = haystack.len();
if needle.is_empty() || needle.len() > n {
return Vec::new();
}
let cores = std::thread::available_parallelism().map_or(1, std::num::NonZero::get);
let width = n.div_ceil(cores * 4).max(needle.len());
let leaves = n.div_ceil(width);
crate::trace::rung("find all across", "leaves", leaves);
let mut found: Vec<Vec<usize>> = vec![Vec::new(); leaves];
let plan = flynnel::JobPlan::set_profile(
0,
u32::try_from(leaves).unwrap_or(u32::MAX),
flynnel::DispatchProfile::Streaming,
);
flynnel::sched::par_iter::for_each_chunk_indexed_min_leaf(&plan, &mut found, 1, |base, slots| {
for (k, slot) in slots.iter_mut().enumerate() {
let lo = (base + k) * width;
let hi = (lo + width).min(n);
let edge = (hi + needle.len() - 1).min(n);
*slot = find_all_one_pass(&haystack[lo..edge], needle)
.into_iter()
.map(|r| lo + r)
.filter(|&at| at < hi)
.collect();
}
});
found.concat()
}
#[must_use]
pub fn find_scalar(haystack: &[u8], needle: &[u8]) -> Option<usize> {
if needle.is_empty() {
return Some(0);
}
if needle.len() > haystack.len() {
return None;
}
let first = needle[0];
let last = haystack.len() - needle.len();
(0..=last).find(|&i| haystack[i] == first && &haystack[i..i + needle.len()] == needle)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn find_avx2(haystack: &[u8], needle: &[u8]) -> Option<usize> {
use core::arch::x86_64::{
_mm256_cmpeq_epi8, _mm256_loadu_si256, _mm256_movemask_epi8, _mm256_set1_epi8,
};
let n = needle.len();
let last = haystack.len() - n;
let first = _mm256_set1_epi8(needle[0] as i8);
let final_byte = _mm256_set1_epi8(needle[n - 1] as i8);
let mut i = 0;
while i + 32 <= haystack.len() {
let chunk = unsafe { _mm256_loadu_si256(haystack.as_ptr().add(i).cast()) };
let mut mask = _mm256_movemask_epi8(_mm256_cmpeq_epi8(chunk, first)) as u32;
if mask != 0 && i + n - 1 + 32 <= haystack.len() {
let tail = unsafe { _mm256_loadu_si256(haystack.as_ptr().add(i + n - 1).cast()) };
mask &= _mm256_movemask_epi8(_mm256_cmpeq_epi8(tail, final_byte)) as u32;
}
while mask != 0 {
let pos = i + mask.trailing_zeros() as usize;
if pos <= last && &haystack[pos..pos + n] == needle {
return Some(pos);
}
mask &= mask - 1;
}
i += 32;
}
while i <= last {
if haystack[i] == needle[0] && &haystack[i..i + n] == needle {
return Some(i);
}
i += 1;
}
None
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse2")]
unsafe fn find_sse2(haystack: &[u8], needle: &[u8]) -> Option<usize> {
use core::arch::x86_64::{
_mm_cmpeq_epi8, _mm_loadu_si128, _mm_movemask_epi8, _mm_set1_epi8,
};
let n = needle.len();
let last = haystack.len() - n;
let first = _mm_set1_epi8(needle[0] as i8);
let final_byte = _mm_set1_epi8(needle[n - 1] as i8);
let mut i = 0;
while i + 16 <= haystack.len() {
let chunk = unsafe { _mm_loadu_si128(haystack.as_ptr().add(i).cast()) };
let mut mask = _mm_movemask_epi8(_mm_cmpeq_epi8(chunk, first)) as u32;
if mask != 0 && i + n - 1 + 16 <= haystack.len() {
let tail = unsafe { _mm_loadu_si128(haystack.as_ptr().add(i + n - 1).cast()) };
mask &= _mm_movemask_epi8(_mm_cmpeq_epi8(tail, final_byte)) as u32;
}
while mask != 0 {
let pos = i + mask.trailing_zeros() as usize;
if pos <= last && &haystack[pos..pos + n] == needle {
return Some(pos);
}
mask &= mask - 1;
}
i += 16;
}
while i <= last {
if haystack[i] == needle[0] && &haystack[i..i + n] == needle {
return Some(i);
}
i += 1;
}
None
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn find_avx512(haystack: &[u8], needle: &[u8]) -> Option<usize> {
use core::arch::x86_64::{_mm512_cmpeq_epi8_mask, _mm512_loadu_si512, _mm512_set1_epi8};
let n = needle.len();
let last = haystack.len() - n;
let first = _mm512_set1_epi8(needle[0] as i8);
let final_byte = _mm512_set1_epi8(needle[n - 1] as i8);
let mut i = 0;
while i + 64 <= haystack.len() {
let chunk = unsafe { _mm512_loadu_si512(haystack.as_ptr().add(i).cast()) };
let mut mask: u64 = _mm512_cmpeq_epi8_mask(chunk, first);
if mask != 0 && i + n - 1 + 64 <= haystack.len() {
let tail = unsafe { _mm512_loadu_si512(haystack.as_ptr().add(i + n - 1).cast()) };
mask &= _mm512_cmpeq_epi8_mask(tail, final_byte);
}
while mask != 0 {
let pos = i + mask.trailing_zeros() as usize;
if pos <= last && &haystack[pos..pos + n] == needle {
return Some(pos);
}
mask &= mask - 1;
}
i += 64;
}
while i <= last {
if haystack[i] == needle[0] && &haystack[i..i + n] == needle {
return Some(i);
}
i += 1;
}
None
}
#[must_use]
pub fn find_avx512_emulated(haystack: &[u8], needle: &[u8]) -> Option<usize> {
if needle.is_empty() {
return Some(0);
}
if needle.len() > haystack.len() {
return None;
}
let n = needle.len();
let last = haystack.len() - n;
let first = needle[0];
let final_byte = needle[n - 1];
let mut i = 0;
while i + 64 <= haystack.len() {
let mut mask: u64 = 0;
for j in 0..64 {
if haystack[i + j] == first {
mask |= 1u64 << j;
}
}
if mask != 0 && i + n - 1 + 64 <= haystack.len() {
let mut tail: u64 = 0;
for j in 0..64 {
if haystack[i + n - 1 + j] == final_byte {
tail |= 1u64 << j;
}
}
mask &= tail;
}
while mask != 0 {
let pos = i + mask.trailing_zeros() as usize;
if pos <= last && &haystack[pos..pos + n] == needle {
return Some(pos);
}
mask &= mask - 1;
}
i += 64;
}
while i <= last {
if haystack[i] == first && &haystack[i..i + n] == needle {
return Some(i);
}
i += 1;
}
None
}
#[must_use]
pub fn nth_byte(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
if n == 0 {
return None;
}
#[cfg(target_arch = "x86_64")]
{
match crate::isa::tier() {
crate::isa::Tier::Avx512 => return unsafe { nth_byte_avx512(haystack, byte, n) },
crate::isa::Tier::Avx2 => return unsafe { nth_byte_avx2(haystack, byte, n) },
crate::isa::Tier::Sse2 => return unsafe { nth_byte_sse2(haystack, byte, n) },
crate::isa::Tier::Scalar => {}
}
}
nth_byte_scalar(haystack, byte, n)
}
#[must_use]
pub fn nth_byte_scalar(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
let skip = n.checked_sub(1)?;
haystack.iter().enumerate().filter(|&(_, &b)| b == byte).nth(skip).map(|(i, _)| i)
}
#[must_use]
pub fn nth_byte_back(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
if n == 0 {
return None;
}
#[cfg(target_arch = "x86_64")]
{
match crate::isa::tier() {
crate::isa::Tier::Avx512 => return unsafe { nth_byte_back_avx512(haystack, byte, n) },
crate::isa::Tier::Avx2 => return unsafe { nth_byte_back_avx2(haystack, byte, n) },
crate::isa::Tier::Sse2 => return unsafe { nth_byte_back_sse2(haystack, byte, n) },
crate::isa::Tier::Scalar => {}
}
}
nth_byte_back_scalar(haystack, byte, n)
}
#[must_use]
pub fn nth_byte_back_scalar(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
let skip = n.checked_sub(1)?;
haystack.iter().enumerate().rev().filter(|&(_, &b)| b == byte).nth(skip).map(|(i, _)| i)
}
#[must_use]
pub fn count_byte(haystack: &[u8], byte: u8) -> usize {
#[cfg(target_arch = "x86_64")]
{
match crate::isa::tier() {
crate::isa::Tier::Avx512 => return unsafe { count_byte_avx512(haystack, byte) },
crate::isa::Tier::Avx2 => return unsafe { count_byte_avx2(haystack, byte) },
crate::isa::Tier::Sse2 => return unsafe { count_byte_sse2(haystack, byte) },
crate::isa::Tier::Scalar => {}
}
}
count_byte_scalar(haystack, byte)
}
#[must_use]
pub fn count_byte_scalar(haystack: &[u8], byte: u8) -> usize {
haystack.iter().filter(|&&b| b == byte).count()
}
#[must_use]
pub fn count_byte_across(haystack: &[u8], byte: u8) -> usize {
let n = haystack.len();
if n == 0 {
return 0;
}
let cores = std::thread::available_parallelism().map_or(1, std::num::NonZero::get);
let width = n.div_ceil(cores * 4).max(1);
let leaves = n.div_ceil(width);
crate::trace::rung("count byte across", "leaves", leaves);
let mut counts = vec![0usize; leaves];
let plan = flynnel::JobPlan::set_profile(
0,
u32::try_from(leaves).expect("the leaves number at most four a core"),
flynnel::DispatchProfile::Streaming,
);
flynnel::sched::par_iter::for_each_chunk_indexed_min_leaf(&plan, &mut counts, 1, |base, slots| {
for (k, slot) in slots.iter_mut().enumerate() {
let lo = (base + k) * width;
let hi = (lo + width).min(n);
*slot = count_byte(&haystack[lo..hi], byte);
}
});
counts.iter().sum()
}
pub const SPLIT_FROM: usize = 4 * 1024 * 1024;
#[must_use]
pub fn count_byte_split(haystack: &[u8], byte: u8) -> usize {
if haystack.len() < SPLIT_FROM { count_byte(haystack, byte) } else { count_byte_across(haystack, byte) }
}
#[inline]
fn nth_set_bit(mut mask: u64, k: usize) -> usize {
for _ in 1..k {
mask &= mask - 1;
}
mask.trailing_zeros() as usize
}
#[inline]
fn nth_set_bit_from_top(mut mask: u64, k: usize) -> usize {
for _ in 1..k {
mask ^= 1u64 << (63 - mask.leading_zeros());
}
(63 - mask.leading_zeros()) as usize
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn nth_byte_avx2(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
use core::arch::x86_64::{_mm256_cmpeq_epi8, _mm256_loadu_si256, _mm256_movemask_epi8, _mm256_set1_epi8};
let wanted = _mm256_set1_epi8(byte as i8);
let mut left = n;
let mut i = 0;
while i + 32 <= haystack.len() {
let chunk = unsafe { _mm256_loadu_si256(haystack.as_ptr().add(i).cast()) };
let mask = u64::from(_mm256_movemask_epi8(_mm256_cmpeq_epi8(chunk, wanted)) as u32);
let here = mask.count_ones() as usize;
if here >= left {
return Some(i + nth_set_bit(mask, left));
}
left -= here;
i += 32;
}
nth_byte_scalar(&haystack[i..], byte, left).map(|at| i + at)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse2")]
unsafe fn nth_byte_sse2(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
use core::arch::x86_64::{_mm_cmpeq_epi8, _mm_loadu_si128, _mm_movemask_epi8, _mm_set1_epi8};
let wanted = _mm_set1_epi8(byte as i8);
let mut left = n;
let mut i = 0;
while i + 16 <= haystack.len() {
let chunk = unsafe { _mm_loadu_si128(haystack.as_ptr().add(i).cast()) };
let mask = u64::from(_mm_movemask_epi8(_mm_cmpeq_epi8(chunk, wanted)) as u32);
let here = mask.count_ones() as usize;
if here >= left {
return Some(i + nth_set_bit(mask, left));
}
left -= here;
i += 16;
}
nth_byte_scalar(&haystack[i..], byte, left).map(|at| i + at)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn nth_byte_avx512(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
use core::arch::x86_64::{_mm512_cmpeq_epi8_mask, _mm512_loadu_si512, _mm512_set1_epi8};
let wanted = _mm512_set1_epi8(byte as i8);
let mut left = n;
let mut i = 0;
while i + 64 <= haystack.len() {
let chunk = unsafe { _mm512_loadu_si512(haystack.as_ptr().add(i).cast()) };
let mask: u64 = _mm512_cmpeq_epi8_mask(chunk, wanted);
let here = mask.count_ones() as usize;
if here >= left {
return Some(i + nth_set_bit(mask, left));
}
left -= here;
i += 64;
}
nth_byte_scalar(&haystack[i..], byte, left).map(|at| i + at)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn nth_byte_back_avx2(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
use core::arch::x86_64::{_mm256_cmpeq_epi8, _mm256_loadu_si256, _mm256_movemask_epi8, _mm256_set1_epi8};
let wanted = _mm256_set1_epi8(byte as i8);
let mut left = n;
let mut end = haystack.len();
while end >= 32 {
let start = end - 32;
let chunk = unsafe { _mm256_loadu_si256(haystack.as_ptr().add(start).cast()) };
let mask = u64::from(_mm256_movemask_epi8(_mm256_cmpeq_epi8(chunk, wanted)) as u32);
let here = mask.count_ones() as usize;
if here >= left {
return Some(start + nth_set_bit_from_top(mask, left));
}
left -= here;
end = start;
}
nth_byte_back_scalar(&haystack[..end], byte, left)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse2")]
unsafe fn nth_byte_back_sse2(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
use core::arch::x86_64::{_mm_cmpeq_epi8, _mm_loadu_si128, _mm_movemask_epi8, _mm_set1_epi8};
let wanted = _mm_set1_epi8(byte as i8);
let mut left = n;
let mut end = haystack.len();
while end >= 16 {
let start = end - 16;
let chunk = unsafe { _mm_loadu_si128(haystack.as_ptr().add(start).cast()) };
let mask = u64::from(_mm_movemask_epi8(_mm_cmpeq_epi8(chunk, wanted)) as u32);
let here = mask.count_ones() as usize;
if here >= left {
return Some(start + nth_set_bit_from_top(mask, left));
}
left -= here;
end = start;
}
nth_byte_back_scalar(&haystack[..end], byte, left)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn nth_byte_back_avx512(haystack: &[u8], byte: u8, n: usize) -> Option<usize> {
use core::arch::x86_64::{_mm512_cmpeq_epi8_mask, _mm512_loadu_si512, _mm512_set1_epi8};
let wanted = _mm512_set1_epi8(byte as i8);
let mut left = n;
let mut end = haystack.len();
while end >= 64 {
let start = end - 64;
let chunk = unsafe { _mm512_loadu_si512(haystack.as_ptr().add(start).cast()) };
let mask: u64 = _mm512_cmpeq_epi8_mask(chunk, wanted);
let here = mask.count_ones() as usize;
if here >= left {
return Some(start + nth_set_bit_from_top(mask, left));
}
left -= here;
end = start;
}
nth_byte_back_scalar(&haystack[..end], byte, left)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn count_byte_avx2(haystack: &[u8], byte: u8) -> usize {
use core::arch::x86_64::{_mm256_cmpeq_epi8, _mm256_loadu_si256, _mm256_movemask_epi8, _mm256_set1_epi8};
let wanted = _mm256_set1_epi8(byte as i8);
let mut count = 0usize;
let mut i = 0;
while i + 32 <= haystack.len() {
let chunk = unsafe { _mm256_loadu_si256(haystack.as_ptr().add(i).cast()) };
count += (_mm256_movemask_epi8(_mm256_cmpeq_epi8(chunk, wanted)) as u32).count_ones() as usize;
i += 32;
}
count + count_byte_scalar(&haystack[i..], byte)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse2")]
unsafe fn count_byte_sse2(haystack: &[u8], byte: u8) -> usize {
use core::arch::x86_64::{_mm_cmpeq_epi8, _mm_loadu_si128, _mm_movemask_epi8, _mm_set1_epi8};
let wanted = _mm_set1_epi8(byte as i8);
let mut count = 0usize;
let mut i = 0;
while i + 16 <= haystack.len() {
let chunk = unsafe { _mm_loadu_si128(haystack.as_ptr().add(i).cast()) };
count += (_mm_movemask_epi8(_mm_cmpeq_epi8(chunk, wanted)) as u32).count_ones() as usize;
i += 16;
}
count + count_byte_scalar(&haystack[i..], byte)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn count_byte_avx512(haystack: &[u8], byte: u8) -> usize {
use core::arch::x86_64::{_mm512_cmpeq_epi8_mask, _mm512_loadu_si512, _mm512_set1_epi8};
let wanted = _mm512_set1_epi8(byte as i8);
let mut count = 0usize;
let mut i = 0;
while i + 64 <= haystack.len() {
let chunk = unsafe { _mm512_loadu_si512(haystack.as_ptr().add(i).cast()) };
count += _mm512_cmpeq_epi8_mask(chunk, wanted).count_ones() as usize;
i += 64;
}
count + count_byte_scalar(&haystack[i..], byte)
}
#[must_use]
#[inline]
pub fn word_run(b: &[u8]) -> usize {
#[cfg(target_arch = "x86_64")]
{
if crate::isa::tier() >= crate::isa::Tier::Avx512 {
return unsafe { word_run_avx512(b) };
}
if crate::isa::tier() >= crate::isa::Tier::Avx2 {
return unsafe { word_run_avx2(b) };
}
}
word_run_scalar(b)
}
#[must_use]
pub fn word_run_scalar(b: &[u8]) -> usize {
b.iter().take_while(|&&c| c == b'_' || c.is_ascii_alphanumeric()).count()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn word_run_avx2(b: &[u8]) -> usize {
use core::arch::x86_64::{
_mm256_and_si256, _mm256_cmpeq_epi8, _mm256_cmpgt_epi8, _mm256_loadu_si256,
_mm256_movemask_epi8, _mm256_or_si256, _mm256_set1_epi8,
};
let n = b.len();
let mut i = 0;
while i + 32 <= n {
let v = unsafe { _mm256_loadu_si256(b.as_ptr().add(i).cast()) };
let lower = _mm256_and_si256(
_mm256_cmpgt_epi8(v, _mm256_set1_epi8(96)),
_mm256_cmpgt_epi8(_mm256_set1_epi8(123), v),
);
let upper = _mm256_and_si256(
_mm256_cmpgt_epi8(v, _mm256_set1_epi8(64)),
_mm256_cmpgt_epi8(_mm256_set1_epi8(91), v),
);
let digit = _mm256_and_si256(
_mm256_cmpgt_epi8(v, _mm256_set1_epi8(47)),
_mm256_cmpgt_epi8(_mm256_set1_epi8(58), v),
);
let under = _mm256_cmpeq_epi8(v, _mm256_set1_epi8(95));
let word = _mm256_or_si256(_mm256_or_si256(lower, upper), _mm256_or_si256(digit, under));
let mask = _mm256_movemask_epi8(word) as u32;
if mask != 0xFFFF_FFFF {
return i + (!mask).trailing_zeros() as usize;
}
i += 32;
}
i + word_run_scalar(&b[i..])
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn word_run_avx512(b: &[u8]) -> usize {
use core::arch::x86_64::{
_mm512_cmpeq_epi8_mask, _mm512_cmpgt_epi8_mask, _mm512_loadu_si512, _mm512_set1_epi8,
};
let n = b.len();
let mut i = 0;
while i + 64 <= n {
let v = unsafe { _mm512_loadu_si512(b.as_ptr().add(i).cast()) };
let lower = _mm512_cmpgt_epi8_mask(v, _mm512_set1_epi8(96))
& _mm512_cmpgt_epi8_mask(_mm512_set1_epi8(123), v);
let upper = _mm512_cmpgt_epi8_mask(v, _mm512_set1_epi8(64))
& _mm512_cmpgt_epi8_mask(_mm512_set1_epi8(91), v);
let digit = _mm512_cmpgt_epi8_mask(v, _mm512_set1_epi8(47))
& _mm512_cmpgt_epi8_mask(_mm512_set1_epi8(58), v);
let under = _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(95));
let word: u64 = lower | upper | digit | under;
if word != u64::MAX {
return i + (!word).trailing_zeros() as usize;
}
i += 64;
}
i + word_run_scalar(&b[i..])
}
#[must_use]
#[inline]
pub fn space_run(b: &[u8]) -> usize {
#[cfg(target_arch = "x86_64")]
{
if crate::isa::tier() >= crate::isa::Tier::Avx512 {
return unsafe { space_run_avx512(b) };
}
if crate::isa::tier() >= crate::isa::Tier::Avx2 {
return unsafe { space_run_avx2(b) };
}
}
space_run_scalar(b)
}
#[must_use]
pub fn space_run_scalar(b: &[u8]) -> usize {
b.iter().take_while(|&&c| c.is_ascii_whitespace()).count()
}
#[must_use]
#[inline]
pub fn nonspace_run(b: &[u8]) -> usize {
#[cfg(target_arch = "x86_64")]
{
if crate::isa::tier() >= crate::isa::Tier::Avx512 {
return unsafe { nonspace_run_avx512(b) };
}
if crate::isa::tier() >= crate::isa::Tier::Avx2 {
return unsafe { nonspace_run_avx2(b) };
}
}
nonspace_run_scalar(b)
}
#[must_use]
pub fn nonspace_run_scalar(b: &[u8]) -> usize {
b.iter().take_while(|&&c| !c.is_ascii_whitespace()).count()
}
#[must_use]
pub fn nonspace_spans_at_least(b: &[u8], min_len: usize) -> Vec<(usize, usize)> {
#[cfg(target_arch = "x86_64")]
{
if crate::isa::tier() >= crate::isa::Tier::Avx512 {
return unsafe { nonspace_spans_at_least_avx512(b, min_len) };
}
if crate::isa::tier() >= crate::isa::Tier::Avx2 {
return unsafe { nonspace_spans_at_least_avx2(b, min_len) };
}
}
nonspace_spans_at_least_scalar(b, min_len)
}
#[must_use]
pub fn nonspace_spans_at_least_scalar(b: &[u8], min_len: usize) -> Vec<(usize, usize)> {
let min_len = min_len.max(1);
let mut spans = Vec::new();
let mut start = 0;
for (j, &c) in b.iter().enumerate() {
if c.is_ascii_whitespace() {
if j - start >= min_len {
spans.push((start, j));
}
start = j + 1;
}
}
if b.len() - start >= min_len {
spans.push((start, b.len()));
}
spans
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn nonspace_spans_at_least_avx2(b: &[u8], min_len: usize) -> Vec<(usize, usize)> {
use core::arch::x86_64::_mm256_loadu_si256;
let min_len = min_len.max(1);
let n = b.len();
let mut spans = Vec::new();
let mut start = 0;
let mut i = 0;
while i + 32 <= n {
let mask = unsafe { whitespace_lanes(_mm256_loadu_si256(b.as_ptr().add(i).cast())) };
if mask != 0 {
let first = mask.trailing_zeros() as usize;
if i + first - start >= min_len {
spans.push((start, i + first));
}
if min_len < 32 && 32 - mask.count_ones() as usize >= min_len {
let mut rest = mask & (mask - 1);
let mut prev = first;
while rest != 0 {
let q = rest.trailing_zeros() as usize;
if q - prev > min_len {
spans.push((i + prev + 1, i + q));
}
prev = q;
rest &= rest - 1;
}
}
start = i + 32 - mask.leading_zeros() as usize;
}
i += 32;
}
for (j, &c) in b.iter().enumerate().skip(i) {
if c.is_ascii_whitespace() {
if j - start >= min_len {
spans.push((start, j));
}
start = j + 1;
}
}
if n - start >= min_len {
spans.push((start, n));
}
spans
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn nonspace_spans_at_least_avx512(b: &[u8], min_len: usize) -> Vec<(usize, usize)> {
use core::arch::x86_64::_mm512_loadu_si512;
let min_len = min_len.max(1);
let n = b.len();
let mut spans = Vec::new();
let mut start = 0;
let mut i = 0;
while i + 64 <= n {
let mask =
unsafe { whitespace_lanes_512(_mm512_loadu_si512(b.as_ptr().add(i).cast())) };
if mask != 0 {
let first = mask.trailing_zeros() as usize;
if i + first - start >= min_len {
spans.push((start, i + first));
}
if min_len < 64 && 64 - mask.count_ones() as usize >= min_len {
let mut rest = mask & (mask - 1);
let mut prev = first;
while rest != 0 {
let q = rest.trailing_zeros() as usize;
if q - prev > min_len {
spans.push((i + prev + 1, i + q));
}
prev = q;
rest &= rest - 1;
}
}
start = i + 64 - mask.leading_zeros() as usize;
}
i += 64;
}
for (j, &c) in b.iter().enumerate().skip(i) {
if c.is_ascii_whitespace() {
if j - start >= min_len {
spans.push((start, j));
}
start = j + 1;
}
}
if n - start >= min_len {
spans.push((start, n));
}
spans
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
#[inline]
unsafe fn whitespace_lanes_512(v: core::arch::x86_64::__m512i) -> u64 {
use core::arch::x86_64::{_mm512_cmpeq_epi8_mask, _mm512_set1_epi8};
_mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(9))
| _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(10))
| _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(12))
| _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(13))
| _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(32))
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn whitespace_lanes(v: core::arch::x86_64::__m256i) -> u32 {
use core::arch::x86_64::{_mm256_cmpeq_epi8, _mm256_movemask_epi8, _mm256_or_si256, _mm256_set1_epi8};
let ws = _mm256_or_si256(
_mm256_or_si256(
_mm256_or_si256(
_mm256_cmpeq_epi8(v, _mm256_set1_epi8(9)),
_mm256_cmpeq_epi8(v, _mm256_set1_epi8(10)),
),
_mm256_or_si256(
_mm256_cmpeq_epi8(v, _mm256_set1_epi8(12)),
_mm256_cmpeq_epi8(v, _mm256_set1_epi8(13)),
),
),
_mm256_cmpeq_epi8(v, _mm256_set1_epi8(32)),
);
_mm256_movemask_epi8(ws) as u32
}
pub struct PairPositions<'a, const X: u8, const Y: u8> {
b: &'a [u8],
tier: crate::isa::Tier,
base: usize,
mask: u64,
}
pub type QuotePositions<'a> = PairPositions<'a, b'"', b'\''>;
pub type CloseOrNewlinePositions<'a> = PairPositions<'a, b'"', b'\n'>;
impl<'a, const X: u8, const Y: u8> PairPositions<'a, X, Y> {
#[must_use]
pub fn new(b: &'a [u8]) -> Self {
let tier = crate::isa::tier();
let mask = pair_lanes_at::<X, Y>(b, 0, tier);
PairPositions { b, tier, base: 0, mask }
}
pub fn next_at_or_after(&mut self, from: usize) -> Option<usize> {
let n = self.b.len();
if from >= n {
return None;
}
if from < self.base || from >= self.base + 64 {
self.base = from & !63;
self.mask = pair_lanes_at::<X, Y>(self.b, self.base, self.tier);
}
let mut mask = self.mask & (!0u64 << (from - self.base));
while mask == 0 {
self.base += 64;
if self.base >= n {
return None;
}
self.mask = pair_lanes_at::<X, Y>(self.b, self.base, self.tier);
mask = self.mask;
}
Some(self.base + mask.trailing_zeros() as usize)
}
}
#[cfg_attr(not(target_arch = "x86_64"), allow(unused_variables))]
fn pair_lanes_at<const X: u8, const Y: u8>(b: &[u8], base: usize, tier: crate::isa::Tier) -> u64 {
let end = (base + 64).min(b.len());
let block = &b[base..end];
#[cfg(target_arch = "x86_64")]
{
if block.len() == 64 {
if tier >= crate::isa::Tier::Avx512 {
return unsafe { pair_lanes_512::<X, Y>(block) };
}
if tier >= crate::isa::Tier::Avx2 {
return unsafe { pair_lanes_avx2::<X, Y>(block) };
}
}
}
let mut mask = 0u64;
for (i, &c) in block.iter().enumerate() {
if c == X || c == Y {
mask |= 1u64 << i;
}
}
mask
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn pair_lanes_512<const X: u8, const Y: u8>(block: &[u8]) -> u64 {
use core::arch::x86_64::{_mm512_cmpeq_epi8_mask, _mm512_loadu_si512, _mm512_set1_epi8};
let v = unsafe { _mm512_loadu_si512(block.as_ptr().cast()) };
_mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(X as i8)) | _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(Y as i8))
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn pair_lanes_avx2<const X: u8, const Y: u8>(block: &[u8]) -> u64 {
use core::arch::x86_64::{_mm256_cmpeq_epi8, _mm256_loadu_si256, _mm256_movemask_epi8, _mm256_or_si256, _mm256_set1_epi8};
let half = |at: usize| {
let v = unsafe { _mm256_loadu_si256(block.as_ptr().add(at).cast()) };
let q = _mm256_or_si256(
_mm256_cmpeq_epi8(v, _mm256_set1_epi8(X as i8)),
_mm256_cmpeq_epi8(v, _mm256_set1_epi8(Y as i8)),
);
_mm256_movemask_epi8(q) as u32 as u64
};
half(0) | (half(32) << 32)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn space_run_avx2(b: &[u8]) -> usize {
use core::arch::x86_64::_mm256_loadu_si256;
let n = b.len();
let mut i = 0;
while i + 32 <= n {
let mask = unsafe { whitespace_lanes(_mm256_loadu_si256(b.as_ptr().add(i).cast())) };
if mask != 0xFFFF_FFFF {
return i + (!mask).trailing_zeros() as usize;
}
i += 32;
}
i + space_run_scalar(&b[i..])
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn space_run_avx512(b: &[u8]) -> usize {
use core::arch::x86_64::_mm512_loadu_si512;
let n = b.len();
let mut i = 0;
while i + 64 <= n {
let mask =
unsafe { whitespace_lanes_512(_mm512_loadu_si512(b.as_ptr().add(i).cast())) };
if mask != u64::MAX {
return i + (!mask).trailing_zeros() as usize;
}
i += 64;
}
i + space_run_scalar(&b[i..])
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn nonspace_run_avx2(b: &[u8]) -> usize {
use core::arch::x86_64::_mm256_loadu_si256;
let n = b.len();
let mut i = 0;
while i + 32 <= n {
let mask = unsafe { whitespace_lanes(_mm256_loadu_si256(b.as_ptr().add(i).cast())) };
if mask != 0 {
return i + mask.trailing_zeros() as usize;
}
i += 32;
}
i + nonspace_run_scalar(&b[i..])
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn nonspace_run_avx512(b: &[u8]) -> usize {
use core::arch::x86_64::_mm512_loadu_si512;
let n = b.len();
let mut i = 0;
while i + 64 <= n {
let mask =
unsafe { whitespace_lanes_512(_mm512_loadu_si512(b.as_ptr().add(i).cast())) };
if mask != 0 {
return i + mask.trailing_zeros() as usize;
}
i += 64;
}
i + nonspace_run_scalar(&b[i..])
}
#[must_use]
#[inline]
pub fn digit_run(b: &[u8]) -> usize {
#[cfg(target_arch = "x86_64")]
{
if crate::isa::tier() >= crate::isa::Tier::Avx512 {
return unsafe { digit_run_avx512(b) };
}
if crate::isa::tier() >= crate::isa::Tier::Avx2 {
return unsafe { digit_run_avx2(b) };
}
}
digit_run_scalar(b)
}
#[must_use]
pub fn digit_run_scalar(b: &[u8]) -> usize {
b.iter().take_while(|c| c.is_ascii_digit()).count()
}
#[must_use]
#[inline]
pub fn digit_find(b: &[u8]) -> Option<usize> {
#[cfg(target_arch = "x86_64")]
{
if crate::isa::tier() >= crate::isa::Tier::Avx512 {
return unsafe { digit_find_avx512(b) };
}
if crate::isa::tier() >= crate::isa::Tier::Avx2 {
return unsafe { digit_find_avx2(b) };
}
}
digit_find_scalar(b)
}
#[must_use]
pub fn digit_find_scalar(b: &[u8]) -> Option<usize> {
b.iter().position(u8::is_ascii_digit)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn digit_find_avx2(b: &[u8]) -> Option<usize> {
use core::arch::x86_64::{
_mm256_and_si256, _mm256_cmpgt_epi8, _mm256_loadu_si256, _mm256_movemask_epi8,
_mm256_set1_epi8,
};
let n = b.len();
let mut i = 0;
while i + 32 <= n {
let v = unsafe { _mm256_loadu_si256(b.as_ptr().add(i).cast()) };
let digit = _mm256_and_si256(
_mm256_cmpgt_epi8(v, _mm256_set1_epi8(47)),
_mm256_cmpgt_epi8(_mm256_set1_epi8(58), v),
);
let mask = _mm256_movemask_epi8(digit) as u32;
if mask != 0 {
return Some(i + mask.trailing_zeros() as usize);
}
i += 32;
}
digit_find_scalar(&b[i..]).map(|k| i + k)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn digit_find_avx512(b: &[u8]) -> Option<usize> {
use core::arch::x86_64::{_mm512_cmpgt_epi8_mask, _mm512_loadu_si512, _mm512_set1_epi8};
let n = b.len();
let mut i = 0;
while i + 64 <= n {
let v = unsafe { _mm512_loadu_si512(b.as_ptr().add(i).cast()) };
let digit: u64 = _mm512_cmpgt_epi8_mask(v, _mm512_set1_epi8(47))
& _mm512_cmpgt_epi8_mask(_mm512_set1_epi8(58), v);
if digit != 0 {
return Some(i + digit.trailing_zeros() as usize);
}
i += 64;
}
digit_find_scalar(&b[i..]).map(|k| i + k)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn digit_run_avx2(b: &[u8]) -> usize {
use core::arch::x86_64::{
_mm256_and_si256, _mm256_cmpgt_epi8, _mm256_loadu_si256, _mm256_movemask_epi8,
_mm256_set1_epi8,
};
let n = b.len();
let mut i = 0;
while i + 32 <= n {
let v = unsafe { _mm256_loadu_si256(b.as_ptr().add(i).cast()) };
let digit = _mm256_and_si256(
_mm256_cmpgt_epi8(v, _mm256_set1_epi8(47)),
_mm256_cmpgt_epi8(_mm256_set1_epi8(58), v),
);
let mask = _mm256_movemask_epi8(digit) as u32;
if mask != 0xFFFF_FFFF {
return i + (!mask).trailing_zeros() as usize;
}
i += 32;
}
i + digit_run_scalar(&b[i..])
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw")]
unsafe fn digit_run_avx512(b: &[u8]) -> usize {
use core::arch::x86_64::{_mm512_cmpgt_epi8_mask, _mm512_loadu_si512, _mm512_set1_epi8};
let n = b.len();
let mut i = 0;
while i + 64 <= n {
let v = unsafe { _mm512_loadu_si512(b.as_ptr().add(i).cast()) };
let digit: u64 = _mm512_cmpgt_epi8_mask(v, _mm512_set1_epi8(47))
& _mm512_cmpgt_epi8_mask(_mm512_set1_epi8(58), v);
if digit != u64::MAX {
return i + (!digit).trailing_zeros() as usize;
}
i += 64;
}
i + digit_run_scalar(&b[i..])
}
#[derive(Clone, Copy, Debug)]
pub struct ClassTables {
lo: [u8; 16],
hi: [u8; 16],
}
impl ClassTables {
#[must_use]
pub fn for_class(class: &crate::byte_nfa::ByteClass) -> Option<ClassTables> {
let mut lo = [0u8; 16];
for b in 0..=255u8 {
if !class.has(b) {
continue;
}
if b >= 0x80 {
return None;
}
lo[(b & 0x0F) as usize] |= 1 << (b >> 4);
}
let mut hi = [0u8; 16];
for (h, slot) in hi.iter_mut().enumerate() {
*slot = if h < 8 { 1 << h } else { 0 };
}
Some(ClassTables { lo, hi })
}
#[must_use]
pub fn has(&self, b: u8) -> bool {
self.lo[(b & 0x0F) as usize] & self.hi[((b >> 4) & 0x0F) as usize] != 0
}
#[must_use]
pub fn has_run_of(&self, input: &[u8], least: usize) -> bool {
if least == 0 {
return true;
}
self.run_scan(input, Some(least)) >= least
}
#[must_use]
pub fn longest_run(&self, input: &[u8]) -> usize {
self.run_scan(input, None)
}
fn run_scan(&self, input: &[u8], stop_at: Option<usize>) -> usize {
#[cfg(target_arch = "x86_64")]
{
if matches!(crate::isa::tier(), crate::isa::Tier::Avx2 | crate::isa::Tier::Avx512) {
return unsafe { self.run_scan_avx2(input, stop_at) };
}
}
self.run_scan_scalar(input, stop_at)
}
#[must_use]
pub fn longest_run_scalar(&self, input: &[u8]) -> usize {
self.run_scan_scalar(input, None)
}
fn run_scan_scalar(&self, input: &[u8], stop_at: Option<usize>) -> usize {
let mut longest = 0usize;
let mut run = 0usize;
for &b in input {
if self.has(b) {
run += 1;
if run > longest {
longest = run;
if stop_at.is_some_and(|n| longest >= n) {
return longest;
}
}
} else {
run = 0;
}
}
longest
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn run_scan_avx2(&self, input: &[u8], stop_at: Option<usize>) -> usize {
use core::arch::x86_64::{
__m256i, _mm256_and_si256, _mm256_broadcastsi128_si256, _mm256_cmpeq_epi8,
_mm256_loadu_si256, _mm256_movemask_epi8, _mm256_set1_epi8, _mm256_setzero_si256,
_mm256_shuffle_epi8, _mm256_srli_epi16, _mm_loadu_si128,
};
let lo_tbl =
_mm256_broadcastsi128_si256(unsafe { _mm_loadu_si128(self.lo.as_ptr().cast()) });
let hi_tbl =
_mm256_broadcastsi128_si256(unsafe { _mm_loadu_si128(self.hi.as_ptr().cast()) });
let low_nibble = _mm256_set1_epi8(0x0F);
let zero = _mm256_setzero_si256();
let mut longest = 0usize;
let mut run = 0usize;
let mut at = 0usize;
while at + 32 <= input.len() {
let v: __m256i = unsafe { _mm256_loadu_si256(input.as_ptr().add(at).cast()) };
let lo_idx = _mm256_and_si256(v, low_nibble);
let hi_idx = _mm256_and_si256(_mm256_srli_epi16::<4>(v), low_nibble);
let met = _mm256_and_si256(
_mm256_shuffle_epi8(lo_tbl, lo_idx),
_mm256_shuffle_epi8(hi_tbl, hi_idx),
);
let absent = _mm256_movemask_epi8(_mm256_cmpeq_epi8(met, zero)) as u32;
let present = !absent;
if present == u32::MAX {
run += 32;
if run > longest {
longest = run;
}
} else {
let head = present.trailing_ones() as usize;
run += head;
if run > longest {
longest = run;
}
let tail = present.leading_ones() as usize;
let mut rest = present >> head;
let mut seen = head;
while rest != 0 {
let gap = rest.trailing_zeros() as usize;
rest >>= gap;
seen += gap;
if rest == 0 {
break;
}
let ones = rest.trailing_ones() as usize;
if seen + ones < 32 && ones > longest {
longest = ones;
}
rest >>= ones;
seen += ones;
}
run = tail;
if run > longest {
longest = run;
}
}
if stop_at.is_some_and(|n| longest >= n) {
return longest;
}
at += 32;
}
for &b in &input[at..] {
if self.has(b) {
run += 1;
if run > longest {
longest = run;
}
} else {
run = 0;
}
}
longest
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pair_positions_are_the_bytes_a_reading_finds() {
fn check<const X: u8, const Y: u8>(text: &[u8]) {
let want: Vec<usize> =
text.iter().enumerate().filter(|&(_, &c)| c == X || c == Y).map(|(i, _)| i).collect();
assert!(want.len() > 100, "the file holds the bytes to find");
let first_at_or_after = |from: usize| want[want.partition_point(|&q| q < from)..].first().copied();
for stride in [1usize, 2, 5, 63, 64, 65, 200, 4097] {
let mut found = PairPositions::<'_, X, Y>::new(text);
let mut from = 0;
let mut handed = 0;
while let Some(p) = found.next_at_or_after(from) {
assert_eq!(Some(p), first_at_or_after(from), "from {from}, stride {stride}");
handed += 1;
from = p + stride;
}
assert_eq!(first_at_or_after(from), None, "stride {stride}: a byte at or past {from} was not handed out");
if stride == 1 {
assert_eq!(handed, want.len());
}
}
let mut found = PairPositions::<'_, X, Y>::new(text);
for from in (0..text.len()).rev().step_by(37).chain([0, 130, 1, 64, 63, 5000, 64, 0]) {
assert_eq!(found.next_at_or_after(from), first_at_or_after(from), "from {from}, asked out of order");
}
assert_eq!(found.next_at_or_after(text.len()), None);
assert_eq!(found.next_at_or_after(text.len() + 100), None);
}
let text: &[u8] = include_bytes!("byte_simd.rs");
check::<b'"', b'\''>(text);
check::<b'"', b'\n'>(text);
assert_eq!(QuotePositions::new(b"").next_at_or_after(0), None);
assert_eq!(QuotePositions::new(b"no quote here").next_at_or_after(0), None);
assert_eq!(CloseOrNewlinePositions::new(b"no close here").next_at_or_after(0), None);
}
#[test]
fn digit_find_matches_scalar() {
let big: Vec<u8> = (0..6000u32)
.map(|i| {
let r = i.wrapping_mul(2_654_435_761) >> 23;
match r % 5 {
0 => (r & 0xFF) as u8,
1 => b'0' + (r % 10) as u8,
2 => b' ',
_ => b'a' + (r % 26) as u8,
}
})
.collect();
for start in [0usize, 1, 31, 32, 33, 63, 64, 65, 1000, 5999] {
let tail = &big[start.min(big.len())..];
assert_eq!(digit_find(tail), digit_find_scalar(tail), "from {start}");
}
assert_eq!(digit_find(b""), None);
assert_eq!(digit_find(&[b'x'; 200]), None);
let mut first = [b'x'; 200];
first[0] = b'7';
assert_eq!(digit_find(&first), Some(0));
let mut block = [b'x'; 200];
block[63] = b'7';
assert_eq!(digit_find(&block), Some(63));
let mut tail = [b'x'; 200];
tail[199] = b'7';
assert_eq!(digit_find(&tail), Some(199));
}
#[test]
fn run_finders_match_scalar() {
let big: Vec<u8> = (0..6000u32)
.map(|i| {
let r = i.wrapping_mul(2_654_435_761) >> 23;
if r % 4 == 0 {
(r & 0xFF) as u8
} else {
b" \t\n0189abcXYZ_"[(r % 13) as usize]
}
})
.collect();
for start in [0usize, 1, 5, 31, 32, 33, 63, 64, 100, 1000, 5990, 5999] {
let s = &big[start..];
assert_eq!(space_run(s), space_run_scalar(s), "space @ {start}");
assert_eq!(digit_run(s), digit_run_scalar(s), "digit @ {start}");
}
assert_eq!(space_run(b" \t\nx"), 5);
assert_eq!(space_run(b"\x0b spaces"), 0); assert_eq!(digit_run(b"0123456789012345678901234567890123!"), 34);
assert_eq!(digit_run(b"12.5"), 2);
}
#[test]
fn word_run_matches_scalar() {
let big: Vec<u8> = (0..5000u32)
.map(|i| {
let r = i.wrapping_mul(2_654_435_761) >> 24;
if r % 5 == 0 { (r & 0xFF) as u8 } else { b"abcXYZ_0189"[(r % 11) as usize] }
})
.collect();
for start in [0usize, 1, 5, 31, 32, 33, 63, 64, 100, 1000, 4990, 4999] {
assert_eq!(
word_run(&big[start..]),
word_run_scalar(&big[start..]),
"mismatch at start {start}"
);
}
assert_eq!(word_run(b""), 0);
assert_eq!(word_run(b" abc"), 0);
assert_eq!(word_run(b"abcdefghijklmnopqrstuvwxyz012345!"), 32);
assert_eq!(word_run(b"name_1 rest"), 6);
for len in [31usize, 32, 33, 63, 64, 65, 127, 128, 129] {
let mut run: Vec<u8> = std::iter::repeat_n(b'w', len).collect();
run.push(b' ');
run.extend_from_slice(b"tail");
assert_eq!(word_run(&run), len, "a run of {len} word bytes");
assert_eq!(word_run_scalar(&run), len, "the scalar reference at {len}");
}
}
#[test]
fn every_run_function_agrees_with_its_reference_at_both_window_widths() {
for len in [0usize, 1, 31, 32, 33, 63, 64, 65, 127, 128, 129, 200] {
for (class, other) in [(b'w', b' '), (b'7', b'x'), (b' ', b'w'), (b'x', b'\n')] {
let mut run: Vec<u8> = std::iter::repeat_n(class, len).collect();
run.push(other);
run.extend_from_slice(b"tail rest");
assert_eq!(word_run(&run), word_run_scalar(&run), "word at {len}");
assert_eq!(digit_run(&run), digit_run_scalar(&run), "digit at {len}");
assert_eq!(space_run(&run), space_run_scalar(&run), "space at {len}");
assert_eq!(
nonspace_run(&run),
nonspace_run_scalar(&run),
"nonspace at {len}"
);
}
}
}
#[test]
fn nonspace_run_matches_scalar() {
let big: Vec<u8> = (0..5000u32)
.map(|i| {
let r = i.wrapping_mul(2_654_435_761) >> 24;
match r % 23 {
0 => b' ',
1 => b'\n',
2 => b'\t',
3 => 11,
_ => b"abcXYZ_0189+/=."[(r % 15) as usize],
}
})
.collect();
for start in [0usize, 1, 5, 31, 32, 33, 63, 64, 100, 1000, 4990, 4999] {
assert_eq!(
nonspace_run(&big[start..]),
nonspace_run_scalar(&big[start..]),
"mismatch at start {start}"
);
}
assert_eq!(nonspace_run(b""), 0);
assert_eq!(nonspace_run(b" abc"), 0);
assert_eq!(nonspace_run(b"abcdefghijklmnopqrstuvwxyz012345 "), 32);
assert_eq!(nonspace_run(b"name_1\x0bx rest"), 8);
assert_eq!(nonspace_run(b"\x0c"), 0);
}
#[test]
fn nonspace_spans_match_scalar() {
let big: Vec<u8> = (0..6000u32)
.map(|i| {
let r = i.wrapping_mul(2_654_435_761) >> 24;
if r % 41 == 0 || i % 1000 == 31 || i % 1000 == 32 {
b' '
} else {
b"abcXYZ_0189+/=."[(r % 15) as usize]
}
})
.collect();
for min_len in [1usize, 2, 16, 31, 32, 33, 48, 63, 64, 65, 100, 1000] {
for start in [0usize, 1, 31, 32, 33, 63, 64, 65, 100, 127, 128, 5990] {
assert_eq!(
nonspace_spans_at_least(&big[start..], min_len),
nonspace_spans_at_least_scalar(&big[start..], min_len),
"min_len {min_len} from {start}"
);
}
}
assert_eq!(nonspace_spans_at_least(b"", 1), Vec::new());
assert_eq!(nonspace_spans_at_least(b" ", 1), Vec::new());
assert_eq!(nonspace_spans_at_least(b"abc", 3), vec![(0, 3)]);
assert_eq!(nonspace_spans_at_least(b"abc def", 3), vec![(0, 3), (4, 7)]);
assert_eq!(nonspace_spans_at_least(b"ab cdef", 3), vec![(3, 7)]);
let long: Vec<u8> =
std::iter::repeat_n(b'x', 70).chain(*b"\n").chain(std::iter::repeat_n(b'y', 47)).collect();
assert_eq!(nonspace_spans_at_least(&long, 48), vec![(0, 70)]);
}
#[test]
fn dispatched_find_matches_scalar_on_small_cases() {
let cases: &[(&[u8], &[u8])] = &[
(b"hello world", b"world"),
(b"hello world", b"xyz"),
(b"aaaaaaab", b"ab"),
(b"abcabcabc", b"cab"),
(b"", b"x"),
(b"x", b""),
(b"needle at end NEEDLE", b"NEEDLE"),
(b"\x00\x01\x02\x03", b"\x02\x03"),
];
for &(h, n) in cases {
assert_eq!(find(h, n), find_scalar(h, n), "h={h:?} n={n:?}");
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn each_simd_path_matches_scalar_on_large_input() {
let big: Vec<u8> =
(0..5000u32).map(|i| (i.wrapping_mul(2_654_435_761) >> 24) as u8).collect();
let needles: Vec<&[u8]> = vec![
&big[0..1],
&big[100..104],
&big[2500..2506],
&big[4996..5000], &big[1000..1100], b"\xff\xff\xff",
b"\xff\xff", ];
if std::is_x86_feature_detected!("avx2") {
for nd in &needles {
assert_eq!(unsafe { find_avx2(&big, nd) }, find_scalar(&big, nd));
}
}
if std::is_x86_feature_detected!("sse2") {
for nd in &needles {
assert_eq!(unsafe { find_sse2(&big, nd) }, find_scalar(&big, nd));
}
}
#[cfg(target_arch = "x86_64")]
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("avx512bw") {
for nd in &needles {
assert_eq!(unsafe { find_avx512(&big, nd) }, find_scalar(&big, nd));
}
}
}
#[test]
fn the_counting_searches_answer_what_the_scalar_baselines_answer() {
let big: Vec<u8> = (0..5003u32).map(|i| (i.wrapping_mul(2_654_435_761) >> 24) as u8).collect();
let texts: Vec<&[u8]> = vec![&big, &big[..1], &big[..15], &big[..31], &big[..63], &big[..65], b"", b"\n\n\n"];
for text in &texts {
for byte in [big[0], b'\n', 0xFF] {
let total = count_byte_scalar(text, byte);
assert_eq!(count_byte(text, byte), total, "count of {byte} over {} bytes", text.len());
assert_eq!(count_byte_across(text, byte), total, "split count of {byte} over {} bytes", text.len());
for n in 0..=total + 1 {
assert_eq!(nth_byte(text, byte, n), nth_byte_scalar(text, byte, n), "{n}th {byte}");
assert_eq!(nth_byte_back(text, byte, n), nth_byte_back_scalar(text, byte, n), "{n}th {byte} from the end");
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn each_counting_path_matches_scalar() {
let big: Vec<u8> = (0..5003u32).map(|i| (i.wrapping_mul(2_654_435_761) >> 24) as u8).collect();
let byte = big[7];
let total = count_byte_scalar(&big, byte);
let check = |name: &str,
nth: &dyn Fn(usize) -> Option<usize>,
back: &dyn Fn(usize) -> Option<usize>,
count: &dyn Fn() -> usize| {
assert_eq!(count(), total, "{name} count");
for n in 1..=total + 1 {
assert_eq!(nth(n), nth_byte_scalar(&big, byte, n), "{name} {n}th");
assert_eq!(back(n), nth_byte_back_scalar(&big, byte, n), "{name} {n}th from the end");
}
};
if std::is_x86_feature_detected!("sse2") {
check(
"sse2",
&|n| unsafe { nth_byte_sse2(&big, byte, n) },
&|n| unsafe { nth_byte_back_sse2(&big, byte, n) },
&|| unsafe { count_byte_sse2(&big, byte) },
);
}
if std::is_x86_feature_detected!("avx2") {
check(
"avx2",
&|n| unsafe { nth_byte_avx2(&big, byte, n) },
&|n| unsafe { nth_byte_back_avx2(&big, byte, n) },
&|| unsafe { count_byte_avx2(&big, byte) },
);
}
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("avx512bw") {
check(
"avx512",
&|n| unsafe { nth_byte_avx512(&big, byte, n) },
&|n| unsafe { nth_byte_back_avx512(&big, byte, n) },
&|| unsafe { count_byte_avx512(&big, byte) },
);
}
}
#[test]
fn avx512_algorithm_is_byte_exact_via_emulation() {
let big: Vec<u8> =
(0..5000u32).map(|i| (i.wrapping_mul(2_654_435_761) >> 24) as u8).collect();
let needles: Vec<&[u8]> = vec![
&big[0..1],
&big[63..67], &big[64..68], &big[100..104],
&big[2500..2506],
&big[4996..5000], &big[1000..1100], b"\xff\xff\xff",
b"\xff\xff", b"",
];
for nd in &needles {
assert_eq!(
find_avx512_emulated(&big, nd),
find_scalar(&big, nd),
"emulated AVX-512 disagrees with scalar on needle {nd:?}"
);
}
let cases: &[(&[u8], &[u8])] = &[
(b"hello world", b"world"),
(b"hello world", b"xyz"),
(b"", b"x"),
(b"x", b""),
(b"\x00\x01\x02\x03", b"\x02\x03"),
];
for &(h, n) in cases {
assert_eq!(find_avx512_emulated(h, n), find_scalar(h, n), "h={h:?} n={n:?}");
}
}
#[test]
fn the_split_find_all_reports_what_one_pass_reports() {
let mut hay = Vec::new();
for i in 0..20_000u32 {
hay.extend_from_slice(format!("let value_{i} = {i} ; call_{i}(alpha) ;\n").as_bytes());
}
hay.extend_from_slice(&b"aaaaaaaaaaaaaaaaaaaaaaaa\n".repeat(400));
for needle in [
&b"let"[..],
&b"alpha"[..],
&b"aa"[..],
&b"aaaa"[..],
&b"zzzqqq"[..],
&b"\n"[..],
&b"value_19999 = 19999"[..],
] {
assert_eq!(
find_all_across(&hay, needle),
find_all_one_pass(&hay, needle),
"needle {needle:?} over {} bytes",
hay.len()
);
}
for n in [0usize, 1, 2, 3, 64, 4096] {
let small = &hay[..n.min(hay.len())];
assert_eq!(find_all_across(small, b"let"), find_all(small, b"let"), "{n} bytes");
}
assert!(find_all_across(&hay, b"").is_empty(), "an empty needle reports nothing");
assert!(find_all(&hay, b"").is_empty(), "and the one-pass form agrees");
}
#[test]
fn the_spanned_walk_reports_what_one_pass_reports() {
let mut hay = Vec::new();
for i in 0..600u32 {
hay.extend_from_slice(format!("let value_{i} = {i} ; call_{i}(alpha) ;\n").as_bytes());
}
hay.extend_from_slice(&b"aaaaaaaaaaaaaaaaaaaaaaaa\n".repeat(40));
for needle in [
&b"let"[..],
&b"alpha"[..],
&b"aa"[..],
&b"aaaa"[..],
&b"zzzqqq"[..],
&b"\n"[..],
&b"a"[..],
] {
let want = find_all_one_pass(&hay, needle);
for span in [1usize, 2, 3, 7, 64, 4096, hay.len(), hay.len() * 2] {
let got: Vec<usize> = occurrences_by_span(&hay, needle, span).collect();
assert_eq!(got, want, "needle {needle:?} at span {span}");
}
}
for span in [1usize, 8, 4096] {
assert!(occurrences_by_span(&hay, b"", span).next().is_none(), "an empty needle");
for n in [0usize, 1, 2, 3, 64] {
let small = &hay[..n.min(hay.len())];
let got: Vec<usize> = occurrences_by_span(small, b"let", span).collect();
assert_eq!(got, find_all_one_pass(small, b"let"), "{n} bytes at span {span}");
}
}
let got: Vec<usize> = occurrences(&hay, b"let").collect();
assert_eq!(got, find_all_one_pass(&hay, b"let"), "the shipped span");
let got: Vec<usize> = occurrences_by_span(&hay, b"let", 0).collect();
assert_eq!(got, find_all_one_pass(&hay, b"let"), "a span of zero");
}
#[test]
fn the_split_walk_and_the_unsplit_walk_report_the_same_positions() {
let mut hay = Vec::new();
for i in 0..900u32 {
hay.extend_from_slice(format!("let value_{i} = {i} ; call_{i}(alpha) ;\n").as_bytes());
}
hay.extend_from_slice(&b"aaaaaaaaaaaaaaaaaaaaaaaa\n".repeat(60));
for needle in [
&b"let"[..],
&b"alpha"[..],
&b"aa"[..],
&b"aaaa"[..],
&b"a"[..],
&b"zzzqqq"[..],
&b"\n"[..],
&b""[..],
] {
let one = find_all_one_pass(&hay, needle);
let spanned: Vec<usize> = occurrences(&hay, needle).collect();
let unsplit: Vec<usize> = occurrences_unsplit(&hay, needle).collect();
assert_eq!(spanned, one, "the span walk against one pass, needle {needle:?}");
assert_eq!(unsplit, one, "the unsplit walk against one pass, needle {needle:?}");
}
for n in [0usize, 1, 2, 3] {
let small = &hay[..n];
let spanned: Vec<usize> = occurrences(small, b"let").collect();
let unsplit: Vec<usize> = occurrences_unsplit(small, b"let").collect();
assert_eq!(spanned, find_all_one_pass(small, b"let"), "{n} bytes, span walk");
assert_eq!(unsplit, find_all_one_pass(small, b"let"), "{n} bytes, unsplit walk");
}
}
fn classes() -> Vec<(&'static str, crate::byte_nfa::ByteClass)> {
use crate::byte_nfa::ByteClass;
let digits = ByteClass::range(b'0', b'9');
let lower = ByteClass::range(b'a', b'z');
let upper = ByteClass::range(b'A', b'Z');
vec![
("digits", digits),
("hex", digits.union(ByteClass::range(b'a', b'f')).union(ByteClass::range(b'A', b'F'))),
(
"base64",
digits
.union(lower)
.union(upper)
.union(ByteClass::just(b'+'))
.union(ByteClass::just(b'/')),
),
("word", digits.union(lower).union(upper).union(ByteClass::just(b'_'))),
("one byte", ByteClass::just(b'q')),
("none", ByteClass::none()),
]
}
#[test]
fn the_nibble_tables_agree_with_the_class_on_every_byte() {
for (name, class) in classes() {
let tables = ClassTables::for_class(&class).expect("an ASCII class compiles");
for b in 0..=255u8 {
assert_eq!(tables.has(b), class.has(b), "{name} disagrees on byte {b:#04x}");
}
}
}
#[test]
fn a_class_above_ascii_has_no_tables() {
use crate::byte_nfa::ByteClass;
assert!(ClassTables::for_class(&ByteClass::any()).is_none());
assert!(ClassTables::for_class(&ByteClass::just(0x80)).is_none());
assert!(ClassTables::for_class(&ByteClass::range(b'a', b'z')).is_some());
}
#[test]
fn the_vector_run_answers_what_the_scalar_run_answers() {
let mut inputs: Vec<Vec<u8>> = vec![
Vec::new(),
b"q".to_vec(),
b"abc def".to_vec(),
b"a".repeat(31),
b"a".repeat(32),
b"a".repeat(33),
b"a".repeat(1000),
[b" ".repeat(32), b"a".repeat(32), b" ".repeat(32)].concat(),
[b"a".repeat(30), b" ".to_vec(), b"a".repeat(40)].concat(),
[b"a".repeat(64), b" ".to_vec()].concat(),
[b" ".to_vec(), b"a".repeat(64)].concat(),
];
for k in 0..40usize {
inputs.push([b" ".repeat(k), b"abcdef".to_vec(), b" ".repeat(40 - k)].concat());
}
for (name, class) in classes() {
let tables = ClassTables::for_class(&class).expect("an ASCII class compiles");
for input in &inputs {
assert_eq!(
tables.longest_run(input),
tables.longest_run_scalar(input),
"{name} over {} bytes",
input.len()
);
}
}
}
#[test]
fn stopping_early_answers_the_threshold_the_full_scan_answers() {
let inputs: Vec<Vec<u8>> = vec![
Vec::new(),
b"a".to_vec(),
b"a".repeat(31),
b"a".repeat(32),
b"a".repeat(33),
[b" ".repeat(40), b"a".repeat(20), b" ".repeat(40)].concat(),
[b"a".repeat(10), b" ".to_vec(), b"a".repeat(20), b" ".to_vec(), b"a".repeat(5)]
.concat(),
[b"a".repeat(200), b" ".to_vec(), b"a".repeat(3)].concat(),
];
for (name, class) in classes() {
let tables = ClassTables::for_class(&class).expect("an ASCII class compiles");
for input in &inputs {
let longest = tables.longest_run(input);
for least in 0..=(longest + 3) {
assert_eq!(
tables.has_run_of(input, least),
longest >= least,
"{name}, {} bytes, threshold {least} against a longest of {longest}",
input.len()
);
}
}
}
}
#[test]
fn the_longest_run_is_what_a_necessary_condition_reads() {
use crate::byte_nfa::ByteClass;
let hex = ByteClass::range(b'0', b'9')
.union(ByteClass::range(b'a', b'f'))
.union(ByteClass::range(b'A', b'F'));
let tables = ClassTables::for_class(&hex).expect("an ASCII class compiles");
assert_eq!(tables.longest_run(&b"a".repeat(31)), 31);
assert_eq!(tables.longest_run(&b"a".repeat(32)), 32);
let split = [b"a".repeat(20), b" ".to_vec(), b"a".repeat(20)].concat();
assert_eq!(tables.longest_run(&split), 20);
assert_eq!(tables.longest_run(b"abcdefzabcdef"), 6);
}
}