use crate::error::{Error, ErrorKind, Result};
use crate::hir::Hir;
use super::super::interpreter::BacktrackingVm;
use super::super::shared::{BudgetExhausted, CaptureSlots, DEFAULT_BACKTRACK_LIMIT};
use dynasmrt::ExecutableBuffer;
enum Halt {
Retry,
Budget,
}
#[cfg(target_arch = "x86_64")]
use super::x86_64::BacktrackingCompiler;
#[cfg(target_arch = "aarch64")]
use super::aarch64::BacktrackingCompiler;
#[cfg(all(target_arch = "x86_64", target_os = "windows"))]
type MatchFn = unsafe extern "win64" fn(*const u8, usize, *mut i64, u64) -> i64;
#[cfg(all(target_arch = "x86_64", not(target_os = "windows")))]
type MatchFn = unsafe extern "sysv64" fn(*const u8, usize, *mut i64, u64) -> i64;
#[cfg(target_arch = "aarch64")]
type MatchFn = unsafe extern "C" fn(*const u8, usize, *mut i64, u64) -> i64;
pub(super) const STACK_EXHAUSTED: i64 = -2;
pub(super) const BUDGET_EXHAUSTED: i64 = -3;
pub struct BacktrackingJit {
#[allow(dead_code)]
pub(super) code: ExecutableBuffer,
pub(super) match_fn: MatchFn,
pub(super) capture_count: u32,
pub(super) vm: BacktrackingVm,
pub(super) needs_left_context: bool,
}
impl BacktrackingJit {
pub fn is_match(&self, input: &[u8]) -> bool {
self.find(input).is_some()
}
fn run(&self, input: &[u8], limit: u64) -> std::result::Result<Option<CaptureSlots>, Halt> {
let num_slots = (self.capture_count as usize + 1) * 2;
let mut buf: Vec<i64> = vec![-1; num_slots];
let result =
unsafe { (self.match_fn)(input.as_ptr(), input.len(), buf.as_mut_ptr(), limit) };
if result == STACK_EXHAUSTED {
return Err(Halt::Retry);
}
if result == BUDGET_EXHAUSTED {
return Err(Halt::Budget);
}
if result < 0 {
return Ok(None);
}
let mut captures = Vec::with_capacity(self.capture_count as usize + 1);
for i in 0..=self.capture_count as usize {
let (start, end) = (buf[i * 2], buf[i * 2 + 1]);
if start >= 0 && end >= 0 {
captures.push(Some((start as usize, end as usize)));
} else {
captures.push(None);
}
}
Ok(Some(captures))
}
pub fn find(&self, input: &[u8]) -> Option<(usize, usize)> {
self.captures(input).and_then(|caps| caps[0])
}
pub fn captures(&self, input: &[u8]) -> Option<Vec<Option<(usize, usize)>>> {
self.try_captures_from(input, 0, DEFAULT_BACKTRACK_LIMIT)
.unwrap_or(None)
}
pub fn try_captures_from(
&self,
input: &[u8],
from: usize,
limit: u64,
) -> std::result::Result<Option<CaptureSlots>, BudgetExhausted> {
if from > input.len() {
return Ok(None);
}
if self.needs_left_context && from > 0 {
return self.vm.try_captures_from(input, from, limit);
}
let (haystack, offset) = if from == 0 {
(input, 0)
} else {
(&input[from..], from)
};
let caps = match self.run(haystack, limit) {
Ok(caps) => caps,
Err(Halt::Retry) => self.vm.try_captures_from(haystack, 0, limit)?,
Err(Halt::Budget) => return Err(BudgetExhausted),
};
Ok(caps.map(|caps| {
caps.into_iter()
.map(|slot| slot.map(|(s, e)| (s + offset, e + offset)))
.collect()
}))
}
pub fn find_at(&self, input: &[u8], start: usize) -> Option<(usize, usize)> {
self.find_from(input, start)
}
pub fn find_from(&self, input: &[u8], from: usize) -> Option<(usize, usize)> {
self.captures_from(input, from).and_then(|caps| caps[0])
}
pub fn captures_from(&self, input: &[u8], from: usize) -> Option<Vec<Option<(usize, usize)>>> {
self.try_captures_from(input, from, DEFAULT_BACKTRACK_LIMIT)
.unwrap_or(None)
}
#[cfg(test)]
pub fn debug_match(&self, input: &[u8]) -> (i64, Vec<i64>) {
let num_slots = (self.capture_count as usize + 1) * 2;
let mut captures: Vec<i64> = vec![-1; num_slots];
let result = unsafe {
(self.match_fn)(
input.as_ptr(),
input.len(),
captures.as_mut_ptr(),
DEFAULT_BACKTRACK_LIMIT,
)
};
(result, captures)
}
}
pub fn compile_backtracking(hir: &Hir) -> Result<BacktrackingJit> {
if crate::hir::has_unbounded_nullable_repeat(&hir.expr) {
return Err(Error::new(
ErrorKind::Jit("unbounded repetition over a nullable body".to_string()),
"",
));
}
let compiler = BacktrackingCompiler::new(hir)?;
let mut jit = compiler.compile()?;
let props = &hir.props;
jit.needs_left_context =
props.has_start_anchor || props.has_multiline_anchors || props.has_word_boundary;
Ok(jit)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hir::translate;
use crate::parser::parse;
fn compile_pattern(pattern: &str) -> Result<BacktrackingJit> {
let ast = parse(pattern)?;
let hir = translate(&ast)?;
compile_backtracking(&hir)
}
#[test]
fn test_literal_debug() {
let jit = compile_pattern("hello").unwrap();
let (result, caps) = jit.debug_match(b"hello");
println!("result: {}, caps: {:?}", result, caps);
assert!(result >= 0, "Expected match, got result={}", result);
}
#[test]
fn test_literal() {
let jit = compile_pattern("hello").unwrap();
assert!(jit.is_match(b"hello"));
assert!(jit.is_match(b"say hello world"));
assert!(!jit.is_match(b"helo"));
}
#[test]
fn test_simple_backref() {
let jit = compile_pattern(r"(a)\1").unwrap();
let (result_aa, caps_aa) = jit.debug_match(b"aa");
println!("(a)\\1 on 'aa': result={}, caps={:?}", result_aa, caps_aa);
let (result_ab, caps_ab) = jit.debug_match(b"ab");
println!("(a)\\1 on 'ab': result={}, caps={:?}", result_ab, caps_ab);
let (result_a, caps_a) = jit.debug_match(b"a");
println!("(a)\\1 on 'a': result={}, caps={:?}", result_a, caps_a);
assert!(jit.is_match(b"aa"), "Should match 'aa'");
assert!(!jit.is_match(b"ab"), "Should NOT match 'ab'");
assert!(!jit.is_match(b"a"), "Should NOT match 'a'");
}
#[test]
fn test_quoted_string() {
let jit = compile_pattern(r#"(['"])[^'"]*\1"#).unwrap();
let (r1, c1) = jit.debug_match(br#""hello""#);
println!(r#"['"][^'"]*\1 on "hello": result={}, caps={:?}"#, r1, c1);
let (r2, c2) = jit.debug_match(b"'world'");
println!(r#"['"][^'"]*\1 on 'world': result={}, caps={:?}"#, r2, c2);
let (r3, c3) = jit.debug_match(br#""mixed'"#);
println!(r#"['"][^'"]*\1 on "mixed': result={}, caps={:?}"#, r3, c3);
let (r4, c4) = jit.debug_match(b"'mixed\"");
println!(r#"['"][^'"]*\1 on 'mixed": result={}, caps={:?}"#, r4, c4);
assert!(jit.is_match(br#""hello""#), "Should match \"hello\"");
assert!(jit.is_match(b"'world'"), "Should match 'world'");
assert!(!jit.is_match(br#""mixed'"#), "Should NOT match \"mixed'");
assert!(!jit.is_match(b"'mixed\""), "Should NOT match 'mixed\"");
}
#[test]
fn test_alternation_backref() {
let jit = compile_pattern(r"(a|b)\1").unwrap();
let (result_aa, caps_aa) = jit.debug_match(b"aa");
println!("(a|b)\\1 on 'aa': result={}, caps={:?}", result_aa, caps_aa);
let (result_bb, caps_bb) = jit.debug_match(b"bb");
println!("(a|b)\\1 on 'bb': result={}, caps={:?}", result_bb, caps_bb);
let (result_ab, caps_ab) = jit.debug_match(b"ab");
println!("(a|b)\\1 on 'ab': result={}, caps={:?}", result_ab, caps_ab);
let (result_ba, caps_ba) = jit.debug_match(b"ba");
println!("(a|b)\\1 on 'ba': result={}, caps={:?}", result_ba, caps_ba);
assert!(jit.is_match(b"aa"), "Should match 'aa'");
assert!(jit.is_match(b"bb"), "Should match 'bb'");
assert!(!jit.is_match(b"ab"), "Should NOT match 'ab'");
assert!(!jit.is_match(b"ba"), "Should NOT match 'ba'");
}
#[test]
fn test_captures() {
let jit = compile_pattern(r"(a)(b)\2\1").unwrap();
let caps = jit.captures(b"abba").unwrap();
assert_eq!(caps[0], Some((0, 4))); assert_eq!(caps[1], Some((0, 1))); assert_eq!(caps[2], Some((1, 2))); }
}