1use std::collections::BTreeMap;
22
23use thiserror::Error;
24
25use crate::atn::parser_atn::ParserAtn;
26use crate::recognizer::{Recognizer, RecognizerData};
27use crate::token::{Token, TokenId, TokenSink, TokenSource, TokenSpec, TokenStoreError};
28use crate::tree::{Node, NodeKind};
29use crate::{BaseParser, CommonTokenStream, TOKEN_EOF};
30
31const MATCH_STACK_RED_ZONE: usize = 1024 * 1024;
32const MATCH_STACK_SIZE: usize = 4 * 1024 * 1024;
33
34#[derive(Clone, Copy, Debug, Eq, PartialEq)]
36enum TagKind {
37 Rule { rule_index: usize, bypass_type: i32 },
41 Token { token_type: i32 },
43}
44
45#[derive(Clone, Debug, Eq, PartialEq)]
53struct TagInfo {
54 kind: TagKind,
55 name: String,
57 label: Option<String>,
59}
60
61impl TagInfo {
62 fn label_keys(&self) -> impl Iterator<Item = &str> {
67 std::iter::once(self.name.as_str()).chain(self.label.as_deref())
68 }
69}
70
71#[derive(Clone, Debug, Eq, PartialEq)]
73enum Chunk {
74 Text(String),
76 Tag { name: String, label: Option<String> },
78}
79
80#[derive(Clone, Debug, Eq, Error, PartialEq)]
82pub enum ParseTreePatternError {
83 #[error("unterminated tag in pattern: {pattern}")]
85 UnterminatedTag { pattern: String },
86 #[error("missing start tag in pattern: {pattern}")]
88 MissingStartTag { pattern: String },
89 #[error("tag delimiters out of order in pattern: {pattern}")]
91 DelimitersOutOfOrder { pattern: String },
92 #[error("empty tag in pattern: {pattern}")]
94 EmptyTag { pattern: String },
95 #[error("unknown token {name} in pattern: {pattern}")]
97 UnknownToken { name: String, pattern: String },
98 #[error("unknown rule {name} in pattern: {pattern}")]
100 UnknownRule { name: String, pattern: String },
101 #[error("invalid tag {tag} in pattern: {pattern}")]
104 InvalidTag { tag: String, pattern: String },
105 #[error("could not tokenize pattern chunk: {message}")]
107 Tokenization { message: String },
108 #[error("start rule did not consume the full pattern: {pattern}")]
110 StartRuleDoesNotConsumeFullPattern { pattern: String },
111 #[error("could not interpret pattern as rule {rule_index}: {message}")]
113 CannotInvokeStartRule { rule_index: usize, message: String },
114 #[error("could not build rule-bypass ATN: {message}")]
116 BypassAtn { message: String },
117 #[error("{which} delimiter cannot be empty")]
120 EmptyDelimiter { which: &'static str },
121}
122
123#[derive(Clone, Debug, Eq, PartialEq)]
129struct Delimiters {
130 start: String,
131 stop: String,
132 escape: String,
133}
134
135impl Default for Delimiters {
136 fn default() -> Self {
137 Self {
138 start: "<".to_owned(),
139 stop: ">".to_owned(),
140 escape: "\\".to_owned(),
141 }
142 }
143}
144
145fn split(pattern: &str, delimiters: &Delimiters) -> Result<Vec<Chunk>, ParseTreePatternError> {
153 let chars: Vec<char> = pattern.chars().collect();
154 let start: Vec<char> = delimiters.start.chars().collect();
155 let stop: Vec<char> = delimiters.stop.chars().collect();
156 let escape: Vec<char> = delimiters.escape.chars().collect();
157
158 let matches_at = |at: usize, needle: &[char]| -> bool {
159 !needle.is_empty() && chars[at..].starts_with(needle)
160 };
161
162 let mut starts = Vec::new();
164 let mut stops = Vec::new();
165 let mut position = 0;
166 while position < chars.len() {
167 if matches_at(position, &escape) && matches_at(position + escape.len(), &start) {
168 position += escape.len() + start.len();
169 } else if matches_at(position, &escape) && matches_at(position + escape.len(), &stop) {
170 position += escape.len() + stop.len();
171 } else if matches_at(position, &start) {
172 starts.push(position);
173 position += start.len();
174 } else if matches_at(position, &stop) {
175 stops.push(position);
176 position += stop.len();
177 } else {
178 position += 1;
179 }
180 }
181
182 if starts.len() > stops.len() {
183 return Err(ParseTreePatternError::UnterminatedTag {
184 pattern: pattern.to_owned(),
185 });
186 }
187 if starts.len() < stops.len() {
188 return Err(ParseTreePatternError::MissingStartTag {
189 pattern: pattern.to_owned(),
190 });
191 }
192 for (open, close) in starts.iter().zip(&stops) {
193 if open >= close {
194 return Err(ParseTreePatternError::DelimitersOutOfOrder {
195 pattern: pattern.to_owned(),
196 });
197 }
198 }
199 for (close, next_open) in stops.iter().zip(starts.iter().skip(1)) {
204 if close + stop.len() > *next_open {
205 return Err(ParseTreePatternError::DelimitersOutOfOrder {
206 pattern: pattern.to_owned(),
207 });
208 }
209 }
210
211 let slice = |from: usize, to: usize| -> String { chars[from..to].iter().collect() };
212
213 let ntags = starts.len();
215 let mut chunks = Vec::new();
216 if ntags == 0 {
217 chunks.push(Chunk::Text(slice(0, chars.len())));
218 } else if starts[0] > 0 {
219 chunks.push(Chunk::Text(slice(0, starts[0])));
220 }
221 for index in 0..ntags {
222 let tag = slice(starts[index] + start.len(), stops[index]);
223 chunks.push(parse_tag(&tag, pattern)?);
224 if index + 1 < ntags {
225 chunks.push(Chunk::Text(slice(
226 stops[index] + stop.len(),
227 starts[index + 1],
228 )));
229 }
230 }
231 if ntags > 0 {
232 let after_last = stops[ntags - 1] + stop.len();
233 if after_last < chars.len() {
234 chunks.push(Chunk::Text(slice(after_last, chars.len())));
235 }
236 }
237
238 if !delimiters.escape.is_empty() {
240 for chunk in &mut chunks {
241 if let Chunk::Text(text) = chunk {
242 *text = strip_escape(text, &delimiters.escape);
243 }
244 }
245 }
246
247 Ok(chunks)
248}
249
250fn strip_escape(text: &str, escape: &str) -> String {
253 let mut out = String::with_capacity(text.len());
254 let mut rest = text;
255 while let Some(at) = rest.find(escape) {
256 out.push_str(&rest[..at]);
257 rest = &rest[at + escape.len()..];
258 }
259 out.push_str(rest);
260 out
261}
262
263fn parse_tag(tag: &str, pattern: &str) -> Result<Chunk, ParseTreePatternError> {
266 let (label, name) = tag.find(':').map_or((None, tag), |colon| {
267 (Some(tag[..colon].to_owned()), &tag[colon + 1..])
268 });
269 if name.is_empty() {
270 return Err(ParseTreePatternError::EmptyTag {
271 pattern: pattern.to_owned(),
272 });
273 }
274 Ok(Chunk::Tag {
275 name: name.to_owned(),
276 label,
277 })
278}
279
280pub trait PatternLexer {
291 fn tokenize_chunk(&mut self, text: &str) -> Result<Vec<TokenSpec>, ParseTreePatternError>;
298}
299
300impl<F> PatternLexer for F
301where
302 F: FnMut(&str) -> Result<Vec<TokenSpec>, ParseTreePatternError>,
303{
304 fn tokenize_chunk(&mut self, text: &str) -> Result<Vec<TokenSpec>, ParseTreePatternError> {
305 self(text)
306 }
307}
308
309pub fn lex_pattern_chunk<L>(
327 text: &str,
328 make_lexer: impl FnOnce(crate::InputStream) -> L,
329) -> Result<Vec<TokenSpec>, ParseTreePatternError>
330where
331 L: TokenSource,
332{
333 let lexer = make_lexer(crate::InputStream::new(text));
334 let mut stream =
335 CommonTokenStream::try_new(lexer).map_err(|error| ParseTreePatternError::Tokenization {
336 message: error.to_string(),
337 })?;
338 stream.fill();
339 if let Some(error) = stream.drain_source_errors().into_iter().next() {
340 return Err(ParseTreePatternError::Tokenization {
341 message: format!("lexer error at {}:{}", error.line, error.column),
342 });
343 }
344 Ok(stream
345 .tokens()
346 .filter(|token| token.token_type() != TOKEN_EOF)
347 .map(|token| {
348 TokenSpec::explicit(token.token_type(), token.text_or_empty())
349 .with_channel(token.channel())
350 })
351 .collect())
352}
353
354#[derive(Debug)]
365pub struct ParseTreePatternMatcher<'a> {
366 bypass_atn: ParserAtn,
367 data: &'a RecognizerData,
368 delimiters: Delimiters,
369}
370
371impl<'a> ParseTreePatternMatcher<'a> {
372 pub fn new(atn: &ParserAtn, data: &'a RecognizerData) -> Result<Self, ParseTreePatternError> {
383 let bypass_atn =
384 atn.with_bypass_alternatives()
385 .map_err(|error| ParseTreePatternError::BypassAtn {
386 message: error.to_string(),
387 })?;
388 Ok(Self {
389 bypass_atn,
390 data,
391 delimiters: Delimiters::default(),
392 })
393 }
394
395 pub fn set_delimiters(
407 &mut self,
408 start: impl Into<String>,
409 stop: impl Into<String>,
410 escape: impl Into<String>,
411 ) -> Result<(), ParseTreePatternError> {
412 let start = start.into();
413 let stop = stop.into();
414 if start.is_empty() {
415 return Err(ParseTreePatternError::EmptyDelimiter { which: "start" });
416 }
417 if stop.is_empty() {
418 return Err(ParseTreePatternError::EmptyDelimiter { which: "stop" });
419 }
420 self.delimiters = Delimiters {
421 start,
422 stop,
423 escape: escape.into(),
424 };
425 Ok(())
426 }
427
428 pub fn compile(
441 &self,
442 pattern: &str,
443 rule_index: usize,
444 lexer: impl PatternLexer,
445 ) -> Result<ParseTreePattern, ParseTreePatternError> {
446 let chunks = split(pattern, &self.delimiters)?;
447 let (specs, tags_by_index) = self.tokenize(&chunks, pattern, lexer)?;
448 let tree = self.interpret(specs, &tags_by_index, rule_index, pattern)?;
449 Ok(ParseTreePattern {
450 pattern: pattern.to_owned(),
451 pattern_rule_index: rule_index,
452 tree,
453 })
454 }
455
456 fn tokenize(
460 &self,
461 chunks: &[Chunk],
462 pattern: &str,
463 mut lexer: impl PatternLexer,
464 ) -> Result<(Vec<TokenSpec>, BTreeMap<usize, TagInfo>), ParseTreePatternError> {
465 let mut specs = Vec::new();
466 let mut tags_by_index = BTreeMap::new();
467 for chunk in chunks {
468 match chunk {
469 Chunk::Tag { name, label } => {
470 let (spec, tag) = self.tag_token(name, label.clone(), pattern)?;
471 tags_by_index.insert(specs.len(), tag);
472 specs.push(spec);
473 }
474 Chunk::Text(text) => {
475 specs.extend(lexer.tokenize_chunk(text)?);
476 }
477 }
478 }
479 if let Some(at) = specs
486 .iter()
487 .position(|spec| spec.token_type == TOKEN_EOF)
488 .filter(|at| at + 1 < specs.len())
489 {
490 return Err(ParseTreePatternError::Tokenization {
491 message: format!(
492 "EOF at pattern token {at} terminates the stream; {} following token(s) \
493 would be ignored",
494 specs.len() - at - 1
495 ),
496 });
497 }
498 Ok((specs, tags_by_index))
499 }
500
501 fn tag_token(
507 &self,
508 name: &str,
509 label: Option<String>,
510 pattern: &str,
511 ) -> Result<(TokenSpec, TagInfo), ParseTreePatternError> {
512 let display = tag_display(name, label.as_deref());
513 let first = name
514 .chars()
515 .next()
516 .ok_or_else(|| ParseTreePatternError::InvalidTag {
517 tag: name.to_owned(),
518 pattern: pattern.to_owned(),
519 })?;
520 if first.is_uppercase() {
521 let token_type = self.data.vocabulary().token_type(name).ok_or_else(|| {
522 ParseTreePatternError::UnknownToken {
523 name: name.to_owned(),
524 pattern: pattern.to_owned(),
525 }
526 })?;
527 let spec = TokenSpec::explicit(token_type, display);
528 let tag = TagInfo {
529 kind: TagKind::Token { token_type },
530 name: name.to_owned(),
531 label,
532 };
533 Ok((spec, tag))
534 } else if first.is_lowercase() {
535 let rule_index =
536 self.rule_index(name)
537 .ok_or_else(|| ParseTreePatternError::UnknownRule {
538 name: name.to_owned(),
539 pattern: pattern.to_owned(),
540 })?;
541 let bypass_type = self
544 .bypass_atn
545 .bypass_token_type(rule_index)
546 .map_err(|error| ParseTreePatternError::BypassAtn {
547 message: error.to_string(),
548 })?;
549 let spec = TokenSpec::explicit(bypass_type, display);
550 let tag = TagInfo {
551 kind: TagKind::Rule {
552 rule_index,
553 bypass_type,
554 },
555 name: name.to_owned(),
556 label,
557 };
558 Ok((spec, tag))
559 } else {
560 Err(ParseTreePatternError::InvalidTag {
561 tag: name.to_owned(),
562 pattern: pattern.to_owned(),
563 })
564 }
565 }
566
567 fn rule_index(&self, name: &str) -> Option<usize> {
570 self.data.rule_names().iter().rposition(|rule| rule == name)
571 }
572
573 fn interpret(
576 &self,
577 specs: Vec<TokenSpec>,
578 tags_by_index: &BTreeMap<usize, TagInfo>,
579 rule_index: usize,
580 pattern: &str,
581 ) -> Result<PatternTree, ParseTreePatternError> {
582 let trailing_eof = specs
583 .last()
584 .is_some_and(|spec| spec.token_type == TOKEN_EOF);
585 let source = PatternTokenSource { specs, index: 0 };
586 let mut parser = BaseParser::new(CommonTokenStream::new(source), self.data.clone());
587 parser.remove_error_listeners();
594 let root = parser
595 .parse_atn_rule(&self.bypass_atn, rule_index)
596 .map_err(|error| ParseTreePatternError::CannotInvokeStartRule {
597 rule_index,
598 message: error.to_string(),
599 })?;
600 if parser.number_of_syntax_errors() > 0 {
601 return Err(ParseTreePatternError::CannotInvokeStartRule {
602 rule_index,
603 message: format!(
604 "pattern is not valid for the rule: {} syntax error(s) during pattern parse",
605 parser.number_of_syntax_errors()
606 ),
607 });
608 }
609
610 if parser.token_stream().la_token(1) != TOKEN_EOF {
613 return Err(ParseTreePatternError::StartRuleDoesNotConsumeFullPattern {
614 pattern: pattern.to_owned(),
615 });
616 }
617
618 let file = parser.into_parsed_file(root);
619 if trailing_eof
625 && !file.tree().descendants().any(|node| {
626 node.as_terminal()
627 .is_some_and(|terminal| terminal.symbol().token_type() == TOKEN_EOF)
628 })
629 {
630 return Err(ParseTreePatternError::StartRuleDoesNotConsumeFullPattern {
631 pattern: pattern.to_owned(),
632 });
633 }
634 let tags = rekey_tags_by_token_id(tags_by_index);
635 Ok(PatternTree { file, tags })
636 }
637}
638
639#[derive(Debug)]
642struct PatternTokenSource {
643 specs: Vec<TokenSpec>,
644 index: usize,
645}
646
647impl TokenSource for PatternTokenSource {
648 fn next_token(&mut self, sink: &mut TokenSink<'_>) -> Result<TokenId, TokenStoreError> {
649 let spec = self
650 .specs
651 .get(self.index)
652 .cloned()
653 .unwrap_or_else(|| TokenSpec::eof(self.index, self.index, 1, self.index));
654 self.index += 1;
655 sink.push(spec)
656 }
657
658 fn line(&self) -> usize {
659 1
660 }
661
662 fn column(&self) -> usize {
663 self.index
664 }
665
666 fn source_name(&self) -> &'static str {
667 "tree-pattern"
668 }
669}
670
671fn rekey_tags_by_token_id(tags_by_index: &BTreeMap<usize, TagInfo>) -> BTreeMap<TokenId, TagInfo> {
678 tags_by_index
679 .iter()
680 .filter_map(|(&index, tag)| Some((TokenId::try_from(index).ok()?, tag.clone())))
681 .collect()
682}
683
684fn tag_display(name: &str, label: Option<&str>) -> String {
686 label.map_or_else(|| format!("<{name}>"), |label| format!("<{label}:{name}>"))
687}
688
689#[derive(Debug)]
696struct PatternTree {
697 file: crate::tree::ParsedFile,
698 tags: BTreeMap<TokenId, TagInfo>,
699}
700
701#[derive(Debug)]
707pub struct ParseTreePattern {
708 pattern: String,
709 pattern_rule_index: usize,
710 tree: PatternTree,
711}
712
713impl ParseTreePattern {
714 #[must_use]
716 pub fn pattern(&self) -> &str {
717 &self.pattern
718 }
719
720 #[must_use]
722 pub const fn pattern_rule_index(&self) -> usize {
723 self.pattern_rule_index
724 }
725
726 #[must_use]
734 pub fn pattern_tree(&self) -> Node<'_> {
735 self.tree.file.tree()
736 }
737
738 #[must_use]
741 pub fn match_tree<'subject>(&self, tree: Node<'subject>) -> ParseTreeMatch<'subject> {
742 let mut labels: BTreeMap<String, Vec<Node<'subject>>> = BTreeMap::new();
743 let pattern_root = self.tree.file.tree();
744 let mismatched = match_impl(tree, pattern_root, &self.tree.tags, &mut labels);
745 ParseTreeMatch {
746 tree,
747 labels,
748 mismatched_node: mismatched,
749 }
750 }
751
752 #[must_use]
754 pub fn matches(&self, tree: Node<'_>) -> bool {
755 self.match_tree(tree).succeeded()
756 }
757
758 pub fn find_all<'subject, R>(
771 &self,
772 tree: Node<'subject>,
773 xpath: &str,
774 recognizer: &R,
775 ) -> Result<Vec<ParseTreeMatch<'subject>>, crate::XPathError>
776 where
777 R: Recognizer + ?Sized,
778 {
779 Ok(crate::XPath::find_all(tree, xpath, recognizer)?
780 .into_iter()
781 .map(|subtree| self.match_tree(subtree))
782 .filter(ParseTreeMatch::succeeded)
783 .collect())
784 }
785}
786
787#[derive(Clone, Debug)]
792pub struct ParseTreeMatch<'subject> {
793 tree: Node<'subject>,
794 labels: BTreeMap<String, Vec<Node<'subject>>>,
795 mismatched_node: Option<Node<'subject>>,
796}
797
798impl<'subject> ParseTreeMatch<'subject> {
799 #[must_use]
801 pub const fn succeeded(&self) -> bool {
802 self.mismatched_node.is_none()
803 }
804
805 #[must_use]
807 pub const fn tree(&self) -> Node<'subject> {
808 self.tree
809 }
810
811 #[must_use]
813 pub const fn mismatched_node(&self) -> Option<Node<'subject>> {
814 self.mismatched_node
815 }
816
817 #[must_use]
822 pub fn get(&self, label: &str) -> Option<Node<'subject>> {
823 self.labels
824 .get(label)
825 .and_then(|nodes| nodes.last().copied())
826 }
827
828 #[must_use]
830 pub fn get_all(&self, label: &str) -> &[Node<'subject>] {
831 self.labels.get(label).map_or(&[], Vec::as_slice)
832 }
833
834 #[must_use]
836 pub const fn labels(&self) -> &BTreeMap<String, Vec<Node<'subject>>> {
837 &self.labels
838 }
839}
840
841fn match_impl<'subject>(
853 tree: Node<'subject>,
854 pattern: Node<'_>,
855 tags: &BTreeMap<TokenId, TagInfo>,
856 labels: &mut BTreeMap<String, Vec<Node<'subject>>>,
857) -> Option<Node<'subject>> {
858 stacker::maybe_grow(MATCH_STACK_RED_ZONE, MATCH_STACK_SIZE, || {
862 match (leaf_kind(tree), leaf_kind(pattern)) {
863 (Some(_), Some(_)) => match_terminals(tree, pattern, tags, labels),
864 (None, None) => match_rules(tree, pattern, tags, labels),
865 _ => Some(tree),
867 }
868 })
869}
870
871fn leaf_kind(node: Node<'_>) -> Option<i32> {
875 match node.kind() {
876 NodeKind::Terminal => node.as_terminal().map(|t| t.symbol().token_type()),
877 NodeKind::Error => node.as_error().map(|e| e.symbol().token_type()),
878 NodeKind::Rule => None,
879 }
880}
881
882fn match_terminals<'subject>(
883 tree: Node<'subject>,
884 pattern: Node<'_>,
885 tags: &BTreeMap<TokenId, TagInfo>,
886 labels: &mut BTreeMap<String, Vec<Node<'subject>>>,
887) -> Option<Node<'subject>> {
888 let tree_type = leaf_kind(tree);
889 let pattern_type = leaf_kind(pattern);
890 if tree_type != pattern_type {
891 return Some(tree);
892 }
893 match pattern_token_tag(pattern, tags) {
895 Some(tag) => {
896 bind(labels, tag, tree);
897 None
898 }
899 None if leaf_text(tree) == leaf_text(pattern) => None,
900 None => Some(tree),
901 }
902}
903
904fn match_rules<'subject>(
905 tree: Node<'subject>,
906 pattern: Node<'_>,
907 tags: &BTreeMap<TokenId, TagInfo>,
908 labels: &mut BTreeMap<String, Vec<Node<'subject>>>,
909) -> Option<Node<'subject>> {
910 let (Some(tree_rule), Some(pattern_rule)) = (tree.as_rule(), pattern.as_rule()) else {
913 return Some(tree);
914 };
915
916 if let Some((tag_rule_index, tag)) = rule_tag_of(pattern, tags) {
918 return if tree_rule.rule_index() == tag_rule_index {
919 bind(labels, tag, tree);
920 None
921 } else {
922 Some(tree)
923 };
924 }
925
926 if tree_rule.child_count() != pattern_rule.child_count() {
927 return Some(tree);
928 }
929 for (tree_child, pattern_child) in tree.children().zip(pattern.children()) {
930 if let Some(mismatch) = match_impl(tree_child, pattern_child, tags, labels) {
931 return Some(mismatch);
932 }
933 }
934 None
935}
936
937fn pattern_token_tag<'a>(
939 pattern: Node<'_>,
940 tags: &'a BTreeMap<TokenId, TagInfo>,
941) -> Option<&'a TagInfo> {
942 let token_id = pattern.as_terminal()?.token_id();
943 let tag = tags.get(&token_id)?;
944 matches!(tag.kind, TagKind::Token { .. }).then_some(tag)
945}
946
947fn rule_tag_of<'a>(
951 pattern: Node<'_>,
952 tags: &'a BTreeMap<TokenId, TagInfo>,
953) -> Option<(usize, &'a TagInfo)> {
954 let rule = pattern.as_rule()?;
955 if rule.child_count() != 1 {
956 return None;
957 }
958 let child = pattern.children().next()?;
959 let token_id = child.as_terminal()?.token_id();
960 let tag = tags.get(&token_id)?;
961 match tag.kind {
962 TagKind::Rule { rule_index, .. } => Some((rule_index, tag)),
963 TagKind::Token { .. } => None,
964 }
965}
966
967fn bind<'subject>(
969 labels: &mut BTreeMap<String, Vec<Node<'subject>>>,
970 tag: &TagInfo,
971 node: Node<'subject>,
972) {
973 for key in tag.label_keys() {
974 labels.entry(key.to_owned()).or_default().push(node);
975 }
976}
977
978fn leaf_text(node: Node<'_>) -> &str {
980 node.as_terminal()
981 .map(crate::tree::TerminalNodeView::text)
982 .or_else(|| node.as_error().map(crate::tree::ErrorNodeView::text))
983 .unwrap_or("")
984}
985
986#[cfg(test)]
987#[allow(clippy::disallowed_methods)] mod tests {
989 use super::*;
990 use crate::token::{TokenSpec, TokenStore};
991 use crate::tree::{NodeId, ParseTreeStorage, ParsedFile, ParserRuleContext};
992
993 const RULE_STAT: usize = 0;
995 const RULE_EXPR: usize = 1;
996 const ASSIGN: i32 = 1;
997 const SEMI: i32 = 2;
998 const ID: i32 = 3;
999 const INT: i32 = 4;
1000 const BYPASS_EXPR: i32 = 6;
1003
1004 fn split_default(pattern: &str) -> Result<Vec<Chunk>, ParseTreePatternError> {
1007 split(pattern, &Delimiters::default())
1008 }
1009
1010 #[test]
1011 fn split_interleaves_text_and_tags() {
1012 let chunks = split_default("<ID> = <expr> ;").expect("valid pattern");
1013 insta::assert_debug_snapshot!("split_interleaves_text_and_tags", chunks);
1014 }
1015
1016 #[test]
1017 fn split_parses_labeled_tags() {
1018 let chunks = split_default("<lhs:ID> = <e:expr>").expect("valid pattern");
1019 insta::assert_debug_snapshot!("split_parses_labeled_tags", chunks);
1020 }
1021
1022 #[test]
1023 fn split_strips_escapes_from_text_only() {
1024 let chunks = split_default(r"a \< b <ID> c \> d").expect("valid pattern");
1026 insta::assert_debug_snapshot!("split_strips_escapes", chunks);
1027 }
1028
1029 #[test]
1030 fn split_no_tags_is_single_text_chunk() {
1031 let chunks = split_default("a = 3 ;").expect("valid pattern");
1032 insta::assert_debug_snapshot!("split_no_tags", chunks);
1033 }
1034
1035 #[test]
1036 fn split_rejects_malformed_patterns() {
1037 let cases = ["<ID", "ID>", "><", "<>", "<a:>", "<a<b>>"];
1038 let errors: Vec<_> = cases
1039 .into_iter()
1040 .map(|pattern| {
1041 (
1042 pattern,
1043 split_default(pattern).expect_err("invalid").to_string(),
1044 )
1045 })
1046 .collect();
1047 insta::assert_debug_snapshot!("split_rejects_malformed", errors);
1048 }
1049
1050 #[test]
1051 fn split_honors_custom_delimiters() {
1052 let delimiters = Delimiters {
1053 start: "[[".to_owned(),
1054 stop: "]]".to_owned(),
1055 escape: "%".to_owned(),
1056 };
1057 let chunks = split("x [[expr]] y", &delimiters).expect("valid pattern");
1058 insta::assert_debug_snapshot!("split_custom_delimiters", chunks);
1059 }
1060
1061 enum Build {
1065 Rule(usize, Vec<Self>),
1066 Token(i32, &'static str),
1068 TokenTag {
1070 token_type: i32,
1071 name: &'static str,
1072 label: Option<&'static str>,
1073 },
1074 RuleTag {
1076 rule_index: usize,
1077 bypass_type: i32,
1078 name: &'static str,
1079 label: Option<&'static str>,
1080 },
1081 }
1082
1083 struct TreeFactory {
1084 tokens: TokenStore,
1085 storage: ParseTreeStorage,
1086 tags: BTreeMap<TokenId, TagInfo>,
1087 }
1088
1089 impl TreeFactory {
1090 fn new() -> Self {
1091 Self {
1092 tokens: TokenStore::new(None, "TreePattern"),
1093 storage: ParseTreeStorage::new(),
1094 tags: BTreeMap::new(),
1095 }
1096 }
1097
1098 fn push_token(&mut self, token_type: i32, text: &str) -> TokenId {
1099 self.tokens
1100 .push(TokenSpec::explicit(token_type, text))
1101 .expect("test token fits")
1102 }
1103
1104 fn build(&mut self, spec: &Build) -> NodeId {
1105 match spec {
1106 Build::Token(token_type, text) => {
1107 let id = self.push_token(*token_type, text);
1108 self.storage.terminal(id)
1109 }
1110 Build::TokenTag {
1111 token_type,
1112 name,
1113 label,
1114 } => {
1115 let id = self.push_token(*token_type, &format!("<{name}>"));
1116 self.tags.insert(
1117 id,
1118 TagInfo {
1119 kind: TagKind::Token {
1120 token_type: *token_type,
1121 },
1122 name: (*name).to_owned(),
1123 label: label.map(str::to_owned),
1124 },
1125 );
1126 self.storage.terminal(id)
1127 }
1128 Build::RuleTag {
1129 rule_index,
1130 bypass_type,
1131 name,
1132 label,
1133 } => {
1134 let id = self.push_token(*bypass_type, &format!("<{name}>"));
1135 self.tags.insert(
1136 id,
1137 TagInfo {
1138 kind: TagKind::Rule {
1139 rule_index: *rule_index,
1140 bypass_type: *bypass_type,
1141 },
1142 name: (*name).to_owned(),
1143 label: label.map(str::to_owned),
1144 },
1145 );
1146 let leaf = self.storage.terminal(id);
1148 let mut context = ParserRuleContext::new(*rule_index, -1);
1149 self.storage.add_child(&mut context, leaf);
1150 self.storage.finish_rule(context)
1151 }
1152 Build::Rule(rule_index, children) => {
1153 let child_ids: Vec<_> = children.iter().map(|c| self.build(c)).collect();
1154 let mut context = ParserRuleContext::new(*rule_index, -1);
1155 for child in child_ids {
1156 self.storage.add_child(&mut context, child);
1157 }
1158 self.storage.finish_rule(context)
1159 }
1160 }
1161 }
1162
1163 fn into_file(self, root: NodeId) -> (ParsedFile, BTreeMap<TokenId, TagInfo>) {
1164 (ParsedFile::new(self.tokens, self.storage, root), self.tags)
1165 }
1166 }
1167
1168 fn subject_tree(spec: &Build) -> ParsedFile {
1170 let mut factory = TreeFactory::new();
1171 let root = factory.build(spec);
1172 factory.into_file(root).0
1173 }
1174
1175 fn pattern_from(rule_index: usize, spec: &Build) -> ParseTreePattern {
1177 let mut factory = TreeFactory::new();
1178 let root = factory.build(spec);
1179 let (file, tags) = factory.into_file(root);
1180 ParseTreePattern {
1181 pattern: "<test>".to_owned(),
1182 pattern_rule_index: rule_index,
1183 tree: PatternTree { file, tags },
1184 }
1185 }
1186
1187 fn subject_x_eq_3() -> ParsedFile {
1189 subject_tree(&Build::Rule(
1190 RULE_STAT,
1191 vec![
1192 Build::Token(ID, "x"),
1193 Build::Token(ASSIGN, "="),
1194 Build::Rule(RULE_EXPR, vec![Build::Token(INT, "3")]),
1195 Build::Token(SEMI, ";"),
1196 ],
1197 ))
1198 }
1199
1200 #[test]
1201 fn matches_rule_tag_and_binds_label() {
1202 let pattern = pattern_from(
1204 RULE_STAT,
1205 &Build::Rule(
1206 RULE_STAT,
1207 vec![
1208 Build::TokenTag {
1209 token_type: ID,
1210 name: "ID",
1211 label: None,
1212 },
1213 Build::Token(ASSIGN, "="),
1214 Build::RuleTag {
1215 rule_index: RULE_EXPR,
1216 bypass_type: BYPASS_EXPR,
1217 name: "expr",
1218 label: Some("e"),
1219 },
1220 Build::Token(SEMI, ";"),
1221 ],
1222 ),
1223 );
1224 let subject = subject_x_eq_3();
1225 let result = pattern.match_tree(subject.tree());
1226
1227 assert!(result.succeeded(), "pattern should match");
1228 assert_eq!(result.get("ID").map(Node::text), Some("x".to_owned()));
1230 assert_eq!(result.get("e").map(Node::text), Some("3".to_owned()));
1231 assert_eq!(result.get("expr").map(Node::text), Some("3".to_owned()));
1232 assert!(result.get("absent").is_none());
1233 }
1234
1235 #[test]
1236 fn literal_mismatch_reports_first_bad_node() {
1237 let pattern = pattern_from(
1239 RULE_STAT,
1240 &Build::Rule(
1241 RULE_STAT,
1242 vec![
1243 Build::Token(ID, "y"),
1244 Build::Token(ASSIGN, "="),
1245 Build::RuleTag {
1246 rule_index: RULE_EXPR,
1247 bypass_type: BYPASS_EXPR,
1248 name: "expr",
1249 label: None,
1250 },
1251 Build::Token(SEMI, ";"),
1252 ],
1253 ),
1254 );
1255 let subject = subject_x_eq_3();
1256 let result = pattern.match_tree(subject.tree());
1257
1258 assert!(!result.succeeded());
1259 assert_eq!(
1260 result.mismatched_node().map(Node::text),
1261 Some("x".to_owned())
1262 );
1263 }
1264
1265 #[test]
1266 fn child_count_mismatch_fails_at_rule() {
1267 let pattern = pattern_from(
1269 RULE_STAT,
1270 &Build::Rule(
1271 RULE_STAT,
1272 vec![
1273 Build::TokenTag {
1274 token_type: ID,
1275 name: "ID",
1276 label: None,
1277 },
1278 Build::Token(ASSIGN, "="),
1279 Build::RuleTag {
1280 rule_index: RULE_EXPR,
1281 bypass_type: BYPASS_EXPR,
1282 name: "expr",
1283 label: None,
1284 },
1285 ],
1286 ),
1287 );
1288 let subject = subject_x_eq_3();
1289 let result = pattern.match_tree(subject.tree());
1290
1291 assert!(!result.succeeded());
1292 assert!(result.mismatched_node().and_then(Node::as_rule).is_some());
1294 }
1295
1296 #[test]
1297 fn rule_tag_type_mismatch_fails() {
1298 let pattern = pattern_from(
1300 RULE_STAT,
1301 &Build::RuleTag {
1302 rule_index: RULE_EXPR,
1303 bypass_type: BYPASS_EXPR,
1304 name: "expr",
1305 label: None,
1306 },
1307 );
1308 let subject = subject_x_eq_3(); let result = pattern.match_tree(subject.tree());
1310 assert!(!result.succeeded());
1311 }
1312
1313 #[test]
1314 fn get_all_returns_every_binding_in_order() {
1315 let pattern = pattern_from(
1317 RULE_STAT,
1318 &Build::Rule(
1319 RULE_STAT,
1320 vec![
1321 Build::RuleTag {
1322 rule_index: RULE_EXPR,
1323 bypass_type: BYPASS_EXPR,
1324 name: "expr",
1325 label: Some("operand"),
1326 },
1327 Build::RuleTag {
1328 rule_index: RULE_EXPR,
1329 bypass_type: BYPASS_EXPR,
1330 name: "expr",
1331 label: Some("operand"),
1332 },
1333 ],
1334 ),
1335 );
1336 let subject = subject_tree(&Build::Rule(
1337 RULE_STAT,
1338 vec![
1339 Build::Rule(RULE_EXPR, vec![Build::Token(INT, "1")]),
1340 Build::Rule(RULE_EXPR, vec![Build::Token(INT, "2")]),
1341 ],
1342 ));
1343 let result = pattern.match_tree(subject.tree());
1344
1345 assert!(result.succeeded());
1346 let operands: Vec<_> = result.get_all("operand").iter().map(|n| n.text()).collect();
1347 assert_eq!(operands, vec!["1".to_owned(), "2".to_owned()]);
1348 assert_eq!(result.get_all("expr").len(), 2);
1350 }
1351
1352 use crate::atn::AtnStateKind;
1355 use crate::atn::parser_atn::{ParserAtn, ParserAtnBuilder, ParserTransitionSpec};
1356 use crate::vocabulary::Vocabulary;
1357
1358 fn stat_expr_atn() -> ParserAtn {
1364 let mut atn = ParserAtnBuilder::new(4);
1365 for (number, kind, rule) in [
1367 (0, AtnStateKind::RuleStart, 0), (1, AtnStateKind::Basic, 0), (2, AtnStateKind::Basic, 0), (3, AtnStateKind::Basic, 0), (4, AtnStateKind::RuleStop, 0), (5, AtnStateKind::RuleStart, 1), (6, AtnStateKind::BlockStart, 1), (7, AtnStateKind::Basic, 1), (8, AtnStateKind::Basic, 1), (9, AtnStateKind::BlockEnd, 1), (10, AtnStateKind::RuleStop, 1), ] {
1379 assert_eq!(
1380 atn.add_state(kind, Some(rule)).expect("state").index(),
1381 number
1382 );
1383 }
1384 atn.set_rule_to_start_state(vec![0, 5]).expect("starts");
1385 atn.set_rule_to_stop_state(vec![4, 10]).expect("stops");
1386 atn.set_end_state(6, 9).expect("expr block end");
1387 atn.add_decision_state(6).expect("decision");
1388
1389 atn.add_transition(
1391 0,
1392 ParserTransitionSpec::Atom {
1393 target: 1,
1394 label: ID,
1395 },
1396 )
1397 .expect("edge");
1398 atn.add_transition(
1399 1,
1400 ParserTransitionSpec::Atom {
1401 target: 2,
1402 label: ASSIGN,
1403 },
1404 )
1405 .expect("edge");
1406 atn.add_transition(
1407 2,
1408 ParserTransitionSpec::Rule {
1409 target: 5,
1410 rule_index: 1,
1411 follow_state: 3,
1412 precedence: 0,
1413 },
1414 )
1415 .expect("edge");
1416 atn.add_transition(
1417 3,
1418 ParserTransitionSpec::Atom {
1419 target: 4,
1420 label: SEMI,
1421 },
1422 )
1423 .expect("edge");
1424 atn.add_transition(10, ParserTransitionSpec::Epsilon { target: 3 })
1427 .expect("edge");
1428
1429 atn.add_transition(5, ParserTransitionSpec::Epsilon { target: 6 })
1431 .expect("edge");
1432 atn.add_transition(6, ParserTransitionSpec::Epsilon { target: 7 })
1433 .expect("edge");
1434 atn.add_transition(6, ParserTransitionSpec::Epsilon { target: 8 })
1435 .expect("edge");
1436 atn.add_transition(
1437 7,
1438 ParserTransitionSpec::Atom {
1439 target: 9,
1440 label: INT,
1441 },
1442 )
1443 .expect("edge");
1444 atn.add_transition(
1445 8,
1446 ParserTransitionSpec::Atom {
1447 target: 9,
1448 label: ID,
1449 },
1450 )
1451 .expect("edge");
1452 atn.add_transition(9, ParserTransitionSpec::Epsilon { target: 10 })
1453 .expect("edge");
1454
1455 atn.finish().expect("valid stat/expr ATN")
1456 }
1457
1458 fn stat_expr_data() -> RecognizerData {
1459 RecognizerData::new(
1460 "StatExpr.g4",
1461 Vocabulary::new(
1462 [None, Some("'='"), Some("';'"), None, None],
1463 [None, Some("ASSIGN"), Some("SEMI"), Some("ID"), Some("INT")],
1464 [None::<&str>, None],
1465 ),
1466 )
1467 .with_rule_names(["stat", "expr"])
1468 }
1469
1470 fn stat_expr_chunk_lexer(text: &str) -> Result<Vec<TokenSpec>, ParseTreePatternError> {
1476 let mut specs = Vec::new();
1477 for word in text.split_whitespace() {
1478 let token_type = match word {
1479 "=" => ASSIGN,
1480 ";" => SEMI,
1481 _ if word.chars().all(|c| c.is_ascii_digit()) => INT,
1482 _ if word.chars().all(|c| c.is_ascii_alphanumeric()) => ID,
1483 other => {
1484 return Err(ParseTreePatternError::Tokenization {
1485 message: format!("unexpected chunk word {other:?}"),
1486 });
1487 }
1488 };
1489 specs.push(TokenSpec::explicit(token_type, word));
1490 }
1491 Ok(specs)
1492 }
1493
1494 fn stat_expr_matcher_and_data() -> (ParserAtn, RecognizerData) {
1495 (stat_expr_atn(), stat_expr_data())
1496 }
1497
1498 #[test]
1499 fn compile_and_match_full_pattern() {
1500 let (atn, data) = stat_expr_matcher_and_data();
1501 let matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1502 let pattern = matcher
1503 .compile("<ID> = <e:expr> ;", RULE_STAT, stat_expr_chunk_lexer)
1504 .expect("compiles");
1505
1506 let mut parser = BaseParser::new(
1508 CommonTokenStream::new(stat_expr_subject("x = 3 ;")),
1509 data.clone(),
1510 );
1511 let root = parser
1512 .parse_atn_rule(&atn, RULE_STAT)
1513 .expect("subject parse");
1514 let subject = parser.into_parsed_file(root);
1515
1516 let result = pattern.match_tree(subject.tree());
1517 assert!(result.succeeded(), "pattern should match `x = 3 ;`");
1518 assert_eq!(result.get("ID").map(Node::text), Some("x".to_owned()));
1519 assert_eq!(result.get("e").map(Node::text), Some("3".to_owned()));
1520 }
1521
1522 #[test]
1523 fn compile_rejects_patterns_that_only_parse_via_recovery() {
1524 let (atn, data) = stat_expr_matcher_and_data();
1529 let matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1530 for pattern in ["<ID> <e:expr> ;", "<ID> = ;", "= <expr> ;", "x 3 ;"] {
1531 let error = matcher
1532 .compile(pattern, RULE_STAT, stat_expr_chunk_lexer)
1533 .expect_err("recovered pattern parse must be rejected");
1534 assert!(
1535 matches!(error, ParseTreePatternError::CannotInvokeStartRule { .. }),
1536 "unexpected error for {pattern:?}: {error}"
1537 );
1538 }
1539 }
1540
1541 #[test]
1542 fn split_rejects_overlapping_tags_without_panicking() {
1543 let error = split_default("<a<b>>").expect_err("overlapping tags");
1546 assert!(matches!(
1547 error,
1548 ParseTreePatternError::DelimitersOutOfOrder { .. }
1549 ));
1550 }
1551
1552 #[test]
1553 fn compile_rejects_tokens_after_an_eof_tag() {
1554 let (atn, data) = stat_expr_matcher_and_data();
1557 let matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1558 let error = matcher
1559 .compile(
1560 "<ID> = <expr> ; <EOF> garbage",
1561 RULE_STAT,
1562 stat_expr_chunk_lexer,
1563 )
1564 .expect_err("tokens after an EOF tag must be rejected");
1565 assert!(
1566 matches!(error, ParseTreePatternError::Tokenization { .. }),
1567 "unexpected error: {error}"
1568 );
1569 }
1570
1571 #[test]
1572 fn compile_rejects_unconsumed_trailing_eof_tag() {
1573 let (atn, data) = stat_expr_matcher_and_data();
1578 let matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1579 let error = matcher
1580 .compile("<ID> = <expr> ; <EOF>", RULE_STAT, stat_expr_chunk_lexer)
1581 .expect_err("unconsumed trailing EOF tag must be rejected");
1582 assert!(
1583 matches!(
1584 error,
1585 ParseTreePatternError::StartRuleDoesNotConsumeFullPattern { .. }
1586 ),
1587 "unexpected error: {error}"
1588 );
1589 }
1590
1591 #[test]
1592 fn compile_rejects_partial_pattern() {
1593 let (atn, data) = stat_expr_matcher_and_data();
1596 let matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1597 let error = matcher
1598 .compile("<ID> = <expr> ; extra", RULE_STAT, stat_expr_chunk_lexer)
1599 .expect_err("trailing token should be rejected");
1600 assert!(
1601 matches!(
1602 error,
1603 ParseTreePatternError::StartRuleDoesNotConsumeFullPattern { .. }
1604 ),
1605 "unexpected error: {error}"
1606 );
1607 }
1608
1609 #[test]
1610 fn compile_rejects_unknown_tag_names() {
1611 let (atn, data) = stat_expr_matcher_and_data();
1612 let matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1613 let unknown_token = matcher
1614 .compile("<NOPE> = <expr> ;", RULE_STAT, stat_expr_chunk_lexer)
1615 .expect_err("unknown token tag");
1616 assert!(matches!(
1617 unknown_token,
1618 ParseTreePatternError::UnknownToken { .. }
1619 ));
1620 let unknown_rule = matcher
1621 .compile("<ID> = <nope> ;", RULE_STAT, stat_expr_chunk_lexer)
1622 .expect_err("unknown rule tag");
1623 assert!(matches!(
1624 unknown_rule,
1625 ParseTreePatternError::UnknownRule { .. }
1626 ));
1627 }
1628
1629 #[test]
1630 fn set_delimiters_validates_and_switches_tag_syntax() {
1631 let (atn, data) = stat_expr_matcher_and_data();
1632 let mut matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1633
1634 assert!(matches!(
1636 matcher.set_delimiters("", ">", "\\"),
1637 Err(ParseTreePatternError::EmptyDelimiter { which: "start" })
1638 ));
1639 assert!(matches!(
1640 matcher.set_delimiters("<", "", "\\"),
1641 Err(ParseTreePatternError::EmptyDelimiter { which: "stop" })
1642 ));
1643
1644 matcher
1647 .set_delimiters("[[", "]]", "%")
1648 .expect("valid delimiters");
1649 matcher
1650 .compile("[[ID]] = [[e:expr]] ;", RULE_STAT, stat_expr_chunk_lexer)
1651 .expect("custom-delimiter pattern compiles");
1652 matcher
1653 .compile("<ID> = <expr> ;", RULE_STAT, stat_expr_chunk_lexer)
1654 .expect_err("old delimiters are literal text now");
1655 }
1656
1657 #[test]
1658 fn compiled_pattern_does_not_match_different_structure() {
1659 let (atn, data) = stat_expr_matcher_and_data();
1660 let matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1661 let pattern = matcher
1663 .compile("y = <expr> ;", RULE_STAT, stat_expr_chunk_lexer)
1664 .expect("compiles");
1665
1666 let mut parser = BaseParser::new(
1667 CommonTokenStream::new(stat_expr_subject("x = 3 ;")),
1668 data.clone(),
1669 );
1670 let root = parser
1671 .parse_atn_rule(&atn, RULE_STAT)
1672 .expect("subject parse");
1673 let subject = parser.into_parsed_file(root);
1674
1675 let result = pattern.match_tree(subject.tree());
1676 assert!(
1677 !result.succeeded(),
1678 "identifier `x` should not match literal `y`"
1679 );
1680 }
1681
1682 fn stat_expr_subject(input: &str) -> PatternTokenSource {
1685 let specs = stat_expr_chunk_lexer(input).expect("valid subject input");
1686 PatternTokenSource { specs, index: 0 }
1687 }
1688
1689 #[derive(Debug)]
1690 struct StatExprRecognizer {
1691 data: RecognizerData,
1692 }
1693
1694 impl Recognizer for StatExprRecognizer {
1695 fn data(&self) -> &RecognizerData {
1696 &self.data
1697 }
1698
1699 fn data_mut(&mut self) -> &mut RecognizerData {
1700 &mut self.data
1701 }
1702 }
1703
1704 #[test]
1705 fn find_all_pairs_xpath_selection_with_pattern_matching() {
1706 let (atn, data) = stat_expr_matcher_and_data();
1707 let matcher = ParseTreePatternMatcher::new(&atn, &data).expect("matcher");
1708 let pattern = matcher
1710 .compile("<INT>", RULE_EXPR, stat_expr_chunk_lexer)
1711 .expect("compiles");
1712
1713 let mut parser = BaseParser::new(
1714 CommonTokenStream::new(stat_expr_subject("x = 3 ;")),
1715 data.clone(),
1716 );
1717 let root = parser
1718 .parse_atn_rule(&atn, RULE_STAT)
1719 .expect("subject parse");
1720 let subject = parser.into_parsed_file(root);
1721 let recognizer = StatExprRecognizer { data };
1722
1723 let matches = pattern
1725 .find_all(subject.tree(), "//expr", &recognizer)
1726 .expect("valid xpath");
1727 assert_eq!(matches.len(), 1);
1728 assert_eq!(matches[0].tree().text(), "3");
1729 let none = pattern
1731 .find_all(subject.tree(), "//stat", &recognizer)
1732 .expect("valid xpath");
1733 assert!(none.is_empty());
1734 }
1735}