1use crate::ast::{Expr, Grammar, GrammarItem, RuleDecl};
4use crate::diagnostic::{Diagnostic, DiagnosticKind};
5use crate::error::RantlrError;
6use crate::fix_hint::FixHint;
7
8pub fn lint(source: &str, grammar: &Grammar) -> Result<(), RantlrError> {
10 let diags = lint_all(source, grammar);
11 if let Some(diag) = diags.into_iter().next() {
12 return Err(RantlrError::from_diagnostic(source, diag));
13 }
14 Ok(())
15}
16
17pub fn lint_all(source: &str, grammar: &Grammar) -> Vec<Diagnostic> {
19 let mut out = Vec::new();
20 for item in &grammar.items {
21 if let GrammarItem::Rule(rule) = item {
22 if let Some(diag) = check_direct_left_recursion(source, rule) {
23 out.push(diag);
24 }
25 }
26 }
27 out
28}
29
30fn check_direct_left_recursion(source: &str, rule: &RuleDecl) -> Option<Diagnostic> {
32 if !leftmost_refs(&rule.body).iter().any(|r| r == &rule.name) {
33 return None;
34 }
35
36 let message = format!(
37 "Rule '{}' is left-recursive. In Rantlr, use 'repeat {{ ... }}' instead of calling the rule at the start to ensure better performance and avoid stack overflow.",
38 rule.name
39 );
40
41 let fix = suggest_left_recursion_fix(source, rule);
42
43 let mut diag = Diagnostic::error(DiagnosticKind::Lint, message, rule.name_span)
44 .with_code("left-recursion")
45 .with_help(format!(
46 "rewrite `{name}` so it starts with a different symbol, then use `repeat {{ ... }}` for the recursive suffix",
47 name = rule.name
48 ));
49
50 if let Some(fix) = fix {
51 diag = diag.with_fix(fix);
52 }
53
54 Some(diag)
55}
56
57fn suggest_left_recursion_fix(source: &str, rule: &RuleDecl) -> Option<FixHint> {
58 let start = rule.span.start.as_usize().min(source.len());
59 let end = rule.span.end.as_usize().min(source.len());
60 if start >= end {
61 return None;
62 }
63 let before = source[start..end].trim_end().to_string();
64
65 let after = match &rule.body {
66 Expr::Seq { items } if matches!(items.first(), Some(Expr::Ref { name, .. }) if name == &rule.name)
68 && items.len() >= 2 =>
69 {
70 let rest = &items[1..];
71 let base = infer_base_symbol(rest, &rule.name)
72 .unwrap_or_else(|| "/* your_base_case */".into());
73 let suffix = rest
74 .iter()
75 .map(expr_to_gr)
76 .collect::<Vec<_>>()
77 .join("\n ");
78 format!(
79 "rule {} {{\n {}\n repeat {{\n {}\n }}\n}}",
80 rule.name, base, suffix
81 )
82 }
83 Expr::Match { arms } => {
85 let non_self: Vec<&Expr> = arms
86 .iter()
87 .filter(|a| !leftmost_refs(a).iter().any(|r| *r == rule.name.as_str()))
88 .collect();
89 if non_self.is_empty() {
90 format!(
91 "rule {} {{\n // TODO: add a base case that does NOT call `{0}` first\n /* base */\n}}",
92 rule.name
93 )
94 } else if non_self.len() == 1 {
95 format!(
96 "rule {} {{\n {}\n}}",
97 rule.name,
98 expr_to_gr(non_self[0])
99 )
100 } else {
101 let arms_txt = non_self
102 .iter()
103 .map(|a| expr_to_gr(a))
104 .collect::<Vec<_>>()
105 .join(" | ");
106 format!("rule {} {{\n match {}\n}}", rule.name, arms_txt)
107 }
108 }
109 _ => format!(
110 "rule {} {{\n // Start with something that is NOT `{0}`, then loop:\n /* base */\n repeat {{\n /* recursive suffix */\n }}\n}}",
111 rule.name
112 ),
113 };
114
115 Some(FixHint::new(
116 "Rewrite left recursion with repeat",
117 format!(
118 "Whoops! Rule `{name}` is calling itself too early — that's a left-recursion loop. Click Fix to turn it into a `repeat` block that starts with a base case first.",
119 name = rule.name
120 ),
121 before,
122 after,
123 rule.span,
124 source,
125 ))
126}
127
128fn infer_base_symbol(rest: &[Expr], self_name: &str) -> Option<String> {
129 for expr in rest.iter().rev() {
131 if let Expr::Ref { name, .. } = expr {
132 if name != self_name {
133 return Some(name.clone());
134 }
135 }
136 }
137 None
138}
139
140fn expr_to_gr(expr: &Expr) -> String {
141 match expr {
142 Expr::Ref { name, .. } => name.clone(),
143 Expr::Literal { value, .. } => format!("\"{}\"", escape_gr_string(value)),
144 Expr::Group { body } => format!("({})", expr_to_gr(body)),
145 Expr::Seq { items } => items.iter().map(expr_to_gr).collect::<Vec<_>>().join(" "),
146 Expr::Match { arms } => {
147 let inner = arms
148 .iter()
149 .map(expr_to_gr)
150 .collect::<Vec<_>>()
151 .join(" | ");
152 format!("match {inner}")
153 }
154 Expr::Alt { alts } => alts
155 .iter()
156 .map(expr_to_gr)
157 .collect::<Vec<_>>()
158 .join(" | "),
159 Expr::Optional { body } => format!("optional {{ {} }}", expr_to_gr(body)),
160 Expr::Repeat { min, max, body } => {
161 let bounds = match (min, max) {
162 (None, None) => String::new(),
163 (Some(a), None) => format!("({a}..)"),
164 (Some(a), Some(b)) => format!("({a}..{b})"),
165 (None, Some(b)) => format!("(..{b})"),
166 };
167 format!("repeat{bounds} {{ {} }}", expr_to_gr(body))
168 }
169 }
170}
171
172fn escape_gr_string(s: &str) -> String {
173 s.chars()
174 .map(|c| match c {
175 '\\' => "\\\\".into(),
176 '"' => "\\\"".into(),
177 '\n' => "\\n".into(),
178 '\t' => "\\t".into(),
179 '\r' => "\\r".into(),
180 c => c.to_string(),
181 })
182 .collect()
183}
184
185fn leftmost_refs(expr: &Expr) -> Vec<&str> {
187 match expr {
188 Expr::Alt { alts } => alts.iter().flat_map(leftmost_refs).collect(),
189 Expr::Seq { items } => {
190 let mut refs = Vec::new();
191 for item in items {
192 refs.extend(leftmost_refs(item));
193 if !is_nullable(item) {
194 break;
195 }
196 }
197 refs
198 }
199 Expr::Match { arms } => arms.iter().flat_map(leftmost_refs).collect(),
200 Expr::Repeat { body, .. } | Expr::Optional { body } | Expr::Group { body } => {
201 leftmost_refs(body)
202 }
203 Expr::Ref { name, .. } => vec![name.as_str()],
204 Expr::Literal { .. } => vec![],
205 }
206}
207
208fn is_nullable(expr: &Expr) -> bool {
209 match expr {
210 Expr::Optional { .. } => true,
211 Expr::Repeat { min, .. } => matches!(min, None | Some(0)),
212 Expr::Group { body } => is_nullable(body),
213 Expr::Alt { alts } => alts.iter().any(is_nullable),
214 Expr::Seq { items } => items.iter().all(is_nullable),
215 Expr::Match { arms } => arms.iter().any(is_nullable),
216 Expr::Ref { .. } | Expr::Literal { .. } => false,
217 }
218}
219
220#[cfg(test)]
221mod tests {
222 use super::*;
223 use crate::parser::parse;
224
225 #[test]
226 fn accepts_calculator() {
227 let src = include_str!("../testdata/calculator.gr");
228 let g = parse(src).unwrap();
229 assert!(lint_all(src, &g).is_empty());
230 }
231
232 #[test]
233 fn flags_direct_left_recursion_with_fix() {
234 let src = r#"
235 grammar Bad;
236 token N = Number;
237 rule expr {
238 expr "+" term
239 }
240 rule term {
241 N
242 }
243 "#;
244 let g = parse(src).unwrap();
245 let diags = lint_all(src, &g);
246 assert_eq!(diags.len(), 1);
247 assert_eq!(diags[0].code.as_deref(), Some("left-recursion"));
248 let fix = diags[0].fix.as_ref().expect("fix hint");
249 assert!(fix.after.contains("repeat"));
250 assert!(fix.after.contains("term"));
251 assert!(fix.after.contains("\"+\""));
252 assert!(fix.after.contains("rule expr {\n term\n"));
254 assert!(!fix.suggestion.is_empty());
255 assert_eq!(fix.suggestion, fix.after);
256 assert!(fix.start < fix.end);
257 assert!(fix.start_line >= 1);
258 assert!(fix.explanation.to_lowercase().contains("left-recursion"));
259 }
260
261 #[test]
262 fn flags_left_recursion_through_match() {
263 let src = r#"
264 grammar Bad;
265 rule expr {
266 match expr | "x"
267 }
268 "#;
269 let g = parse(src).unwrap();
270 let diags = lint_all(src, &g);
271 assert_eq!(diags.len(), 1);
272 assert!(diags[0].message.contains("left-recursive"));
273 let fix = diags[0].fix.as_ref().unwrap();
274 assert!(fix.after.contains("\"x\""));
275 }
276
277 #[test]
278 fn allows_right_recursion_and_repeat() {
279 let src = r#"
280 grammar Ok;
281 token N = Number;
282 rule expr {
283 term
284 repeat {
285 match "+" | "-"
286 term
287 }
288 }
289 rule term {
290 N
291 }
292 rule list {
293 item
294 optional { list }
295 }
296 rule item { N }
297 "#;
298 let g = parse(src).unwrap();
299 assert!(lint_all(src, &g).is_empty());
300 }
301}