use std::collections::{HashMap, HashSet};
use crate::capture::{add_capture, Capture};
use crate::query::{NegativeQuery, QueryTree};
use crate::util::parse_number_literal;
use crate::RegexMap;
use colored::Colorize;
use tree_sitter::{Node, TreeCursor};
pub fn build_query_tree(
source: &str,
cursor: &mut TreeCursor,
is_cpp: bool,
regex_constraints: Option<RegexMap>,
) -> QueryTree {
_build_query_tree(source, cursor, 0, is_cpp, false, false, regex_constraints)
}
fn _build_query_tree(
source: &str,
c: &mut TreeCursor,
id: usize,
is_cpp: bool,
is_multi_pattern: bool,
strict_mode: bool,
regex_constraints: Option<RegexMap>,
) -> QueryTree {
let mut b = QueryBuilder {
query_source: source.to_string(),
captures: Vec::new(),
negations: Vec::new(),
id,
cpp: is_cpp,
regex_constraints: match regex_constraints {
Some(r) => r,
None => RegexMap::new(HashMap::new()),
},
};
if c.node().kind() == "translation_unit" {
debug!("query cursor specifies translation_unit");
c.goto_first_child();
}
let mut variables = HashSet::new();
let sexp = if !is_multi_pattern {
let needs_anchor = c.node().kind() == "compound_statement" && id == 0;
debug!("query needs anchor: {}", needs_anchor);
let mut s = b.build(c, 0, strict_mode);
if !needs_anchor {
s += "@";
s += &add_capture(&mut b.captures, Capture::Display);
}
s += &process_captures(&b.captures, 0, &mut variables);
if needs_anchor {
let capture = Capture::Display;
format!(
"(function_definition body: {}) @{}",
s,
&add_capture(&mut b.captures, capture)
)
} else {
"(".to_string() + &s + ")"
}
} else {
assert!(c.goto_first_child());
assert!(c.goto_next_sibling());
let mut s = String::new();
loop {
let child = c.node();
if !c.goto_next_sibling() {
break;
}
let before = b.captures.len();
let mut cursor = child.walk();
let child_sexp = b.build(&mut cursor, 0, strict_mode);
let captures = &process_captures(&b.captures, before, &mut variables);
if !child_sexp.is_empty() {
s += &format!("({} {})", child_sexp, captures);
}
}
s
};
debug!("tree_sitter query {}: {}", id, sexp);
QueryTree::new(
crate::ts_query(&sexp, is_cpp),
b.captures,
variables,
b.negations,
id,
)
}
fn process_captures(
captures: &[Capture],
offset: usize,
variables: &mut HashSet<String>,
) -> String {
let mut vars: HashMap<String, Vec<usize>> = HashMap::new();
let mut sexp = String::new();
for (i, c) in captures.iter().skip(offset).enumerate() {
match c {
Capture::Display => (),
Capture::Check(s) => {
sexp += &format!(r#"(#eq? @{} "{}")"#, (i + offset).to_string(), s);
}
Capture::Variable(var, _) => {
vars.entry(var.clone())
.or_insert_with(Vec::new)
.push(i + offset);
variables.insert(var.clone());
}
_ => (),
}
}
for (_, vec) in vars.iter() {
if vec.len() > 1 {
let a = vec[0].to_string();
for capture in vec.iter().skip(1) {
let b = capture.to_string();
sexp += &format!(r#"(#eq? @{} @{})"#, a, b);
}
}
}
sexp
}
struct QueryBuilder {
query_source: String,
captures: Vec<Capture>, negations: Vec<NegativeQuery>, id: usize, cpp: bool, regex_constraints: RegexMap,
}
impl QueryBuilder {
fn get_text(&self, n: &tree_sitter::Node) -> &str {
&self.query_source[n.byte_range()]
}
fn is_subexpr_wildcard(&self, query: Node) -> bool {
if query.kind() != "call_expression" {
return false;
}
let f = query.child_by_field_name("function").unwrap();
if f.utf8_text(self.query_source.as_bytes()).unwrap() == "_" {
return true;
}
false
}
fn is_comparison_binary_exp(&self, n: Node) -> bool {
assert!(n.kind() == "binary_expression");
if let Some(op) = n.child(1) {
[">", "<", "<=", ">="].contains(&op.kind())
} else {
false
}
}
fn is_commutative_binary_exp(&self, n: Node) -> bool {
assert!(n.kind() == "binary_expression");
if let Some(op) = n.child(1) {
["+", "*", "&", "|", "==", "!="].contains(&op.kind())
} else {
false
}
}
fn is_transformable_binary_exp(&self, n: Node) -> bool {
self.is_comparison_binary_exp(n) || self.is_commutative_binary_exp(n)
}
fn build(&mut self, c: &mut TreeCursor, depth: usize, strict_mode: bool) -> String {
if !c.node().is_named() {
return format!(r#""{}""#, c.node().kind());
}
let kind = c.node().kind();
match kind {
"binary_expression" if self.is_transformable_binary_exp(c.node()) => {
assert!(c.goto_first_child());
let left = self.build(c, depth + 1, strict_mode);
assert!(c.goto_next_sibling());
let op = c.node().kind();
let alt_op = match op {
">" => "<",
"<" => ">",
"<=" => ">=",
">=" => "<=",
_ => op,
};
assert!(c.goto_next_sibling());
let right = self.build(c, depth + 1, strict_mode);
c.goto_parent();
return format! {"[(binary_expression left: {0} operator: \"{1}\" right: {2})
(binary_expression left: {2} operator: \"{3}\" right: {0})]", left, op, right, alt_op};
}
"labeled_statement" => {
let label = c.node().child(0).unwrap();
if self.get_text(&label).to_uppercase() == "NOT" {
self.build_negative_query(c);
return "".to_string();
} else if self.get_text(&label).to_uppercase() == "STRICT" {
if let Some(child) = c.node().named_child(1) {
return self.build(&mut child.walk(), depth, true);
} else {
return "".to_string();
}
}
}
"compound_statement" if c.node().named_child_count() > 0 => {
self.id += 1;
let mut c = c.node().walk();
let capture = Capture::Subquery(Box::new(_build_query_tree(
&self.query_source,
&mut c,
self.id,
self.cpp,
true,
false, Some(self.regex_constraints.clone()),
)));
return "(compound_statement) @".to_string()
+ &add_capture(&mut self.captures, capture);
}
"identifier"
| "type_identifier"
| "field_identifier"
| "sized_type_specifier"
| "primitive_type"
| "namespace_identifier" => return self.build_identifier(c),
"assignment_expression" => return self.build_assignment(c, depth, strict_mode),
"call_expression" => match self.build_call_expr(c, depth, strict_mode) {
Some(s) => return s,
_ => (),
},
"expression_statement" => {
if let Some(child) = c.node().named_child(0) {
if !strict_mode || self.is_subexpr_wildcard(child) {
if self.get_text(&child) != "_" {
c.goto_first_child();
return self.build(c, depth, strict_mode);
}
}
}
}
"number_literal" => {
let pattern = self.get_text(&c.node());
let capture = if let Some(num) = parse_number_literal(pattern) {
Capture::Number(num)
} else {
warn! {"Could not parse {} as a number. Forcing string matching", pattern}
Capture::Check(pattern.to_string())
};
return format! {"(number_literal) @{}", &add_capture(&mut self.captures, capture)};
}
"string_literal" => {
let pattern = self.get_text(&c.node());
let unquoted = &pattern[1..pattern.len() - 1];
if unquoted.starts_with('$') {
let c = Capture::Variable(
unquoted.to_string(),
self.regex_constraints.get(unquoted),
);
return format! {"(string_literal) @{}", &add_capture(&mut self.captures, c)};
}
}
_ => (),
}
let anchoring = kind == "argument_list" && c.node().named_child_count() > 1;
let is_funcdef = kind == "function_definition";
let mut result = format!("({}", c.node().kind());
if !c.goto_first_child() {
if !c.node().is_named() {
return format!(r#""{}""#, c.node().kind());
}
return result + ")";
}
loop {
let name = c.field_name();
if let Some(n) = name {
result += &format!(" {}:", n);
let t = self.build(c, depth + 1, strict_mode);
if n == "declarator" && is_funcdef {
result += &format!("([(_ {}) ({})])", t, t);
} else {
result += &t
}
} else if c.node().is_named() {
if anchoring {
result += " .";
}
result += " ";
result += &self.build(c, depth + 1, strict_mode);
} else {
let sexp = self.build(c, depth + 1, strict_mode);
if sexp.chars().all(|c| char::is_alphanumeric(c) || c == '"') && sexp != "\"\"\"" {
result += &format!(
" {} @{}",
sexp,
&add_capture(&mut self.captures, Capture::Display)
);
}
}
if !c.goto_next_sibling() {
break;
}
}
c.goto_parent();
debug!("generated query: {}", result);
result + ")"
}
fn build_negative_query(&mut self, c: &mut TreeCursor) {
let negated_query = c.node().child(2).unwrap();
let before = self.captures.len() as i64 - 1;
self.id += 1;
self.negations.push(NegativeQuery {
qt: Box::new(_build_query_tree(
&self.query_source,
&mut negated_query.walk(),
self.id,
self.cpp,
false,
false, Some(self.regex_constraints.clone()),
)),
previous_capture_index: before,
});
}
fn build_identifier(&mut self, c: &mut TreeCursor) -> String {
let pattern = self.get_text(&c.node());
let kind = c.node().kind();
if pattern == "_" {
return "(_)".to_string();
}
let mut result = if kind == "type_identifier" {
"[ (type_identifier) (sized_type_specifier) (primitive_type)]".to_string()
} else if kind == "identifier" && pattern.starts_with('$') {
if self.cpp {
"[(identifier) (field_expression) (field_identifier) (qualified_identifier) (this)]"
.to_string()
} else {
"[(identifier) (field_expression) (field_identifier)]".to_string()
}
} else {
format!("({})", kind)
};
let capture = if pattern.starts_with('$') {
Capture::Variable(pattern.to_string(), self.regex_constraints.get(pattern))
} else {
Capture::Check(pattern.to_string())
};
result += " @";
result += &add_capture(&mut self.captures, capture);
result
}
fn build_call_expr(
&mut self,
c: &mut TreeCursor,
depth: usize,
strict_mode: bool,
) -> Option<String> {
if self.is_subexpr_wildcard(c.node()) {
let mut arg = c.node().child_by_field_name("arguments").unwrap().walk();
arg.goto_first_child();
arg.goto_next_sibling();
let mut copy = arg.clone();
copy.goto_next_sibling();
if copy.goto_next_sibling() {
warn! {"sub expression '{}' with multiple arguments is not supported.
Do you want to match on a function call '$foo()' instead?",
self.get_text(&c.node()).to_string().red()};
warn! {"converting to function call..."};
return None;
}
if depth == 0 {
return Some(self.build(&mut arg, depth, strict_mode));
}
self.id += 1;
let capture = Capture::Subquery(Box::new(_build_query_tree(
&self.query_source,
&mut arg,
self.id,
self.cpp,
false,
strict_mode,
Some(self.regex_constraints.clone()),
)));
return Some("_ @".to_string() + &add_capture(&mut self.captures, capture));
}
let function = c.node().child_by_field_name("function").unwrap();
let arguments = c.node().child_by_field_name("arguments").unwrap();
if function.kind() == "identifier" {
let pattern = self.get_text(&function);
if !pattern.starts_with('$') {
let capture = Capture::Check(pattern.to_string());
let capture_str = "@".to_string() + &add_capture(&mut self.captures, capture);
let a = self.build(&mut arguments.walk(), depth + 1, false);
let fs = if strict_mode {
format! {"(identifier) {}",capture_str}
} else {
if self.cpp {
format! {"[(field_expression field: (field_identifier){0})
(qualified_identifier name: (identifier){0})
(qualified_identifier name: (qualified_identifier (identifier){0}))
(qualified_identifier name: (qualified_identifier (qualified_identifier (identifier){0})))
(qualified_identifier name: (qualified_identifier (qualified_identifier
(qualified_identifier (identifier){0}))))
(identifier) {0}]",capture_str}
} else {
format! {"[(field_expression field: (field_identifier){0})
(identifier) {0}]",capture_str}
}
};
let result = format! {"(call_expression function: {} arguments: {})", fs, a};
return Some(result);
}
}
None
}
fn build_assignment(&mut self, c: &mut TreeCursor, depth: usize, strict_mode: bool) -> String {
assert!(c.goto_first_child());
let left = self.build(c, depth + 1, strict_mode);
let left_is_identifier = c.node().kind() == "identifier";
assert!(c.goto_next_sibling());
let optional_cast = |r: String| format! {"[(cast_expression value: {}) {}]", r, r};
let result = if c.node().kind() != "=" || !left_is_identifier {
let operator = self.build(c, depth + 1, strict_mode);
assert!(c.goto_next_sibling());
let right = optional_cast(self.build(c, depth + 1, strict_mode));
format! {"(assignment_expression left: {} {} right: {})" , left, operator, right}
} else {
assert!(c.goto_next_sibling());
let right = optional_cast(self.build(c, depth + 1, strict_mode));
format! {r"[(assignment_expression left: {0} right: {1})
(init_declarator declarator: {0} value: {1})
(init_declarator declarator:(pointer_declarator declarator: {0}) value: {1})]", left,right}
};
c.goto_parent();
return result;
}
}