use std::borrow::Cow;
use std::sync::OnceLock;
use std::sync::atomic::{AtomicUsize, Ordering};
use super::ast::{Ast, CharClass, ClassAtom, LookKind, ParsedRegex};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RequiredLiterals<'a> {
None,
One(Cow<'a, str>),
Any(Vec<Cow<'a, str>>),
}
impl RequiredLiterals<'_> {
fn is_empty(&self) -> bool {
matches!(self, Self::None)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Prefilter {
None,
Byte(u8),
ByteSet {
bytes: Vec<u8>,
bitmap: [u64; 4],
},
Literal(String),
Any {
literals: Vec<String>,
ascii_case_insensitive: bool,
mixed_width_fold_mask: u8,
finder: Option<MultiLiteralFinder>,
},
Factor {
factor: RequiredFactor,
literals: Box<Prefilter>,
},
}
impl Prefilter {
pub fn from_regex(parsed: &ParsedRegex) -> Self {
parsed.prefilter().clone()
}
pub fn from_pattern(ast: &Ast) -> Self {
Self::from_required(required_literals(ast), false)
}
pub(crate) fn from_required(
required: RequiredLiterals<'_>,
ascii_case_insensitive: bool,
) -> Self {
if ascii_case_insensitive {
let literals = match required {
RequiredLiterals::None => return Self::None,
RequiredLiterals::One(literal) => vec![literal],
RequiredLiterals::Any(literals) => literals,
};
if literals.is_empty()
|| literals
.iter()
.any(|literal| literal.is_empty() || !literal.is_ascii())
{
return Self::None;
}
let mixed_width_fold_mask = literals.iter().fold(0u8, |mask, literal| {
literal.bytes().fold(mask, |mask, byte| {
mask | u8::from(byte.eq_ignore_ascii_case(&b's'))
| (u8::from(byte.eq_ignore_ascii_case(&b'k')) << 1)
})
});
let finder = MultiLiteralFinder::for_literals_ignore_ascii_case(&literals);
let literals = if finder.is_some() {
Vec::new()
} else {
literals.into_iter().map(Cow::into_owned).collect()
};
return Self::Any {
finder,
literals,
ascii_case_insensitive: true,
mixed_width_fold_mask,
};
}
match required {
RequiredLiterals::None => Self::None,
RequiredLiterals::One(literal) => prefilter_one(literal.into_owned()),
RequiredLiterals::Any(literals) if literals.is_empty() => Self::None,
RequiredLiterals::Any(literals) if literals.len() == 1 => prefilter_one(
literals
.into_iter()
.next()
.expect("one literal")
.into_owned(),
),
RequiredLiterals::Any(literals)
if literals
.iter()
.all(|literal| literal.len() == 1 && literal.is_ascii()) =>
{
let bytes = literals
.into_iter()
.map(|literal| literal.as_bytes()[0])
.collect::<Vec<_>>();
let mut bitmap = [0u64; 4];
for &byte in &bytes {
bitmap[byte as usize >> 6] |= 1u64 << (byte & 63);
}
Self::ByteSet { bytes, bitmap }
}
RequiredLiterals::Any(literals) => {
let finder = MultiLiteralFinder::for_literals(&literals);
let literals = if finder.is_some() {
Vec::new()
} else {
literals.into_iter().map(Cow::into_owned).collect()
};
Self::Any {
finder,
literals,
ascii_case_insensitive: false,
mixed_width_fold_mask: 0,
}
}
}
}
pub fn may_match(&self, haystack: &str, from: usize) -> bool {
if !haystack.is_char_boundary(from) {
return false;
}
let Some(slice) = haystack.get(from..) else {
return false;
};
match self {
Self::None => true,
Self::Byte(byte) => find_byte(slice.as_bytes(), *byte).is_some(),
Self::ByteSet { bytes, bitmap } => {
find_byte_set(slice.as_bytes(), bytes, bitmap).is_some()
}
Self::Literal(literal) => find_literal(slice, literal).is_some(),
Self::Factor { factor, literals } => {
literals.may_match(haystack, from) && factor.find(slice.as_bytes()).is_some()
}
Self::Any {
literals,
ascii_case_insensitive: false,
finder,
..
} => finder.as_ref().map_or_else(
|| {
literals
.iter()
.any(|literal| find_literal(slice, literal).is_some())
},
|finder| finder.find(slice.as_bytes()).is_some(),
),
Self::Any {
literals,
ascii_case_insensitive: true,
mixed_width_fold_mask,
finder,
} => {
finder.as_ref().map_or_else(
|| {
literals
.iter()
.any(|literal| contains_ignore_ascii_case(slice, literal))
},
|finder| finder.find(slice.as_bytes()).is_some(),
) || first_ascii_case_fold_candidate(slice, *mixed_width_fold_mask).is_some()
}
}
}
pub fn next_occurrence(&self, haystack: &str, from: usize) -> Option<usize> {
if !haystack.is_char_boundary(from) {
return None;
}
let slice = haystack.get(from..)?;
match self {
Self::None => Some(from),
Self::Byte(byte) => find_byte(slice.as_bytes(), *byte).map(|pos| from + pos),
Self::ByteSet { bytes, bitmap } => {
find_byte_set(slice.as_bytes(), bytes, bitmap).map(|pos| from + pos)
}
Self::Literal(literal) => find_literal(slice, literal).map(|pos| from + pos),
Self::Factor { factor, literals } => {
let factor = factor.find(slice.as_bytes())? + from;
if literals.is_enabled() {
Some(literals.next_occurrence(haystack, from)?.min(factor))
} else {
Some(factor)
}
}
Self::Any {
literals,
ascii_case_insensitive: false,
finder,
..
} => finder.as_ref().map_or_else(
|| {
literals
.iter()
.filter_map(|literal| find_literal(slice, literal))
.min()
.map(|pos| from + pos)
},
|finder| finder.find(slice.as_bytes()).map(|pos| from + pos),
),
Self::Any {
literals,
ascii_case_insensitive: true,
mixed_width_fold_mask,
finder,
} => {
let literal = match finder {
Some(finder) => finder.find(slice.as_bytes()),
None => literals
.iter()
.filter_map(|literal| find_ignore_ascii_case(slice, literal))
.min(),
};
literal
.into_iter()
.chain(first_ascii_case_fold_candidate(
slice,
*mixed_width_fold_mask,
))
.min()
.map(|pos| from + pos)
}
}
}
pub fn is_enabled(&self) -> bool {
!matches!(self, Self::None)
}
}
#[doc(hidden)]
#[derive(Debug)]
pub struct MultiLiteralFinder {
literals: SortedLiterals,
work: AtomicUsize,
automaton: OnceLock<LiteralAutomaton>,
}
impl Clone for MultiLiteralFinder {
fn clone(&self) -> Self {
Self {
literals: self.literals.clone(),
work: AtomicUsize::new(self.work.load(Ordering::Relaxed)),
automaton: self.automaton.clone(),
}
}
}
impl PartialEq for MultiLiteralFinder {
fn eq(&self, other: &Self) -> bool {
self.literals == other.literals
}
}
impl Eq for MultiLiteralFinder {}
impl MultiLiteralFinder {
fn for_literals<S: AsRef<str>>(literals: &[S]) -> Option<Self> {
Self::worthwhile(literals).then(|| Self::new(literals, false))
}
fn for_literals_ignore_ascii_case<S: AsRef<str>>(literals: &[S]) -> Option<Self> {
Self::worthwhile(literals).then(|| Self::new(literals, true))
}
fn worthwhile<S: AsRef<str>>(literals: &[S]) -> bool {
let total_bytes = literals
.iter()
.map(|literal| literal.as_ref().len())
.sum::<usize>();
literals.len() >= multi_literal_min_literals()
&& total_bytes >= multi_literal_min_total_bytes()
}
fn new<S: AsRef<str>>(literals: &[S], fold_ascii_case: bool) -> Self {
Self {
literals: SortedLiterals::new(literals, fold_ascii_case),
work: AtomicUsize::new(0),
automaton: OnceLock::new(),
}
}
fn find(&self, haystack: &[u8]) -> Option<usize> {
if let Some(automaton) = self.automaton.get() {
return automaton.find(haystack, self.literals.fold_ascii_case);
}
let mut work = 0usize;
let found = self.literals.find(haystack, &mut work);
let total = self
.work
.fetch_add(work, Ordering::Relaxed)
.saturating_add(work);
if total >= self.automaton_work_threshold() {
self.automaton
.get_or_init(|| LiteralAutomaton::new(&self.literals));
}
found
}
fn automaton_work_threshold(&self) -> usize {
self.literals
.bytes
.len()
.saturating_mul(multi_literal_automaton_work_per_byte())
}
#[cfg(test)]
fn find_with_automaton(&self, haystack: &[u8]) -> Option<usize> {
self.automaton
.get_or_init(|| LiteralAutomaton::new(&self.literals))
.find(haystack, self.literals.fold_ascii_case)
}
#[cfg(test)]
fn find_sorted(&self, haystack: &[u8]) -> Option<usize> {
self.literals.find(haystack, &mut 0)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct SortedLiterals {
bytes: Box<[u8]>,
ends: Box<[u32]>,
first: Box<[u32; 257]>,
fold_ascii_case: bool,
}
impl SortedLiterals {
fn new<S: AsRef<str>>(literals: &[S], fold_ascii_case: bool) -> Self {
let already_sorted = !fold_ascii_case
&& literals
.windows(2)
.all(|pair| pair[0].as_ref().as_bytes() < pair[1].as_ref().as_bytes());
let total_bytes = literals
.iter()
.map(|literal| literal.as_ref().len())
.sum::<usize>();
let mut bytes = Vec::with_capacity(total_bytes);
let mut ends = Vec::with_capacity(literals.len());
if already_sorted {
for literal in literals {
bytes.extend_from_slice(literal.as_ref().as_bytes());
ends.push(u32::try_from(bytes.len()).expect("prefilter literals exceed u32"));
}
} else {
let mut folded = Vec::with_capacity(total_bytes);
let mut spans = Vec::with_capacity(literals.len());
for literal in literals {
let start = folded.len();
folded.extend_from_slice(literal.as_ref().as_bytes());
if fold_ascii_case {
folded[start..].make_ascii_lowercase();
}
spans.push(start..folded.len());
}
spans.sort_unstable_by(|left, right| folded[left.clone()].cmp(&folded[right.clone()]));
spans.dedup_by(|right, left| folded[right.clone()] == folded[left.clone()]);
for span in spans {
bytes.extend_from_slice(&folded[span]);
ends.push(u32::try_from(bytes.len()).expect("prefilter literals exceed u32"));
}
}
let mut first = Box::new([0u32; 257]);
let mut start = 0usize;
for (index, &end) in ends.iter().enumerate() {
debug_assert!(end as usize > start, "prefilter literals are non-empty");
first[bytes[start] as usize + 1] = index as u32 + 1;
start = end as usize;
}
for byte in 1..first.len() {
first[byte] = first[byte].max(first[byte - 1]);
}
Self {
bytes: bytes.into_boxed_slice(),
ends: ends.into_boxed_slice(),
first,
fold_ascii_case,
}
}
fn len(&self) -> usize {
self.ends.len()
}
fn start(&self, index: usize) -> usize {
index
.checked_sub(1)
.map_or(0, |previous| self.ends[previous] as usize)
}
fn get(&self, index: usize) -> &[u8] {
&self.bytes[self.start(index)..self.ends[index] as usize]
}
#[inline]
fn fold(&self, byte: u8) -> u8 {
if self.fold_ascii_case {
byte.to_ascii_lowercase()
} else {
byte
}
}
fn find(&self, haystack: &[u8], work: &mut usize) -> Option<usize> {
for (start, &byte) in haystack.iter().enumerate() {
let byte = self.fold(byte) as usize;
let (low, high) = (self.first[byte] as usize, self.first[byte + 1] as usize);
if low != high && self.matches_at(haystack, start, low, high, work) {
*work = work.saturating_add(start + 1);
return Some(start);
}
}
*work = work.saturating_add(haystack.len());
None
}
fn matches_at(
&self,
haystack: &[u8],
start: usize,
mut low: usize,
mut high: usize,
work: &mut usize,
) -> bool {
let mut depth = 1usize;
loop {
if self.ends[low] as usize - self.start(low) == depth {
return true;
}
let Some(&input) = haystack.get(start + depth) else {
return false;
};
let input = self.fold(input);
let byte_at = |index: usize| self.bytes[self.start(index) + depth];
let mut lower = low;
let mut upper = high;
while lower < upper {
*work += 1;
let middle = lower + (upper - lower) / 2;
if byte_at(middle) < input {
lower = middle + 1;
} else {
upper = middle;
}
}
if lower == high || byte_at(lower) != input {
return false;
}
low = lower;
upper = high;
while lower < upper {
*work += 1;
let middle = lower + (upper - lower) / 2;
if byte_at(middle) <= input {
lower = middle + 1;
} else {
upper = middle;
}
}
high = lower;
depth += 1;
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct LiteralAutomaton {
nodes: Vec<FinderNode>,
edge_bytes: Vec<u8>,
edge_targets: Vec<u32>,
root_edges: Box<[u32; 256]>,
max_literal_len: usize,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
struct FinderNode {
edge_start: u32,
edge_len: u32,
failure: u32,
output_len: u32,
}
impl LiteralAutomaton {
fn new(literals: &SortedLiterals) -> Self {
let capacity = literals.bytes.len().saturating_add(1);
let mut nodes = Vec::with_capacity(capacity);
let mut parents = Vec::with_capacity(capacity);
let mut node_bytes = Vec::with_capacity(capacity);
nodes.push(FinderNode::default());
parents.push(0u32);
node_bytes.push(0u8);
let mut path = vec![0u32];
let mut previous: &[u8] = &[];
let mut max_literal_len = 0usize;
for index in 0..literals.len() {
let literal = literals.get(index);
debug_assert!(!literal.is_empty());
max_literal_len = max_literal_len.max(literal.len());
let common = previous
.iter()
.zip(literal)
.take_while(|(left, right)| left == right)
.count();
path.truncate(common + 1);
for &byte in &literal[common..] {
let parent = *path.last().expect("trie path keeps the root");
let node = u32::try_from(nodes.len()).expect("prefilter trie exceeds u32");
nodes[parent as usize].edge_len += 1;
nodes.push(FinderNode::default());
parents.push(parent);
node_bytes.push(byte);
path.push(node);
}
let state = *path.last().expect("trie path keeps the root") as usize;
let output_len = u32::try_from(literal.len()).expect("prefilter literal exceeds u32");
nodes[state].output_len = nodes[state].output_len.max(output_len);
previous = literal;
}
let mut edge_start = 0u32;
for node in &mut nodes {
node.edge_start = edge_start;
node.failure = edge_start;
edge_start += node.edge_len;
}
let edge_count = edge_start as usize;
let mut edge_bytes = vec![0u8; edge_count];
let mut edge_targets = vec![0u32; edge_count];
for child in 1..nodes.len() {
let parent = parents[child] as usize;
let slot = nodes[parent].failure as usize;
nodes[parent].failure += 1;
edge_bytes[slot] = node_bytes[child];
edge_targets[slot] = child as u32;
}
drop(parents);
drop(node_bytes);
for node in &mut nodes {
node.failure = 0;
}
let mut automaton = Self {
nodes,
edge_bytes,
edge_targets,
root_edges: Box::new([u32::MAX; 256]),
max_literal_len,
};
let root = automaton.nodes[0];
for index in root.edge_start as usize..(root.edge_start + root.edge_len) as usize {
automaton.root_edges[automaton.edge_bytes[index] as usize] =
automaton.edge_targets[index];
}
let mut queue = Vec::with_capacity(automaton.nodes.len());
queue.push(0u32);
let mut cursor = 0usize;
while let Some(&state) = queue.get(cursor) {
cursor += 1;
let node = automaton.nodes[state as usize];
for index in node.edge_start as usize..(node.edge_start + node.edge_len) as usize {
let byte = automaton.edge_bytes[index];
let child = automaton.edge_targets[index];
queue.push(child);
if state == 0 {
continue;
}
let mut failure = node.failure;
let target = loop {
if let Some(next) = automaton.goto(failure, byte) {
break next;
}
if failure == 0 {
break 0;
}
failure = automaton.nodes[failure as usize].failure;
};
let failure_output = automaton.nodes[target as usize].output_len;
let child = &mut automaton.nodes[child as usize];
child.failure = target;
child.output_len = child.output_len.max(failure_output);
}
}
automaton
}
fn find(&self, haystack: &[u8], fold_ascii_case: bool) -> Option<usize> {
let mut state = 0u32;
let mut best = None;
for (index, byte) in haystack.iter().copied().enumerate() {
let byte = if fold_ascii_case {
byte.to_ascii_lowercase()
} else {
byte
};
state = self.step(state, byte);
let output_len = self.nodes[state as usize].output_len as usize;
if output_len != 0 {
let start = index + 1 - output_len;
best = Some(best.map_or(start, |current: usize| current.min(start)));
}
if let Some(best) = best
&& index.saturating_add(2) >= best.saturating_add(self.max_literal_len)
{
break;
}
}
best
}
fn step(&self, mut state: u32, byte: u8) -> u32 {
loop {
if let Some(next) = self.goto(state, byte) {
return next;
}
if state == 0 {
return 0;
}
state = self.nodes[state as usize].failure;
}
}
#[inline]
fn goto(&self, state: u32, byte: u8) -> Option<u32> {
if state == 0 {
let next = self.root_edges[byte as usize];
return (next != u32::MAX).then_some(next);
}
let node = &self.nodes[state as usize];
let start = node.edge_start as usize;
let end = start + node.edge_len as usize;
let bytes = &self.edge_bytes[start..end];
let index = if bytes.len() <= 16 {
bytes.iter().position(|candidate| *candidate == byte)?
} else {
bytes.binary_search(&byte).ok()?
};
Some(self.edge_targets[start + index])
}
}
fn multi_literal_min_literals() -> usize {
4
}
fn multi_literal_min_total_bytes() -> usize {
32
}
fn multi_literal_automaton_work_per_byte() -> usize {
8
}
fn prefilter_one(literal: String) -> Prefilter {
if literal.len() == 1 && literal.is_ascii() {
Prefilter::Byte(literal.as_bytes()[0])
} else {
Prefilter::Literal(literal)
}
}
fn find_byte(haystack: &[u8], needle: u8) -> Option<usize> {
memchr::memchr(needle, haystack)
}
#[inline]
pub(crate) fn find_byte_set(haystack: &[u8], bytes: &[u8], bitmap: &[u64; 4]) -> Option<usize> {
match bytes {
[] => None,
[byte] => memchr::memchr(*byte, haystack),
[a, b] => memchr::memchr2(*a, *b, haystack),
[a, b, c] => memchr::memchr3(*a, *b, *c, haystack),
_ => find_byte_set_bitmap(haystack, bitmap),
}
}
#[inline]
fn byte_in_set(bitmap: &[u64; 4], byte: u8) -> bool {
bitmap[byte as usize >> 6] & (1u64 << (byte & 63)) != 0
}
pub(crate) fn find_byte_set_bitmap(haystack: &[u8], bitmap: &[u64; 4]) -> Option<usize> {
let mut index = 0;
let len = haystack.len();
while index + 8 <= len {
if byte_in_set(bitmap, haystack[index]) {
return Some(index);
}
if byte_in_set(bitmap, haystack[index + 1]) {
return Some(index + 1);
}
if byte_in_set(bitmap, haystack[index + 2]) {
return Some(index + 2);
}
if byte_in_set(bitmap, haystack[index + 3]) {
return Some(index + 3);
}
if byte_in_set(bitmap, haystack[index + 4]) {
return Some(index + 4);
}
if byte_in_set(bitmap, haystack[index + 5]) {
return Some(index + 5);
}
if byte_in_set(bitmap, haystack[index + 6]) {
return Some(index + 6);
}
if byte_in_set(bitmap, haystack[index + 7]) {
return Some(index + 7);
}
index += 8;
}
while index < len {
if byte_in_set(bitmap, haystack[index]) {
return Some(index);
}
index += 1;
}
None
}
fn find_literal(haystack: &str, needle: &str) -> Option<usize> {
memchr::memmem::find(haystack.as_bytes(), needle.as_bytes())
}
fn contains_ignore_ascii_case(haystack: &str, needle: &str) -> bool {
find_ignore_ascii_case(haystack, needle).is_some()
}
fn first_ascii_case_fold_candidate(haystack: &str, mask: u8) -> Option<usize> {
if mask == 0 {
return None;
}
let bytes = haystack.as_bytes();
let mut from = 0usize;
loop {
let relative = match mask {
1 => memchr::memchr(0xc5, bytes.get(from..)?),
2 => memchr::memchr(0xe2, bytes.get(from..)?),
_ => memchr::memchr2(0xc5, 0xe2, bytes.get(from..)?),
}?;
let offset = from + relative;
if (mask & 1 != 0 && matches!(bytes.get(offset..), Some([0xc5, 0xbf, ..])))
|| (mask & 2 != 0 && matches!(bytes.get(offset..), Some([0xe2, 0x84, 0xaa, ..])))
{
return Some(offset);
}
from = offset + 1;
}
}
fn find_ignore_ascii_case(haystack: &str, needle: &str) -> Option<usize> {
if needle.is_empty() {
return Some(0);
}
let hay = haystack.as_bytes();
let needle = needle.as_bytes();
if needle.len() > hay.len() {
return None;
}
let first = needle[0];
let lower = first.to_ascii_lowercase();
let upper = first.to_ascii_uppercase();
let mut from = 0usize;
loop {
let rest = hay.get(from..)?;
let relative = if lower == upper {
memchr::memchr(first, rest)?
} else {
memchr::memchr2(lower, upper, rest)?
};
let pos = from + relative;
if hay
.get(pos..pos + needle.len())
.is_some_and(|window| window.eq_ignore_ascii_case(needle))
{
return Some(pos);
}
from = pos + 1;
}
}
const MAX_PREFILTER_CURSOR_SLOTS: usize = 1024;
#[derive(Debug, Clone, Default)]
pub(crate) struct PrefilterCursors {
line_ptr: usize,
line_len: usize,
generation: u64,
slots: Vec<CursorSlot>,
overflow_slots: Vec<OverflowCursorSlot>,
}
#[derive(Debug, Clone, Copy, Default)]
struct CursorSlot {
generation: u64,
searched_from: usize,
next_occurrence: Option<usize>,
}
#[derive(Debug, Clone, Copy, Default)]
struct OverflowCursorSlot {
matcher_slot: u32,
cursor: CursorSlot,
}
impl PrefilterCursors {
pub(crate) fn begin_line(&mut self, line: &str) {
self.line_ptr = line.as_ptr() as usize;
self.line_len = line.len();
self.generation = self.generation.wrapping_add(1).max(1);
}
pub(crate) fn may_match(
&mut self,
slot: u32,
prefilter: &Prefilter,
line: &str,
start: usize,
) -> bool {
if !prefilter.is_enabled() {
return true;
}
self.next_occurrence(slot, prefilter, line, start).is_some()
}
pub(crate) fn next_occurrence(
&mut self,
slot: u32,
prefilter: &Prefilter,
line: &str,
start: usize,
) -> Option<usize> {
if !prefilter.is_enabled() {
return Some(start);
}
if slot == u32::MAX {
return prefilter.next_occurrence(line, start);
}
let (line_ptr, line_len) = (line.as_ptr() as usize, line.len());
if self.line_ptr != line_ptr || self.line_len != line_len {
self.line_ptr = line_ptr;
self.line_len = line_len;
self.generation = self.generation.wrapping_add(1);
}
if self.generation == 0 {
self.generation = 1;
}
let generation = self.generation;
if (slot as usize) < MAX_PREFILTER_CURSOR_SLOTS {
let slot = slot as usize;
if slot >= self.slots.len() {
self.slots.resize(slot + 1, CursorSlot::default());
}
return cached_next_occurrence(
&mut self.slots[slot],
true,
generation,
prefilter,
line,
start,
);
}
let overflow_index = slot as usize % MAX_PREFILTER_CURSOR_SLOTS;
if overflow_index >= self.overflow_slots.len() {
self.overflow_slots
.resize(overflow_index + 1, OverflowCursorSlot::default());
}
let entry = &mut self.overflow_slots[overflow_index];
let same_matcher = entry.matcher_slot == slot;
entry.matcher_slot = slot;
cached_next_occurrence(
&mut entry.cursor,
same_matcher,
generation,
prefilter,
line,
start,
)
}
}
#[inline]
fn cached_next_occurrence(
entry: &mut CursorSlot,
same_matcher: bool,
generation: u64,
prefilter: &Prefilter,
line: &str,
start: usize,
) -> Option<usize> {
let stale = !same_matcher || entry.generation != generation || entry.searched_from > start;
if !stale {
match entry.next_occurrence {
None => return None,
Some(occurrence) if occurrence >= start => return Some(occurrence),
Some(_) => {}
}
}
let next = prefilter.next_occurrence(line, start);
*entry = CursorSlot {
generation,
searched_from: start,
next_occurrence: next,
};
next
}
pub fn required_literal(pattern: &str) -> Option<String> {
let parsed = super::ast::parse(pattern);
match required_literals(&parsed.ast) {
RequiredLiterals::One(literal) => Some(literal.into_owned()),
RequiredLiterals::Any(literals) => literals
.into_iter()
.max_by_key(|literal| literal.len())
.map(Cow::into_owned),
RequiredLiterals::None => literal_prefix(pattern),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RequiredFactor {
items: Box<[FactorItem]>,
first_byte: Option<u8>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct FactorItem {
set: [u64; 4],
min: u32,
max: u32,
}
const FACTOR_UNBOUNDED: u32 = u32::MAX;
const FACTOR_SELECTIVE_SET_BYTES: u32 = 8;
const FACTOR_MAX_ITEMS: usize = 64;
impl FactorItem {
fn run_end(&self, haystack: &[u8], start: usize, limit: usize) -> usize {
let mut end = start;
while end - start < limit && haystack.get(end).is_some_and(|byte| self.contains(*byte)) {
end += 1;
}
end
}
fn fixed(byte: u8) -> Self {
let mut set = [0u64; 4];
set[byte as usize >> 6] |= 1u64 << (byte & 63);
Self {
set,
min: 1,
max: 1,
}
}
#[inline]
fn contains(&self, byte: u8) -> bool {
byte_in_set(&self.set, byte)
}
fn is_variable(&self) -> bool {
self.min != self.max
}
fn disjoint(&self, other: &Self) -> bool {
self.set.iter().zip(&other.set).all(|(a, b)| a & b == 0)
}
fn set_len(&self) -> u32 {
self.set.iter().map(|word| word.count_ones()).sum()
}
fn single_byte(&self) -> Option<u8> {
(self.set_len() == 1).then(|| {
let word = self
.set
.iter()
.position(|word| *word != 0)
.expect("one set bit");
(word * 64) as u8 + self.set[word].trailing_zeros() as u8
})
}
}
impl RequiredFactor {
fn new(items: Vec<FactorItem>) -> Self {
let first_byte = items.first().and_then(FactorItem::single_byte);
Self {
items: items.into_boxed_slice(),
first_byte,
}
}
pub(crate) fn score(&self) -> usize {
self.items
.iter()
.filter(|item| item.set_len() <= FACTOR_SELECTIVE_SET_BYTES)
.map(|item| item.min as usize)
.sum()
}
fn find(&self, haystack: &[u8]) -> Option<usize> {
let first = self.items.first()?;
let mut runs: Option<Box<[ItemRun]>> = None;
let mut scanned_total = 0usize;
let per_byte = 4 + 2 * self.items.len();
let budget = haystack.len().saturating_mul(per_byte).saturating_add(64);
let mut from = 0usize;
while from < haystack.len() {
let rest = &haystack[from..];
let relative = match self.first_byte {
Some(byte) => memchr::memchr(byte, rest)?,
None => find_byte_set_bitmap(rest, &first.set)?,
};
let start = from + relative;
let (matched, scanned) = match runs.as_deref_mut() {
None => self.matches_at(haystack, start),
Some(runs) => self.matches_at_with_runs(haystack, start, runs),
};
if matched {
return Some(start);
}
scanned_total = scanned_total.saturating_add(scanned);
if scanned_total >= budget {
return Some(start);
}
if runs.is_none() && scanned_total > haystack.len() {
runs = Some(vec![ItemRun::EMPTY; self.items.len()].into_boxed_slice());
}
from = start + 1;
}
None
}
fn matches_at(&self, haystack: &[u8], start: usize) -> (bool, usize) {
let mut position = start;
for item in &self.items {
let mut count = 0u32;
while count < item.max
&& haystack
.get(position)
.is_some_and(|byte| item.contains(*byte))
{
count += 1;
position += 1;
}
if count < item.min {
return (false, position - start + 1);
}
}
(true, position - start)
}
fn matches_at_with_runs(
&self,
haystack: &[u8],
start: usize,
runs: &mut [ItemRun],
) -> (bool, usize) {
let mut position = start;
let mut scanned = 0usize;
for (item, run) in self.items.iter().zip(runs) {
if !(run.start <= position && position <= run.end) {
let end = item.run_end(haystack, position, usize::MAX);
scanned += end - position + 1;
*run = ItemRun {
start: position,
end,
};
}
let count = (run.end - position).min(item.max as usize);
if count < item.min as usize {
return (false, scanned);
}
position += count;
}
(true, scanned)
}
}
#[derive(Clone, Copy)]
struct ItemRun {
start: usize,
end: usize,
}
impl ItemRun {
const EMPTY: Self = Self {
start: usize::MAX,
end: 0,
};
}
pub(crate) fn required_factor(ast: &Ast) -> Option<RequiredFactor> {
let mut best = None;
collect_required_factors(ast, &mut best);
best
}
fn collect_required_factors(ast: &Ast, best: &mut Option<RequiredFactor>) {
match ast {
Ast::Concat(nodes) => {
let mut run = Vec::new();
for node in nodes {
let checkpoint = run.len();
if factor_items(node, &mut run) {
if let Ast::Look {
kind: LookKind::Ahead,
child,
} = node
{
collect_required_factors(child, best);
}
continue;
}
run.truncate(checkpoint);
finish_factor_run(std::mem::take(&mut run), best);
collect_required_factors(node, best);
}
finish_factor_run(run, best);
}
Ast::Group { child, .. }
| Ast::Look {
kind: LookKind::Ahead,
child,
} => collect_required_factors(child, best),
Ast::Flags { flags, child } if !flags.case_insensitive => {
collect_required_factors(child, best);
}
Ast::Repeat { node, min, .. } if *min > 0 => collect_required_factors(node, best),
_ => {
let mut run = Vec::new();
if factor_items(ast, &mut run) {
finish_factor_run(run, best);
}
}
}
}
fn factor_items(ast: &Ast, out: &mut Vec<FactorItem>) -> bool {
if out.len() > FACTOR_MAX_ITEMS {
return false;
}
match ast {
Ast::Empty | Ast::Anchor(_) | Ast::Look { .. } => true,
Ast::Literal(literal) => {
out.extend(literal.bytes().map(FactorItem::fixed));
true
}
Ast::Class(class) => {
out.push(class_factor_item(class));
true
}
Ast::Group { child, .. } => factor_items(child, out),
Ast::Flags { flags, child } if !flags.case_insensitive => factor_items(child, out),
Ast::Concat(nodes) => nodes.iter().all(|node| factor_items(node, out)),
Ast::Repeat { node, min, max, .. } => {
let mut inner = Vec::new();
if !factor_items(node, &mut inner) || inner.len() > 1 {
return false;
}
let Some(item) = inner.pop() else {
return true;
};
let Ok(min) = u32::try_from(*min) else {
return false;
};
let max = match max {
None => FACTOR_UNBOUNDED,
Some(max) => match u32::try_from(*max) {
Ok(max) => item.max.saturating_mul(max).min(FACTOR_UNBOUNDED - 1),
Err(_) => return false,
},
};
out.push(FactorItem {
set: item.set,
min: item.min.saturating_mul(min),
max,
});
true
}
_ => false,
}
}
fn class_factor_item(class: &CharClass) -> FactorItem {
let ascii = super::bytecode::ascii_class_masks(class).0;
if class_may_match_non_ascii(class) {
FactorItem {
set: [ascii[0], ascii[1], u64::MAX, u64::MAX],
min: 1,
max: 4,
}
} else {
FactorItem {
set: [ascii[0], ascii[1], 0, 0],
min: 1,
max: 1,
}
}
}
fn class_may_match_non_ascii(class: &CharClass) -> bool {
class.negated
|| class.atoms.iter().any(|atom| match atom {
ClassAtom::Char(ch) => !ch.is_ascii(),
ClassAtom::Range(_, end) => !end.is_ascii(),
ClassAtom::Nested(nested) => class_may_match_non_ascii(nested),
ClassAtom::Perl(_) | ClassAtom::Posix { .. } | ClassAtom::Unicode { .. } => true,
})
}
fn finish_factor_run(run: Vec<FactorItem>, best: &mut Option<RequiredFactor>) {
let mut segment: Vec<FactorItem> = Vec::new();
let mut items = run.into_iter().peekable();
while let Some(mut item) = items.next() {
if segment.is_empty() {
item.max = item.min;
}
if item.min == 0 && item.max == 0 {
continue;
}
let greedy_exact = !item.is_variable()
|| items
.peek()
.is_some_and(|next| next.min > 0 && item.disjoint(next));
if greedy_exact {
segment.push(item);
continue;
}
if item.min > 0 {
item.max = item.min;
segment.push(item);
}
offer_factor(std::mem::take(&mut segment), best);
}
offer_factor(segment, best);
}
fn offer_factor(mut items: Vec<FactorItem>, best: &mut Option<RequiredFactor>) {
if let Some(last) = items.last_mut() {
last.max = last.min;
}
while items.last().is_some_and(|item| item.min == 0) {
items.pop();
}
if items.len() < 2
|| items
.iter()
.all(|item| !item.is_variable() && item.set_len() == 1)
{
return;
}
let candidate = RequiredFactor::new(items);
let score = candidate.score();
if score >= 2 && best.as_ref().is_none_or(|best| score > best.score()) {
*best = Some(candidate);
}
}
pub fn required_literals(ast: &Ast) -> RequiredLiterals<'_> {
match collect_required_literals(ast) {
RequiredLiterals::Any(mut literals) => {
literals.sort_unstable();
literals.dedup();
RequiredLiterals::Any(literals)
}
required => required,
}
}
fn collect_required_literals(ast: &Ast) -> RequiredLiterals<'_> {
if let Some(literal) = exact_literal(ast).filter(|literal| !literal.is_empty()) {
return RequiredLiterals::One(literal);
}
match ast {
Ast::Concat(nodes) => sequence_required_literals(nodes),
Ast::Alternation(branches) => alternation_required_literals(branches),
Ast::Group { child, .. } | Ast::Flags { child, .. } => collect_required_literals(child),
Ast::Look {
kind: LookKind::Ahead,
child,
} => collect_required_literals(child),
Ast::Repeat { node, min, .. } if *min > 0 => collect_required_literals(node),
Ast::Class(class) => class_required_literals(class),
_ => RequiredLiterals::None,
}
}
fn sequence_required_literals(nodes: &[Ast]) -> RequiredLiterals<'_> {
let mut best = RequiredLiterals::None;
let mut run = Cow::Borrowed("");
for node in nodes {
let run_len = run.len();
if append_exact_literal(node, &mut run) {
continue;
}
truncate_run(&mut run, run_len);
if !run.is_empty() {
best = choose_more_selective(
best,
RequiredLiterals::One(std::mem::replace(&mut run, Cow::Borrowed(""))),
);
}
let candidate = collect_required_literals(node);
best = choose_more_selective(best, candidate);
}
if !run.is_empty() {
best = choose_more_selective(best, RequiredLiterals::One(run));
}
best
}
fn truncate_run(run: &mut Cow<'_, str>, len: usize) {
match run {
Cow::Borrowed(literal) => *literal = &literal[..len],
Cow::Owned(literal) => literal.truncate(len),
}
}
fn exact_literal(ast: &Ast) -> Option<Cow<'_, str>> {
match ast {
Ast::Empty => Some(Cow::Borrowed("")),
Ast::Literal(literal) => Some(Cow::Borrowed(literal)),
Ast::Concat(nodes) => {
if !nodes.iter().all(is_exact_literal) {
return None;
}
let mut out = Cow::Borrowed("");
for node in nodes {
append_exact_literal(node, &mut out);
}
Some(out)
}
Ast::Group { child, .. } | Ast::Flags { child, .. } => exact_literal(child),
_ => None,
}
}
fn is_exact_literal(ast: &Ast) -> bool {
match ast {
Ast::Empty | Ast::Literal(_) => true,
Ast::Concat(nodes) => nodes.iter().all(is_exact_literal),
Ast::Group { child, .. } | Ast::Flags { child, .. } => is_exact_literal(child),
_ => false,
}
}
fn append_exact_literal<'a>(ast: &'a Ast, out: &mut Cow<'a, str>) -> bool {
match ast {
Ast::Empty => true,
Ast::Literal(literal) => {
if out.is_empty() {
*out = Cow::Borrowed(literal);
} else {
out.to_mut().push_str(literal);
}
true
}
Ast::Concat(nodes) => nodes.iter().all(|node| append_exact_literal(node, out)),
Ast::Group { child, .. } | Ast::Flags { child, .. } => append_exact_literal(child, out),
_ => false,
}
}
fn alternation_required_literals(branches: &[Ast]) -> RequiredLiterals<'_> {
let mut literals = Vec::new();
for branch in branches {
match collect_required_literals(branch) {
RequiredLiterals::One(literal) => literals.push(literal),
RequiredLiterals::Any(mut branch_literals) => literals.append(&mut branch_literals),
RequiredLiterals::None => return RequiredLiterals::None,
}
}
RequiredLiterals::Any(literals)
}
fn choose_more_selective<'a>(
left: RequiredLiterals<'a>,
right: RequiredLiterals<'a>,
) -> RequiredLiterals<'a> {
if left.is_empty() {
return right;
}
if right.is_empty() {
return left;
}
let left_len = max_literal_len(&left);
let right_len = max_literal_len(&right);
if right_len > left_len {
return right;
}
if right_len < left_len {
return left;
}
if literal_cardinality(&right) < literal_cardinality(&left) {
right
} else {
left
}
}
fn max_literal_len(literals: &RequiredLiterals) -> usize {
match literals {
RequiredLiterals::None => 0,
RequiredLiterals::One(literal) => literal.len(),
RequiredLiterals::Any(literals) => literals
.iter()
.map(|literal| literal.len())
.max()
.unwrap_or(0),
}
}
fn literal_cardinality(literals: &RequiredLiterals) -> usize {
match literals {
RequiredLiterals::None => usize::MAX,
RequiredLiterals::One(_) => 1,
RequiredLiterals::Any(literals) => {
let mut distinct = literals
.iter()
.map(|literal| &**literal)
.collect::<Vec<_>>();
distinct.sort_unstable();
distinct.dedup();
distinct.len()
}
}
}
fn class_required_literals(class: &CharClass) -> RequiredLiterals<'_> {
if class.negated || class.atoms.is_empty() {
return RequiredLiterals::None;
}
let mut literals = Vec::new();
for atom in &class.atoms {
match atom {
ClassAtom::Char(ch) => literals.push(Cow::Owned(ch.to_string())),
ClassAtom::Range(..)
| ClassAtom::Perl(_)
| ClassAtom::Posix { .. }
| ClassAtom::Unicode { .. }
| ClassAtom::Nested(_) => return RequiredLiterals::None,
}
}
literals.sort();
literals.dedup();
match literals.len() {
0 => RequiredLiterals::None,
1 => RequiredLiterals::One(literals.remove(0)),
_ => RequiredLiterals::Any(literals),
}
}
fn literal_prefix(pattern: &str) -> Option<String> {
let mut literal = String::new();
let mut escaped = false;
for ch in pattern.chars() {
if escaped {
if ch.is_ascii_alphanumeric() {
return None;
}
literal.push(ch);
escaped = false;
continue;
}
match ch {
'\\' => escaped = true,
'(' | ')' | '[' | ']' | '{' | '}' | '|' | '?' | '*' | '+' | '.' | '^' | '$' => break,
ch => literal.push(ch),
}
}
(!literal.is_empty()).then_some(literal)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LiteralSet {
literals: Vec<String>,
trie: Vec<LiteralTrieNode>,
empty_pattern: Option<usize>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
struct LiteralTrieNode {
edges: Vec<(u8, usize)>,
terminal_patterns: Vec<usize>,
}
impl LiteralSet {
pub(crate) fn retained_heap_bytes(&self) -> usize {
let mut bytes = self
.literals
.capacity()
.saturating_mul(std::mem::size_of::<String>())
.saturating_add(
self.trie
.capacity()
.saturating_mul(std::mem::size_of::<LiteralTrieNode>()),
);
for literal in &self.literals {
bytes = bytes.saturating_add(literal.capacity());
}
for node in &self.trie {
bytes = bytes
.saturating_add(
node.edges
.capacity()
.saturating_mul(std::mem::size_of::<(u8, usize)>()),
)
.saturating_add(
node.terminal_patterns
.capacity()
.saturating_mul(std::mem::size_of::<usize>()),
);
}
bytes
}
pub fn new(literals: Vec<String>) -> Self {
let mut trie = vec![LiteralTrieNode::default()];
let mut empty_pattern = None;
for (pattern, literal) in literals.iter().enumerate() {
if literal.is_empty() {
empty_pattern =
Some(empty_pattern.map_or(pattern, |best: usize| best.min(pattern)));
continue;
}
let mut node = 0usize;
for byte in literal.bytes() {
let next = trie[node]
.edges
.iter()
.find_map(|(edge, next)| (*edge == byte).then_some(*next));
node = if let Some(next) = next {
next
} else {
let next = trie.len();
trie.push(LiteralTrieNode::default());
trie[node].edges.push((byte, next));
next
};
}
trie[node].terminal_patterns.push(pattern);
}
Self {
literals,
trie,
empty_pattern,
}
}
pub fn literals(&self) -> &[String] {
&self.literals
}
pub fn find(&self, haystack: &str, from: usize) -> Option<(usize, usize, usize)> {
if !haystack.is_char_boundary(from) {
return None;
}
for start in haystack[from..]
.char_indices()
.map(|(offset, _)| from + offset)
.chain(std::iter::once(haystack.len()))
{
let mut best = self.empty_pattern.map(|pattern| (pattern, start));
let mut node = 0usize;
let mut end = start;
while let Some(byte) = haystack.as_bytes().get(end) {
let Some(next) = self.trie[node]
.edges
.iter()
.find_map(|(edge, next)| (*edge == *byte).then_some(*next))
else {
break;
};
node = next;
end += 1;
for pattern in &self.trie[node].terminal_patterns {
if best.is_none_or(|(best, _)| *pattern < best) {
best = Some((*pattern, end));
}
}
}
if let Some((pattern, end)) = best {
return Some((pattern, start, end));
}
}
None
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::engine::regex::ast::parse;
#[test]
fn required_factor_rejects_subscripts_for_empty_array_declarators() {
let parsed = parse(r"(\w+)\s*(\[) *(])\s*(=)");
let prefilter = parsed.prefilter();
assert!(
matches!(prefilter, Prefilter::Factor { .. }),
"{prefilter:?}"
);
assert!(!prefilter.may_match("values[index] = values[index] * 2;", 0));
assert!(prefilter.may_match("int values[ ] = {1};", 0));
assert!(prefilter.may_match("int values[] = {1};", 0));
assert!(!prefilter.may_match("int values[] = {1};", 11));
assert_eq!(prefilter.next_occurrence("a[i] = b[i]", 0), None);
assert!(prefilter.next_occurrence("a[i] b[ ] =", 0).is_some());
}
#[test]
fn required_factor_is_not_derived_under_case_folding() {
for pattern in [r"(?i)\[ *]=", r"(?i:\[ *])=", r"x(?i)\[ *]"] {
let parsed = parse(pattern);
assert!(
!matches!(parsed.prefilter(), Prefilter::Factor { .. }),
"{pattern}: {:?}",
parsed.prefilter()
);
}
}
#[test]
fn required_factor_stops_where_greedy_verification_would_be_inexact() {
assert_eq!(required_factor(&parse("xa*a").ast), None);
assert_eq!(required_factor(&parse("x *y?z").ast), None);
assert!(required_factor(&parse(r"x *\(").ast).is_some());
}
#[test]
fn required_factor_search_stays_linear() {
let factor = required_factor(&parse("<[^>]*>").ast).expect("byte-run factor");
assert_eq!(factor.find("<".repeat(100_000).as_bytes()), None);
let bounded = required_factor(&parse("<[^>]{0,20}>").ast).expect("byte-run factor");
assert_eq!(bounded.find("<".repeat(100_000).as_bytes()), None);
assert_eq!(factor.find(b"x<y>"), Some(1));
assert_eq!(factor.find(b"x<y"), None);
}
#[test]
fn required_factor_search_matches_a_direct_scan() {
fn direct(factor: &RequiredFactor, haystack: &[u8]) -> Option<usize> {
(0..haystack.len()).find(|&start| {
let mut position = start;
factor.items.iter().all(|item| {
let mut count = 0u32;
while count < item.max
&& haystack
.get(position)
.is_some_and(|byte| item.contains(*byte))
{
count += 1;
position += 1;
}
count >= item.min
})
})
}
let mut seed = 0x2545_f491_u32;
let mut next = |bound: u32| {
seed ^= seed << 13;
seed ^= seed >> 17;
seed ^= seed << 5;
seed % bound
};
let mut checked = 0;
for pattern in [
"<[^>]*>",
"a{2,3}b",
r"x *\(",
"ab{0,2}c",
"[ab]+c{2}",
r"<[^>]*>\s*x",
"<[^>]{0,3}>",
"a[^b]{0,2}b[^c]*c",
] {
let Some(factor) = required_factor(&parse(pattern).ast) else {
continue;
};
checked += 1;
for _ in 0..500 {
let haystack = (0..next(24))
.map(|_| b"<>abcxy("[next(8) as usize])
.collect::<Vec<_>>();
assert_eq!(
factor.find(&haystack),
direct(&factor, &haystack),
"{pattern} on {:?}",
String::from_utf8_lossy(&haystack)
);
}
}
assert!(
checked >= 3,
"only {checked} patterns have byte-run factors"
);
}
#[test]
fn required_factor_prefilter_never_rejects_a_matching_start() {
use crate::engine::regex::AnchorContext;
use crate::engine::regex::backtrack::StepBudget;
use crate::engine::regex::bytecode::{BytecodeScratch, Program};
let patterns = [
r"(\[) *(])\s*=",
r"x *\(",
r"xa*a",
r"x[ab]*b",
r"[a-z]+ *\(",
r"\w+ *\(",
r"é+ ?x",
r"[^\]]* *\]x",
r"(?:ab)+c d",
r"a{2,3} b",
r"a(?=b *c)",
r"a*+a b",
r"(?<=x)y *z",
r"b\s*+=\s*+c",
r"[ab]{2} *c",
r"\[\s*\]",
r"(?:\[ *]|\( *\))=",
r"(a)\1 *b",
r"(?<n>a) *\k<n>",
r"(\( *\)|\g<1>x) *=",
r"𝒳 *[^a]",
r"[^x] +\t",
r"\h+ *\]",
];
let alphabet = [
"a", "b", "c", "x", "y", "z", " ", " ", "[", "]", "(", ")", "=", "é", "𝒳", "\t", "\n",
];
let mut seed = 0x9e37_79b9_7f4a_7c15u64;
let mut next = move || {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
seed
};
let mut matched_starts = 0usize;
for pattern in patterns {
let parsed = parse(pattern);
let program =
Program::compile_captures(&parsed, &(0..=parsed.capture_count).collect::<Vec<_>>())
.unwrap_or_else(|error| panic!("{pattern}: {error:?}"));
let prefilter = Prefilter::from_regex(&parsed);
let mut scratch = BytecodeScratch::default();
for _ in 0..20_000 {
let len = (next() % 14) as usize;
let line = (0..len)
.map(|_| alphabet[(next() % alphabet.len() as u64) as usize])
.collect::<String>();
for start in (0..=line.len()).filter(|index| line.is_char_boundary(*index)) {
let mut budget = StepBudget::new(1_000_000);
let matched = program
.execute(
&line,
start,
AnchorContext::line_start(),
&mut budget,
&mut scratch,
)
.expect("budget");
if matched.is_some() {
matched_starts += 1;
assert!(
prefilter.may_match(&line, start),
"{pattern} matches {line:?} at {start} but {prefilter:?} rejects it"
);
}
}
}
}
assert!(matched_starts > 15_000, "{matched_starts}");
}
#[test]
fn extracts_safe_literal_prefix() {
assert_eq!(required_literal("foo|bar"), Some("foo".to_owned()));
assert_eq!(required_literal(r"\w+"), None);
}
#[test]
fn extracts_alternation_literals() {
let parsed = parse("foo|bar");
assert_eq!(
required_literals(&parsed.ast),
RequiredLiterals::Any(vec!["bar".into(), "foo".into()])
);
}
#[test]
fn extracts_positive_lookahead_literals() {
let parsed = parse(r"(?<=return)\s*(?=(<)\s*([A-Za-z]+))");
assert_eq!(
required_literals(&parsed.ast),
RequiredLiterals::One("<".into())
);
let parsed = parse(r"(?<!\\)(?=;)");
assert_eq!(
required_literals(&parsed.ast),
RequiredLiterals::One(";".into())
);
let parsed = parse(r"(?<=return)");
assert_eq!(required_literals(&parsed.ast), RequiredLiterals::None);
}
#[test]
fn extracts_positive_class_literals() {
let parsed = parse(r"(?=[;)])(?<!\\)");
assert_eq!(
required_literals(&parsed.ast),
RequiredLiterals::Any(vec![")".into(), ";".into()])
);
let parsed = parse(r"(?=[A-Z])");
assert_eq!(required_literals(&parsed.ast), RequiredLiterals::None);
}
#[test]
fn enables_ascii_literal_prefilter_for_case_insensitive_patterns() {
let parsed = parse(r"(?i)foo");
let prefilter = Prefilter::from_regex(&parsed);
assert!(prefilter.is_enabled());
assert!(prefilter.may_match("xxFOO", 0));
assert!(!prefilter.may_match("xxbar", 0));
let parsed = parse(r"(?i)k");
let prefilter = Prefilter::from_regex(&parsed);
assert!(prefilter.may_match("K", 0));
assert_eq!(prefilter.next_occurrence("xxK", 0), Some(2));
let parsed = parse(r"(?i)café");
assert!(!Prefilter::from_regex(&parsed).is_enabled());
let parsed = parse(r"foo");
assert!(Prefilter::from_regex(&parsed).is_enabled());
}
#[test]
fn prefilter_uses_byte_scan_for_single_byte() {
let parsed = parse("x+");
let prefilter = Prefilter::from_pattern(&parsed.ast);
assert!(prefilter.may_match("abcx", 0));
assert!(!prefilter.may_match("abc", 0));
}
#[test]
fn byte_set_prefilter_finds_first_of_many_bytes() {
let parsed = parse("a|e|i|o|u");
let prefilter = Prefilter::from_pattern(&parsed.ast);
assert_eq!(prefilter.next_occurrence("xxxyz o", 0), Some(6));
assert_eq!(prefilter.next_occurrence("xxxyz o", 6), Some(6));
assert_eq!(prefilter.next_occurrence("xxxyz o", 7), None);
assert!(!prefilter.may_match("bcdfg", 0));
}
#[test]
fn ignore_ascii_case_search_finds_first_folded_needle() {
assert_eq!(find_ignore_ascii_case("xxSELECT", "select"), Some(2));
assert_eq!(find_ignore_ascii_case("Select", "SELECT"), Some(0));
assert_eq!(find_ignore_ascii_case("nope", "SELECT"), None);
assert_eq!(find_ignore_ascii_case("ssssSELECT", "select"), Some(4));
}
#[test]
fn multi_literal_finder_preserves_leftmost_and_failure_outputs() {
let literals = [
"bc", "abcd", "suffix", "hers", "his", "she", "he", "keyword",
]
.into_iter()
.map(str::to_owned)
.collect::<Vec<_>>();
let finder = MultiLiteralFinder::new(&literals, false);
assert_eq!(finder.find(b"zabcd"), Some(1));
assert_eq!(finder.find(b"ushers"), Some(1));
assert_eq!(finder.find(b"nothing"), None);
}
#[test]
fn multi_literal_finder_matches_naive_leftmost_search() {
let mut seed = 0x9e37_79b9_7f4a_7c15u64;
let mut next = |bound: u64| {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
seed % bound
};
for round in 0..200 {
let alphabet: &[u8] = if round % 3 == 0 {
b"ab"
} else if round % 3 == 1 {
b"abcd-"
} else {
b"abcdefghijklmnopqrstuvwxyz0123456789"
};
let count = 1 + next(40) as usize;
let mut literals = (0..count)
.map(|_| {
let len = 1 + next(6) as usize;
(0..len)
.map(|_| alphabet[next(alphabet.len() as u64) as usize] as char)
.collect::<String>()
})
.collect::<Vec<_>>();
if round % 3 == 2 {
literals.extend(alphabet.iter().map(|byte| format!("a{}", *byte as char)));
}
let fold = round % 2 == 1;
if fold {
for literal in literals.iter_mut().step_by(2) {
literal.make_ascii_uppercase();
}
}
let finder = MultiLiteralFinder::new(&literals, fold);
let tiered = MultiLiteralFinder::new(&literals, fold);
for _ in 0..20 {
let len = next(24) as usize;
let mut haystack = (0..len)
.map(|_| alphabet[next(alphabet.len() as u64) as usize])
.collect::<Vec<_>>();
if fold {
for byte in haystack.iter_mut().step_by(3) {
byte.make_ascii_uppercase();
}
}
let naive = (0..=haystack.len()).find(|&start| {
literals.iter().any(|literal| {
haystack[start..]
.get(..literal.len())
.is_some_and(|window| {
if fold {
window.eq_ignore_ascii_case(literal.as_bytes())
} else {
window == literal.as_bytes()
}
})
})
});
let context = format!(
"literals {literals:?} haystack {:?}",
String::from_utf8_lossy(&haystack)
);
assert_eq!(finder.find_sorted(&haystack), naive, "{context}");
assert_eq!(finder.find_with_automaton(&haystack), naive, "{context}");
assert_eq!(tiered.find(&haystack), naive, "{context}");
}
}
}
#[test]
fn multi_literal_finder_builds_its_automaton_only_after_enough_queries() {
let literals = (0..64)
.map(|index| format!("keyword_{index:02}"))
.collect::<Vec<_>>();
let finder = MultiLiteralFinder::new(&literals, false);
assert!(finder.automaton.get().is_none(), "construction stays lazy");
let line = format!("{} keyword_42", "x".repeat(40));
assert_eq!(finder.find(line.as_bytes()), Some(41));
assert!(
finder.automaton.get().is_none(),
"one short query stays sorted"
);
let mut queries = 0;
while finder.automaton.get().is_none() {
assert_eq!(finder.find(line.as_bytes()), Some(41));
queries += 1;
assert!(queries < 10_000, "automaton threshold is reachable");
}
assert_eq!(finder.find(line.as_bytes()), Some(41));
assert_eq!(finder.find(b"keyword_4 keyword_7"), None);
assert_eq!(finder, MultiLiteralFinder::new(&literals, false));
}
#[test]
fn high_process_global_slots_use_a_bounded_collision_safe_memo() {
let keyword = parse("keyword");
let keyword = Prefilter::from_pattern(&keyword.ast);
let missing = parse("missing");
let missing = Prefilter::from_pattern(&missing.ast);
let mut cursors = PrefilterCursors::default();
let line = "keyword";
let high_slot = 1_000_000;
let colliding_slot = high_slot + MAX_PREFILTER_CURSOR_SLOTS as u32;
cursors.begin_line(line);
assert!(cursors.may_match(high_slot, &keyword, line, 0));
assert!(!cursors.may_match(colliding_slot, &missing, line, 0));
assert!(cursors.may_match(high_slot, &keyword, line, 0));
assert!(cursors.slots.is_empty());
assert!(!cursors.overflow_slots.is_empty());
assert!(cursors.overflow_slots.len() <= MAX_PREFILTER_CURSOR_SLOTS);
}
#[test]
fn explicit_line_boundary_invalidates_reused_string_storage() {
let parsed = parse("keyword");
let prefilter = Prefilter::from_pattern(&parsed.ast);
let mut cursors = PrefilterCursors::default();
let mut line = String::from("no-match");
cursors.begin_line(&line);
assert!(!cursors.may_match(7, &prefilter, &line, 0));
line.clear();
line.push_str("keyword!");
cursors.begin_line(&line);
assert!(cursors.may_match(7, &prefilter, &line, 0));
}
#[test]
fn large_any_prefilter_uses_compiled_finder_without_changing_answers() {
let parsed = parse(concat!(
"alpha_long|beta_long|gamma_long|delta_long|epsilon_long|zeta_long|theta_long|keyword_long|",
"iota_long|kappa_long|lambda_long|mu_long_value|nu_long_value|xi_long_value|omicron_long|pi_long_value",
));
let prefilter = Prefilter::from_pattern(&parsed.ast);
let Prefilter::Any { finder, .. } = &prefilter else {
panic!("expected Any prefilter");
};
assert!(finder.is_some());
assert_eq!(
prefilter.next_occurrence("xx keyword_long alpha", 0),
Some(3)
);
assert!(prefilter.may_match("xx keyword_long alpha", 3));
assert!(!prefilter.may_match("xx keyword_long alpha", 4));
assert!(!prefilter.may_match("unrelated", 0));
}
#[test]
fn case_insensitive_finder_matches_per_literal_search() {
let literals: Vec<String> = [
"abstract",
"accept",
"accepting",
"add",
"add-corresponding",
"Alias",
"SELECT",
"sKip",
"kind",
"he",
"she",
"hers",
"Z_9",
]
.into_iter()
.map(str::to_owned)
.collect();
let prefilter = Prefilter::from_required(
RequiredLiterals::Any(literals.iter().cloned().map(Cow::Owned).collect()),
true,
);
let Prefilter::Any {
finder,
mixed_width_fold_mask,
..
} = &prefilter
else {
panic!("expected Any prefilter");
};
assert!(finder.is_some());
let reference = |slice: &str| {
literals
.iter()
.filter_map(|literal| find_ignore_ascii_case(slice, literal))
.chain(first_ascii_case_fold_candidate(
slice,
*mixed_width_fold_mask,
))
.min()
};
for text in [
"",
"nothing here",
" ADD-CORRESPONDING x",
"xaccEPTing",
"uSHErs",
"ſkip and \u{212a}ind",
"Ä add é SELECT",
"z_9 Z_9",
"ACCEPT",
"aDd",
] {
for from in (0..=text.len()).filter(|from| text.is_char_boundary(*from)) {
let expected = reference(&text[from..]).map(|pos| from + pos);
assert_eq!(
prefilter.next_occurrence(text, from),
expected,
"{text:?} from {from}"
);
assert_eq!(
prefilter.may_match(text, from),
expected.is_some(),
"{text:?} from {from}"
);
}
}
}
#[test]
fn literal_set_leftmost_lowest_index() {
let set = LiteralSet::new(vec!["bb".into(), "b".into(), "a".into()]);
assert_eq!(set.find("abb", 0), Some((2, 0, 1)));
assert_eq!(set.find("abb", 1), Some((0, 1, 3)));
}
#[test]
fn literal_trie_preserves_empty_prefix_and_utf8_order() {
let set = LiteralSet::new(vec!["éx".into(), "é".into(), "".into()]);
assert_eq!(set.find("zéx", 1), Some((0, 1, 4)));
let set = LiteralSet::new(vec!["".into(), "é".into()]);
assert_eq!(set.find("é", 0), Some((0, 0, 0)));
assert_eq!(set.find("é", 1), None);
}
pub(super) fn reference_required_literals(ast: &Ast) -> RequiredLiterals<'_> {
fn exact(ast: &Ast, out: &mut String) -> bool {
match ast {
Ast::Empty => true,
Ast::Literal(literal) => {
out.push_str(literal);
true
}
Ast::Concat(nodes) => nodes.iter().all(|node| exact(node, out)),
Ast::Group { child, .. } | Ast::Flags { child, .. } => exact(child, out),
_ => false,
}
}
fn more_selective<'a>(
left: RequiredLiterals<'a>,
right: RequiredLiterals<'a>,
) -> RequiredLiterals<'a> {
let len = |literals: &RequiredLiterals| match literals {
RequiredLiterals::None => 0,
RequiredLiterals::One(literal) => literal.len(),
RequiredLiterals::Any(literals) => literals
.iter()
.map(|literal| literal.len())
.max()
.unwrap_or(0),
};
let count = |literals: &RequiredLiterals| match literals {
RequiredLiterals::None => usize::MAX,
RequiredLiterals::One(_) => 1,
RequiredLiterals::Any(literals) => literals.len(),
};
if left.is_empty() {
return right;
}
if right.is_empty() || len(&right) < len(&left) {
return left;
}
if len(&right) > len(&left) || count(&right) < count(&left) {
right
} else {
left
}
}
let mut whole = String::new();
if exact(ast, &mut whole) && !whole.is_empty() {
return RequiredLiterals::One(Cow::Owned(whole));
}
match ast {
Ast::Concat(nodes) => {
let mut best = RequiredLiterals::None;
let mut run = String::new();
for node in nodes {
let run_len = run.len();
if exact(node, &mut run) {
continue;
}
run.truncate(run_len);
if !run.is_empty() {
best = more_selective(
best,
RequiredLiterals::One(Cow::Owned(std::mem::take(&mut run))),
);
}
best = more_selective(best, reference_required_literals(node));
}
if !run.is_empty() {
best = more_selective(best, RequiredLiterals::One(Cow::Owned(run)));
}
best
}
Ast::Alternation(branches) => {
let mut literals = Vec::new();
for branch in branches {
match reference_required_literals(branch) {
RequiredLiterals::One(literal) => literals.push(literal),
RequiredLiterals::Any(mut nested) => literals.append(&mut nested),
RequiredLiterals::None => return RequiredLiterals::None,
}
}
literals.sort();
literals.dedup();
RequiredLiterals::Any(literals)
}
Ast::Group { child, .. } | Ast::Flags { child, .. } => {
reference_required_literals(child)
}
Ast::Look {
kind: LookKind::Ahead,
child,
} => reference_required_literals(child),
Ast::Repeat { node, min, .. } if *min > 0 => reference_required_literals(node),
Ast::Class(class) => class_required_literals(class),
_ => RequiredLiterals::None,
}
}
#[test]
fn required_literals_match_per_level_normalization() {
use crate::engine::grammar::load_dev_grammar_from_str;
use crate::engine::state::GrammarId;
use crate::grammars::registry::CORE_ASSETS;
let mut patterns = vec![
r"(&)(?=[A-Za-z])((a(s(ymp(eq)?|cr|t)|n(d|g)?)|A(s(sign|cr)|nd|MP))|(b(s|ig)|B(e|scr)))"
.to_owned(),
r"x(?:ab|cd|ab)y(?:ab|cd)z(?:e|f|e|f)".to_owned(),
r"(?:(?:ab|ab)|cd)(?:ef|gh|ef)".to_owned(),
r"(?i)k(?:ey|ind)|(?=[;)])s(?:ab|ab)+".to_owned(),
];
for (index, asset) in CORE_ASSETS.iter().enumerate() {
let grammar = load_dev_grammar_from_str(GrammarId(index as u16), asset.source)
.expect("core grammar parses");
patterns.extend(grammar.patterns.iter().map(|pattern| pattern.to_string()));
}
for pattern in &patterns {
let parsed = parse(pattern);
assert_eq!(
required_literals(&parsed.ast),
reference_required_literals(&parsed.ast),
"{pattern:?}"
);
}
}
}