use super::ast::*;
use super::lexer::{EscapeKind, TokenKind};
use super::state::Parser;
use crate::error::{Error, ErrorKind, Result};
use crate::hir::unicode_data;
use crate::hir::HORIZONTAL_WHITESPACE;
impl Parser<'_> {
pub(super) fn parse_class(&mut self) -> Result<Expr> {
let start_span = self.current.span;
self.lexer.set_in_class(true);
self.advance()?;
let negated = if matches!(self.current.kind, TokenKind::Caret) {
self.advance()?;
true
} else {
false
};
let mut ranges = Vec::new();
if matches!(self.current.kind, TokenKind::CloseBracket) {
ranges.push(ClassRange::single(']'));
self.advance()?;
} else if matches!(self.current.kind, TokenKind::Hyphen) && self.peek_set_op().is_none() {
ranges.push(ClassRange::single('-'));
self.advance()?;
}
ranges.extend(self.parse_class_term()?);
while let Some(op) = self.peek_set_op() {
let op_span = self.current.span;
if ranges.is_empty() {
return Err(Error::with_span(
ErrorKind::InvalidClassSetOp,
self.pattern,
op_span,
));
}
self.advance()?; self.advance()?;
let rhs = self.parse_class_term()?;
if rhs.is_empty() {
return Err(Error::with_span(
ErrorKind::InvalidClassSetOp,
self.pattern,
op_span,
));
}
ranges = op.apply(&ranges, &rhs);
}
self.lexer.set_in_class(false);
if matches!(self.current.kind, TokenKind::Eof) {
return Err(Error::with_span(
ErrorKind::UnmatchedOpenBracket,
self.pattern,
start_span,
));
}
self.advance()?;
if ranges.is_empty() {
return Err(Error::with_span(
ErrorKind::EmptyClass,
self.pattern,
start_span,
));
}
Ok(Expr::Class(Box::new(Class::new(ranges, negated))))
}
fn peek_set_op(&self) -> Option<SetOp> {
let next = self.pattern.get(self.current.span.end..)?.chars().next();
match (&self.current.kind, next) {
(TokenKind::Literal('&'), Some('&')) => Some(SetOp::Intersection),
(TokenKind::Hyphen, Some('-')) => Some(SetOp::Difference),
(TokenKind::Literal('~'), Some('~')) => Some(SetOp::SymmetricDifference),
_ => None,
}
}
fn parse_class_term(&mut self) -> Result<Vec<ClassRange>> {
let mut ranges = Vec::new();
while !matches!(self.current.kind, TokenKind::CloseBracket | TokenKind::Eof)
&& self.peek_set_op().is_none()
{
if matches!(self.current.kind, TokenKind::OpenBracket) {
if let Expr::Class(c) = self.parse_class()? {
if c.negated {
ranges.extend(complement_ranges(&c.ranges));
} else {
ranges.extend(c.ranges);
}
}
continue;
}
let item = self.parse_class_item()?;
match item {
ClassItem::Char(start_char) => {
if matches!(self.current.kind, TokenKind::Hyphen)
&& self.peek_set_op().is_none()
{
self.advance()?;
if matches!(self.current.kind, TokenKind::CloseBracket) {
ranges.push(ClassRange::single(start_char));
ranges.push(ClassRange::single('-'));
break;
}
let end_item = self.parse_class_item()?;
match end_item {
ClassItem::Char(end_char) => {
if start_char > end_char {
return Err(Error::with_span(
ErrorKind::InvalidClassRange {
start: start_char,
end: end_char,
},
self.pattern,
self.current.span,
));
}
ranges.push(ClassRange::new(start_char, end_char));
}
ClassItem::Ranges(r) => {
ranges.push(ClassRange::single(start_char));
ranges.push(ClassRange::single('-'));
ranges.extend(r);
}
ClassItem::UnicodeProperty { name, negated: _ } => {
ranges.push(ClassRange::single(start_char));
ranges.push(ClassRange::single('-'));
if let Some(code_point_ranges) = unicode_data::get_property(&name) {
for &(start, end) in code_point_ranges {
if (0xD800..=0xDFFF).contains(&start) {
continue;
}
let start = start.min(0x10FFFF);
let end = end.min(0x10FFFF);
if let (Some(s), Some(e)) =
(char::from_u32(start), char::from_u32(end))
{
ranges.push(ClassRange::new(s, e));
}
}
} else {
return Err(Error::with_span(
ErrorKind::UnknownUnicodeProperty(name.clone()),
self.pattern,
self.current.span,
));
}
}
}
} else {
ranges.push(ClassRange::single(start_char));
}
}
ClassItem::Ranges(r) => {
ranges.extend(r);
}
ClassItem::UnicodeProperty { name, negated } => {
if let Some(code_point_ranges) = unicode_data::get_property(&name) {
for &(start, end) in code_point_ranges {
if (0xD800..=0xDFFF).contains(&start) {
continue;
}
let start = start.min(0x10FFFF);
let end = end.min(0x10FFFF);
if start < 0xD800 && end > 0xDFFF {
if let (Some(s), Some(e)) =
(char::from_u32(start), char::from_u32(0xD7FF))
{
ranges.push(ClassRange::new(s, e));
}
if let (Some(s), Some(e)) =
(char::from_u32(0xE000), char::from_u32(end))
{
ranges.push(ClassRange::new(s, e));
}
} else if start <= 0xD7FF && (0xD800..=0xDFFF).contains(&end) {
if let (Some(s), Some(e)) =
(char::from_u32(start), char::from_u32(0xD7FF))
{
ranges.push(ClassRange::new(s, e));
}
} else if (0xD800..=0xDFFF).contains(&start) && end > 0xDFFF {
if let (Some(s), Some(e)) =
(char::from_u32(0xE000), char::from_u32(end))
{
ranges.push(ClassRange::new(s, e));
}
} else if let (Some(s), Some(e)) =
(char::from_u32(start), char::from_u32(end))
{
ranges.push(ClassRange::new(s, e));
}
}
if negated {
let count_to_drain = code_point_ranges
.iter()
.filter(|&&(start, _)| {
!(0xD800..=0xDFFF).contains(&start)
})
.map(|&(start, end)| {
let start = start.min(0x10FFFF);
let end = end.min(0x10FFFF);
if start < 0xD800 && end > 0xDFFF {
2
} else {
1
}
})
.sum::<usize>();
let drain_start = ranges.len().saturating_sub(count_to_drain);
let positive_ranges: Vec<ClassRange> =
ranges.drain(drain_start..).collect();
ranges.extend(complement_ranges(&positive_ranges));
}
} else {
return Err(Error::with_span(
ErrorKind::UnknownUnicodeProperty(name.clone()),
self.pattern,
self.current.span,
));
}
}
}
}
Ok(ranges)
}
fn parse_class_item(&mut self) -> Result<ClassItem> {
match &self.current.kind {
TokenKind::Escape(esc) => {
match esc {
EscapeKind::Digit => {
self.advance()?;
Ok(ClassItem::Ranges(vec![ClassRange::new('0', '9')]))
}
EscapeKind::NotDigit => {
self.advance()?;
Ok(ClassItem::Ranges(vec![
ClassRange::new('\x00', '/'), ClassRange::new(':', '\u{10FFFF}'), ]))
}
EscapeKind::Word => {
self.advance()?;
Ok(ClassItem::Ranges(vec![
ClassRange::new('a', 'z'),
ClassRange::new('A', 'Z'),
ClassRange::new('0', '9'),
ClassRange::single('_'),
]))
}
EscapeKind::NotWord => {
self.advance()?;
Ok(ClassItem::Ranges(vec![
ClassRange::new('\x00', '/'), ClassRange::new(':', '@'), ClassRange::new('[', '^'), ClassRange::single('`'), ClassRange::new('{', '\u{10FFFF}'), ]))
}
EscapeKind::Whitespace => {
self.advance()?;
Ok(ClassItem::Ranges(unicode_whitespace_ranges()))
}
EscapeKind::NotWhitespace => {
self.advance()?;
Ok(ClassItem::Ranges(complement_ranges(
&unicode_whitespace_ranges(),
)))
}
EscapeKind::HorizontalWhitespace => {
self.advance()?;
Ok(ClassItem::Ranges(horizontal_whitespace_ranges()))
}
EscapeKind::NotHorizontalWhitespace => {
self.advance()?;
Ok(ClassItem::Ranges(complement_ranges(
&horizontal_whitespace_ranges(),
)))
}
EscapeKind::Literal(c) => {
let c = *c;
self.advance()?;
Ok(ClassItem::Char(c))
}
EscapeKind::Newline => {
self.advance()?;
Ok(ClassItem::Char('\n'))
}
EscapeKind::CarriageReturn => {
self.advance()?;
Ok(ClassItem::Char('\r'))
}
EscapeKind::Tab => {
self.advance()?;
Ok(ClassItem::Char('\t'))
}
EscapeKind::FormFeed => {
self.advance()?;
Ok(ClassItem::Char('\x0C'))
}
EscapeKind::VerticalTab => {
self.advance()?;
Ok(ClassItem::Char('\x0B'))
}
EscapeKind::Null => {
self.advance()?;
Ok(ClassItem::Char('\0'))
}
EscapeKind::Hex(c) | EscapeKind::Unicode(c) => {
let c = *c;
self.advance()?;
Ok(ClassItem::Char(c))
}
EscapeKind::WordBoundary => {
self.advance()?;
Ok(ClassItem::Char('\u{8}'))
}
EscapeKind::UnicodeProperty(name) => {
let name = name.clone();
self.advance()?;
Ok(ClassItem::UnicodeProperty {
name,
negated: false,
})
}
EscapeKind::NotUnicodeProperty(name) => {
let name = name.clone();
self.advance()?;
Ok(ClassItem::UnicodeProperty {
name,
negated: true,
})
}
_ => Err(Error::with_span(
ErrorKind::EscapeNotAllowedInClass(self.current_text().to_string()),
self.pattern,
self.current.span,
)),
}
}
TokenKind::PosixClass { name, negated } => {
let name = name.clone();
let negated = *negated;
let ranges = posix_class_ranges(&name).ok_or_else(|| {
Error::with_span(
ErrorKind::UnknownPosixClass(name.clone()),
self.pattern,
self.current.span,
)
})?;
let ranges = if negated {
complement_ranges(ranges)
} else {
ranges.to_vec()
};
self.advance()?;
Ok(ClassItem::Ranges(ranges))
}
TokenKind::Literal(c) => {
let c = *c;
self.advance()?;
Ok(ClassItem::Char(c))
}
TokenKind::Caret => {
self.advance()?;
Ok(ClassItem::Char('^'))
}
_ => Err(Error::with_span(
ErrorKind::UnexpectedChar(self.current_char().unwrap_or('?')),
self.pattern,
self.current.span,
)),
}
}
}
enum ClassItem {
Char(char),
Ranges(Vec<ClassRange>),
UnicodeProperty { name: String, negated: bool },
}
fn unicode_whitespace_ranges() -> Vec<ClassRange> {
codepoints_to_ranges(unicode_data::PERL_SPACE)
}
fn horizontal_whitespace_ranges() -> Vec<ClassRange> {
codepoints_to_ranges(HORIZONTAL_WHITESPACE)
}
fn codepoints_to_ranges(table: &[(u32, u32)]) -> Vec<ClassRange> {
table
.iter()
.filter_map(|&(s, e)| Some(ClassRange::new(char::from_u32(s)?, char::from_u32(e)?)))
.collect()
}
fn posix_class_ranges(name: &str) -> Option<&'static [ClassRange]> {
const ALPHA: &[ClassRange] = &[
ClassRange {
start: 'A',
end: 'Z',
},
ClassRange {
start: 'a',
end: 'z',
},
];
const DIGIT: &[ClassRange] = &[ClassRange {
start: '0',
end: '9',
}];
const ALNUM: &[ClassRange] = &[
ClassRange {
start: 'A',
end: 'Z',
},
ClassRange {
start: 'a',
end: 'z',
},
ClassRange {
start: '0',
end: '9',
},
];
const UPPER: &[ClassRange] = &[ClassRange {
start: 'A',
end: 'Z',
}];
const LOWER: &[ClassRange] = &[ClassRange {
start: 'a',
end: 'z',
}];
const SPACE: &[ClassRange] = &[
ClassRange {
start: '\t',
end: '\r',
},
ClassRange {
start: ' ',
end: ' ',
},
];
const BLANK: &[ClassRange] = &[
ClassRange {
start: '\t',
end: '\t',
},
ClassRange {
start: ' ',
end: ' ',
},
];
const CNTRL: &[ClassRange] = &[
ClassRange {
start: '\x00',
end: '\x1F',
},
ClassRange {
start: '\x7F',
end: '\x7F',
},
];
const PRINT: &[ClassRange] = &[ClassRange {
start: '\x20',
end: '\x7E',
}];
const GRAPH: &[ClassRange] = &[ClassRange {
start: '\x21',
end: '\x7E',
}];
const PUNCT: &[ClassRange] = &[
ClassRange {
start: '\x21',
end: '\x2F',
},
ClassRange {
start: '\x3A',
end: '\x40',
},
ClassRange {
start: '\x5B',
end: '\x60',
},
ClassRange {
start: '\x7B',
end: '\x7E',
},
];
const XDIGIT: &[ClassRange] = &[
ClassRange {
start: '0',
end: '9',
},
ClassRange {
start: 'A',
end: 'F',
},
ClassRange {
start: 'a',
end: 'f',
},
];
const WORD: &[ClassRange] = &[
ClassRange {
start: 'A',
end: 'Z',
},
ClassRange {
start: 'a',
end: 'z',
},
ClassRange {
start: '0',
end: '9',
},
ClassRange {
start: '_',
end: '_',
},
];
const ASCII: &[ClassRange] = &[ClassRange {
start: '\x00',
end: '\x7F',
}];
match name {
"alpha" => Some(ALPHA),
"digit" => Some(DIGIT),
"alnum" => Some(ALNUM),
"upper" => Some(UPPER),
"lower" => Some(LOWER),
"space" => Some(SPACE),
"blank" => Some(BLANK),
"cntrl" => Some(CNTRL),
"print" => Some(PRINT),
"graph" => Some(GRAPH),
"punct" => Some(PUNCT),
"xdigit" => Some(XDIGIT),
"word" => Some(WORD),
"ascii" => Some(ASCII),
_ => None,
}
}
fn complement_ranges(ranges: &[ClassRange]) -> Vec<ClassRange> {
let mut pts: Vec<(u32, u32)> = ranges
.iter()
.map(|r| (r.start as u32, r.end as u32))
.collect();
pts.sort_by_key(|r| r.0);
let mut out = Vec::new();
let mut next = 0u32;
for (s, e) in pts {
if s > next {
push_scalar_range(&mut out, next, s - 1);
}
if e + 1 > next {
next = e + 1;
}
}
if next <= 0x10FFFF {
push_scalar_range(&mut out, next, 0x10FFFF);
}
out
}
fn push_scalar_range_checked(out: &mut Vec<ClassRange>, start: u32, end: u32) {
if let (Some(s), Some(e)) = (char::from_u32(start), char::from_u32(end)) {
out.push(ClassRange::new(s, e));
}
}
fn push_scalar_range(out: &mut Vec<ClassRange>, start: u32, end: u32) {
const SUR_LO: u32 = 0xD800;
const SUR_HI: u32 = 0xDFFF;
if start > end {
return;
}
if end < SUR_LO || start > SUR_HI {
push_scalar_range_checked(out, start, end);
} else {
if start < SUR_LO {
push_scalar_range_checked(out, start, SUR_LO - 1);
}
if end > SUR_HI {
push_scalar_range_checked(out, SUR_HI + 1, end);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum SetOp {
Intersection,
Difference,
SymmetricDifference,
}
impl SetOp {
fn apply(self, lhs: &[ClassRange], rhs: &[ClassRange]) -> Vec<ClassRange> {
match self {
SetOp::Intersection => intersect_ranges(lhs, rhs),
SetOp::Difference => difference_ranges(lhs, rhs),
SetOp::SymmetricDifference => symmetric_difference_ranges(lhs, rhs),
}
}
}
type Interval = (u32, u32);
fn sorted_merged(mut intervals: Vec<Interval>) -> Vec<Interval> {
intervals.sort_by_key(|r| r.0);
let mut merged: Vec<Interval> = Vec::new();
for (start, end) in intervals {
if let Some(last) = merged.last_mut() {
if start <= last.1.saturating_add(1) {
last.1 = last.1.max(end);
continue;
}
}
merged.push((start, end));
}
merged
}
fn to_intervals(ranges: &[ClassRange]) -> Vec<Interval> {
sorted_merged(
ranges
.iter()
.map(|r| (r.start as u32, r.end as u32))
.collect(),
)
}
fn intervals_to_ranges(intervals: &[Interval]) -> Vec<ClassRange> {
let mut out = Vec::new();
for &(start, end) in intervals {
push_scalar_range(&mut out, start, end);
}
out
}
fn intersect_intervals(a: &[Interval], b: &[Interval]) -> Vec<Interval> {
let mut out = Vec::new();
let mut i = 0;
let mut j = 0;
while i < a.len() && j < b.len() {
let (a_start, a_end) = a[i];
let (b_start, b_end) = b[j];
let start = a_start.max(b_start);
let end = a_end.min(b_end);
if start <= end {
out.push((start, end));
}
if a_end < b_end {
i += 1;
} else {
j += 1;
}
}
out
}
fn difference_intervals(a: &[Interval], b: &[Interval]) -> Vec<Interval> {
let mut out = Vec::new();
let mut j = 0usize;
for &(a_start, a_end) in a {
let mut cur = a_start;
while j < b.len() && b[j].1 < cur {
j += 1;
}
let mut k = j;
while cur <= a_end && k < b.len() && b[k].0 <= a_end {
let (b_start, b_end) = b[k];
if b_start > cur {
out.push((cur, b_start - 1));
}
if b_end >= a_end {
cur = a_end + 1;
break;
}
cur = b_end + 1;
k += 1;
}
if cur <= a_end {
out.push((cur, a_end));
}
j = k;
}
out
}
fn intersect_ranges(a: &[ClassRange], b: &[ClassRange]) -> Vec<ClassRange> {
let a = to_intervals(a);
let b = to_intervals(b);
intervals_to_ranges(&sorted_merged(intersect_intervals(&a, &b)))
}
fn difference_ranges(a: &[ClassRange], b: &[ClassRange]) -> Vec<ClassRange> {
let a = to_intervals(a);
let b = to_intervals(b);
intervals_to_ranges(&sorted_merged(difference_intervals(&a, &b)))
}
fn symmetric_difference_ranges(a: &[ClassRange], b: &[ClassRange]) -> Vec<ClassRange> {
let a = to_intervals(a);
let b = to_intervals(b);
let mut sym = difference_intervals(&a, &b);
sym.extend(difference_intervals(&b, &a));
intervals_to_ranges(&sorted_merged(sym))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn recognizes_all_14_posix_class_names() {
for name in [
"alpha", "digit", "alnum", "upper", "lower", "space", "blank", "cntrl", "print",
"graph", "punct", "xdigit", "word", "ascii",
] {
assert!(
posix_class_ranges(name).is_some(),
"{name} should be a recognized POSIX class"
);
}
}
#[test]
fn rejects_unknown_names() {
assert!(posix_class_ranges("bogus").is_none());
assert!(posix_class_ranges("").is_none());
assert!(posix_class_ranges("Alpha").is_none());
}
#[test]
fn alpha_is_ascii_letters_only() {
let ranges = posix_class_ranges("alpha").unwrap();
assert!(ranges.iter().any(|r| r.contains('a')));
assert!(ranges.iter().any(|r| r.contains('X')));
assert!(!ranges.iter().any(|r| r.contains('9')));
assert!(!ranges.iter().any(|r| r.contains('_')));
assert!(!ranges.iter().any(|r| r.contains('é')));
}
#[test]
fn digit_is_0_to_9() {
let ranges = posix_class_ranges("digit").unwrap();
assert!(ranges.iter().any(|r| r.contains('5')));
assert!(!ranges.iter().any(|r| r.contains('a')));
}
#[test]
fn space_matches_space_and_tab_but_not_a() {
let ranges = posix_class_ranges("space").unwrap();
assert!(ranges.iter().any(|r| r.contains(' ')));
assert!(ranges.iter().any(|r| r.contains('\t')));
assert!(!ranges.iter().any(|r| r.contains('a')));
}
#[test]
fn word_matches_underscore() {
let ranges = posix_class_ranges("word").unwrap();
assert!(ranges.iter().any(|r| r.contains('_')));
}
fn ranges_contain_cp(ranges: &[ClassRange], cp: u32) -> bool {
match char::from_u32(cp) {
Some(c) => ranges.iter().any(|r| r.contains(c)),
None => false,
}
}
#[test]
fn set_ops_match_naive_membership_oracle() {
const WINDOW: std::ops::RangeInclusive<u32> = 0u32..=0x2FF;
let cases: &[(&str, Vec<ClassRange>, Vec<ClassRange>)] = &[
(
"disjoint",
vec![ClassRange::new('a', 'f')],
vec![ClassRange::new('x', 'z')],
),
(
"b fully contains a",
vec![ClassRange::new('m', 'p')],
vec![ClassRange::new('a', 'z')],
),
(
"a fully contains b",
vec![ClassRange::new('a', 'z')],
vec![ClassRange::new('m', 'p')],
),
(
"partial overlap, a starts first",
vec![ClassRange::new('a', 'm')],
vec![ClassRange::new('g', 'z')],
),
(
"partial overlap, b starts first",
vec![ClassRange::new('g', 'z')],
vec![ClassRange::new('a', 'm')],
),
(
"identical sets",
vec![ClassRange::new('a', 'z'), ClassRange::new('0', '9')],
vec![ClassRange::new('a', 'z'), ClassRange::new('0', '9')],
),
("a empty", vec![], vec![ClassRange::new('a', 'z')]),
("b empty", vec![ClassRange::new('a', 'z')], vec![]),
("both empty", vec![], vec![]),
(
"adjacent but not overlapping",
vec![ClassRange::new('a', 'c')],
vec![ClassRange::new('d', 'f')],
),
(
"multiple holes punched in one a interval",
vec![ClassRange::new('\0', '\u{FF}')],
vec![
ClassRange::new('\u{10}', '\u{1F}'),
ClassRange::new('\u{40}', '\u{5F}'),
ClassRange::new('\u{90}', '\u{9F}'),
],
),
];
for (label, a, b) in cases {
let inter = intersect_ranges(a, b);
let diff = difference_ranges(a, b);
let sym = symmetric_difference_ranges(a, b);
for cp in WINDOW {
let in_a = ranges_contain_cp(a, cp);
let in_b = ranges_contain_cp(b, cp);
let expect_inter = in_a && in_b;
let expect_diff = in_a && !in_b;
let expect_sym = in_a != in_b;
assert_eq!(
ranges_contain_cp(&inter, cp),
expect_inter,
"{label}: intersect mismatch at U+{cp:04X}"
);
assert_eq!(
ranges_contain_cp(&diff, cp),
expect_diff,
"{label}: difference mismatch at U+{cp:04X}"
);
assert_eq!(
ranges_contain_cp(&sym, cp),
expect_sym,
"{label}: symmetric_difference mismatch at U+{cp:04X}"
);
}
}
}
}