use memchr::memmem::Finder;
use crate::hir::{
compute_capture_count, CodepointClass, Hir, HirClass, HirExpr, HirLookaroundKind, HirProps,
HirRepeat,
};
use crate::literal::{byte_class_set, extract_literals};
use crate::nfa::{ByteClass, ByteRange};
use crate::vm::pike::shared::decode_utf8_codepoint;
#[derive(Debug, Clone)]
pub(crate) enum RunClass {
Bytes(ByteClass),
Codepoints(CodepointClass),
}
#[derive(Debug, Clone)]
pub(crate) struct ReverseSuffixPlan {
pub(crate) literal: Vec<u8>,
pub(crate) class: RunClass,
pub(crate) allows_empty_run: bool,
}
#[derive(Debug)]
pub(crate) struct ReverseSuffixSearch {
plan: ReverseSuffixPlan,
finder: Finder<'static>,
}
impl ReverseSuffixSearch {
pub(crate) fn new(hir: &Hir) -> Option<Self> {
let plan = reverse_suffix_plan(hir)?;
let finder = Finder::new(&plan.literal).into_owned();
Some(Self { plan, finder })
}
pub(crate) fn find(
&self,
input: &[u8],
from: usize,
mut confirm: impl FnMut(usize) -> Option<(usize, usize)>,
) -> Option<(usize, usize)> {
let mut pos = from;
loop {
let p = pos + self.finder.find(input.get(pos..)?)?;
let mut start = run_start(input, p, &self.plan.class, from);
while start < p && !crate::nfa::is_utf8_boundary(input, start) {
start += 1;
}
if start == p && !self.plan.allows_empty_run {
pos = p + 1;
continue;
}
if let Some(found) = confirm(start) {
return Some(found);
}
pos = (p + 1).max(run_end(input, p, &self.plan.class));
}
}
}
pub(crate) fn reverse_suffix_plan(hir: &Hir) -> Option<ReverseSuffixPlan> {
if hir.props.has_backrefs || hir.props.has_non_greedy {
return None;
}
let HirExpr::Concat(parts) = &hir.expr else {
return None;
};
let [head, tail] = parts.as_slice() else {
return None;
};
let HirExpr::Repeat(repeat) = head else {
return None;
};
let class = run_class(repeat)?;
let HirExpr::Lookaround(look) = tail else {
return None;
};
if !matches!(look.kind, HirLookaroundKind::PositiveLookahead) {
return None;
}
if count_lookarounds(&hir.expr) != 1 {
return None;
}
let literal = leading_literal(&look.expr)?;
if literal.len() < 2 {
return None;
}
Some(ReverseSuffixPlan {
literal,
class,
allows_empty_run: repeat.min == 0,
})
}
fn run_class(repeat: &HirRepeat) -> Option<RunClass> {
if !repeat.greedy || repeat.max.is_some() || repeat.min > 1 {
return None;
}
if compute_capture_count(&repeat.expr) > 0 {
return None;
}
match &repeat.expr {
HirExpr::Class(class) => Some(RunClass::Bytes(byte_class(class))),
HirExpr::UnicodeCpClass(cpclass) => Some(RunClass::Codepoints(cpclass.clone())),
_ => None,
}
}
fn byte_class(class: &HirClass) -> ByteClass {
let members = byte_class_set(&class.ranges, class.negated);
let is_member = |byte: usize| members.get(byte).is_some_and(|&flag| flag != 0);
let mut ranges = Vec::new();
let mut byte = 0usize;
while byte < 256 {
if !is_member(byte) {
byte += 1;
continue;
}
let start = byte as u8;
while byte < 256 && is_member(byte) {
byte += 1;
}
ranges.push(ByteRange::new(start, (byte - 1) as u8));
}
ByteClass::new(ranges)
}
fn leading_literal(inner: &HirExpr) -> Option<Vec<u8>> {
let probe = Hir {
expr: inner.clone(),
props: HirProps::default(),
};
let literals = extract_literals(&probe);
if literals.prefix_offset != 0 {
return None;
}
literals.single_prefix().map(|prefix| prefix.to_vec())
}
fn count_lookarounds(expr: &HirExpr) -> usize {
match expr {
HirExpr::Empty
| HirExpr::Literal(_)
| HirExpr::Class(_)
| HirExpr::UnicodeCpClass(_)
| HirExpr::Anchor(_)
| HirExpr::Backref(_) => 0,
HirExpr::Concat(exprs) | HirExpr::Alt(exprs) => exprs.iter().map(count_lookarounds).sum(),
HirExpr::Repeat(rep) => count_lookarounds(&rep.expr),
HirExpr::Capture(cap) => count_lookarounds(&cap.expr),
HirExpr::Lookaround(look) => 1 + count_lookarounds(&look.expr),
}
}
pub(crate) fn run_start(input: &[u8], p: usize, class: &RunClass, floor: usize) -> usize {
let mut start = p.min(input.len());
let floor = floor.min(start);
match class {
RunClass::Bytes(bytes) => {
while start > floor {
let Some(&byte) = input.get(start - 1) else {
break;
};
if !bytes.contains(byte) {
break;
}
start -= 1;
}
}
RunClass::Codepoints(cpclass) => {
while start > floor {
let boundary = prev_boundary(input, start);
if boundary < floor {
break;
}
let Some(bytes) = input.get(boundary..start) else {
break;
};
match decode_utf8_codepoint(bytes) {
Some((cp, len)) if len == start - boundary && cpclass.contains(cp) => {
start = boundary;
}
_ => break,
}
}
}
}
start
}
pub(crate) fn run_end(input: &[u8], p: usize, class: &RunClass) -> usize {
let mut end = p.min(input.len());
match class {
RunClass::Bytes(bytes) => {
while input.get(end).is_some_and(|&byte| bytes.contains(byte)) {
end += 1;
}
}
RunClass::Codepoints(cpclass) => {
while let Some((cp, len)) = input.get(end..).and_then(decode_utf8_codepoint) {
if len == 0 || !cpclass.contains(cp) {
break;
}
end += len;
}
}
}
end
}
fn prev_boundary(input: &[u8], end: usize) -> usize {
let mut i = end.saturating_sub(1);
while i > 0 && end - i < 4 && input.get(i).is_some_and(|&byte| (byte & 0xC0) == 0x80) {
i -= 1;
}
i
}
#[cfg(test)]
mod tests {
use super::*;
fn plan(pattern: &str) -> Option<ReverseSuffixPlan> {
let ast = crate::parser::parse(pattern).unwrap();
let hir = crate::hir::translate(&ast).unwrap();
reverse_suffix_plan(&hir)
}
fn search(pattern: &str) -> ReverseSuffixSearch {
let ast = crate::parser::parse(pattern).unwrap();
let hir = crate::hir::translate(&ast).unwrap();
ReverseSuffixSearch::new(&hir).expect("gate accepts this shape")
}
fn confirms(
search: &ReverseSuffixSearch,
input: &[u8],
from: usize,
matches: Option<(usize, usize)>,
) -> (Option<(usize, usize)>, Vec<usize>) {
let mut asked = Vec::new();
let found = search.find(input, from, |start| {
asked.push(start);
matches.filter(|&(match_start, _)| match_start == start)
});
(found, asked)
}
#[test]
fn test_search_confirms_one_start_per_run() {
let ing = search(r"\w+(?=ing\b)");
let (found, asked) = confirms(&ing, b"singing", 0, Some((0, 4)));
assert_eq!(found, Some((0, 4)));
assert_eq!(asked, vec![0]);
let (found, asked) = confirms(&ing, b"singing", 0, None);
assert_eq!(found, None);
assert_eq!(asked, vec![0]);
let (found, asked) = confirms(&ing, b"singing ringing", 0, Some((8, 12)));
assert_eq!(found, Some((8, 12)));
assert_eq!(asked, vec![0, 8]);
}
#[test]
fn test_search_skips_an_empty_run_only_when_the_repeat_needs_one() {
let ing = search(r"\w+(?=ing\b)");
assert_eq!(confirms(&ing, b"ings", 0, None), (None, vec![]));
assert_eq!(confirms(&ing, b" ing", 0, None), (None, vec![]));
let xy = search(r"[a-z]+(?=xy)");
assert_eq!(
confirms(&xy, b" xyxy", 0, Some((1, 3))),
(Some((1, 3)), vec![1])
);
let star = search(r"\w*(?=ing)");
assert_eq!(
confirms(&star, b"ings", 0, Some((0, 0))),
(Some((0, 0)), vec![0])
);
}
#[test]
fn test_search_never_confirms_a_start_before_the_resume_point() {
let ing = search(r"\w+(?=ing\b)");
assert_eq!(confirms(&ing, b"singing", 4, None), (None, vec![]));
assert_eq!(
confirms(&ing, b"singing", 2, Some((2, 4))),
(Some((2, 4)), vec![2])
);
}
fn bytes_class(ranges: &[(u8, u8)]) -> RunClass {
RunClass::Bytes(byte_class(&HirClass::new(ranges.to_vec(), false)))
}
fn codepoint_class(ranges: &[(u32, u32)]) -> RunClass {
RunClass::Codepoints(CodepointClass::new(ranges.to_vec(), false))
}
#[test]
fn test_plan_accepts_a_class_run_before_a_literal_lookahead() {
let accepted = plan(r"\w+(?=ing\b)").expect("class run before a literal lookahead");
assert_eq!(accepted.literal, b"ing".to_vec());
assert!(!accepted.allows_empty_run);
let RunClass::Bytes(class) = &accepted.class else {
panic!("ASCII \\w is a byte class");
};
assert!(class.contains(b'a'));
assert!(class.contains(b'_'));
assert!(!class.contains(b' '));
let accepted = plan(r"\w*(?=ing)").expect("a zero-or-more run is still one run");
assert!(accepted.allows_empty_run);
let accepted = plan(r"[a-z]+(?=xy)").expect("an explicit class is the same shape");
assert_eq!(accepted.literal, b"xy".to_vec());
}
#[test]
fn test_plan_accepts_a_codepoint_class_run() {
let accepted = plan(r"(?u:\w+(?=ing))").expect("a codepoint run is still one run");
assert!(matches!(accepted.class, RunClass::Codepoints(_)));
assert_eq!(accepted.literal, b"ing".to_vec());
}
#[test]
fn test_plan_refuses_everything_that_is_not_one_class_run() {
assert!(plan(r"(?:abcde|c)(?=d)").is_none());
assert!(plan(r"\w+(?=ing|ed)").is_none());
assert!(plan(r"\w+(?=\w*ing)").is_none());
assert!(plan(r"\w+\s+error\s+\w+").is_none());
assert!(plan(r"\w+?(?=ing)").is_none());
assert!(plan(r"\w{2,5}(?=ing)").is_none());
assert!(plan(r"(\w)+(?=ing)").is_none());
assert!(plan(r"(?i:\w+(?=ing))").is_none());
assert!(plan(r"(?i)\w+(?=ing)").is_none());
assert!(plan(r"\w+(?=(i)ng\1)").is_none());
assert!(plan(r"\w+(?=x)").is_none());
assert!(plan(r"\w+(?!ing)").is_none());
assert!(plan(r"\w+(?=ing(?=s))").is_none());
assert!(plan(r"abc(?=ing)").is_none());
assert!(plan(r"\w+ing").is_none());
}
#[test]
fn test_run_scan_over_a_byte_class() {
let class = bytes_class(&[(b'a', b'z')]);
let input = b"foo bar";
assert_eq!(run_start(input, 3, &class, 0), 0);
assert_eq!(run_end(input, 0, &class), 3);
assert_eq!(run_start(input, input.len(), &class, 0), 4);
assert_eq!(run_end(input, input.len(), &class), input.len());
assert_eq!(run_start(input, 0, &class, 0), 0);
assert_eq!(run_start(input, 4, &class, 0), 4);
assert_eq!(run_end(input, 3, &class), 3);
let whole = b"foobar";
assert_eq!(run_start(whole, whole.len(), &class, 0), 0);
assert_eq!(run_end(whole, 0, &class), whole.len());
assert_eq!(run_start(whole, 99, &class, 0), 0);
assert_eq!(run_end(whole, 99, &class), whole.len());
}
#[test]
fn test_run_start_never_walks_past_the_floor() {
let class = bytes_class(&[(b'a', b'z')]);
let whole = b"foobar";
assert_eq!(run_start(whole, whole.len(), &class, 3), 3);
assert_eq!(run_start(whole, 3, &class, 3), 3);
assert_eq!(run_start(whole, 3, &class, 99), 3);
let letters = codepoint_class(&[(0x61, 0x7A), (0xE9, 0xE9), (0x4E16, 0x4E16)]);
let input = "aé世".as_bytes(); assert_eq!(input.len(), 6);
assert_eq!(run_start(input, 6, &letters, 0), 0);
assert_eq!(run_start(input, 6, &letters, 3), 3);
assert_eq!(run_start(input, 6, &letters, 2), 3);
}
#[test]
fn test_run_scan_over_a_codepoint_class() {
let letters = codepoint_class(&[(0x61, 0x7A), (0xE9, 0xE9), (0x4E16, 0x4E16)]);
let input = "aé世 b".as_bytes();
assert_eq!(input.len(), 8);
assert_eq!(run_start(input, 6, &letters, 0), 0);
assert_eq!(run_end(input, 0, &letters), 6);
assert_eq!(run_start(input, input.len(), &letters, 0), 7);
assert_eq!(run_end(input, input.len(), &letters), input.len());
assert_eq!(run_start(input, 0, &letters, 0), 0);
assert_eq!(run_start(input, 7, &letters, 0), 7);
assert_eq!(run_end(input, 6, &letters), 6);
let whole = "aé世".as_bytes();
assert_eq!(run_start(whole, whole.len(), &letters, 0), 0);
assert_eq!(run_end(whole, 0, &letters), whole.len());
let other = "a漢éb".as_bytes();
assert_eq!(other.len(), 7);
assert_eq!(run_start(other, other.len(), &letters, 0), 4);
assert_eq!(run_end(other, 4, &letters), other.len());
assert_eq!(run_end(other, 0, &letters), 1);
assert_eq!(run_start(whole, 99, &letters, 0), 0);
assert_eq!(run_end(whole, 99, &letters), whole.len());
}
}