use tree_sitter::{Node, Parser, Tree};
const MAX_ITERATIONS: u32 = 8;
const C_KEYWORDS: &[&str] = &[
"auto",
"break",
"case",
"char",
"const",
"continue",
"default",
"do",
"double",
"else",
"enum",
"extern",
"float",
"for",
"goto",
"if",
"inline",
"int",
"long",
"register",
"restrict",
"return",
"short",
"signed",
"sizeof",
"static",
"struct",
"switch",
"typedef",
"union",
"unsigned",
"void",
"volatile",
"while",
"_Alignas",
"_Alignof",
"_Atomic",
"_Bool",
"_Complex",
"_Generic",
"_Imaginary",
"_Noreturn",
"_Static_assert",
"_Thread_local",
];
fn is_bare_identifier(text: &str) -> bool {
if C_KEYWORDS.contains(&text) {
return false;
}
let mut chars = text.chars();
match chars.next() {
Some(c) if c.is_ascii_alphabetic() || c == '_' => {}
_ => return false,
}
chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
}
fn find_blankable_identifier_error(node: &Node, source: &str) -> Option<(usize, usize)> {
if node.kind() == "ERROR" {
let start = node.start_byte();
let end = node.end_byte();
let text = &source[start..end];
let trimmed = text.trim();
if !trimmed.is_empty() && is_bare_identifier(trimmed) {
let leading = text.len() - text.trim_start().len();
let trailing = text.len() - text.trim_end().len();
return Some((start + leading, end - trailing));
}
}
for i in 0..node.child_count() {
if let Some(child) = node.child(i) {
if let Some(found) = find_blankable_identifier_error(&child, source) {
return Some(found);
}
}
}
None
}
fn blank_range(source: &str, start: usize, end: usize) -> String {
let mut bytes = source.as_bytes().to_vec();
for b in bytes.iter_mut().take(end).skip(start) {
if *b != b'\n' && *b != b'\r' {
*b = b' ';
}
}
String::from_utf8(bytes).unwrap_or_else(|_| source.to_string())
}
fn blank_or_mark_noreturn(source: &str, start: usize, end: usize) -> String {
let trimmed = source[start..end].trim();
if crate::analyze::noreturn::NORETURN_ATTRIBUTE_MACRO_NAMES.contains(&trimmed) {
if let Some(marked) = crate::analyze::noreturn::write_marker(source, start, end) {
return marked;
}
}
blank_range(source, start, end)
}
fn find_blankable_preproc_brace_error(
node: &Node,
source: &str,
) -> Option<((usize, usize), (usize, usize))> {
if matches!(
node.kind(),
"preproc_if" | "preproc_ifdef" | "preproc_elif" | "preproc_elifdef"
) {
let children: Vec<Node> = (0..node.child_count())
.filter_map(|i| node.child(i))
.collect();
if let Some(endif_idx) = children
.iter()
.position(|c| !c.is_named() && c.kind() == "#endif")
{
let mut idx = endif_idx;
let mut error_child = None;
while idx > 0 {
idx -= 1;
let c = &children[idx];
if c.kind() == "comment" {
continue;
}
if c.is_error() {
let trimmed = source[c.start_byte()..c.end_byte()].trim();
if trimmed == "{" || trimmed == "}" {
error_child = Some(*c);
}
}
break;
}
if let Some(err) = error_child {
return Some((
(node.start_byte(), err.start_byte()),
(
children[endif_idx].start_byte(),
children[endif_idx].end_byte(),
),
));
}
}
}
for i in 0..node.child_count() {
if let Some(child) = node.child(i) {
if let Some(found) = find_blankable_preproc_brace_error(&child, source) {
return Some(found);
}
}
}
None
}
pub fn parse_with_recovery(parser: &mut Parser, source: String) -> Option<(Tree, String)> {
let mut text = source;
let mut tree = parser.parse(&text, None)?;
for _ in 0..MAX_ITERATIONS {
if !tree.root_node().has_error() {
break;
}
if let Some((start, end)) = find_blankable_identifier_error(&tree.root_node(), &text) {
text = blank_or_mark_noreturn(&text, start, end);
} else if let Some(((s1, e1), (s2, e2))) =
find_blankable_preproc_brace_error(&tree.root_node(), &text)
{
text = blank_range(&text, s1, e1);
text = blank_range(&text, s2, e2);
} else {
break;
}
tree = parser.parse(&text, None)?;
}
Some((tree, text))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parser::c_language;
fn recover(src: &str) -> (bool, String) {
let mut parser = Parser::new();
parser.set_language(&c_language()).unwrap();
let (tree, text) = parse_with_recovery(&mut parser, src.to_string()).unwrap();
(tree.root_node().has_error(), text)
}
#[test]
fn recovers_unknown_calling_convention_in_funcptr_typedef() {
let src = "typedef void (GL_APIENTRYP PFNGLFOO)(int x);\nint y = 1;\n";
let (has_error, _) = recover(src);
assert!(!has_error);
}
#[test]
fn preserves_byte_length() {
let src = "typedef void (GL_APIENTRYP PFNGLFOO)(int x);\nint y = 1;\n";
let (_, text) = recover(src);
assert_eq!(text.len(), src.len());
let pos_orig = src.find("int y = 1;").unwrap();
let pos_out = text.find("int y = 1;").unwrap();
assert_eq!(pos_orig, pos_out);
}
#[test]
fn clean_file_untouched_and_no_reparse_cost() {
let src = "int main(void) { return 0; }\n";
let (has_error, text) = recover(src);
assert!(!has_error);
assert_eq!(text, src);
}
#[test]
fn never_blanks_a_c_keyword_stranded_in_an_error_node() {
let src = "#define STATIC static\nSTATIC void f(void) {\n}\n";
let (_, text) = recover(src);
assert!(text.contains("void f(void)"));
}
#[test]
fn does_not_blank_a_meaningful_enum_like_macro() {
let src = "#define MAX_SIZE 32\nint arr[MAX_SIZE];\n";
let (has_error, text) = recover(src);
assert!(!has_error);
assert_eq!(text, src);
}
#[test]
fn recovers_cplusplus_guarded_extern_c_brace() {
let src = "int x;\n#if defined(__cplusplus)\n}\n#endif\n";
let (_, text) = recover(src);
assert!(!text.contains("__cplusplus"));
assert!(text.contains('}'));
}
#[test]
fn does_not_corrupt_whole_file_header_guard_on_switch_in_macro_error() {
let src = concat!(
"#ifndef UTHASH_H\n",
"#define UTHASH_H\n",
"#define HASH_SFH(key,keylen,hashv) \\\n",
"do { \\\n",
" unsigned const char *_sfh_key=(unsigned const char*)(key); \\\n",
" uint32_t _sfh_tmp, _sfh_len = (uint32_t)keylen; \\\n",
" unsigned _sfh_rem = _sfh_len & 3U; \\\n",
" _sfh_len >>= 2; \\\n",
" hashv = 0xcafebabeu; \\\n",
" for (;_sfh_len > 0U; _sfh_len--) { \\\n",
" hashv += get16bits (_sfh_key); \\\n",
" _sfh_tmp = ((uint32_t)(get16bits (_sfh_key+2)) << 11) ^ hashv; \\\n",
" hashv = (hashv << 16) ^ _sfh_tmp; \\\n",
" _sfh_key += 2U*sizeof (uint16_t); \\\n",
" hashv += hashv >> 11; \\\n",
" } \\\n",
" switch (_sfh_rem) { \\\n",
" case 3: hashv += get16bits (_sfh_key); \\\n",
" hashv ^= hashv << 16; \\\n",
" hashv ^= (uint32_t)(_sfh_key[sizeof (uint16_t)]) << 18; \\\n",
" hashv += hashv >> 11; \\\n",
" break; \\\n",
" case 2: hashv += get16bits (_sfh_key); \\\n",
" hashv ^= hashv << 11; \\\n",
" hashv += hashv >> 17; \\\n",
" break; \\\n",
" case 1: hashv += *_sfh_key; \\\n",
" hashv ^= hashv << 10; \\\n",
" hashv += hashv >> 1; \\\n",
" break; \\\n",
" default: ; \\\n",
" } \\\n",
" hashv ^= hashv << 3; \\\n",
" hashv += hashv >> 5; \\\n",
" hashv ^= hashv << 4; \\\n",
" hashv += hashv >> 17; \\\n",
" hashv ^= hashv << 25; \\\n",
" hashv += hashv >> 6; \\\n",
"} while (0)\n",
"int tail;\n",
"#endif\n",
);
let (_, text) = recover(src);
assert_eq!(text, src, "recovery must not alter this file at all");
}
#[test]
fn preproc_brace_recovery_preserves_byte_length_and_line_count() {
let src = "int x;\n#if defined(__cplusplus)\n}\n#endif\nint y;\n";
let (_, text) = recover(src);
assert_eq!(text.len(), src.len());
assert_eq!(text.matches('\n').count(), src.matches('\n').count());
let pos_orig = src.find("int y;").unwrap();
let pos_out = text.find("int y;").unwrap();
assert_eq!(pos_orig, pos_out);
}
#[test]
fn stops_after_max_iterations_without_hanging() {
let mut src = String::new();
for i in 0..20 {
src.push_str(&format!(
"typedef void (UNKNOWNCONV{i} PFNFOO{i})(int x);\n"
));
}
let mut parser = Parser::new();
parser.set_language(&c_language()).unwrap();
let (_, text) = parse_with_recovery(&mut parser, src.clone()).unwrap();
assert_eq!(text.len(), src.len());
}
}