use tree_sitter::Node;
use crate::declarations::node_text;
use brokk_bifrost_core::hash::{HashMap, HashSet};
const TRIVIAL_TYPE_KINDS: &[&str] = &["primitive_type", "sized_type_specifier"];
#[derive(Default)]
struct NameFacts {
declarations: usize,
trivial_signature: bool,
poisoned: bool,
}
pub struct CppTemporaryFreeCallIndex<'a> {
source: &'a str,
facts: Option<HashMap<&'a str, NameFacts>>,
visited_nodes: usize,
}
impl<'a> CppTemporaryFreeCallIndex<'a> {
pub fn build(source: &'a str, root: Node<'a>) -> Self {
let mut facts: HashMap<&'a str, NameFacts> = HashMap::default();
let mut declaration_names: HashSet<usize> = HashSet::default();
let mut provable_file = true;
let mut visited_nodes = 0usize;
let mut stack = vec![root];
while let Some(node) = stack.pop() {
visited_nodes += 1;
let kind = node.kind();
if kind.starts_with("preproc_") {
provable_file = false;
}
if let Some(candidate) = free_function_candidate(node) {
declaration_names.insert(candidate.name.id());
let name = node_text(candidate.name, source);
let entry = facts.entry(name).or_default();
entry.declarations += 1;
entry.trivial_signature = candidate.trivial_signature;
}
if matches!(
kind,
"identifier" | "field_identifier" | "type_identifier" | "namespace_identifier"
) && !declaration_names.contains(&node.id())
&& !(kind == "identifier" && is_callee_position(node))
{
facts.entry(node_text(node, source)).or_default().poisoned = true;
}
let mut cursor = node.walk();
stack.extend(node.named_children(&mut cursor));
}
Self {
source,
facts: provable_file.then_some(facts),
visited_nodes,
}
}
pub fn visited_nodes(&self) -> usize {
self.visited_nodes
}
pub fn call_is_provably_temporary_free(&self, call: Node<'_>) -> bool {
debug_assert_eq!(call.kind(), "call_expression");
let Some(facts) = &self.facts else {
return false;
};
let Some(callee) = call.child_by_field_name("function") else {
return false;
};
if callee.kind() != "identifier" {
return false;
}
let Some(name) = facts.get(node_text(callee, self.source)) else {
return false;
};
if name.declarations != 1 || !name.trivial_signature || name.poisoned {
return false;
}
let Some(arguments) = call.child_by_field_name("arguments") else {
return false;
};
let mut cursor = arguments.walk();
arguments
.named_children(&mut cursor)
.all(argument_is_trivially_shaped)
}
}
struct FreeFunctionCandidate<'a> {
name: Node<'a>,
trivial_signature: bool,
}
fn free_function_candidate(node: Node<'_>) -> Option<FreeFunctionCandidate<'_>> {
if !matches!(node.kind(), "function_definition" | "declaration") {
return None;
}
if node.parent()?.kind() != "translation_unit" {
return None;
}
let type_node = node.child_by_field_name("type")?;
let mut declarator = node.child_by_field_name("declarator")?;
let mut indirect_return = false;
while matches!(
declarator.kind(),
"pointer_declarator" | "reference_declarator"
) {
indirect_return = true;
declarator = declarator.child_by_field_name("declarator")?;
}
if declarator.kind() != "function_declarator" {
return None;
}
let name = declarator.child_by_field_name("declarator")?;
if name.kind() != "identifier" {
return None;
}
let trivial_return = indirect_return || TRIVIAL_TYPE_KINDS.contains(&type_node.kind());
let trivial_signature = trivial_return
&& declarator
.child_by_field_name("parameters")
.is_some_and(|parameters| {
let mut cursor = parameters.walk();
parameters
.named_children(&mut cursor)
.all(parameter_is_provably_trivial)
});
Some(FreeFunctionCandidate {
name,
trivial_signature,
})
}
fn parameter_is_provably_trivial(parameter: Node<'_>) -> bool {
if parameter.kind() != "parameter_declaration" {
return false;
}
let Some(type_node) = parameter.child_by_field_name("type") else {
return false;
};
if TRIVIAL_TYPE_KINDS.contains(&type_node.kind()) {
return true;
}
parameter
.child_by_field_name("declarator")
.is_some_and(declarator_contains_pointer)
}
fn declarator_contains_pointer(declarator: Node<'_>) -> bool {
let mut stack = vec![declarator];
while let Some(node) = stack.pop() {
if matches!(
node.kind(),
"pointer_declarator" | "abstract_pointer_declarator"
) {
return true;
}
let mut cursor = node.walk();
stack.extend(node.named_children(&mut cursor));
}
false
}
fn is_callee_position(node: Node<'_>) -> bool {
node.parent().is_some_and(|parent| {
parent.kind() == "call_expression"
&& parent
.child_by_field_name("function")
.is_some_and(|function| function.id() == node.id())
})
}
fn argument_is_trivially_shaped(argument: Node<'_>) -> bool {
let mut node = argument;
loop {
match node.kind() {
"parenthesized_expression" => {
let mut cursor = node.walk();
let mut children = node.named_children(&mut cursor);
let (Some(inner), None) = (children.next(), children.next()) else {
return false;
};
node = inner;
}
"identifier"
| "number_literal"
| "char_literal"
| "string_literal"
| "concatenated_string"
| "true"
| "false"
| "null"
| "nullptr" => return true,
"call_expression" => return true,
_ => return false,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tree_sitter::{Parser, Tree};
fn parse_cpp(source: &str) -> Tree {
let mut parser = Parser::new();
parser
.set_language(&tree_sitter_cpp::LANGUAGE.into())
.expect("cpp language");
parser.parse(source, None).expect("cpp tree")
}
fn classified_calls(source: &str) -> Vec<(String, bool)> {
let tree = parse_cpp(source);
let index = CppTemporaryFreeCallIndex::build(source, tree.root_node());
let mut calls = Vec::new();
let mut stack = vec![tree.root_node()];
while let Some(node) = stack.pop() {
if node.kind() == "call_expression" {
calls.push((
node_text(node, source).to_string(),
index.call_is_provably_temporary_free(node),
));
}
let mut cursor = node.walk();
stack.extend(node.named_children(&mut cursor));
}
calls.sort();
calls
}
fn assert_calls(source: &str, expected: &[(&str, bool)]) {
let mut expected = expected
.iter()
.map(|(text, provable)| (text.to_string(), *provable))
.collect::<Vec<_>>();
expected.sort();
assert_eq!(classified_calls(source), expected);
}
#[test]
fn exact_local_free_function_calls_are_provable() {
assert_calls(
r#"
const char *dfb_source() {
return "tainted";
}
void dfb_sink(const char *value) {}
void run() {
dfb_sink(dfb_source());
}
"#,
&[("dfb_source()", true), ("dfb_sink(dfb_source())", true)],
);
}
#[test]
fn discarded_result_and_literal_argument_are_provable() {
assert_calls(
r#"
const char *dfb_source() {
return "tainted";
}
void dfb_sink(const char *value) {}
void run() {
dfb_source();
dfb_sink("clean");
}
"#,
&[("dfb_source()", true), ("dfb_sink(\"clean\")", true)],
);
}
#[test]
fn local_prototype_is_provable() {
assert_calls(
r#"
const char *dfb_source();
void run() {
dfb_source();
}
"#,
&[("dfb_source()", true)],
);
}
#[test]
fn overloaded_functions_stay_unproven() {
assert_calls(
r#"
const char *dfb_source() { return "a"; }
const char *dfb_source(int selector) { return "b"; }
void run() {
dfb_source();
}
"#,
&[("dfb_source()", false)],
);
}
#[test]
fn function_pointer_calls_stay_unproven() {
assert_calls(
r#"
const char *real_source() { return "a"; }
void run() {
const char *(*dfb_source)() = real_source;
dfb_source();
}
"#,
&[("dfb_source()", false)],
);
}
#[test]
fn virtual_member_calls_stay_unproven() {
assert_calls(
r#"
struct Producer {
virtual const char *dfb_source() { return "a"; }
};
void run(Producer *producer) {
producer->dfb_source();
}
"#,
&[("producer->dfb_source()", false)],
);
}
#[test]
fn member_declaration_poisons_unqualified_calls() {
assert_calls(
r#"
const char *dfb_source() { return "a"; }
struct Wrapper {
const char *dfb_source() { return "b"; }
const char *read() { return dfb_source(); }
};
"#,
&[("dfb_source()", false)],
);
}
#[test]
fn template_functions_stay_unproven() {
assert_calls(
r#"
template <typename T>
T dfb_source() { return T(); }
void run() {
dfb_source<const char *>();
}
"#,
&[("dfb_source<const char *>()", false), ("T()", false)],
);
}
#[test]
fn class_return_types_stay_unproven() {
assert_calls(
r#"
struct Token {};
Token dfb_source() { return Token{}; }
void run() {
dfb_source();
}
"#,
&[("dfb_source()", false)],
);
}
#[test]
fn class_reference_parameters_stay_unproven() {
assert_calls(
r#"
struct Token {};
void dfb_sink(const Token &value) {}
void run() {
dfb_sink(Token{});
}
"#,
&[("dfb_sink(Token{})", false)],
);
}
#[test]
fn default_arguments_stay_unproven() {
assert_calls(
r#"
void dfb_sink(const char *value = "d") {}
void run() {
dfb_sink();
}
"#,
&[("dfb_sink()", false)],
);
}
#[test]
fn complex_argument_expressions_stay_unproven() {
assert_calls(
r#"
void dfb_sink(const char *value) {}
void run(const char *left) {
dfb_sink(left + 1);
}
"#,
&[("dfb_sink(left + 1)", false)],
);
}
#[test]
fn preprocessor_content_makes_the_file_unprovable() {
assert_calls(
r#"
#include <string>
const char *dfb_source() { return "a"; }
void run() {
dfb_source();
}
"#,
&[("dfb_source()", false)],
);
}
#[test]
fn address_taken_names_stay_unproven() {
assert_calls(
r#"
const char *dfb_source() { return "a"; }
void keep(const char *(*pointer)()) {}
void run() {
keep(&dfb_source);
dfb_source();
}
"#,
&[("dfb_source()", false), ("keep(&dfb_source)", false)],
);
}
#[test]
fn pointer_indirection_over_class_types_is_provable() {
assert_calls(
r#"
struct Token;
Token *dfb_source() { return nullptr; }
void dfb_sink(Token *value) {}
void run() {
dfb_sink((dfb_source()));
}
"#,
&[("dfb_source()", true), ("dfb_sink((dfb_source()))", true)],
);
}
}