use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
use serde::Serialize;
use crate::ast::{Expr, Grammar, GrammarItem};
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct RuleNode {
pub id: String,
pub label: String,
pub self_recursive: bool,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct RuleEdge {
pub from: String,
pub to: String,
pub recursive: bool,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct RuleGraph {
pub nodes: Vec<RuleNode>,
pub edges: Vec<RuleEdge>,
}
pub fn rule_graph(grammar: &Grammar) -> RuleGraph {
let rule_names: BTreeSet<String> = grammar
.items
.iter()
.filter_map(|i| match i {
GrammarItem::Rule(r) => Some(r.name.clone()),
_ => None,
})
.collect();
let mut adj: BTreeMap<String, BTreeSet<String>> = BTreeMap::new();
let mut order = Vec::new();
for item in &grammar.items {
if let GrammarItem::Rule(rule) = item {
order.push(rule.name.clone());
adj.entry(rule.name.clone()).or_default();
let mut refs = BTreeSet::new();
collect_rule_refs(&rule.body, &rule_names, &mut refs);
for r in refs {
adj.get_mut(&rule.name).unwrap().insert(r);
}
}
}
let cyclic_edges = find_cyclic_edges(&adj);
let nodes: Vec<RuleNode> = order
.iter()
.map(|name| {
let self_recursive = adj.get(name).map(|s| s.contains(name)).unwrap_or(false);
RuleNode {
id: name.clone(),
label: name.clone(),
self_recursive,
}
})
.collect();
let mut edges = Vec::new();
for (from, tos) in &adj {
for to in tos {
let recursive =
cyclic_edges.contains(&(from.clone(), to.clone())) || from == to;
edges.push(RuleEdge {
from: from.clone(),
to: to.clone(),
recursive,
});
}
}
RuleGraph { nodes, edges }
}
fn collect_rule_refs(expr: &Expr, rules: &BTreeSet<String>, out: &mut BTreeSet<String>) {
match expr {
Expr::Ref { name, .. } => {
if rules.contains(name) {
out.insert(name.clone());
}
}
Expr::Literal { .. } => {}
Expr::Group { body } | Expr::Optional { body } | Expr::Repeat { body, .. } => {
collect_rule_refs(body, rules, out);
}
Expr::Seq { items } => {
for item in items {
collect_rule_refs(item, rules, out);
}
}
Expr::Match { arms } | Expr::Alt { alts: arms } => {
for arm in arms {
collect_rule_refs(arm, rules, out);
}
}
}
}
fn find_cyclic_edges(adj: &BTreeMap<String, BTreeSet<String>>) -> HashSet<(String, String)> {
let mut cyclic = HashSet::new();
let mut reach_cache: HashMap<String, HashSet<String>> = HashMap::new();
for n in adj.keys() {
reach_cache.insert(n.clone(), reachable(n, adj));
}
for (u, tos) in adj {
for v in tos {
if reach_cache.get(v).map(|s| s.contains(u)).unwrap_or(false) {
cyclic.insert((u.clone(), v.clone()));
}
}
}
cyclic
}
fn reachable(start: &str, adj: &BTreeMap<String, BTreeSet<String>>) -> HashSet<String> {
let mut seen = HashSet::new();
let mut stack = vec![start.to_string()];
while let Some(n) = stack.pop() {
if !seen.insert(n.clone()) {
continue;
}
if let Some(next) = adj.get(&n) {
for t in next {
stack.push(t.clone());
}
}
}
seen
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parser::parse;
#[test]
fn calculator_call_edges_and_paren_cycle() {
let src = include_str!("../testdata/calculator.gr");
let g = parse(src).unwrap();
let graph = rule_graph(&g);
assert!(graph.nodes.iter().any(|n| n.id == "expr"));
assert!(graph.edges.iter().any(|e| e.from == "expr" && e.to == "term"));
assert!(graph.edges.iter().any(|e| e.from == "term" && e.to == "factor"));
assert!(graph.edges.iter().any(|e| e.from == "factor" && e.to == "expr" && e.recursive));
}
#[test]
fn left_recursive_self_edge() {
let src = include_str!("../testdata/left_recursive.gr");
let g = parse(src).unwrap();
let graph = rule_graph(&g);
let expr = graph.nodes.iter().find(|n| n.id == "expr").unwrap();
assert!(expr.self_recursive);
assert!(graph.edges.iter().any(|e| e.from == "expr" && e.to == "expr" && e.recursive));
}
}