diff --git a/src/parser.rs b/src/parser.rs index 4d7c35e0..0f1b1fad 100644 --- a/src/parser.rs +++ b/src/parser.rs @@ -1255,6 +1255,9 @@ pub struct BaseParser { /// selected rule spans after recognition, avoiding many speculative /// nodes that are thrown away with losing paths. fast_token_nodes_enabled: bool, + /// Whether fast recognition should retain private/public rule alternatives + /// in deferred tree metadata. + fast_track_alt_numbers: bool, /// Parser-owned append-only storage for speculative recognition output. /// Each public interpreted-rule entry clears lengths while retaining /// bounded backing capacities for parser reuse. @@ -1387,6 +1390,10 @@ struct FastDeferredRuleId(u32); enum FastDeferredNode { Fragment(NodeSeqId), Rule(FastDeferredRuleId), + Alternative(u32), + LeftRecursiveBoundary { + rule_index: u32, + }, Concat { prefix: FastDeferredNodeId, suffix: FastDeferredNodeId, @@ -1596,6 +1603,18 @@ impl RecognitionArena { self.push_deferred_node(FastDeferredNode::Rule(rule)) } + fn deferred_alternative(&mut self, alt_number: usize) -> FastDeferredNodeId { + self.push_deferred_node(FastDeferredNode::Alternative( + u32::try_from(alt_number).expect("alternative number fits in u32"), + )) + } + + fn deferred_left_recursive_boundary(&mut self, rule_index: usize) -> FastDeferredNodeId { + self.push_deferred_node(FastDeferredNode::LeftRecursiveBoundary { + rule_index: u32::try_from(rule_index).expect("rule index fits in u32"), + }) + } + fn concat_deferred_nodes( &mut self, prefix: FastDeferredNodeId, @@ -1681,6 +1700,16 @@ impl RecognitionArena { self.nodes[id.0 as usize] } + fn set_boundary_alt_number(&mut self, id: RecognizedNodeId, alt_number: u32) { + let ArenaRecognizedNode::LeftRecursiveBoundary { + alt_number: stored, .. + } = &mut self.nodes[id.0 as usize] + else { + unreachable!("deferred boundary must materialize as a boundary node"); + }; + *stored = alt_number; + } + fn extra(&self, id: RecognitionExtraId) -> &RecognitionExtra { &self.extras[id.0 as usize] } @@ -3962,7 +3991,6 @@ fn atn_has_predicate_transitions(atn: &Atn) -> bool { fn can_use_fast_predicate_recognizer(atn: &Atn, options: &ParserRuntimeOptions<'_>) -> bool { options.init_action_rules.is_empty() && !options.track_alt_numbers - && !options.track_context_alt_numbers && options .predicates .iter() @@ -4075,6 +4103,18 @@ struct FastPredicateContext<'a> { member_values: &'a BTreeMap, } +#[derive(Clone, Copy, Debug, Default)] +struct AltNumberTracking { + public: bool, + context: bool, +} + +impl AltNumberTracking { + const fn any(self) -> bool { + self.public || self.context + } +} + struct FastRecognizeScratch<'a, 'b> { predicate_context: Option>, visiting: &'b mut FxHashSet, @@ -4761,6 +4801,7 @@ where fast_first_set_prefilter: true, fast_recovery_enabled: true, fast_token_nodes_enabled: true, + fast_track_alt_numbers: false, recognition_arena: RecognitionArena::default(), last_recognition_arena_root: NodeSeqId::EMPTY, last_recognition_arena_diagnostics: DiagnosticSeqId::EMPTY, @@ -4799,6 +4840,7 @@ where self.fast_first_set_prefilter = true; self.fast_recovery_enabled = true; self.fast_token_nodes_enabled = self.build_parse_trees; + self.fast_track_alt_numbers = false; self.reset_recognition_arena(); } @@ -6624,7 +6666,13 @@ where rule_index: usize, precedence: i32, ) -> Result { - self.parse_atn_rule_with_precedence_inner(atn, rule_index, precedence, None) + self.parse_atn_rule_with_precedence_inner( + atn, + rule_index, + precedence, + None, + AltNumberTracking::default(), + ) } fn parse_atn_rule_with_precedence_inner( @@ -6633,6 +6681,7 @@ where rule_index: usize, precedence: i32, predicate_context: Option>, + alt_tracking: AltNumberTracking, ) -> Result { let start_state = atn.rule_to_start_state().get(rule_index).ok_or_else(|| { AntlrError::Unsupported(format!("rule {rule_index} has no start state")) @@ -6652,6 +6701,7 @@ where let caller_follow_state = self.pending_invoking_follow_state(atn); self.fast_recovery_enabled = false; self.fast_token_nodes_enabled = false; + self.fast_track_alt_numbers = alt_tracking.any(); let top_request = FastRecognizeTopRequest { start_state, stop_state, @@ -6663,7 +6713,7 @@ where self.fast_token_nodes_enabled = self.build_parse_trees; let needs_tree_retry = matches!( &first_pass, - Ok((outcome, _)) + Ok((outcome, _, _)) if self.build_parse_trees && self .recognition_arena @@ -6683,9 +6733,9 @@ where // boundaries also need the token-node pass; otherwise the fold has // no concrete left operand to wrap into ANTLR's recursive context. Err(_) => true, - Ok((outcome, _)) => !outcome.diagnostics.is_empty() || needs_tree_retry, + Ok((outcome, _, _)) => !outcome.diagnostics.is_empty() || needs_tree_retry, }; - let (outcome, _expected) = if needs_retry { + let (outcome, _expected, alt_number) = if needs_retry { self.fast_first_set_prefilter = false; self.fast_recovery_enabled = false; let clean_retry = self.fast_recognize_top(atn, top_request, predicate_context); @@ -6698,7 +6748,7 @@ where select_better_top_outcome(first_pass, clean_retry, &self.recognition_arena) }; let selected = if clean_selected.is_err() - || matches!(&clean_selected, Ok((outcome, _)) if !outcome.diagnostics.is_empty()) + || matches!(&clean_selected, Ok((outcome, _, _)) if !outcome.diagnostics.is_empty()) { self.fast_recovery_enabled = true; let recovery_retry = self.fast_recognize_top(atn, top_request, predicate_context); @@ -6742,6 +6792,12 @@ where 0 }, ); + if alt_tracking.public { + context.set_alt_number(alt_number); + } + if alt_tracking.context { + context.set_context_alt_number(alt_number); + } if let Some(token) = self.token_id_at(start_index) { self.set_context_start(&mut context, token); } @@ -6762,7 +6818,11 @@ where { let mut cursor = live_root; while let Some(link) = self.recognition_arena.link(cursor) { - let child = self.arena_recognized_node_tree(link.head, false, false)?; + let child = self.arena_recognized_node_tree( + link.head, + alt_tracking.public, + alt_tracking.context, + )?; self.tree.add_child(&mut context, child); cursor = link.tail; } @@ -6772,6 +6832,7 @@ where start_index, stop_index, live_root, + alt_tracking, )?; } } @@ -6806,7 +6867,7 @@ where atn: &Atn, request: FastRecognizeTopRequest, predicate_context: Option>, - ) -> Result<(FastRecognizeOutcome, ExpectedTokens), ExpectedTokens> { + ) -> Result<(FastRecognizeOutcome, ExpectedTokens, usize), ExpectedTokens> { let FastRecognizeTopRequest { start_state, stop_state, @@ -6870,10 +6931,12 @@ where }; match selected { Some(mut outcome) => { - if self.build_parse_trees { - self.materialize_fast_outcome_nodes(&mut outcome); - } - Ok((outcome, expected)) + let alt_number = if self.build_parse_trees || self.fast_track_alt_numbers { + self.materialize_fast_outcome_nodes(&mut outcome) + } else { + 0 + }; + Ok((outcome, expected, alt_number)) } None => Err(expected), } @@ -6968,12 +7031,14 @@ where fn arena_recognized_node_tree_with_implicit_tokens( &mut self, node_id: RecognizedNodeId, + alt_tracking: AltNumberTracking, ) -> Result { let node = self.recognition_arena.node(node_id); match node { ArenaRecognizedNode::Rule { rule_index, invoking_state, + alt_number, start_index, stop_index, children, @@ -6984,6 +7049,12 @@ where invoking_state as isize, self.recognition_arena.sequence_len(children), ); + if alt_tracking.public { + context.set_alt_number(alt_number as usize); + } + if alt_tracking.context { + context.set_context_alt_number(alt_number as usize); + } if let Some(token) = self.token_id_at(start_index as usize) { self.set_context_start(&mut context, token); } @@ -6998,10 +7069,13 @@ where start_index as usize, stop_index.map(|index| index as usize), children, + alt_tracking, )?; Ok(self.rule_node(context)) } - _ => self.arena_recognized_node_tree(node_id, false, false), + _ => { + self.arena_recognized_node_tree(node_id, alt_tracking.public, alt_tracking.context) + } } } @@ -7011,18 +7085,21 @@ where start_index: usize, stop_index: Option, mut children: NodeSeqId, + alt_tracking: AltNumberTracking, ) -> Result<(), AntlrError> { let mut cursor = Some(start_index); while let Some(link) = self.recognition_arena.link(children) { if let Some((child_start, child_stop)) = self.recognition_arena.node_span(link.head) { self.add_visible_terminals_before(context, &mut cursor, child_start)?; - let child = self.arena_recognized_node_tree_with_implicit_tokens(link.head)?; + let child = + self.arena_recognized_node_tree_with_implicit_tokens(link.head, alt_tracking)?; self.tree.add_child(context, child); if let Some(child_stop) = child_stop { cursor = self.next_visible_after_token(child_stop); } } else { - let child = self.arena_recognized_node_tree_with_implicit_tokens(link.head)?; + let child = + self.arena_recognized_node_tree_with_implicit_tokens(link.head, alt_tracking)?; self.tree.add_child(context, child); } children = link.tail; @@ -7203,6 +7280,10 @@ where semantics, member_values: &member_values, }), + AltNumberTracking { + public: track_alt_numbers, + context: track_context_alt_numbers, + }, ) .map(|tree| (tree, Vec::new())); if self.unknown_predicate_hits.is_empty() && self.unhandled_action_hits.is_empty() { @@ -7924,13 +8005,37 @@ where .concat_deferred_nodes(fragment, outcome.deferred_nodes); } + fn defer_fast_outcome_alternative( + &mut self, + outcome: &mut FastRecognizeOutcome, + alt_number: usize, + ) { + let alternative = self.recognition_arena.deferred_alternative(alt_number); + outcome.deferred_nodes = self + .recognition_arena + .concat_deferred_nodes(alternative, outcome.deferred_nodes); + } + + fn defer_fast_outcome_boundary( + &mut self, + outcome: &mut FastRecognizeOutcome, + rule_index: usize, + ) { + let boundary = self + .recognition_arena + .deferred_left_recursive_boundary(rule_index); + outcome.deferred_nodes = self + .recognition_arena + .concat_deferred_nodes(boundary, outcome.deferred_nodes); + } + fn materialize_fast_deferred_nodes( &mut self, root: FastDeferredNodeId, initial_suffix: NodeSeqId, - ) -> NodeSeqId { + ) -> (NodeSeqId, usize) { if root.is_empty() { - return initial_suffix; + return (initial_suffix, 0); } enum Frame { @@ -7939,10 +8044,17 @@ where FinishRule { rule: FastDeferredRule, parent_suffix: NodeSeqId, + parent_alt_number: u32, + parent_pending_boundary: Option, }, } let mut result = initial_suffix; + // The rope is visited suffix-first while nodes are prepended. Later + // alternatives arrive first, so earlier markers overwrite them; a + // boundary redirects those earlier markers to the wrapped context. + let mut alt_number = 0; + let mut pending_boundary = None; let mut pending = Vec::with_capacity(16); pending.push(Frame::Visit(root)); let mut fragment_nodes = Vec::new(); @@ -7964,13 +8076,32 @@ where FastDeferredNode::Rule(rule) => { let rule = self.recognition_arena.deferred_rule(rule); let parent_suffix = result; + let parent_alt_number = alt_number; + let parent_pending_boundary = pending_boundary; result = rule.children; + alt_number = 0; + pending_boundary = None; pending.push(Frame::FinishRule { rule, parent_suffix, + parent_alt_number, + parent_pending_boundary, }); pending.push(Frame::Visit(rule.deferred_children)); } + FastDeferredNode::Alternative(selected) => { + if let Some(boundary) = pending_boundary { + self.recognition_arena + .set_boundary_alt_number(boundary, selected); + } else { + alt_number = selected; + } + } + FastDeferredNode::LeftRecursiveBoundary { rule_index } => { + let boundary = self.arena_boundary_node(rule_index as usize, 0); + self.arena_prepend(&mut result, boundary); + pending_boundary = Some(boundary); + } FastDeferredNode::Concat { prefix, suffix: deferred_suffix, @@ -7984,11 +8115,13 @@ where Frame::FinishRule { rule, parent_suffix, + parent_alt_number, + parent_pending_boundary, } => { let node = self.recognition_arena.push_node(ArenaRecognizedNode::Rule { rule_index: rule.rule_index, invoking_state: rule.invoking_state, - alt_number: 0, + alt_number, start_index: rule.start_index, stop_index: rule.stop_index, return_values: None, @@ -7996,15 +8129,20 @@ where }); result = parent_suffix; self.arena_prepend(&mut result, node); + alt_number = parent_alt_number; + pending_boundary = parent_pending_boundary; } } } - result + (result, alt_number as usize) } - fn materialize_fast_outcome_nodes(&mut self, outcome: &mut FastRecognizeOutcome) { + fn materialize_fast_outcome_nodes(&mut self, outcome: &mut FastRecognizeOutcome) -> usize { let deferred_nodes = std::mem::take(&mut outcome.deferred_nodes); - outcome.nodes = self.materialize_fast_deferred_nodes(deferred_nodes, outcome.nodes); + let (nodes, alt_number) = + self.materialize_fast_deferred_nodes(deferred_nodes, outcome.nodes); + outcome.nodes = nodes; + alt_number } /// Walks one ordinary `*`/`+` repetition at a time so input length grows @@ -8033,6 +8171,17 @@ where } else { None }; + let (enter_alt_number, exit_alt_number) = if self.fast_track_alt_numbers { + let state = atn + .state(request.state_number) + .expect("repetition request state must exist"); + ( + next_alt_number(state, 2, shape.enter_transition_index, 0, true), + next_alt_number(state, 2, shape.exit_transition_index, 0, true), + ) + } else { + (0, 0) + }; let mut work = Vec::with_capacity(2); push_fast_repetition_work( &mut work, @@ -8054,6 +8203,15 @@ where if !coordinates.insert_entered(path) { continue; } + let path_nodes = if enter_alt_number == 0 { + path.deferred_nodes + } else { + let alternative = self + .recognition_arena + .deferred_alternative(enter_alt_number); + self.recognition_arena + .concat_deferred_nodes(path.deferred_nodes, alternative) + }; let body_outcomes = self.recognize_state_fast( atn, FastRecognizeRequest { @@ -8088,7 +8246,7 @@ where .concat_deferred_nodes(body.deferred_nodes, body_fragment); let deferred_nodes = self .recognition_arena - .concat_deferred_nodes(path.deferred_nodes, body_nodes); + .concat_deferred_nodes(path_nodes, body_nodes); let next_path = FastRepetitionPath { index: body.index, deferred_nodes, @@ -8111,6 +8269,14 @@ where if !coordinates.insert_exited(path) { continue; } + let path_nodes = if exit_alt_number == 0 { + path.deferred_nodes + } else { + let alternative = + self.recognition_arena.deferred_alternative(exit_alt_number); + self.recognition_arena + .concat_deferred_nodes(path.deferred_nodes, alternative) + }; let suffixes = self.recognize_state_fast( atn, FastRecognizeRequest { @@ -8135,7 +8301,7 @@ where for mut outcome in suffixes { outcome.deferred_nodes = self .recognition_arena - .concat_deferred_nodes(path.deferred_nodes, outcome.deferred_nodes); + .concat_deferred_nodes(path_nodes, outcome.deferred_nodes); outcome.diagnostics = self .recognition_arena .concat_diagnostics(path.diagnostics, outcome.diagnostics); @@ -8586,13 +8752,50 @@ where continue; } let target = transition.target(); + let outcomes_before_transition = outcomes.len(); + let left_recursive_boundary = match transition_kind { + ParserTransitionKind::Epsilon + | ParserTransitionKind::Action + | ParserTransitionKind::Predicate + | ParserTransitionKind::Precedence => left_recursive_boundary(atn, state, target), + ParserTransitionKind::Atom + | ParserTransitionKind::Range + | ParserTransitionKind::Set + | ParserTransitionKind::NotSet + | ParserTransitionKind::Wildcard + | ParserTransitionKind::Rule => None, + }; match transition_kind { ParserTransitionKind::Epsilon | ParserTransitionKind::Action => { #[cfg(feature = "perf-counters")] perf_counters::inc(&perf_counters::EPSILON_TRANSITIONS, 1); - let boundary = left_recursive_boundary(atn, state, target); - outcomes.extend( - self.recognize_state_fast( + outcomes.extend(self.recognize_state_fast( + atn, + FastRecognizeRequest { + state_number: target, + stop_state, + index, + rule_start_index, + decision_start_index: next_decision_start_index, + precedence, + depth: depth + 1, + recovery_symbols: Rc::clone(&epsilon_recovery_symbols), + recovery_state: epsilon_recovery_state, + }, + FastRecognizeScratch { + predicate_context, + visiting, + memo, + expected, + native_depth: native_depth + 1, + }, + )); + } + ParserTransitionKind::Predicate => { + #[cfg(feature = "perf-counters")] + perf_counters::inc(&perf_counters::EPSILON_TRANSITIONS, 1); + if self.fast_parser_predicate_matches(predicate_context, transition, index) { + outcomes.extend(self.recognize_state_fast( atn, FastRecognizeRequest { state_number: target, @@ -8612,53 +8815,7 @@ where expected, native_depth: native_depth + 1, }, - ) - .into_iter() - .map(|mut outcome| { - if let Some(rule_index) = boundary { - let boundary = self.arena_boundary_node(rule_index, 0); - self.defer_fast_outcome_node(&mut outcome, boundary); - } - outcome - }), - ); - } - ParserTransitionKind::Predicate => { - #[cfg(feature = "perf-counters")] - perf_counters::inc(&perf_counters::EPSILON_TRANSITIONS, 1); - if self.fast_parser_predicate_matches(predicate_context, transition, index) { - let boundary = left_recursive_boundary(atn, state, target); - outcomes.extend( - self.recognize_state_fast( - atn, - FastRecognizeRequest { - state_number: target, - stop_state, - index, - rule_start_index, - decision_start_index: next_decision_start_index, - precedence, - depth: depth + 1, - recovery_symbols: Rc::clone(&epsilon_recovery_symbols), - recovery_state: epsilon_recovery_state, - }, - FastRecognizeScratch { - predicate_context, - visiting, - memo, - expected, - native_depth: native_depth + 1, - }, - ) - .into_iter() - .map(|mut outcome| { - if let Some(rule_index) = boundary { - let boundary = self.arena_boundary_node(rule_index, 0); - self.defer_fast_outcome_node(&mut outcome, boundary); - } - outcome - }), - ); + )); } else { record_predicate_no_viable(expected, next_decision_start_index, index); } @@ -8666,38 +8823,27 @@ where ParserTransitionKind::Precedence => { let transition_precedence = packed_i32(transition.arg0()); if transition_precedence >= precedence { - let boundary = left_recursive_boundary(atn, state, target); - outcomes.extend( - self.recognize_state_fast( - atn, - FastRecognizeRequest { - state_number: target, - stop_state, - index, - rule_start_index, - decision_start_index: next_decision_start_index, - precedence, - depth: depth + 1, - recovery_symbols: Rc::clone(&epsilon_recovery_symbols), - recovery_state: epsilon_recovery_state, - }, - FastRecognizeScratch { - predicate_context, - visiting, - memo, - expected, - native_depth: native_depth + 1, - }, - ) - .into_iter() - .map(|mut outcome| { - if let Some(rule_index) = boundary { - let boundary = self.arena_boundary_node(rule_index, 0); - self.defer_fast_outcome_node(&mut outcome, boundary); - } - outcome - }), - ); + outcomes.extend(self.recognize_state_fast( + atn, + FastRecognizeRequest { + state_number: target, + stop_state, + index, + rule_start_index, + decision_start_index: next_decision_start_index, + precedence, + depth: depth + 1, + recovery_symbols: Rc::clone(&epsilon_recovery_symbols), + recovery_state: epsilon_recovery_state, + }, + FastRecognizeScratch { + predicate_context, + visiting, + memo, + expected, + native_depth: native_depth + 1, + }, + )); } } ParserTransitionKind::Rule => { @@ -8996,6 +9142,23 @@ where } } } + let alt_number = next_alt_number( + state, + transition_count, + transition_index, + 0, + self.fast_track_alt_numbers, + ); + if alt_number != 0 || left_recursive_boundary.is_some() { + for outcome in &mut outcomes[outcomes_before_transition..] { + if alt_number != 0 { + self.defer_fast_outcome_alternative(outcome, alt_number); + } + if let Some(rule_index) = left_recursive_boundary { + self.defer_fast_outcome_boundary(outcome, rule_index); + } + } + } } if has_inserted_cycle_guard { @@ -11716,10 +11879,10 @@ fn state_is_left_recursive_rule(atn: &Atn, state: AtnState<'_>) -> bool { /// rules. If both passes failed, the second pass's expected-token snapshot /// is returned so the caller renders the same diagnostic ANTLR would. fn select_better_top_outcome( - first: Result<(FastRecognizeOutcome, ExpectedTokens), ExpectedTokens>, - second: Result<(FastRecognizeOutcome, ExpectedTokens), ExpectedTokens>, + first: Result<(FastRecognizeOutcome, ExpectedTokens, usize), ExpectedTokens>, + second: Result<(FastRecognizeOutcome, ExpectedTokens, usize), ExpectedTokens>, arena: &RecognitionArena, -) -> Result<(FastRecognizeOutcome, ExpectedTokens), ExpectedTokens> { +) -> Result<(FastRecognizeOutcome, ExpectedTokens, usize), ExpectedTokens> { match (first, second) { (Ok(first), Ok(second)) => { if arena.diagnostics(first.0.diagnostics).next().is_none() { @@ -13066,6 +13229,49 @@ mod tests { finish_atn(atn) } + fn labeled_left_recursive_operator_atn() -> Atn { + let mut atn = ParserAtnBuilder::new(4); + for (state, kind) in [ + (0, AtnStateKind::RuleStart), + (1, AtnStateKind::BlockStart), + (2, AtnStateKind::StarLoopEntry), + (3, AtnStateKind::StarBlockStart), + (4, AtnStateKind::Basic), + (5, AtnStateKind::Basic), + (6, AtnStateKind::Basic), + (7, AtnStateKind::StarLoopBack), + (8, AtnStateKind::LoopEnd), + (9, AtnStateKind::RuleStop), + ] { + assert_eq!(atn.add_state(kind, Some(0)).expect("state").index(), state); + } + atn.set_left_recursive_rule(0) + .expect("left-recursive rule start"); + atn.set_precedence_rule_decision(2) + .expect("precedence decision"); + atn.set_loop_back_state(8, 7).expect("loop-back state"); + atn.set_rule_to_start_state(vec![0]) + .expect("rule start states"); + atn.set_rule_to_stop_state(vec![9]) + .expect("rule stop states"); + for state in [1, 2, 3] { + atn.add_decision_state(state).expect("decision state"); + } + for (source, target) in [(0, 1), (2, 3), (2, 8), (7, 2), (8, 9)] { + atn.add_transition(source, ParserTransitionSpec::Epsilon { target }) + .expect("epsilon transition"); + } + for (source, target, label) in [(1, 2, 1), (1, 2, 2), (4, 6, 4), (5, 6, 3), (6, 7, 1)] { + atn.add_transition(source, ParserTransitionSpec::Atom { target, label }) + .expect("token transition"); + } + for (target, precedence) in [(4, 2), (5, 1)] { + atn.add_transition(3, ParserTransitionSpec::Precedence { target, precedence }) + .expect("operator precedence"); + } + finish_atn(atn) + } + fn parser_inside_left_recursive_callee(symbol: i32) -> BaseParser { let mut parser = mini_parser(vec![ TestToken::new(symbol).with_text("lookahead"), @@ -15681,7 +15887,9 @@ mod tests { }); } - let mut children = parser.materialize_fast_deferred_nodes(root, NodeSeqId::EMPTY); + let (mut children, alt_number) = + parser.materialize_fast_deferred_nodes(root, NodeSeqId::EMPTY); + assert_eq!(alt_number, 0); for expected_rule in (0..DEPTH).rev() { let mut nodes = parser.recognition_arena.iter(children); let node = nodes.next().expect("nested rule node"); @@ -15704,6 +15912,113 @@ mod tests { .expect("deferred rules should materialize without recursion"); } + #[test] + fn deferred_alternatives_preserve_left_recursive_contexts() { + let mut parser = mini_parser(vec![ + TestToken::new(1).with_text("1"), + TestToken::new(2).with_text("+"), + TestToken::new(1).with_text("2"), + TestToken::eof("parser-test", 3, 1, 3), + ]); + let base = parser.arena_token_node(0, false); + let operator = parser.arena_token_node(1, false); + let right = parser.arena_token_node(2, false); + + let base = parser.recognition_arena.prepend(NodeSeqId::EMPTY, base); + let base = parser.recognition_arena.deferred_fragment(base); + let operator = parser.recognition_arena.prepend(NodeSeqId::EMPTY, operator); + let operator = parser.recognition_arena.deferred_fragment(operator); + let right = parser.recognition_arena.prepend(NodeSeqId::EMPTY, right); + let right = parser.recognition_arena.deferred_fragment(right); + let base_alt = parser.recognition_arena.deferred_alternative(1); + let boundary = parser.recognition_arena.deferred_left_recursive_boundary(0); + let operator_alt = parser.recognition_arena.deferred_alternative(6); + + let mut deferred = FastDeferredNodeId::EMPTY; + for fragment in [base_alt, base, boundary, operator_alt, operator, right] { + deferred = parser + .recognition_arena + .concat_deferred_nodes(deferred, fragment); + } + let (nodes, root_alt_number) = + parser.materialize_fast_deferred_nodes(deferred, NodeSeqId::EMPTY); + let nodes = parser + .recognition_arena + .fold_left_recursive_boundaries(nodes); + + let mut root = ParserRuleContext::new(0, -1); + root.set_context_alt_number(root_alt_number); + let mut cursor = nodes; + while let Some(link) = parser.recognition_arena.link(cursor) { + let child = parser + .arena_recognized_node_tree(link.head, false, true) + .expect("materialized child should become a public tree"); + parser.tree.add_child(&mut root, child); + cursor = link.tail; + } + let tree = parser.rule_node(root); + let contexts = parser + .node(tree) + .descendants() + .filter_map(Node::as_rule) + .map(|rule| { + ( + rule.rule_index(), + rule.alt_number(), + rule.context_alt_number(), + rule.text(), + ) + }) + .collect::>(); + + insta::assert_debug_snapshot!( + "deferred_alternatives_preserve_left_recursive_contexts", + contexts + ); + } + + #[test] + fn fast_recognizer_preserves_labeled_left_recursive_operator_context() { + let atn = labeled_left_recursive_operator_atn(); + let mut parser = mini_parser(vec![ + TestToken::new(1).with_text("a"), + TestToken::new(3).with_text("+"), + TestToken::new(1).with_text("b"), + TestToken::eof("parser-test", 3, 1, 3), + ]); + + let (tree, _) = parser + .parse_atn_rule_with_runtime_options( + &atn, + 0, + ParserRuntimeOptions { + track_context_alt_numbers: true, + ..ParserRuntimeOptions::default() + }, + ) + .expect("labeled left-recursive addition should parse"); + let contexts = parser + .node(tree) + .descendants() + .filter_map(Node::as_rule) + .map(|rule| { + let operator = rule + .children() + .next() + .and_then(Node::as_rule) + .is_some_and(|child| child.rule_index() == rule.rule_index()); + (operator, rule.context_alt_number(), rule.text()) + }) + .collect::>(); + + insta::assert_debug_snapshot!( + "fast_recognizer_preserves_labeled_left_recursive_operator_context", + contexts + ); + assert!(!parser.recognition_arena.deferred_nodes.is_empty()); + assert_eq!(parser.number_of_syntax_errors(), 0); + } + #[test] fn deeply_nested_rule_calls_grow_the_stack() { const DEPTH: usize = 4_096; @@ -16400,7 +16715,7 @@ mod tests { } #[test] - fn predicate_gated_same_lookahead_uses_viable_alternative() { + fn private_context_alt_tracking_keeps_fast_predicate_recognition() { let atn = predicate_gated_same_lookahead_atn([0, 1]); let mut parser = mini_parser(vec![ TestToken::new(1).with_text("x"), @@ -16416,12 +16731,17 @@ mod tests { (0, 0, ParserPredicate::False), (0, 1, ParserPredicate::True), ], + track_context_alt_numbers: true, ..ParserRuntimeOptions::default() }, ) .expect("the second predicate-gated alternative should match"); - assert_eq!(parser.node(tree).text(), "x"); + let root = parser.node(tree).as_rule().expect("entry result is a rule"); + insta::assert_debug_snapshot!( + "private_context_alt_tracking_keeps_fast_predicate_recognition", + (root.alt_number(), root.context_alt_number(), root.text()) + ); assert_eq!(parser.number_of_syntax_errors(), 0); assert_eq!(parser.fast_predicate_cache.get(&(0, 0, 0)), Some(&false)); assert_eq!(parser.fast_predicate_cache.get(&(0, 0, 1)), Some(&true)); diff --git a/src/snapshots/antlr4_runtime__parser__tests__deferred_alternatives_preserve_left_recursive_contexts.snap b/src/snapshots/antlr4_runtime__parser__tests__deferred_alternatives_preserve_left_recursive_contexts.snap new file mode 100644 index 00000000..93afba49 --- /dev/null +++ b/src/snapshots/antlr4_runtime__parser__tests__deferred_alternatives_preserve_left_recursive_contexts.snap @@ -0,0 +1,18 @@ +--- +source: src/parser.rs +expression: contexts +--- +[ + ( + 0, + 0, + 6, + "1+2", + ), + ( + 0, + 0, + 1, + "1", + ), +] diff --git a/src/snapshots/antlr4_runtime__parser__tests__fast_recognizer_preserves_labeled_left_recursive_operator_context.snap b/src/snapshots/antlr4_runtime__parser__tests__fast_recognizer_preserves_labeled_left_recursive_operator_context.snap new file mode 100644 index 00000000..92faefe2 --- /dev/null +++ b/src/snapshots/antlr4_runtime__parser__tests__fast_recognizer_preserves_labeled_left_recursive_operator_context.snap @@ -0,0 +1,16 @@ +--- +source: src/parser.rs +expression: contexts +--- +[ + ( + true, + 2, + "a+b", + ), + ( + false, + 1, + "a", + ), +] diff --git a/src/snapshots/antlr4_runtime__parser__tests__private_context_alt_tracking_keeps_fast_predicate_recognition.snap b/src/snapshots/antlr4_runtime__parser__tests__private_context_alt_tracking_keeps_fast_predicate_recognition.snap new file mode 100644 index 00000000..f309ab8c --- /dev/null +++ b/src/snapshots/antlr4_runtime__parser__tests__private_context_alt_tracking_keeps_fast_predicate_recognition.snap @@ -0,0 +1,9 @@ +--- +source: src/parser.rs +expression: "(root.alt_number(), root.context_alt_number(), root.text())" +--- +( + 0, + 2, + "x", +) diff --git a/tools/parse-bench/README.md b/tools/parse-bench/README.md index dfda2f67..9a5bc644 100644 --- a/tools/parse-bench/README.md +++ b/tools/parse-bench/README.md @@ -68,6 +68,19 @@ The script regenerates parsers into `target/parse-bench`, builds: - a Go ANTLR runner using `github.com/antlr4-go/antlr/v4`, - a tree-sitter runner using `tree-sitter-language-pack`. +Rust generation keeps the pinned `JavaParser.g4` unchanged, including its +`JavaParserBase` option and both semantic predicates. The generated benchmark +parser uses the source compiler's historical `assume-true` policy for those +target-language helpers, matching downstream generation while retaining the +predicate metadata that controls runtime routing. Python and Go still use the +equivalent portable rewrite because that grammars-v4 revision does not provide +`JavaParserBase` implementations for those targets. The +`issue-174-return-expression.java` fixture guards the resulting Java +method-body performance path. JSON rows include a benchmark-variant tag, so +the comparator skips only method changes such as this legacy-to-predicate +transition and resumes Java regression checks once both reports use the same +variant. + The output table reports `min` and `avg` parse time per fixture and a relative ratio against `rust-antlr` for the same fixture. Use `--rust-generated-only` for Adaptive LL delivery evidence so the Rust diff --git a/tools/parse-bench/compare.py b/tools/parse-bench/compare.py index 6619bb04..a44d97c9 100755 --- a/tools/parse-bench/compare.py +++ b/tools/parse-bench/compare.py @@ -9,6 +9,9 @@ from pathlib import Path +DEFAULT_BENCHMARK_VARIANT = "default" + + def result_key(result: dict) -> tuple[str, str, str]: return ( str(result["language"]), @@ -28,6 +31,10 @@ def load_results(path: Path) -> dict[tuple[str, str, str], dict]: return indexed +def result_variant(result: dict) -> str: + return str(result.get("benchmark_variant", DEFAULT_BENCHMARK_VARIANT)) + + def parse_speedup_requirement(value: str) -> tuple[str, str, str, float]: parts = value.split(":") if len(parts) != 4: @@ -131,11 +138,20 @@ def main() -> int: runtimes = set(args.runtime or ["rust-antlr"]) regression_failures: list[str] = [] + variant_mismatches: list[tuple[tuple[str, str, str], str, str]] = [] + compared = 0 for key, head in sorted(current.items()): language, fixture, runtime = key - if runtime not in runtimes or key not in baseline: + base = baseline.get(key) + if runtime not in runtimes or base is None: + continue + base_variant = result_variant(base) + head_variant = result_variant(head) + if base_variant != head_variant: + variant_mismatches.append((key, base_variant, head_variant)) continue - base_avg = float(baseline[key]["avg_ns"]) + compared += 1 + base_avg = float(base["avg_ns"]) head_avg = float(head["avg_ns"]) if base_avg <= 0: continue @@ -148,17 +164,22 @@ def main() -> int: ) speedup_failures = check_speedup_requirements(current, args.require_speedup) - compared = sum( - 1 - for key in current - if key in baseline and key[2] in runtimes - ) + if variant_mismatches: + print( + "parse benchmark compare skipped " + f"{len(variant_mismatches)} result(s) with changed benchmark variants:" + ) + for (language, fixture, runtime), base_variant, head_variant in variant_mismatches: + print( + f" {language}/{fixture} {runtime}: " + f"{base_variant} -> {head_variant}" + ) if compared == 0: message = ( "parse benchmark compare found no matching baseline/current " f"result pairs for runtime(s): {', '.join(sorted(runtimes))}" ) - if args.allow_empty: + if args.allow_empty or variant_mismatches: print(f"{message}; skipping regression comparison") else: print(message, file=sys.stderr) diff --git a/tools/parse-bench/fixtures/java/issue-174-return-expression.java b/tools/parse-bench/fixtures/java/issue-174-return-expression.java new file mode 100644 index 00000000..1e8bfcbf --- /dev/null +++ b/tools/parse-bench/fixtures/java/issue-174-return-expression.java @@ -0,0 +1,5 @@ +class C { + int m() { + return 1; + } +} diff --git a/tools/parse-bench/fixtures/manifest.json b/tools/parse-bench/fixtures/manifest.json index 6cb70133..8299641a 100644 --- a/tools/parse-bench/fixtures/manifest.json +++ b/tools/parse-bench/fixtures/manifest.json @@ -92,6 +92,16 @@ "lex" ] }, + { + "language": "java", + "path": "java/issue-174-return-expression.java", + "source": "repository regression benchmark for issue #174", + "license": "BSD-3-Clause", + "description": "Minimal Java method-body return expression that exposes predicate-aware interpreted-rule routing regressions.", + "phases": [ + "parse" + ] + }, { "language": "java", "path": "java/mojang-data-result.java", diff --git a/tools/parse-bench/run.py b/tools/parse-bench/run.py index f62d71c4..2924c74f 100755 --- a/tools/parse-bench/run.py +++ b/tools/parse-bench/run.py @@ -25,6 +25,8 @@ # ANTLR 4.13.2 still generates Go imports for github.com/antlr4-go/antlr/v4, # whose latest published module tag is v4.13.1. GO_ANTLR_RUNTIME = "v4.13.1" +DEFAULT_BENCHMARK_VARIANT = "default" +JAVA_RUST_PREDICATE_VARIANT = "java-upstream-parser-predicates-v1" @dataclasses.dataclass(frozen=True) @@ -144,6 +146,7 @@ class Measurement: language: str fixture: str runtime: str + benchmark_variant: str min_ns: int avg_ns: int bytes: int @@ -257,7 +260,7 @@ def predicate_replacement(match: re.Match[str]) -> str: return grammar.replace("this.", "l.") -def transform_java(grammar_dir: Path) -> None: +def transform_java_for_portable_target(grammar_dir: Path) -> None: parser = grammar_dir / "JavaParser.g4" text = parser.read_text() text = text.replace(" superClass = JavaParserBase;\n", "") @@ -274,6 +277,31 @@ def transform_java(grammar_dir: Path) -> None: parser.write_text(text) +def prepare_runtime_grammar( + spec: LanguageSpec, + grammars_v4: Path, + target: Path, + runtime: str, +) -> None: + copy_grammar(spec, grammars_v4, target) + if spec.name == "csharp": + if runtime == "python-antlr": + transform_csharp_python(target) + elif runtime == "go-antlr": + transform_csharp_go(target) + elif spec.name == "java" and runtime != "rust-antlr": + # The pinned grammars-v4 revision has no Python or Go JavaParserBase. + # Keep their equivalent portable rewrite, but retain the untouched + # predicates for Rust so this benchmark covers runtime semantic routing. + transform_java_for_portable_target(target) + + +def benchmark_variant(language: str, runtime: str, phase: str) -> str: + if phase == "parse" and language == "java" and runtime == "rust-antlr": + return JAVA_RUST_PREDICATE_VARIANT + return DEFAULT_BENCHMARK_VARIANT + + def generate_antlr( antlr_jar: Path, spec: LanguageSpec, @@ -1217,9 +1245,7 @@ def prepare_work( for spec in specs: if "rust-antlr" in runtimes: base_grammar = work_dir / "grammars" / spec.name / "base" - copy_grammar(spec, args.grammars_v4, base_grammar) - if spec.name == "java": - transform_java(base_grammar) + prepare_runtime_grammar(spec, args.grammars_v4, base_grammar, "rust-antlr") generate_rust_modules( spec, base_grammar, @@ -1230,22 +1256,14 @@ def prepare_work( if "python-antlr" in runtimes: py_grammar = work_dir / "grammars" / spec.name / "python" - copy_grammar(spec, args.grammars_v4, py_grammar) - if spec.name == "csharp": - transform_csharp_python(py_grammar) - if spec.name == "java": - transform_java(py_grammar) + prepare_runtime_grammar(spec, args.grammars_v4, py_grammar, "python-antlr") py_lang_gen = py_gen generate_antlr(args.antlr_jar, spec, py_grammar, py_lang_gen, "Python3") prepare_python_support(spec, args.grammars_v4, py_lang_gen) if "go-antlr" in runtimes: go_grammar = work_dir / "grammars" / spec.name / "go" - copy_grammar(spec, args.grammars_v4, go_grammar) - if spec.name == "csharp": - transform_csharp_go(go_grammar) - if spec.name == "java": - transform_java(go_grammar) + prepare_runtime_grammar(spec, args.grammars_v4, go_grammar, "go-antlr") go_gen = work_dir / "generated" / spec.name / "go" package_name = go_package_name(spec) generate_antlr(args.antlr_jar, spec, go_grammar, go_gen, "Go", package_name) @@ -1361,6 +1379,7 @@ def measure_fixture( language=fixture.language, fixture=fixture.name, runtime=runtime, + benchmark_variant=benchmark_variant(fixture.language, runtime, args.phase), min_ns=int(match.group("min")), avg_ns=int(match.group("avg")), bytes=fixture.abs_path.stat().st_size, diff --git a/tools/parse-bench/test_run.py b/tools/parse-bench/test_run.py index adb4ccb9..886bfbe1 100644 --- a/tools/parse-bench/test_run.py +++ b/tools/parse-bench/test_run.py @@ -1,8 +1,12 @@ +import contextlib import importlib.util +import io +import json import sys import tempfile import unittest from pathlib import Path +from unittest import mock RUN_PATH = Path(__file__).with_name("run.py") @@ -12,6 +16,13 @@ sys.modules[SPEC.name] = RUN SPEC.loader.exec_module(RUN) +COMPARE_PATH = Path(__file__).with_name("compare.py") +COMPARE_SPEC = importlib.util.spec_from_file_location("parse_bench_compare", COMPARE_PATH) +assert COMPARE_SPEC is not None and COMPARE_SPEC.loader is not None +COMPARE = importlib.util.module_from_spec(COMPARE_SPEC) +sys.modules[COMPARE_SPEC.name] = COMPARE +COMPARE_SPEC.loader.exec_module(COMPARE) + class DumpTreeHelpersTests(unittest.TestCase): def generated_sources(self) -> tuple[str, str]: @@ -126,6 +137,174 @@ def test_allows_disjoint_and_nested_work_directories(self) -> None: self.assertTrue(marker.exists()) +class RuntimeGrammarPreparationTests(unittest.TestCase): + JAVA_PARSER = """parser grammar JavaParser; +options { + tokenVocab = JavaLexer; + superClass = JavaParserBase; +} +annotationFieldValue: + { this.IsNotIdentifierAssign() }? annotationValue + | identifier '=' annotationValue + ; +recordComponentList + : recordComponent (',' recordComponent)* { this.DoLastRecordComponent() }? + ; +""" + + def prepare_java(self, root: Path, runtime: str) -> str: + grammar_source = root / "grammars-v4" / "java" / "java" + grammar_source.mkdir(parents=True) + (grammar_source / "JavaLexer.g4").write_text("lexer grammar JavaLexer;\n") + (grammar_source / "JavaParser.g4").write_text(self.JAVA_PARSER) + target = root / runtime + RUN.prepare_runtime_grammar( + RUN.LANGUAGES["java"], + root / "grammars-v4", + target, + runtime, + ) + return (target / "JavaParser.g4").read_text() + + def test_rust_java_grammar_preserves_superclass_and_predicates(self) -> None: + with tempfile.TemporaryDirectory() as temp: + parser = self.prepare_java(Path(temp), "rust-antlr") + + self.assertEqual(parser, self.JAVA_PARSER) + + def test_portable_java_targets_remove_unavailable_base_predicates(self) -> None: + for runtime in ("python-antlr", "go-antlr"): + with self.subTest(runtime=runtime), tempfile.TemporaryDirectory() as temp: + parser = self.prepare_java(Path(temp), runtime) + + self.assertNotIn("superClass = JavaParserBase", parser) + self.assertNotIn("IsNotIdentifierAssign", parser) + self.assertNotIn("DoLastRecordComponent", parser) + self.assertIn("identifier '=' annotationValue", parser) + + +class BenchmarkVariantTests(unittest.TestCase): + @staticmethod + def result( + language: str, + fixture: str, + avg_ns: int, + variant: str | None = None, + ) -> dict[str, object]: + result: dict[str, object] = { + "language": language, + "fixture": fixture, + "runtime": "rust-antlr", + "avg_ns": avg_ns, + } + if variant is not None: + result["benchmark_variant"] = variant + return result + + def compare( + self, + baseline: list[dict[str, object]], + current: list[dict[str, object]], + ) -> tuple[int, str, str]: + with tempfile.TemporaryDirectory() as temp: + root = Path(temp) + baseline_path = root / "baseline.json" + current_path = root / "current.json" + baseline_path.write_text(json.dumps({"results": baseline})) + current_path.write_text(json.dumps({"results": current})) + stdout = io.StringIO() + stderr = io.StringIO() + argv = [ + "compare.py", + "--baseline", + str(baseline_path), + "--current", + str(current_path), + "--max-regression", + "1.15", + ] + with ( + mock.patch.object(sys, "argv", argv), + contextlib.redirect_stdout(stdout), + contextlib.redirect_stderr(stderr), + ): + status = COMPARE.main() + return status, stdout.getvalue(), stderr.getvalue() + + def test_java_rust_parse_results_use_predicate_grammar_variant(self) -> None: + self.assertEqual( + RUN.benchmark_variant("java", "rust-antlr", "parse"), + RUN.JAVA_RUST_PREDICATE_VARIANT, + ) + self.assertEqual( + RUN.benchmark_variant("java", "rust-antlr", "lex"), + RUN.DEFAULT_BENCHMARK_VARIANT, + ) + self.assertEqual( + RUN.benchmark_variant("java", "go-antlr", "parse"), + RUN.DEFAULT_BENCHMARK_VARIANT, + ) + + def test_compare_skips_only_changed_benchmark_variants(self) -> None: + baseline = [ + self.result("java", "Example.java", 100), + self.result("kotlin", "Example.kt", 100), + ] + current = [ + self.result( + "java", + "Example.java", + 1_000, + RUN.JAVA_RUST_PREDICATE_VARIANT, + ), + self.result("kotlin", "Example.kt", 100), + ] + + status, stdout, stderr = self.compare(baseline, current) + + self.assertEqual(status, 0) + self.assertIn("skipped 1 result(s)", stdout) + self.assertIn("passed: 1 result(s)", stdout) + self.assertEqual(stderr, "") + + def test_compare_allows_only_changed_benchmark_variants(self) -> None: + baseline = [self.result("java", "Example.java", 100)] + current = [ + self.result( + "java", + "Example.java", + 1_000, + RUN.JAVA_RUST_PREDICATE_VARIANT, + ) + ] + + status, stdout, stderr = self.compare(baseline, current) + + self.assertEqual(status, 0) + self.assertIn("skipped 1 result(s)", stdout) + self.assertIn("skipping regression comparison", stdout) + self.assertEqual(stderr, "") + + def test_compare_rejects_unrelated_result_sets(self) -> None: + baseline = [self.result("java", "Example.java", 100)] + current = [self.result("kotlin", "Example.kt", 100)] + + status, _, stderr = self.compare(baseline, current) + + self.assertEqual(status, 1) + self.assertIn("no matching baseline/current result pairs", stderr) + + def test_compare_enforces_threshold_for_matching_variants(self) -> None: + variant = RUN.JAVA_RUST_PREDICATE_VARIANT + baseline = [self.result("java", "Example.java", 100, variant)] + current = [self.result("java", "Example.java", 1_000, variant)] + + status, _, stderr = self.compare(baseline, current) + + self.assertEqual(status, 1) + self.assertIn("10.00x", stderr) + + class RustCodegenFlagsTests(unittest.TestCase): def test_combines_native_and_profile_generation(self) -> None: self.assertEqual(