rantlr_core/
rule_graph.rs1use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
4
5use serde::Serialize;
6
7use crate::ast::{Expr, Grammar, GrammarItem};
8
9#[derive(Debug, Clone, Serialize)]
10#[serde(rename_all = "camelCase")]
11pub struct RuleNode {
12 pub id: String,
13 pub label: String,
14 pub self_recursive: bool,
16}
17
18#[derive(Debug, Clone, Serialize)]
19#[serde(rename_all = "camelCase")]
20pub struct RuleEdge {
21 pub from: String,
22 pub to: String,
23 pub recursive: bool,
25}
26
27#[derive(Debug, Clone, Serialize)]
28#[serde(rename_all = "camelCase")]
29pub struct RuleGraph {
30 pub nodes: Vec<RuleNode>,
31 pub edges: Vec<RuleEdge>,
32}
33
34pub fn rule_graph(grammar: &Grammar) -> RuleGraph {
36 let rule_names: BTreeSet<String> = grammar
37 .items
38 .iter()
39 .filter_map(|i| match i {
40 GrammarItem::Rule(r) => Some(r.name.clone()),
41 _ => None,
42 })
43 .collect();
44
45 let mut adj: BTreeMap<String, BTreeSet<String>> = BTreeMap::new();
46 let mut order = Vec::new();
47
48 for item in &grammar.items {
49 if let GrammarItem::Rule(rule) = item {
50 order.push(rule.name.clone());
51 adj.entry(rule.name.clone()).or_default();
52 let mut refs = BTreeSet::new();
53 collect_rule_refs(&rule.body, &rule_names, &mut refs);
54 for r in refs {
55 adj.get_mut(&rule.name).unwrap().insert(r);
56 }
57 }
58 }
59
60 let cyclic_edges = find_cyclic_edges(&adj);
61
62 let nodes: Vec<RuleNode> = order
63 .iter()
64 .map(|name| {
65 let self_recursive = adj.get(name).map(|s| s.contains(name)).unwrap_or(false);
66 RuleNode {
67 id: name.clone(),
68 label: name.clone(),
69 self_recursive,
70 }
71 })
72 .collect();
73
74 let mut edges = Vec::new();
75 for (from, tos) in &adj {
76 for to in tos {
77 let recursive =
78 cyclic_edges.contains(&(from.clone(), to.clone())) || from == to;
79 edges.push(RuleEdge {
80 from: from.clone(),
81 to: to.clone(),
82 recursive,
83 });
84 }
85 }
86
87 RuleGraph { nodes, edges }
88}
89
90fn collect_rule_refs(expr: &Expr, rules: &BTreeSet<String>, out: &mut BTreeSet<String>) {
91 match expr {
92 Expr::Ref { name, .. } => {
93 if rules.contains(name) {
94 out.insert(name.clone());
95 }
96 }
97 Expr::Literal { .. } => {}
98 Expr::Group { body } | Expr::Optional { body } | Expr::Repeat { body, .. } => {
99 collect_rule_refs(body, rules, out);
100 }
101 Expr::Seq { items } => {
102 for item in items {
103 collect_rule_refs(item, rules, out);
104 }
105 }
106 Expr::Match { arms } | Expr::Alt { alts: arms } => {
107 for arm in arms {
108 collect_rule_refs(arm, rules, out);
109 }
110 }
111 }
112}
113
114fn find_cyclic_edges(adj: &BTreeMap<String, BTreeSet<String>>) -> HashSet<(String, String)> {
115 let mut cyclic = HashSet::new();
116 let mut reach_cache: HashMap<String, HashSet<String>> = HashMap::new();
117 for n in adj.keys() {
118 reach_cache.insert(n.clone(), reachable(n, adj));
119 }
120
121 for (u, tos) in adj {
122 for v in tos {
123 if reach_cache.get(v).map(|s| s.contains(u)).unwrap_or(false) {
124 cyclic.insert((u.clone(), v.clone()));
125 }
126 }
127 }
128 cyclic
129}
130
131fn reachable(start: &str, adj: &BTreeMap<String, BTreeSet<String>>) -> HashSet<String> {
132 let mut seen = HashSet::new();
133 let mut stack = vec![start.to_string()];
134 while let Some(n) = stack.pop() {
135 if !seen.insert(n.clone()) {
136 continue;
137 }
138 if let Some(next) = adj.get(&n) {
139 for t in next {
140 stack.push(t.clone());
141 }
142 }
143 }
144 seen
145}
146
147#[cfg(test)]
148mod tests {
149 use super::*;
150 use crate::parser::parse;
151
152 #[test]
153 fn calculator_call_edges_and_paren_cycle() {
154 let src = include_str!("../testdata/calculator.gr");
155 let g = parse(src).unwrap();
156 let graph = rule_graph(&g);
157 assert!(graph.nodes.iter().any(|n| n.id == "expr"));
158 assert!(graph.edges.iter().any(|e| e.from == "expr" && e.to == "term"));
159 assert!(graph.edges.iter().any(|e| e.from == "term" && e.to == "factor"));
160 assert!(graph.edges.iter().any(|e| e.from == "factor" && e.to == "expr" && e.recursive));
162 }
163
164 #[test]
165 fn left_recursive_self_edge() {
166 let src = include_str!("../testdata/left_recursive.gr");
167 let g = parse(src).unwrap();
168 let graph = rule_graph(&g);
169 let expr = graph.nodes.iter().find(|n| n.id == "expr").unwrap();
170 assert!(expr.self_recursive);
171 assert!(graph.edges.iter().any(|e| e.from == "expr" && e.to == "expr" && e.recursive));
172 }
173}