use super::ast::{AnchorKind, Ast, CharClass, ClassAtom, LookKind, ParsedRegex, PerlClassKind};
use super::is_unicode_word_char;
pub(crate) const START_CLASS_ALL: u8 = 0b1111;
const PREV_WORD: u8 = 0b1100;
const PREV_NONWORD: u8 = 0b0011;
const CUR_WORD: u8 = 0b1010;
const CUR_NONWORD: u8 = 0b0101;
const BOUNDARY: u8 = 0b0110;
const NOT_BOUNDARY: u8 = 0b1001;
pub(crate) fn start_class_mask(parsed: &ParsedRegex) -> u8 {
let (mask, continuation) = node_mask(&parsed.ast, START_CLASS_ALL);
let mask = mask | continuation.unwrap_or(0);
if mask == 0 {
START_CLASS_ALL
} else {
mask
}
}
fn bits_prev(word: bool) -> u8 {
if word { PREV_WORD } else { PREV_NONWORD }
}
fn bits_cur(word: bool) -> u8 {
if word { CUR_WORD } else { CUR_NONWORD }
}
fn sides_cur_bits(word: bool, nonword: bool) -> u8 {
let mut bits = 0;
if word {
bits |= CUR_WORD;
}
if nonword {
bits |= CUR_NONWORD;
}
bits
}
fn sides_prev_bits(word: bool, nonword: bool) -> u8 {
let mut bits = 0;
if word {
bits |= PREV_WORD;
}
if nonword {
bits |= PREV_NONWORD;
}
bits
}
fn node_mask(ast: &Ast, constraint: u8) -> (u8, Option<u8>) {
match ast {
Ast::Empty => (0, Some(constraint)),
Ast::Literal(literal) => match literal.chars().next() {
Some(ch) => (constraint & bits_cur(is_unicode_word_char(ch)), None),
None => (0, Some(constraint)),
},
Ast::Class(class) => {
let (word, nonword) = class_word_sides(class);
(constraint & sides_cur_bits(word, nonword), None)
}
Ast::Dot | Ast::Grapheme => (constraint, None),
Ast::Anchor(kind) => {
let narrowed = match kind {
AnchorKind::LineStart | AnchorKind::TextStart => constraint & bits_prev(false),
AnchorKind::LineEnd | AnchorKind::TextEnd | AnchorKind::TextEndOrFinalNewline => {
constraint & bits_cur(false)
}
AnchorKind::Continuation => constraint,
AnchorKind::WordBoundary => constraint & BOUNDARY,
AnchorKind::NotWordBoundary => constraint & NOT_BOUNDARY,
};
(0, Some(narrowed))
}
Ast::Look { kind, child } => {
let narrowed = match kind {
LookKind::Ahead => {
let sides = first_char_sides(child);
if sides.nullable {
constraint
} else {
constraint & sides_cur_bits(sides.word, sides.nonword)
}
}
LookKind::NotAhead => {
if negated_look_excludes_word(child, false) {
constraint & bits_cur(false)
} else {
constraint
}
}
LookKind::Behind => {
let sides = last_char_sides(child);
if sides.nullable {
constraint
} else {
constraint & sides_prev_bits(sides.word, sides.nonword)
}
}
LookKind::NotBehind => {
if negated_look_excludes_word(child, true) {
constraint & bits_prev(false)
} else {
constraint
}
}
};
(0, Some(narrowed))
}
Ast::Concat(nodes) => {
let mut mask = 0;
let mut carried = constraint;
for node in nodes {
let (node_bits, continuation) = node_mask(node, carried);
mask |= node_bits;
match continuation {
Some(narrowed) => carried = narrowed,
None => return (mask, None),
}
}
(mask, Some(carried))
}
Ast::Alternation(branches) => {
let mut mask = 0;
let mut continuation: Option<u8> = None;
for branch in branches {
let (branch_bits, branch_continuation) = node_mask(branch, constraint);
mask |= branch_bits;
if let Some(narrowed) = branch_continuation {
continuation = Some(continuation.unwrap_or(0) | narrowed);
}
}
(mask, continuation)
}
Ast::Repeat { node, min, max, .. } => {
if *max == Some(0) {
return (0, Some(constraint));
}
let (mask, continuation) = node_mask(node, constraint);
if *min == 0 {
(mask, Some(constraint | continuation.unwrap_or(0)))
} else {
(mask, continuation)
}
}
Ast::Group { child, .. } | Ast::Flags { child, .. } => node_mask(child, constraint),
Ast::Backref(_) | Ast::Conditional { .. } | Ast::Subroutine(_) | Ast::Unsupported(_) => {
(constraint, Some(constraint))
}
}
}
#[derive(Clone, Copy)]
struct CharSides {
word: bool,
nonword: bool,
nullable: bool,
}
impl CharSides {
const UNKNOWN: Self = Self {
word: true,
nonword: true,
nullable: true,
};
const ZERO_WIDTH: Self = Self {
word: false,
nonword: false,
nullable: true,
};
}
fn first_char_sides(ast: &Ast) -> CharSides {
char_sides(ast, false)
}
fn last_char_sides(ast: &Ast) -> CharSides {
char_sides(ast, true)
}
fn char_sides(ast: &Ast, from_end: bool) -> CharSides {
match ast {
Ast::Empty => CharSides::ZERO_WIDTH,
Ast::Literal(literal) => {
let ch = if from_end {
literal.chars().next_back()
} else {
literal.chars().next()
};
match ch {
Some(ch) => {
let word = is_unicode_word_char(ch);
CharSides {
word,
nonword: !word,
nullable: false,
}
}
None => CharSides::ZERO_WIDTH,
}
}
Ast::Class(class) => {
let (word, nonword) = class_word_sides(class);
CharSides {
word,
nonword,
nullable: false,
}
}
Ast::Dot | Ast::Grapheme => CharSides {
word: true,
nonword: true,
nullable: false,
},
Ast::Anchor(_) | Ast::Look { .. } => CharSides::ZERO_WIDTH,
Ast::Concat(nodes) => {
let mut word = false;
let mut nonword = false;
let mut iterate = |node: &Ast| -> bool {
let sides = char_sides(node, from_end);
word |= sides.word;
nonword |= sides.nonword;
sides.nullable
};
let nullable = if from_end {
nodes.iter().rev().all(&mut iterate)
} else {
nodes.iter().all(&mut iterate)
};
CharSides {
word,
nonword,
nullable,
}
}
Ast::Alternation(branches) => {
let mut word = false;
let mut nonword = false;
let mut nullable = false;
for branch in branches {
let sides = char_sides(branch, from_end);
word |= sides.word;
nonword |= sides.nonword;
nullable |= sides.nullable;
}
CharSides {
word,
nonword,
nullable,
}
}
Ast::Repeat { node, min, max, .. } => {
if *max == Some(0) {
return CharSides::ZERO_WIDTH;
}
let mut sides = char_sides(node, from_end);
sides.nullable |= *min == 0;
sides
}
Ast::Group { child, .. } | Ast::Flags { child, .. } => char_sides(child, from_end),
Ast::Backref(_) | Ast::Conditional { .. } | Ast::Subroutine(_) | Ast::Unsupported(_) => {
CharSides::UNKNOWN
}
}
}
fn negated_look_excludes_word(child: &Ast, from_end: bool) -> bool {
match child {
Ast::Class(class) => class_covers_all_ascii_word(class),
Ast::Group { child, .. } | Ast::Flags { child, .. } => {
negated_look_excludes_word(child, from_end)
}
Ast::Concat(nodes) => {
let (adjacent, rest) = if from_end {
match nodes.split_last() {
Some((last, rest)) => (last, rest),
None => return false,
}
} else {
match nodes.split_first() {
Some((first, rest)) => (first, rest),
None => return false,
}
};
negated_look_excludes_word(adjacent, from_end)
&& rest.iter().all(matches_empty_unconditionally)
}
Ast::Alternation(branches) => branches
.iter()
.any(|branch| negated_look_excludes_word(branch, from_end)),
Ast::Repeat { node, min, max, .. } => {
*min <= 1
&& max.is_none_or(|max| max >= 1)
&& negated_look_excludes_word(node, from_end)
}
_ => false,
}
}
fn matches_empty_unconditionally(ast: &Ast) -> bool {
match ast {
Ast::Empty => true,
Ast::Literal(literal) => literal.is_empty(),
Ast::Repeat { min, .. } => *min == 0,
Ast::Group { child, .. } | Ast::Flags { child, .. } => matches_empty_unconditionally(child),
Ast::Concat(nodes) => nodes.iter().all(matches_empty_unconditionally),
Ast::Alternation(branches) => branches.iter().any(matches_empty_unconditionally),
_ => false,
}
}
fn class_covers_all_ascii_word(class: &CharClass) -> bool {
if class.negated || !class.intersections.is_empty() {
return false;
}
atoms_cover_all_ascii_word(&class.atoms)
}
fn atoms_cover_all_ascii_word(atoms: &[ClassAtom]) -> bool {
let mut covered = [false; 128];
for atom in atoms {
match atom {
ClassAtom::Perl(PerlClassKind::Word) => return true,
ClassAtom::Char(ch) if ch.is_ascii() => covered[*ch as usize] = true,
ClassAtom::Range(start, end) if start.is_ascii() && end.is_ascii() => {
let (start, end) = (*start.min(end) as usize, *start.max(end) as usize);
for slot in &mut covered[start..=end] {
*slot = true;
}
}
ClassAtom::Posix {
name,
negated: false,
} => match name.as_str() {
"word" => return true,
"alnum" => {
for ch in ('0'..='9').chain('A'..='Z').chain('a'..='z') {
covered[ch as usize] = true;
}
}
"alpha" => {
for ch in ('A'..='Z').chain('a'..='z') {
covered[ch as usize] = true;
}
}
"digit" => {
for ch in '0'..='9' {
covered[ch as usize] = true;
}
}
_ => {}
},
_ => {}
}
}
('0'..='9')
.chain('A'..='Z')
.chain('a'..='z')
.chain(std::iter::once('_'))
.all(|ch| covered[ch as usize])
}
fn class_word_sides(class: &CharClass) -> (bool, bool) {
if class.atoms.is_empty() {
return (true, true);
}
let mut word = false;
let mut nonword = false;
for atom in &class.atoms {
let (atom_word, atom_nonword) = atom_word_sides(atom);
word |= atom_word;
nonword |= atom_nonword;
if word && nonword {
break;
}
}
if class.negated {
let covers_word =
class.intersections.is_empty() && atoms_cover_all_ascii_word(&class.atoms);
let covers_nonword = class.intersections.is_empty()
&& class
.atoms
.iter()
.any(|atom| matches!(atom, ClassAtom::Perl(PerlClassKind::NotWord)));
(!covers_word, !covers_nonword)
} else {
(word, nonword)
}
}
fn atom_word_sides(atom: &ClassAtom) -> (bool, bool) {
match atom {
ClassAtom::Char(ch) => {
let word = is_unicode_word_char(*ch);
(word, !word)
}
ClassAtom::Range(start, end) => {
if start.is_ascii() && end.is_ascii() {
let (start, end) = (*start.min(end), *start.max(end));
let mut word = false;
let mut nonword = false;
for ch in start..=end {
if is_unicode_word_char(ch) {
word = true;
} else {
nonword = true;
}
if word && nonword {
break;
}
}
(word, nonword)
} else {
(true, true)
}
}
ClassAtom::Perl(kind) => match kind {
PerlClassKind::Digit => (true, false),
PerlClassKind::Word => (true, false),
PerlClassKind::HorizontalSpace => (true, false),
PerlClassKind::NotWord => (false, true),
PerlClassKind::Space | PerlClassKind::VerticalSpace => (false, true),
PerlClassKind::NotDigit
| PerlClassKind::NotSpace
| PerlClassKind::NotHorizontalSpace
| PerlClassKind::NotVerticalSpace
| PerlClassKind::NotNewline => (true, true),
},
ClassAtom::Posix { name, negated } => {
if *negated {
return (true, true);
}
match name.as_str() {
"alpha" | "alnum" | "digit" | "xdigit" | "upper" | "lower" | "word" => {
(true, false)
}
"space" | "blank" | "cntrl" => (false, true),
_ => (true, true),
}
}
ClassAtom::Unicode { .. } => (true, true),
ClassAtom::Nested(class) => class_word_sides(class),
}
}
#[cfg(test)]
mod tests {
use super::super::ast::parse;
use super::*;
fn mask(pattern: &str) -> u8 {
start_class_mask(&parse(pattern))
}
const MID_WORD: u8 = 0b1000;
const WORD_START: u8 = 0b0010;
const WORD_END: u8 = 0b0100;
const GAP: u8 = 0b0001;
#[test]
fn keyword_with_lookbehind_only_starts_at_word_starts() {
assert_eq!(mask(r"(?<!\w)this(?!\w)"), WORD_START);
assert_eq!(mask(r"\bwhile\b"), WORD_START);
}
#[test]
fn separator_prefixed_rules_exclude_mid_word() {
let separator =
r"((?:\s*+/\*(?:[^*]++|\*+(?!/))*+\*/\s*+)+|\s++|(?<=\W)|(?=\W)|^|\n?$|\A|\Z)";
assert_eq!(mask(&format!("{separator}(#)\\s*pragma\\b")) & MID_WORD, 0);
assert_eq!(
mask(&format!("{separator}((?<!\\w)this(?!\\w))")) & MID_WORD,
0
);
}
#[test]
fn identifier_patterns_allow_all_word_positions() {
assert_eq!(mask(r"[A-Za-z_]\w*"), WORD_START | MID_WORD);
assert_eq!(mask(r"\w+"), WORD_START | MID_WORD);
}
#[test]
fn punctuation_and_anchor_patterns() {
assert_eq!(mask(r"\{"), GAP | WORD_END);
assert_eq!(mask(r"^\s*#"), GAP);
assert_eq!(mask(r"$"), GAP | WORD_END);
assert_eq!(mask(r"\G\w"), WORD_START | MID_WORD);
}
#[test]
fn negated_lookbehind_with_extra_atoms_still_excludes_word_prev() {
assert_eq!(mask(r"(?<![\w$])if\b") & (MID_WORD | WORD_END), 0);
}
#[test]
fn conservative_constructs_keep_all_classes() {
assert_eq!(mask(r"(a)\1"), WORD_START | MID_WORD);
assert_eq!(mask(r"\1x"), START_CLASS_ALL);
assert_eq!(mask(r".*"), START_CLASS_ALL);
assert_eq!(mask(r"x|.|^"), START_CLASS_ALL);
}
#[test]
fn nullable_first_element_unions_with_following_element() {
assert_eq!(mask(r"\s*#") & (WORD_START | MID_WORD), 0);
assert_ne!(mask(r"\s+#") & CUR_NONWORD, 0);
}
#[test]
fn empty_mask_falls_back_to_all() {
assert_eq!(mask(r"\b\B"), START_CLASS_ALL);
}
}