use lang_parsing_substrate::query;
use std::collections::HashSet;
use tree_sitter::{Node, Parser, Tree};
use crate::analyze::context::ProjectContext;
#[derive(Debug, Default, Clone)]
pub struct RepairMacros {
pub object_macros: HashSet<String>,
pub unused_attribute_macros: HashSet<String>,
}
impl RepairMacros {
pub fn from_context(context: &ProjectContext) -> Self {
Self {
object_macros: HashSet::clone(&context.defined_macro_names),
unused_attribute_macros: HashSet::clone(&context.unused_attribute_macros),
}
}
}
pub const UNUSED_ATTRIBUTE_MARKER: &str = "/*U*/";
const MAX_ITERATIONS: u32 = 8;
use crate::utility::cert_c::ast_utils::is_c_keyword;
fn is_bare_identifier(text: &str) -> bool {
if is_c_keyword(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)> {
let found = query::find_first_descendant(*node, |n| {
n.kind() == "ERROR" && {
let trimmed = source[n.start_byte()..n.end_byte()].trim();
!trimmed.is_empty() && is_bare_identifier(trimmed)
}
})?;
let start = found.start_byte();
let end = found.end_byte();
let text = &source[start..end];
let leading = text.len() - text.trim_start().len();
let trailing = text.len() - text.trim_end().len();
Some((start + leading, end - trailing))
}
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(source: &str, start: usize, end: usize, macros: &RepairMacros) -> 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;
}
}
if macros.unused_attribute_macros.contains(trimmed) {
if let Some(marked) = write_padded_marker(source, start, end, UNUSED_ATTRIBUTE_MARKER) {
return marked;
}
}
blank_range(source, start, end)
}
fn write_padded_marker(source: &str, start: usize, end: usize, marker: &str) -> Option<String> {
let len = end - start;
if len < marker.len() {
return None;
}
let mut out = String::with_capacity(source.len());
out.push_str(&source[..start]);
out.push_str(marker);
out.push_str(&" ".repeat(len - marker.len()));
out.push_str(&source[end..]);
Some(out)
}
fn preceding_macro_tokens(
source: &str,
start: usize,
macros: &RepairMacros,
) -> Vec<(usize, usize)> {
let bytes = source.as_bytes();
let mut found = Vec::new();
let mut i = start;
for _ in 0..MAX_PRECEDING_TOKENS {
let scan_from = i;
while i > 0 && (bytes[i - 1] as char).is_ascii_whitespace() {
i -= 1;
}
let token_end = i;
if token_end == scan_from {
break;
}
while i > 0 && ((bytes[i - 1] as char).is_ascii_alphanumeric() || bytes[i - 1] == b'_') {
i -= 1;
}
let token_start = i;
if token_start == token_end {
break;
}
let token = &source[token_start..token_end];
if is_bare_identifier(token)
&& macros.object_macros.contains(token)
&& !line_is_preprocessor_directive(source, token_start)
{
found.push((token_start, token_end));
}
}
found
}
const MAX_PRECEDING_TOKENS: usize = 4;
fn line_is_preprocessor_directive(source: &str, byte: usize) -> bool {
let line_start = source[..byte].rfind('\n').map_or(0, |i| i + 1);
source[line_start..].trim_start().starts_with('#')
}
fn error_node_count(node: &Node) -> usize {
query::find_descendants(*node, |n| n.is_error() || n.is_missing()).len()
}
fn find_blankable_preproc_brace_error(
node: &Node,
source: &str,
) -> Option<((usize, usize), (usize, usize))> {
let conditionals = query::find_descendants_of_kinds(
*node,
&[
"preproc_if",
"preproc_ifdef",
"preproc_elif",
"preproc_elifdef",
],
);
conditionals
.into_iter()
.find_map(|node| lone_brace_guard_ranges(&node, source))
}
fn lone_brace_guard_ranges(node: &Node, source: &str) -> Option<((usize, usize), (usize, usize))> {
{
let mut cursor = node.walk();
let children: Vec<Node> = node.children(&mut cursor).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(),
),
));
}
}
}
None
}
pub fn parse_with_recovery(
parser: &mut Parser,
source: String,
macros: &RepairMacros,
) -> 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) {
let (repaired, repaired_tree) = choose_repair(parser, &text, start, end, macros)?;
text = repaired;
tree = repaired_tree;
continue;
} 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))
}
fn choose_repair(
parser: &mut Parser,
source: &str,
start: usize,
end: usize,
macros: &RepairMacros,
) -> Option<(String, Tree)> {
let stranded_text = blank_or_mark(source, start, end, macros);
let stranded_tree = parser.parse(&stranded_text, None)?;
if !macros.object_macros.contains(&source[start..end]) {
let stranded_errors = error_node_count(&stranded_tree.root_node());
for (macro_start, macro_end) in preceding_macro_tokens(source, start, macros) {
let macro_text = blank_or_mark(source, macro_start, macro_end, macros);
if let Some(macro_tree) = parser.parse(¯o_text, None) {
if error_node_count(¯o_tree.root_node()) <= stranded_errors {
return Some((macro_text, macro_tree));
}
}
}
}
Some((stranded_text, stranded_tree))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parser::c_language;
fn recover(src: &str) -> (bool, String) {
recover_with(src, &RepairMacros::default())
}
fn recover_with(src: &str, macros: &RepairMacros) -> (bool, String) {
let mut parser = Parser::new();
parser.set_language(&c_language()).unwrap();
let (tree, text) = parse_with_recovery(&mut parser, src.to_string(), macros).unwrap();
(tree.root_node().has_error(), text)
}
fn macros(object: &[&str], unused_attr: &[&str]) -> RepairMacros {
RepairMacros {
object_macros: object.iter().map(|s| (*s).to_string()).collect(),
unused_attribute_macros: unused_attr.iter().map(|s| (*s).to_string()).collect(),
}
}
#[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 blanks_the_macro_not_the_stranded_type_when_the_table_knows_it() {
let src = "CURL_EXTERN CURLcode curl_easy_setopt(int o);\n";
let (has_error, text) = recover_with(src, ¯os(&["CURL_EXTERN"], &[]));
assert!(!has_error);
assert!(
text.contains("CURLcode curl_easy_setopt"),
"the declaration's real type must survive: {text:?}"
);
assert!(
!text.contains("CURL_EXTERN"),
"the macro is what goes: {text:?}"
);
assert_eq!(text.len(), src.len());
}
#[test]
fn without_a_macro_table_the_stranded_token_is_still_the_one_blanked() {
let src = "CURL_EXTERN CURLcode curl_easy_setopt(int o);\n";
let (_, text) = recover(src);
assert!(text.contains("CURL_EXTERN"));
assert!(!text.contains("CURLcode"));
}
#[test]
fn never_blanks_a_macro_named_on_its_own_define_line() {
let src = "#define X\nX GLuint counter;\n";
let (_, text) = recover_with(src, ¯os(&["X", "GLuint"], &[]));
assert!(
text.contains("#define X"),
"definition must survive: {text:?}"
);
}
#[test]
fn leaves_a_marker_where_a_trailing_unused_attribute_macro_was() {
let src =
"typedef unsigned long word_t;\nvoid f(void) {\n word_t totalObjectSize UNUSED;\n}\n";
let (has_error, text) = recover_with(src, ¯os(&["UNUSED"], &["UNUSED"]));
assert!(!has_error);
assert!(text.contains(UNUSED_ATTRIBUTE_MARKER), "{text:?}");
assert_eq!(text.len(), src.len());
}
#[test]
fn leaves_a_marker_where_a_leading_unused_attribute_macro_was() {
let src = "typedef unsigned long pptr_t;\nvoid f(void) {\n UNUSED pptr_t vaddr = 1;\n}\n";
let (has_error, text) = recover_with(src, ¯os(&["UNUSED"], &["UNUSED"]));
assert!(!has_error);
assert!(text.contains(UNUSED_ATTRIBUTE_MARKER), "{text:?}");
assert!(text.contains("pptr_t vaddr = 1;"), "{text:?}");
assert_eq!(text.len(), src.len());
}
#[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(), &RepairMacros::default()).unwrap();
assert_eq!(text.len(), src.len());
}
}