use super::ast::{
AnchorKind, Ast, CharClass, ClassAtom, LookKind, ParsedRegex, PerlClassKind, RegexFlags,
};
use super::is_unicode_word_char;
pub(crate) const START_CLASS_ALL: u8 = 0xff;
const NIBBLE_ALL: u8 = 0b1111;
pub(crate) const UNICODE_SHIFT: u8 = 4;
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;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Mode {
Ascii,
Unicode { case_insensitive: bool },
}
impl Mode {
fn with_flags(self, flags: RegexFlags) -> Self {
match self {
Self::Ascii => Self::Ascii,
Self::Unicode { .. } => Self::Unicode {
case_insensitive: flags.case_insensitive,
},
}
}
fn is_unicode(self) -> bool {
matches!(self, Self::Unicode { .. })
}
}
pub(crate) fn start_class_mask(parsed: &ParsedRegex) -> u8 {
let unicode = Mode::Unicode {
case_insensitive: parsed.flags.case_insensitive,
};
nibble_mask(parsed, Mode::Ascii) | (nibble_mask(parsed, unicode) << UNICODE_SHIFT)
}
fn scalar_word_sides(ch: char, mode: Mode) -> (bool, bool) {
if mode
== (Mode::Unicode {
case_insensitive: true,
})
&& !ch.is_ascii()
{
return (true, true);
}
let word = is_unicode_word_char(ch);
(word, !word)
}
fn nibble_mask(parsed: &ParsedRegex, mode: Mode) -> u8 {
let (mask, continuation) = node_mask(&parsed.ast, NIBBLE_ALL, mode);
let mask = mask | continuation.unwrap_or(0);
if mask == 0 {
NIBBLE_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, mode: Mode) -> (u8, Option<u8>) {
match ast {
Ast::Empty => (0, Some(constraint)),
Ast::Literal(literal) => match literal.chars().next() {
Some(ch) => {
let (word, nonword) = scalar_word_sides(ch, mode);
(constraint & sides_cur_bits(word, nonword), None)
}
None => (0, Some(constraint)),
},
Ast::Class(class) => {
let (word, nonword) = class_word_sides(class, mode);
(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, mode);
if sides.nullable {
constraint
} else {
constraint & sides_cur_bits(sides.word, sides.nonword)
}
}
LookKind::NotAhead => {
if negated_look_excludes_word(child, false, mode) {
constraint & bits_cur(false)
} else {
constraint
}
}
LookKind::Behind => match behind_prev_bits(child, mode) {
Some(bits) => constraint & bits,
None => constraint,
},
LookKind::NotBehind => {
if negated_look_excludes_word(child, true, mode) {
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, mode);
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, mode);
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, mode);
if *min == 0 {
(mask, Some(constraint | continuation.unwrap_or(0)))
} else {
(mask, continuation)
}
}
Ast::Group { child, .. } => node_mask(child, constraint, mode),
Ast::Flags { flags, child } => node_mask(child, constraint, mode.with_flags(*flags)),
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 behind_prev_bits(child: &Ast, mode: Mode) -> Option<u8> {
match child {
Ast::Anchor(AnchorKind::LineStart | AnchorKind::TextStart) => Some(bits_prev(false)),
Ast::Group { child, .. } => behind_prev_bits(child, mode),
Ast::Flags { flags, child } => behind_prev_bits(child, mode.with_flags(*flags)),
Ast::Alternation(branches) => branches.iter().try_fold(0, |bits, branch| {
Some(bits | behind_prev_bits(branch, mode)?)
}),
_ => {
let sides = last_char_sides(child, mode);
(!sides.nullable).then(|| sides_prev_bits(sides.word, sides.nonword))
}
}
}
fn first_char_sides(ast: &Ast, mode: Mode) -> CharSides {
char_sides(ast, false, mode)
}
fn last_char_sides(ast: &Ast, mode: Mode) -> CharSides {
char_sides(ast, true, mode)
}
fn char_sides(ast: &Ast, from_end: bool, mode: Mode) -> 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, nonword) = scalar_word_sides(ch, mode);
CharSides {
word,
nonword,
nullable: false,
}
}
None => CharSides::ZERO_WIDTH,
}
}
Ast::Class(class) => {
let (word, nonword) = class_word_sides(class, mode);
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, mode);
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, mode);
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, mode);
sides.nullable |= *min == 0;
sides
}
Ast::Group { child, .. } => char_sides(child, from_end, mode),
Ast::Flags { flags, child } => char_sides(child, from_end, mode.with_flags(*flags)),
Ast::Backref(_) | Ast::Conditional { .. } | Ast::Subroutine(_) | Ast::Unsupported(_) => {
CharSides::UNKNOWN
}
}
}
fn negated_look_excludes_word(child: &Ast, from_end: bool, mode: Mode) -> bool {
match child {
Ast::Class(class) => class_covers_all_word(class, mode),
Ast::Group { child, .. } => negated_look_excludes_word(child, from_end, mode),
Ast::Flags { flags, child } => {
negated_look_excludes_word(child, from_end, mode.with_flags(*flags))
}
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, mode)
&& rest.iter().all(matches_empty_unconditionally)
}
Ast::Alternation(branches) => branches
.iter()
.any(|branch| negated_look_excludes_word(branch, from_end, mode)),
Ast::Repeat { node, min, max, .. } => {
*min <= 1
&& max.is_none_or(|max| max >= 1)
&& negated_look_excludes_word(node, from_end, mode)
}
_ => 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_word(class: &CharClass, mode: Mode) -> bool {
if class.negated || !class.intersections.is_empty() {
return false;
}
atoms_cover_all_word(&class.atoms, mode)
}
fn atoms_cover_all_word(atoms: &[ClassAtom], mode: Mode) -> bool {
match mode {
Mode::Ascii => atoms_cover_all_ascii_word(atoms),
Mode::Unicode { .. } => atoms.iter().any(|atom| match atom {
ClassAtom::Perl(PerlClassKind::Word) => true,
ClassAtom::Posix {
name,
negated: false,
} => name == "word",
_ => false,
}),
}
}
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, mode: Mode) -> (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, mode);
word |= atom_word;
nonword |= atom_nonword;
if word && nonword {
break;
}
}
if class.negated {
let covers_word =
class.intersections.is_empty() && atoms_cover_all_word(&class.atoms, mode);
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, mode: Mode) -> (bool, bool) {
match atom {
ClassAtom::Char(ch) => scalar_word_sides(*ch, mode),
ClassAtom::Range(start, end) if mode.is_unicode() => unicode_range_word_sides(*start, *end),
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);
}
if mode.is_unicode() {
return match name.as_str() {
"digit" | "xdigit" | "word" => (true, false),
"space" | "blank" | "cntrl" => (false, true),
_ => (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, mode),
}
}
fn unicode_range_word_sides(start: char, end: char) -> (bool, bool) {
if !start.is_ascii() || !end.is_ascii() {
return (true, true);
}
let mut word = false;
let mut nonword = false;
for ch in start.min(end)..=start.max(end) {
if is_unicode_word_char(ch) {
word = true;
} else {
nonword = true;
}
}
(word, nonword)
}
#[cfg(test)]
mod tests {
use super::super::ast::parse;
use super::*;
fn mask(pattern: &str) -> u8 {
start_class_mask(&parse(pattern)) & NIBBLE_ALL
}
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"), NIBBLE_ALL);
assert_eq!(mask(r".*"), NIBBLE_ALL);
assert_eq!(mask(r"x|.|^"), NIBBLE_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"), NIBBLE_ALL);
assert_eq!(unicode_mask(r"\b\B"), NIBBLE_ALL);
}
fn unicode_mask(pattern: &str) -> u8 {
start_class_mask(&parse(pattern)) >> UNICODE_SHIFT
}
#[test]
fn line_start_lookbehind_branch_forces_nonword_previous() {
let pattern = r"(?i:(?<=[^.а-яё\w]|^)(Если|If)(?=[^.а-яё\w]|$))";
assert_eq!(mask(pattern), WORD_START);
assert_eq!(unicode_mask(pattern), GAP | WORD_START);
assert_eq!(mask(r"(?<=\W|x?)if") & MID_WORD, MID_WORD);
}
#[test]
fn unicode_masks_drop_ascii_only_claims() {
assert_eq!(mask(r"(?<=[^A-Za-z0-9_])x") & (MID_WORD | WORD_END), 0);
assert_eq!(unicode_mask(r"(?<=[^A-Za-z0-9_])x") & MID_WORD, MID_WORD);
assert_eq!(unicode_mask(r"(?<![A-Za-z0-9_])x") & MID_WORD, MID_WORD);
assert_eq!(mask(r"[[:alpha:]]") & (GAP | WORD_END), 0);
assert_ne!(unicode_mask(r"[[:alpha:]]") & (GAP | WORD_END), 0);
assert_eq!(unicode_mask(r"(?<!\w)this(?!\w)"), WORD_START);
assert_eq!(unicode_mask(r"\bwhile\b"), WORD_START);
assert_eq!(unicode_mask(r"(?i)[a-z]+") & (GAP | WORD_END), 0);
}
#[test]
fn ascii_case_partners_share_word_ness() {
for ch in (0..=0x10_ffff).filter_map(char::from_u32) {
let word = is_unicode_word_char(ch);
for head in [ch.to_lowercase().next(), ch.to_uppercase().next()]
.into_iter()
.flatten()
.filter(char::is_ascii)
{
assert_eq!(is_unicode_word_char(head), word, "{ch:?} -> {head:?}");
}
}
}
#[test]
fn case_insensitive_non_ascii_scalars_stay_conservative() {
assert_eq!(unicode_mask(r"(?i)ꟓ"), NIBBLE_ALL);
assert_eq!(unicode_mask(r"(?<=(?i:ꟓ))x"), WORD_START | MID_WORD);
assert_eq!(unicode_mask(r"(?<=ꟓ)x"), MID_WORD);
assert_eq!(unicode_mask(r"(?i)(?<!k)x") & MID_WORD, MID_WORD);
assert_eq!(unicode_mask(r"(?i)k"), WORD_START | MID_WORD);
}
}