#![cfg_attr(not(feature = "std"), no_std)]
#![warn(missing_docs)]
#![cfg_attr(docsrs, feature(doc_cfg))]
#[cfg(test)]
extern crate std;
#[cfg(test)]
use std::vec::Vec;
pub const MAX_PATTERN_LEN: usize = 64;
#[derive(Clone, Debug)]
pub struct BitParallelSearcher {
pattern: *const u8,
pattern_len: usize,
masks: [u64; 256],
match_mask: u64,
}
unsafe impl Send for BitParallelSearcher {}
unsafe impl Sync for BitParallelSearcher {}
impl BitParallelSearcher {
#[inline]
pub fn new(pattern: &[u8]) -> Self {
assert!(!pattern.is_empty(), "Pattern cannot be empty");
let pattern_len = pattern.len();
let mut masks = [!0u64; 256];
for (i, &byte) in pattern.iter().enumerate().take(64) {
masks[byte as usize] &= !(1u64 << i);
}
Self {
pattern: pattern.as_ptr(),
pattern_len,
masks,
match_mask: if pattern_len <= 64 {
1u64 << (pattern_len - 1)
} else {
0 },
}
}
#[inline]
pub fn find_in(&self, text: &[u8]) -> Option<usize> {
if self.pattern_len > text.len() {
return None;
}
if self.pattern_len <= MAX_PATTERN_LEN {
self.find_bit_parallel(text)
} else {
self.find_naive(text)
}
}
#[inline(always)]
fn find_bit_parallel(&self, text: &[u8]) -> Option<usize> {
let mut state = !0u64;
let match_mask = self.match_mask;
for (i, &byte) in text.iter().enumerate() {
state = (state << 1) | self.masks[byte as usize];
if (state & match_mask) == 0 {
return Some(i + 1 - self.pattern_len);
}
}
None
}
#[cold]
fn find_naive(&self, text: &[u8]) -> Option<usize> {
if self.pattern_len > text.len() {
return None;
}
let pattern = unsafe {
core::slice::from_raw_parts(self.pattern, self.pattern_len)
};
(0..=text.len() - self.pattern_len)
.find(|&i| &text[i..i + self.pattern_len] == pattern)
}
#[cfg(feature = "std")]
#[cfg_attr(docsrs, doc(cfg(feature = "std")))]
pub fn find_all_in<'t>(&self, text: &'t [u8]) -> impl Iterator<Item = usize> + 't {
FindAllIter {
searcher: self.clone(),
text,
pos: 0,
}
}
#[inline]
pub fn count_in(&self, text: &[u8]) -> usize {
if self.pattern_len > text.len() || self.pattern_len > MAX_PATTERN_LEN {
return self.count_naive(text);
}
let mut count = 0;
let mut state = !0u64;
let match_mask = self.match_mask;
for &byte in text {
state = (state << 1) | self.masks[byte as usize];
if (state & match_mask) == 0 {
count += 1;
}
}
count
}
#[cold]
fn count_naive(&self, text: &[u8]) -> usize {
let pattern = unsafe {
core::slice::from_raw_parts(self.pattern, self.pattern_len)
};
let mut count = 0;
for i in 0..text.len().saturating_sub(self.pattern_len - 1) {
if &text[i..i + self.pattern_len] == pattern {
count += 1;
}
}
count
}
#[inline]
pub fn exists_in(&self, text: &[u8]) -> bool {
self.find_in(text).is_some()
}
#[inline]
pub fn pattern_len(&self) -> usize {
self.pattern_len
}
}
#[cfg(feature = "std")]
#[cfg_attr(docsrs, doc(cfg(feature = "std")))]
pub struct FindAllIter<'t> {
searcher: BitParallelSearcher,
text: &'t [u8],
pos: usize,
}
#[cfg(feature = "std")]
impl<'t> Iterator for FindAllIter<'t> {
type Item = usize;
fn next(&mut self) -> Option<Self::Item> {
if self.pos >= self.text.len() {
return None;
}
let remaining = &self.text[self.pos..];
self.searcher.find_in(remaining).map(|offset| {
let match_pos = self.pos + offset;
self.pos = match_pos + 1; match_pos
})
}
}
#[inline]
pub fn find(text: &[u8], pattern: &[u8]) -> Option<usize> {
if pattern.is_empty() {
return Some(0);
}
BitParallelSearcher::new(pattern).find_in(text)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_search() {
let searcher = BitParallelSearcher::new(b"fox");
assert_eq!(searcher.find_in(b"The quick brown fox"), Some(16));
assert_eq!(searcher.find_in(b"no match here"), None);
}
#[test]
fn test_edge_cases() {
let searcher = BitParallelSearcher::new(b"a");
assert_eq!(searcher.find_in(b"a"), Some(0));
assert_eq!(searcher.find_in(b"ba"), Some(1));
assert_eq!(searcher.find_in(b""), None);
}
#[test]
fn test_repeated_pattern() {
let searcher = BitParallelSearcher::new(b"aa");
assert_eq!(searcher.find_in(b"aaaa"), Some(0));
assert_eq!(searcher.count_in(b"aaaa"), 3); }
#[test]
#[should_panic(expected = "Pattern cannot be empty")]
fn test_empty_pattern() {
BitParallelSearcher::new(b"");
}
#[test]
fn test_pattern_at_boundaries() {
let searcher = BitParallelSearcher::new(b"abc");
assert_eq!(searcher.find_in(b"abc"), Some(0));
assert_eq!(searcher.find_in(b"xabc"), Some(1));
assert_eq!(searcher.find_in(b"xyabc"), Some(2));
assert_eq!(searcher.find_in(b"xyzabc"), Some(3));
}
#[cfg(feature = "std")]
#[test]
fn test_find_all() {
let searcher = BitParallelSearcher::new(b"ab");
let matches: Vec<_> = searcher.find_all_in(b"ababab").collect();
assert_eq!(matches, vec![0, 2, 4]);
}
#[test]
fn test_long_pattern_fallback() {
let pattern = b"a".repeat(65);
let mut text = b"x".to_vec();
text.extend_from_slice(&pattern);
let searcher = BitParallelSearcher::new(&pattern);
assert_eq!(searcher.find_in(&text), Some(1));
}
}