use super::{SearchResult, MatchType};
use std::path::Path;
use tree_sitter::{Language, Parser, Query, QueryCursor, Node};
use walkdir::WalkDir;
use std::fs;
fn get_language(lang: &str) -> Option<Language> {
match lang {
"rust" => Some(tree_sitter_rust::language()),
"javascript" => Some(tree_sitter_javascript::language()),
"typescript" => Some(tree_sitter_typescript::language_typescript()),
"python" => Some(tree_sitter_python::language()),
"go" => Some(tree_sitter_go::language()),
"java" => Some(tree_sitter_java::language()),
"cpp" => Some(tree_sitter_cpp::language()),
"c" => Some(tree_sitter_c::language()),
_ => None,
}
}
pub struct AstSearcher {
_init: bool,
}
impl AstSearcher {
pub fn new() -> Self {
Self { _init: true }
}
fn create_parser(&self, lang: &str) -> Option<Parser> {
if let Some(language) = get_language(lang) {
let mut parser = Parser::new();
if parser.set_language(language).is_ok() {
return Some(parser);
}
}
None
}
pub async fn search(
&self,
pattern: &str,
path: &Path,
language: Option<&str>,
max_results: usize,
) -> Result<Vec<SearchResult>, Box<dyn std::error::Error>> {
let mut results = Vec::new();
for entry in WalkDir::new(path)
.follow_links(true)
.into_iter()
.filter_map(|e| e.ok())
.filter(|e| e.file_type().is_file())
{
let file_path = entry.path();
let lang = language.unwrap_or_else(|| detect_language(file_path));
if lang == "text" {
continue;
}
if let Some(mut parser) = self.create_parser(lang) {
if let Ok(source) = fs::read_to_string(file_path) {
if let Some(tree) = parser.parse(&source, None) {
let file_results = self.search_tree(
&tree,
&source,
pattern,
file_path,
lang,
);
results.extend(file_results);
if results.len() >= max_results {
break;
}
}
}
}
}
results.truncate(max_results);
Ok(results)
}
fn search_tree(
&self,
tree: &tree_sitter::Tree,
source: &str,
pattern: &str,
file_path: &Path,
language: &str,
) -> Vec<SearchResult> {
let mut results = Vec::new();
let root_node = tree.root_node();
let query_str = build_query_string(pattern, language);
if let Some(lang) = get_language(language) {
if let Ok(query) = Query::new(lang, &query_str) {
let mut cursor = QueryCursor::new();
let matches = cursor.matches(&query, root_node, source.as_bytes());
for match_ in matches {
for capture in match_.captures {
let node = capture.node;
let start = node.start_position();
let match_text = source[node.byte_range()].to_string();
let (context_before, context_after) = get_context(source, node, 3);
results.push(SearchResult {
file_path: file_path.to_path_buf(),
line_number: start.row + 1,
column: start.column,
match_text,
context_before,
context_after,
match_type: MatchType::Ast,
score: 0.95,
node_type: Some(node.kind().to_string()),
semantic_context: Some(get_semantic_context(node, source)),
});
}
}
} else {
results.extend(search_nodes_by_text(
root_node,
source,
pattern,
file_path,
));
}
}
results
}
}
impl Default for AstSearcher {
fn default() -> Self {
Self::new()
}
}
fn detect_language(path: &Path) -> &'static str {
match path.extension().and_then(|s| s.to_str()) {
Some("rs") => "rust",
Some("js") | Some("mjs") => "javascript",
Some("ts") | Some("tsx") => "typescript",
Some("py") => "python",
Some("go") => "go",
Some("java") => "java",
Some("cpp") | Some("cc") | Some("cxx") => "cpp",
Some("c") | Some("h") => "c",
_ => "text",
}
}
fn build_query_string(pattern: &str, language: &str) -> String {
if pattern.starts_with("function ") {
let name = pattern.trim_start_matches("function ").trim();
match language {
"rust" => format!("(function_item name: (identifier) @fn (#eq? @fn \"{name}\"))"),
"javascript" | "typescript" => {
format!("(function_declaration name: (identifier) @fn (#eq? @fn \"{name}\"))")
}
"python" => format!("(function_definition name: (identifier) @fn (#eq? @fn \"{name}\"))"),
_ => format!("(identifier) @id (#eq? @id \"{name}\")"),
}
} else if pattern.starts_with("class ") {
let name = pattern.trim_start_matches("class ").trim();
match language {
"rust" => format!("(struct_item name: (type_identifier) @struct (#eq? @struct \"{name}\"))"),
"javascript" | "typescript" => {
format!("(class_declaration name: (identifier) @class (#eq? @class \"{name}\"))")
}
"python" => format!("(class_definition name: (identifier) @class (#eq? @class \"{name}\"))"),
_ => format!("(identifier) @id (#eq? @id \"{name}\")"),
}
} else {
format!("(identifier) @id (#match? @id \"{pattern}\")")
}
}
fn search_nodes_by_text(
node: Node,
source: &str,
pattern: &str,
file_path: &Path,
) -> Vec<SearchResult> {
let mut results = Vec::new();
let mut cursor = node.walk();
loop {
let current = cursor.node();
let node_text = source[current.byte_range()].to_string();
if node_text.contains(pattern) {
let start = current.start_position();
results.push(SearchResult {
file_path: file_path.to_path_buf(),
line_number: start.row + 1,
column: start.column,
match_text: node_text,
context_before: vec![],
context_after: vec![],
match_type: MatchType::Ast,
score: 0.9,
node_type: Some(current.kind().to_string()),
semantic_context: Some(get_semantic_context(current, source)),
});
}
if cursor.goto_first_child() {
continue;
}
if cursor.goto_next_sibling() {
continue;
}
loop {
if !cursor.goto_parent() {
return results;
}
if cursor.goto_next_sibling() {
break;
}
}
}
}
fn get_context(source: &str, node: Node, context_lines: usize) -> (Vec<String>, Vec<String>) {
let lines: Vec<&str> = source.lines().collect();
let start_line = node.start_position().row;
let end_line = node.end_position().row;
let before_start = start_line.saturating_sub(context_lines);
let after_end = std::cmp::min(end_line + context_lines + 1, lines.len());
let context_before = lines[before_start..start_line]
.iter()
.map(|s| s.to_string())
.collect();
let context_after = if end_line + 1 < lines.len() {
lines[(end_line + 1)..after_end]
.iter()
.map(|s| s.to_string())
.collect()
} else {
vec![]
};
(context_before, context_after)
}
fn get_semantic_context(node: Node, source: &str) -> String {
let mut context = format!("{} at {}:{}",
node.kind(),
node.start_position().row + 1,
node.start_position().column
);
if let Some(parent) = find_parent_context(node) {
let parent_text = source[parent.byte_range()].lines().next().unwrap_or("");
context.push_str(&format!(" in {parent_text}"));
}
context
}
fn find_parent_context(mut node: Node) -> Option<Node> {
while let Some(parent) = node.parent() {
match parent.kind() {
"function_item" | "function_declaration" | "function_definition" |
"method_definition" | "struct_item" | "class_declaration" |
"class_definition" | "impl_item" => {
return Some(parent);
}
_ => node = parent,
}
}
None
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_ast_search() {
let searcher = AstSearcher::new();
let results = searcher.search(
"function test",
Path::new("."),
Some("rust"),
10,
).await;
assert!(results.is_ok());
}
#[test]
fn test_detect_language() {
assert_eq!(detect_language(Path::new("test.rs")), "rust");
assert_eq!(detect_language(Path::new("test.js")), "javascript");
assert_eq!(detect_language(Path::new("test.ts")), "typescript");
assert_eq!(detect_language(Path::new("test.py")), "python");
assert_eq!(detect_language(Path::new("test.go")), "go");
}
}