use crate::hir::{Hir, HirExpr, HirLookaroundKind};
#[derive(Debug, Clone, Default)]
pub struct Literals {
pub prefixes: Vec<Vec<u8>>,
pub suffixes: Vec<Vec<u8>>,
pub prefix_complete: bool,
pub starts_with_digit: bool,
pub leading_bytes: Vec<u8>,
}
impl Literals {
pub fn is_empty(&self) -> bool {
self.prefixes.is_empty() && self.suffixes.is_empty()
}
pub fn single_prefix(&self) -> Option<&[u8]> {
if self.prefixes.len() == 1 {
Some(&self.prefixes[0])
} else {
None
}
}
pub fn has_multiple_prefixes(&self) -> bool {
self.prefixes.len() > 1
}
pub fn prefix_count(&self) -> usize {
self.prefixes.len()
}
}
pub fn extract_literals(hir: &Hir) -> Literals {
let mut extractor = LiteralExtractor::new();
let result = extractor.extract(&hir.expr);
let prefix_complete = result.complete
&& !hir.props.has_backrefs
&& !hir.props.has_lookaround
&& !hir.props.has_word_boundary
&& !hir.props.has_anchors;
let starts_with_digit = result.prefixes.is_empty() && starts_with_digit_class(&hir.expr);
let leading_bytes = if result.prefixes.is_empty() && !starts_with_digit {
leading_byte_set(&hir.expr).unwrap_or_default()
} else {
Vec::new()
};
Literals {
prefixes: result.prefixes,
suffixes: vec![],
prefix_complete,
starts_with_digit,
leading_bytes,
}
}
fn leading_byte_set(expr: &HirExpr) -> Option<Vec<u8>> {
match expr {
HirExpr::Class(class) if !class.negated => {
let mut bytes = Vec::new();
for &(lo, hi) in &class.ranges {
if bytes.len() + (hi as usize - lo as usize + 1) > MAX_LEADING_BYTES {
return None;
}
bytes.extend(lo..=hi);
}
(!bytes.is_empty()).then_some(bytes)
}
HirExpr::Literal(bytes) => bytes.first().map(|&b| vec![b]),
HirExpr::Concat(exprs) => exprs
.iter()
.find(|e| !is_zero_width(e))
.and_then(leading_byte_set),
HirExpr::Capture(capture) => leading_byte_set(&capture.expr),
HirExpr::Repeat(repeat) if repeat.min >= 1 => leading_byte_set(&repeat.expr),
_ => None,
}
}
const MAX_LEADING_BYTES: usize = 3;
pub type ByteSet = [u8; 256];
pub fn first_byte_set(hir: &Hir) -> Option<ByteSet> {
expr_first_byte_set(&hir.expr).filter(|set| set.contains(&0))
}
pub fn expr_first_byte_set(expr: &HirExpr) -> Option<ByteSet> {
let mut set = [0u8; 256];
add_first_bytes(expr, &mut set).then_some(set)
}
pub fn byte_class_set(ranges: &[(u8, u8)], negated: bool) -> ByteSet {
let mut set = [u8::from(negated); 256];
for &(lo, hi) in ranges {
for byte in lo..=hi {
if let Some(entry) = set.get_mut(byte as usize) {
*entry = u8::from(!negated);
}
}
}
set
}
pub fn single_byte_run_set(expr: &HirExpr) -> Option<ByteSet> {
match expr {
HirExpr::Class(class) => Some(byte_class_set(&class.ranges, class.negated)),
HirExpr::Alt(branches) => {
let mut run = [0u8; 256];
let mut others = [0u8; 256];
let mut found = false;
for branch in branches {
if let HirExpr::Class(class) = branch {
if !class.negated {
for (entry, member) in
run.iter_mut().zip(byte_class_set(&class.ranges, false))
{
*entry |= member;
}
found = true;
continue;
}
}
let first = expr_first_byte_set(branch)?;
for (entry, member) in others.iter_mut().zip(first) {
*entry |= member;
}
}
let disjoint = run
.iter()
.zip(others)
.all(|(member, other)| *member == 0 || other == 0);
(found && disjoint).then_some(run)
}
_ => None,
}
}
fn add_first_bytes(expr: &HirExpr, set: &mut ByteSet) -> bool {
match expr {
HirExpr::Class(class) => {
let members = byte_class_set(&class.ranges, class.negated);
for (entry, member) in set.iter_mut().zip(members) {
*entry |= member;
}
true
}
HirExpr::Literal(bytes) => match bytes.first() {
Some(&byte) => {
if let Some(entry) = set.get_mut(byte as usize) {
*entry = 1;
}
true
}
None => false,
},
HirExpr::Concat(exprs) => match exprs.iter().find(|expr| !is_zero_width(expr)) {
Some(expr) => add_first_bytes(expr, set),
None => false,
},
HirExpr::Alt(branches) => {
!branches.is_empty() && branches.iter().all(|branch| add_first_bytes(branch, set))
}
HirExpr::Capture(capture) => add_first_bytes(&capture.expr, set),
HirExpr::Repeat(repeat) if repeat.min >= 1 => add_first_bytes(&repeat.expr, set),
_ => false,
}
}
fn starts_with_digit_class(expr: &HirExpr) -> bool {
match expr {
HirExpr::Class(class) => {
!class.negated
&& !class.ranges.is_empty()
&& class
.ranges
.iter()
.all(|(lo, hi)| *lo >= b'0' && *hi <= b'9')
}
HirExpr::Concat(exprs) => {
for e in exprs {
if is_zero_width(e) {
continue;
}
return starts_with_digit_class(e);
}
false
}
HirExpr::Repeat(rep) if rep.min > 0 => starts_with_digit_class(&rep.expr),
HirExpr::Capture(cap) => starts_with_digit_class(&cap.expr),
_ => false,
}
}
fn is_zero_width(expr: &HirExpr) -> bool {
matches!(expr, HirExpr::Anchor(_) | HirExpr::Lookaround(_))
}
#[derive(Debug, Clone, Default)]
struct ExtractionResult {
prefixes: Vec<Vec<u8>>,
complete: bool,
has_nullable_suffix: bool,
}
struct LiteralExtractor {
max_prefixes: usize,
max_prefix_len: usize,
}
impl LiteralExtractor {
fn new() -> Self {
Self {
max_prefixes: 8, max_prefix_len: 8, }
}
fn extract(&mut self, expr: &HirExpr) -> ExtractionResult {
match expr {
HirExpr::Literal(bytes) => {
let truncated = bytes.len() > self.max_prefix_len;
let prefix = if truncated {
bytes[..self.max_prefix_len].to_vec()
} else {
bytes.clone()
};
ExtractionResult {
prefixes: vec![prefix],
complete: !truncated,
has_nullable_suffix: false,
}
}
HirExpr::Concat(exprs) => {
if exprs.is_empty() {
return ExtractionResult::default();
}
let mut start_idx = 0;
while start_idx < exprs.len() && is_zero_width(&exprs[start_idx]) {
start_idx += 1;
}
if start_idx >= exprs.len() {
return ExtractionResult::default();
}
let mut result = self.extract(&exprs[start_idx]);
let mut all_literals_so_far = matches!(&exprs[start_idx], HirExpr::Literal(_));
if result.complete && !result.has_nullable_suffix {
for expr in &exprs[start_idx + 1..] {
if is_zero_width(expr) {
continue;
}
if let HirExpr::Literal(bytes) = expr {
for prefix in &mut result.prefixes {
let remaining = self.max_prefix_len.saturating_sub(prefix.len());
if remaining > 0 {
let extend_len = bytes.len().min(remaining);
prefix.extend_from_slice(&bytes[..extend_len]);
if extend_len < bytes.len() {
result.complete = false;
}
} else {
result.complete = false;
}
}
} else {
all_literals_so_far = false;
result.complete = false;
let sub = self.extract(expr);
if sub.has_nullable_suffix {
result.has_nullable_suffix = true;
}
break;
}
}
} else {
if start_idx + 1 < exprs.len() {
result.complete = false;
}
}
if let Some(last) = exprs.last() {
let actual_last = exprs
.iter()
.rev()
.find(|e| !matches!(e, HirExpr::Anchor(_)))
.unwrap_or(last);
let last_result = self.extract(actual_last);
if last_result.has_nullable_suffix {
result.has_nullable_suffix = true;
}
if !last_result.complete || !matches!(actual_last, HirExpr::Literal(_)) {
if !all_literals_so_far {
result.complete = false;
}
}
}
result
}
HirExpr::Alt(exprs) => {
let mut all_prefixes: Vec<Vec<u8>> = Vec::new();
let mut all_complete = true;
let mut any_nullable_suffix = false;
for expr in exprs {
let sub_result = self.extract(expr);
if sub_result.prefixes.is_empty() {
return self.extract_common_prefix(exprs);
}
all_complete = all_complete && sub_result.complete;
any_nullable_suffix = any_nullable_suffix || sub_result.has_nullable_suffix;
all_prefixes.extend(sub_result.prefixes);
if all_prefixes.len() > self.max_prefixes {
return self.extract_common_prefix(exprs);
}
}
all_prefixes.sort();
all_prefixes.dedup();
ExtractionResult {
prefixes: all_prefixes,
complete: all_complete,
has_nullable_suffix: any_nullable_suffix,
}
}
HirExpr::Repeat(rep) => {
if rep.min > 0 {
let mut result = self.extract(&rep.expr);
result.has_nullable_suffix = true;
result.complete = false;
result
} else {
ExtractionResult {
has_nullable_suffix: true,
..Default::default()
}
}
}
HirExpr::Capture(cap) => self.extract(&cap.expr),
HirExpr::Class(_) => {
ExtractionResult::default()
}
_ => ExtractionResult::default(),
}
}
fn extract_common_prefix(&mut self, exprs: &[HirExpr]) -> ExtractionResult {
let mut all_prefixes: Vec<Vec<u8>> = Vec::new();
for expr in exprs {
let sub_result = self.extract(expr);
if sub_result.prefixes.is_empty() {
return ExtractionResult::default();
}
all_prefixes.extend(sub_result.prefixes);
}
if let Some(common) = find_common_prefix(&all_prefixes) {
if !common.is_empty() {
return ExtractionResult {
prefixes: vec![common],
complete: false, has_nullable_suffix: false,
};
}
}
ExtractionResult::default()
}
}
fn find_common_prefix(seqs: &[Vec<u8>]) -> Option<Vec<u8>> {
if seqs.is_empty() {
return None;
}
let first = &seqs[0];
let mut prefix_len = first.len();
for seq in &seqs[1..] {
let common_len = first
.iter()
.zip(seq.iter())
.take_while(|(a, b)| a == b)
.count();
prefix_len = prefix_len.min(common_len);
}
if prefix_len == 0 {
None
} else {
Some(first[..prefix_len].to_vec())
}
}
pub fn required_literal(hir: &Hir) -> Option<Vec<u8>> {
fn keep_longer(best: &mut Option<Vec<u8>>, candidate: Vec<u8>) {
if candidate.len() > best.as_ref().map_or(0, Vec::len) {
*best = Some(candidate);
}
}
fn walk(expr: &HirExpr) -> Option<Vec<u8>> {
match expr {
HirExpr::Lookaround(look) => match look.kind {
HirLookaroundKind::PositiveLookahead | HirLookaroundKind::PositiveLookbehind => {
let mut extractor = LiteralExtractor::new();
let inner = extractor.extract(&look.expr);
inner.prefixes.first().filter(|l| l.len() >= 2).cloned()
}
_ => None,
},
HirExpr::Literal(bytes) => (!bytes.is_empty()).then(|| bytes.clone()),
HirExpr::Concat(exprs) => {
let mut best: Option<Vec<u8>> = None;
let mut run: Vec<u8> = Vec::new();
for expr in exprs {
if let HirExpr::Literal(bytes) = expr {
run.extend_from_slice(bytes);
continue;
}
keep_longer(&mut best, std::mem::take(&mut run));
if let Some(found) = walk(expr) {
keep_longer(&mut best, found);
}
}
keep_longer(&mut best, run);
best
}
HirExpr::Capture(capture) => walk(&capture.expr),
HirExpr::Repeat(repeat) if repeat.min >= 1 => walk(&repeat.expr),
_ => None,
}
}
walk(&hir.expr)
}
#[cfg(test)]
mod tests {
#[test]
fn test_small_leading_class_yields_a_byte_set() {
assert_eq!(literals(r#"(['"])[^'"]*\1"#).leading_bytes, b"\"'".to_vec());
assert_eq!(literals(r"[abc]xyz").leading_bytes, b"abc".to_vec());
assert_eq!(literals(r"(?:[ab])+z").leading_bytes, b"ab".to_vec());
assert!(literals(r"[a-z]xyz").leading_bytes.is_empty());
assert!(literals(r"[^ab]xyz").leading_bytes.is_empty());
assert!(literals(r"abc[de]").leading_bytes.is_empty());
}
fn literals(pattern: &str) -> Literals {
let ast = crate::parser::parse(pattern).unwrap();
let hir = crate::hir::translate(&ast).unwrap();
extract_literals(&hir)
}
#[test]
fn test_leading_lookaround_does_not_hide_the_prefilter() {
assert!(literals(r"(?<!\$)\d+").starts_with_digit);
assert!(literals(r"(?<=\$)\d+").starts_with_digit);
assert!(literals(r"(?!x)\d+").starts_with_digit);
assert_eq!(literals(r"(?<!a)bcd").prefixes, vec![b"bcd".to_vec()]);
assert_eq!(literals(r"(?=x)abc").prefixes, vec![b"abc".to_vec()]);
assert!(!literals(r"(?<!a)bcd").prefix_complete);
}
fn required(pattern: &str) -> Option<Vec<u8>> {
let ast = crate::parser::parse(pattern).unwrap();
let hir = crate::hir::translate(&ast).unwrap();
required_literal(&hir)
}
#[test]
fn test_required_literal_from_positive_lookahead() {
assert_eq!(required(r"\w+(?=ing\b)"), Some(b"ing".to_vec()));
assert_eq!(required(r"(\w+)(?=ing\b)"), Some(b"ing".to_vec()));
assert_eq!(required(r"a(?=bcd)"), Some(b"bcd".to_vec()));
}
#[test]
fn test_required_literal_from_the_pattern_itself() {
assert_eq!(required(r"(\w+)@\1"), Some(b"@".to_vec()));
assert_eq!(required(r"\d+-->\d+"), Some(b"-->".to_vec()));
assert_eq!(required(r"a\d+bcde\d+"), Some(b"bcde".to_vec()));
assert_eq!(required(r"(?:xy)+\d"), Some(b"xy".to_vec()));
}
#[test]
fn test_no_required_literal_when_a_branch_may_not_need_it() {
assert_eq!(required(r"\w+(?=ing\b)|zzz"), None);
assert_eq!(required(r"(?:a(?=bcd)|q)"), None);
assert_eq!(required(r"abc|def"), None);
assert_eq!(required(r"\w+(?!ing)"), None);
assert_eq!(required(r"\w(?:(?=bcd))?"), None);
assert_eq!(required(r"\d+(?:xy)*"), None);
assert_eq!(required(r"\w+"), None);
}
use super::*;
use crate::hir::translate;
use crate::parser::parse;
fn get_literals(pattern: &str) -> Literals {
let ast = parse(pattern).unwrap();
let hir = translate(&ast).unwrap();
extract_literals(&hir)
}
#[test]
fn test_simple_literal() {
let lits = get_literals("hello");
assert_eq!(lits.prefixes.len(), 1);
assert_eq!(lits.prefixes[0], b"hello");
assert!(lits.prefix_complete);
}
#[test]
fn test_long_literal_truncated() {
let lits = get_literals("helloworld123");
assert_eq!(lits.prefixes.len(), 1);
assert_eq!(lits.prefixes[0], b"hellowor"); assert!(!lits.prefix_complete);
}
#[test]
fn test_no_prefix() {
let lits = get_literals(".*hello");
assert!(lits.prefixes.is_empty());
}
#[test]
fn test_alternation_multi_prefix() {
let lits = get_literals("hello|world");
assert_eq!(lits.prefixes.len(), 2);
assert!(lits.prefixes.contains(&b"hello".to_vec()));
assert!(lits.prefixes.contains(&b"world".to_vec()));
}
#[test]
fn test_alternation_common_prefix() {
let lits = get_literals("hello|help");
assert_eq!(lits.prefixes.len(), 2);
assert!(lits.prefixes.contains(&b"hello".to_vec()));
assert!(lits.prefixes.contains(&b"help".to_vec()));
}
#[test]
fn test_concat_extends_prefix() {
let lits = get_literals("ab");
assert_eq!(lits.prefixes.len(), 1);
assert_eq!(lits.prefixes[0], b"ab");
}
#[test]
fn test_class_no_prefix() {
let lits = get_literals("[abc]hello");
assert!(lits.prefixes.is_empty());
}
#[test]
fn test_repeat_one_or_more() {
let lits = get_literals("a+b");
assert_eq!(lits.prefixes.len(), 1);
assert_eq!(lits.prefixes[0], b"a");
}
#[test]
fn test_repeat_zero_or_more_no_prefix() {
let lits = get_literals("a*b");
assert!(lits.prefixes.is_empty());
}
#[test]
fn test_too_many_alternations() {
let lits = get_literals("a|b|c|d|e|f|g|h|i|j");
assert!(lits.prefixes.is_empty());
}
#[test]
fn test_nested_alternation() {
let lits = get_literals("(cat|dog)food");
assert_eq!(lits.prefixes.len(), 2);
assert!(lits.prefixes.contains(&b"catfood".to_vec()));
assert!(lits.prefixes.contains(&b"dogfood".to_vec()));
}
#[test]
fn test_literal_then_star() {
let lits = get_literals("hello.*world");
assert_eq!(lits.prefixes.len(), 1);
assert_eq!(lits.prefixes[0], b"hello");
}
#[test]
fn test_literal_then_class() {
let lits = get_literals("hello[0-9]+");
assert_eq!(lits.prefixes.len(), 1);
assert_eq!(lits.prefixes[0], b"hello");
}
}