use crate::languages::get_language_info;
use crate::types::{ImplTraitInfo, SemanticAnalysis};
use std::cell::RefCell;
use std::collections::HashMap;
use std::path::Path;
use std::sync::LazyLock;
use thiserror::Error;
use tracing::instrument;
use tree_sitter::{Parser, Query, QueryCursor, StreamingIterator};
use crate::parser_elements::{
extract_calls, extract_def_use, extract_elements, extract_impl_methods,
extract_impl_traits_from_tree, extract_imports, extract_references,
};
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum ParserError {
#[error("Unsupported language: {0}")]
UnsupportedLanguage(String),
#[error("Failed to parse file: {0}")]
ParseError(String),
#[error("Invalid UTF-8 in file")]
InvalidUtf8,
#[error("Query error: {0}")]
QueryError(String),
#[error("Parse timeout exceeded: {0} microseconds")]
Timeout(u64),
}
#[derive(Clone, Copy)]
pub(crate) struct TimeoutConfig {
pub deadline: Option<std::time::Instant>,
pub micros: u64,
}
impl TimeoutConfig {
fn new(timeout_micros: Option<u64>) -> Self {
let deadline = timeout_micros
.map(|us| std::time::Instant::now() + std::time::Duration::from_micros(us));
Self {
deadline,
micros: timeout_micros.unwrap_or(0),
}
}
pub(crate) fn is_exceeded(self) -> bool {
self.deadline
.is_some_and(|d| std::time::Instant::now() >= d)
}
}
pub(crate) struct CompiledQueries {
pub element: Query,
pub call: Query,
pub import: Option<Query>,
pub impl_block: Option<Query>,
pub reference: Option<Query>,
pub impl_trait: Option<Query>,
pub defuse: Option<Query>,
}
fn build_compiled_queries(
lang_info: &crate::languages::LanguageInfo,
) -> Result<CompiledQueries, ParserError> {
let element = Query::new(&lang_info.language, lang_info.element_query).map_err(|e| {
ParserError::QueryError(format!(
"Failed to compile element query for {}: {}",
lang_info.name, e
))
})?;
let call = Query::new(&lang_info.language, lang_info.call_query).map_err(|e| {
ParserError::QueryError(format!(
"Failed to compile call query for {}: {}",
lang_info.name, e
))
})?;
let import = if let Some(import_query_str) = lang_info.import_query {
Some(
Query::new(&lang_info.language, import_query_str).map_err(|e| {
ParserError::QueryError(format!(
"Failed to compile import query for {}: {}",
lang_info.name, e
))
})?,
)
} else {
None
};
let impl_block = if let Some(impl_query_str) = lang_info.impl_query {
Some(
Query::new(&lang_info.language, impl_query_str).map_err(|e| {
ParserError::QueryError(format!(
"Failed to compile impl query for {}: {}",
lang_info.name, e
))
})?,
)
} else {
None
};
let reference = if let Some(reference_query_str) = lang_info.reference_query {
Some(
Query::new(&lang_info.language, reference_query_str).map_err(|e| {
ParserError::QueryError(format!(
"Failed to compile reference query for {}: {}",
lang_info.name, e
))
})?,
)
} else {
None
};
let impl_trait = if let Some(impl_trait_query_str) = lang_info.impl_trait_query {
Some(
Query::new(&lang_info.language, impl_trait_query_str).map_err(|e| {
ParserError::QueryError(format!(
"Failed to compile impl_trait query for {}: {}",
lang_info.name, e
))
})?,
)
} else {
None
};
let defuse = if let Some(defuse_query_str) = lang_info.defuse_query {
Some(
Query::new(&lang_info.language, defuse_query_str).map_err(|e| {
ParserError::QueryError(format!(
"Failed to compile defuse query for {}: {}",
lang_info.name, e
))
})?,
)
} else {
None
};
Ok(CompiledQueries {
element,
call,
import,
impl_block,
reference,
impl_trait,
defuse,
})
}
#[cfg_attr(coverage_nightly, coverage(off))]
fn init_query_cache() -> HashMap<&'static str, CompiledQueries> {
let mut cache = HashMap::new();
for lang_name in crate::lang::supported_languages() {
if let Some(lang_info) = get_language_info(lang_name) {
match build_compiled_queries(&lang_info) {
Ok(compiled) => {
cache.insert(*lang_name, compiled);
}
Err(e) => {
tracing::error!(
"Failed to compile queries for language {}: {}",
lang_name,
e
);
}
}
}
}
cache
}
static QUERY_CACHE: LazyLock<HashMap<&'static str, CompiledQueries>> =
LazyLock::new(init_query_cache);
fn get_compiled_queries(language: &str) -> Result<&'static CompiledQueries, ParserError> {
QUERY_CACHE
.get(language)
.ok_or_else(|| ParserError::UnsupportedLanguage(language.to_string()))
}
thread_local! {
pub(crate) static PARSER: RefCell<Parser> = RefCell::new(Parser::new());
pub(crate) static QUERY_CURSOR: RefCell<QueryCursor> = RefCell::new(QueryCursor::new());
}
pub struct ElementExtractor;
impl ElementExtractor {
#[instrument(skip_all, fields(language))]
pub fn extract_with_depth(source: &str, language: &str) -> Result<(usize, usize), ParserError> {
let lang_info = get_language_info(language)
.ok_or_else(|| ParserError::UnsupportedLanguage(language.to_string()))?;
let tree = PARSER.with(|p| {
let mut parser = p.borrow_mut();
parser
.set_language(&lang_info.language)
.map_err(|e| ParserError::ParseError(format!("Failed to set language: {e}")))?;
parser
.parse(source, None)
.ok_or_else(|| ParserError::ParseError("Failed to parse".to_string()))
})?;
let compiled = get_compiled_queries(language)?;
let (function_count, class_count) = QUERY_CURSOR.with(|c| {
let mut cursor = c.borrow_mut();
cursor.set_max_start_depth(None);
let mut function_count = 0;
let mut class_count = 0;
let mut matches =
cursor.matches(&compiled.element, tree.root_node(), source.as_bytes());
while let Some(mat) = matches.next() {
for capture in mat.captures {
let capture_name = compiled.element.capture_names()[capture.index as usize];
match capture_name {
"function" => function_count += 1,
"class" => class_count += 1,
_ => {}
}
}
}
(function_count, class_count)
});
tracing::debug!(language = %language, functions = function_count, classes = class_count, "parse complete");
Ok((function_count, class_count))
}
}
pub struct SemanticExtractor;
impl SemanticExtractor {
#[instrument(skip_all, fields(language))]
pub fn extract(
source: &str,
language: &str,
ast_recursion_limit: Option<usize>,
timeout_micros: Option<u64>,
) -> Result<SemanticAnalysis, ParserError> {
let tc = TimeoutConfig::new(timeout_micros);
if tc.is_exceeded() {
return Err(ParserError::Timeout(tc.micros));
}
let lang_info = get_language_info(language)
.ok_or_else(|| ParserError::UnsupportedLanguage(language.to_string()))?;
let tree = PARSER.with(|p| {
let mut parser = p.borrow_mut();
parser
.set_language(&lang_info.language)
.map_err(|e| ParserError::ParseError(format!("Failed to set language: {e}")))?;
parser
.parse(source, None)
.ok_or_else(|| ParserError::ParseError("Failed to parse".to_string()))
})?;
if tc.is_exceeded() {
return Err(ParserError::Timeout(tc.micros));
}
let compiled = get_compiled_queries(language)?;
let root = tree.root_node();
let max_depth: Option<u32> = ast_recursion_limit
.filter(|&limit| limit > 0)
.and_then(|limit| u32::try_from(limit).ok());
let mut functions = Vec::new();
let mut classes = Vec::new();
let mut imports = Vec::new();
let mut references = Vec::new();
let mut calls = Vec::new();
let mut call_frequency = HashMap::new();
extract_elements(
source,
compiled,
root,
max_depth,
&mut functions,
&mut classes,
tc,
&lang_info,
)?;
if tc.is_exceeded() {
return Err(ParserError::Timeout(tc.micros));
}
extract_calls(
source,
compiled,
root,
max_depth,
&mut calls,
&mut call_frequency,
tc,
)?;
extract_imports(source, compiled, root, max_depth, &mut imports, tc)?;
extract_impl_methods(source, compiled, root, max_depth, &mut classes, tc)?;
extract_references(source, compiled, root, max_depth, &mut references, tc)?;
let impl_traits = if language == "rust" {
extract_impl_traits_from_tree(source, compiled, root, tc)?
} else {
vec![]
};
tracing::debug!(language = %language, functions = functions.len(), classes = classes.len(), imports = imports.len(), references = references.len(), calls = calls.len(), impl_traits = impl_traits.len(), "extraction complete");
Ok(SemanticAnalysis {
functions,
classes,
imports,
references,
call_frequency,
calls,
impl_traits,
def_use_sites: Vec::new(),
})
}
#[instrument(skip_all, fields(language))]
pub fn extract_module_info(
source: &str,
language: &str,
timeout_micros: Option<u64>,
) -> Result<crate::types::ModuleInfo, ParserError> {
let tc = TimeoutConfig::new(timeout_micros);
if tc.is_exceeded() {
return Err(ParserError::Timeout(tc.micros));
}
let lang_info = get_language_info(language)
.ok_or_else(|| ParserError::UnsupportedLanguage(language.to_string()))?;
let tree = PARSER.with(|p| {
let mut parser = p.borrow_mut();
parser
.set_language(&lang_info.language)
.map_err(|e| ParserError::ParseError(format!("Failed to set language: {e}")))?;
parser
.parse(source, None)
.ok_or_else(|| ParserError::ParseError("Failed to parse".to_string()))
})?;
if tc.is_exceeded() {
return Err(ParserError::Timeout(tc.micros));
}
let compiled = get_compiled_queries(language)?;
let root = tree.root_node();
let mut functions = Vec::new();
let mut classes = Vec::new();
let mut imports = Vec::new();
extract_elements(
source,
compiled,
root,
None,
&mut functions,
&mut classes,
tc,
&lang_info,
)?;
if tc.is_exceeded() {
return Err(ParserError::Timeout(tc.micros));
}
extract_imports(source, compiled, root, None, &mut imports, tc)?;
if tc.is_exceeded() {
return Err(ParserError::Timeout(tc.micros));
}
let module_functions = functions
.into_iter()
.map(|f| crate::types::ModuleFunctionInfo {
name: f.name,
line: f.line,
})
.collect();
let module_imports = imports
.into_iter()
.map(|i| crate::types::ModuleImportInfo {
module: i.module,
items: i.items,
})
.collect();
let line_count = source.lines().count();
Ok(crate::types::ModuleInfo::new(
String::new(), line_count,
language.to_string(),
module_functions,
module_imports,
))
}
pub(crate) fn extract_def_use_for_file(
source: &str,
language: &str,
symbol: &str,
file_path: &str,
ast_recursion_limit: Option<usize>,
) -> Vec<crate::types::DefUseSite> {
let Some(lang_info) = get_language_info(language) else {
return vec![];
};
let Ok(compiled) = get_compiled_queries(language) else {
return vec![];
};
if compiled.defuse.is_none() {
return vec![];
}
let tree = match PARSER.with(|p| {
let mut parser = p.borrow_mut();
if parser.set_language(&lang_info.language).is_err() {
return None;
}
parser.parse(source, None)
}) {
Some(t) => t,
None => return vec![],
};
let root = tree.root_node();
let max_depth: Option<u32> = ast_recursion_limit
.filter(|&limit| limit > 0)
.and_then(|limit| u32::try_from(limit).ok());
extract_def_use(source, compiled, root, symbol, file_path, max_depth)
}
}
#[must_use]
pub fn extract_impl_traits(source: &str, path: &Path) -> Vec<ImplTraitInfo> {
let Some(lang_info) = get_language_info("rust") else {
return vec![];
};
let Ok(compiled) = get_compiled_queries("rust") else {
return vec![];
};
let Some(query) = &compiled.impl_trait else {
return vec![];
};
let Some(tree) = PARSER.with(|p| {
let mut parser = p.borrow_mut();
let _ = parser.set_language(&lang_info.language);
parser.parse(source, None)
}) else {
return vec![];
};
let root = tree.root_node();
let mut results = Vec::new();
QUERY_CURSOR.with(|c| {
let mut cursor = c.borrow_mut();
cursor.set_max_start_depth(None);
let mut matches = cursor.matches(query, root, source.as_bytes());
while let Some(mat) = matches.next() {
let mut trait_name = String::new();
let mut impl_type = String::new();
let mut line = 0usize;
for capture in mat.captures {
let capture_name = query.capture_names()[capture.index as usize];
let node = capture.node;
let text = source[node.start_byte()..node.end_byte()].to_string();
match capture_name {
"trait_name" => {
trait_name = text;
line = node.start_position().row + 1;
}
"impl_type" => {
impl_type = text;
}
_ => {}
}
}
if !trait_name.is_empty() && !impl_type.is_empty() {
results.push(ImplTraitInfo {
trait_name,
impl_type,
path: path.to_path_buf(),
line,
});
}
}
});
results
}
pub(crate) fn execute_query_impl(
language: &str,
source: &str,
query_str: &str,
) -> Result<Vec<crate::QueryCapture>, ParserError> {
let ts_language = crate::languages::get_ts_language(language)
.ok_or_else(|| ParserError::UnsupportedLanguage(language.to_string()))?;
let mut parser = Parser::new();
parser
.set_language(&ts_language)
.map_err(|e| ParserError::QueryError(e.to_string()))?;
let tree = parser
.parse(source.as_bytes(), None)
.ok_or_else(|| ParserError::QueryError("failed to parse source".to_string()))?;
let query =
Query::new(&ts_language, query_str).map_err(|e| ParserError::QueryError(e.to_string()))?;
let source_bytes = source.as_bytes();
let mut captures = Vec::new();
QUERY_CURSOR.with(|c| {
let mut cursor = c.borrow_mut();
cursor.set_max_start_depth(None);
let mut matches = cursor.matches(&query, tree.root_node(), source_bytes);
while let Some(m) = matches.next() {
for cap in m.captures {
let node = cap.node;
let capture_name = query.capture_names()[cap.index as usize].to_string();
let text = node.utf8_text(source_bytes).unwrap_or("").to_string();
captures.push(crate::QueryCapture {
capture_name,
text,
start_line: node.start_position().row,
end_line: node.end_position().row,
start_byte: node.start_byte(),
end_byte: node.end_byte(),
});
}
}
});
Ok(captures)
}
#[cfg(test)]
mod tests_rust {
use super::*;
use crate::types::CallInfo;
#[test]
fn test_ast_recursion_limit_zero_is_unlimited() {
let source = r#"fn hello() -> u32 { 42 }"#;
let result = SemanticExtractor::extract(source, "rust", Some(0), None);
assert!(result.is_ok(), "extract with limit=0 should succeed");
let analysis = result.unwrap();
assert_eq!(
analysis.functions.len(),
1,
"should find exactly one function"
);
}
#[test]
fn test_rust_use_as_imports() {
let source = "use std::io as stdio;\n";
let result = SemanticExtractor::extract(source, "rust", None, None).unwrap();
let stdio_import = result
.imports
.iter()
.find(|imp| imp.items.iter().any(|i| i == "stdio"));
assert!(
stdio_import.is_some(),
"expected import with alias 'stdio' in {:?}",
result.imports
);
}
#[test]
fn test_rust_use_as_clause_plain_identifier() {
let source = "use io as stdio;\n";
let result = SemanticExtractor::extract(source, "rust", None, None).unwrap();
let alias_import = result
.imports
.iter()
.find(|imp| imp.items.iter().any(|i| i == "stdio"));
assert!(
alias_import.is_some(),
"expected import with alias 'stdio' in {:?}",
result.imports
);
}
#[test]
fn test_rust_scoped_use_with_prefix() {
let source = "use std::{io, fs};\n";
let result = SemanticExtractor::extract(source, "rust", None, None).unwrap();
let has_io = result
.imports
.iter()
.any(|imp| imp.items.iter().any(|i| i == "io"));
let has_fs = result
.imports
.iter()
.any(|imp| imp.items.iter().any(|i| i == "fs"));
assert!(has_io, "expected import 'io' in {:?}", result.imports);
assert!(has_fs, "expected import 'fs' in {:?}", result.imports);
}
#[test]
fn test_rust_scoped_use_imports() {
let source = "use std::{io, fs};\n";
let result = SemanticExtractor::extract(source, "rust", None, None).unwrap();
assert!(
!result.imports.is_empty(),
"expected imports in {:?}",
result.imports
);
}
#[test]
fn test_rust_wildcard_imports() {
let source = "use std::*;\n";
let result = SemanticExtractor::extract(source, "rust", None, None).unwrap();
let wildcard = result
.imports
.iter()
.find(|imp| imp.items.iter().any(|i| i == "*"));
assert!(
wildcard.is_some(),
"expected wildcard import in {:?}",
result.imports
);
}
#[test]
fn test_extract_impl_traits_standalone() {
let source = r#"
trait MyTrait {
fn method(&self);
}
impl MyTrait for MyType {
fn method(&self) {}
}
"#;
let result = extract_impl_traits(source, Path::new("test.rs"));
assert!(
!result.is_empty(),
"expected impl trait in result, got {:?}",
result
);
}
#[test]
fn test_ast_recursion_limit_overflow() {
let source = r#"fn hello() -> u32 { 42 }"#;
let result = SemanticExtractor::extract(source, "rust", Some(usize::MAX), None);
assert!(
result.is_ok(),
"extract with limit=usize::MAX should succeed"
);
}
#[test]
fn test_ast_recursion_limit_some() {
let source = r#"fn hello() -> u32 { 42 }"#;
let result = SemanticExtractor::extract(source, "rust", Some(10), None);
assert!(result.is_ok(), "extract with limit=10 should succeed");
let analysis = result.unwrap();
assert_eq!(
analysis.functions.len(),
1,
"should find exactly one function"
);
}
#[test]
fn test_extract_def_use_for_file_finds_write_and_read() {
let source = r#"
fn test() {
let mut x = 5;
x = 10;
let y = x;
}
"#;
let result =
SemanticExtractor::extract_def_use_for_file(source, "rust", "x", "test.rs", None);
let has_write = result
.iter()
.any(|s| s.kind == crate::types::DefUseKind::Write);
let has_read = result
.iter()
.any(|s| s.kind == crate::types::DefUseKind::Read);
assert!(has_write, "expected write site for 'x'");
assert!(has_read, "expected read site for 'x'");
}
#[test]
fn test_extract_def_use_for_file_no_match_returns_empty() {
let source = r#"
fn test() {
let x = 5;
}
"#;
let result = SemanticExtractor::extract_def_use_for_file(
source,
"rust",
"nonexistent",
"test.rs",
None,
);
assert!(
result.is_empty(),
"expected empty result for nonexistent symbol"
);
}
#[test]
fn extract_calls_does_not_panic_on_function_calls() {
let src = r#"
fn foo() {}
fn bar() {
foo();
}
"#;
let result = SemanticExtractor::extract(src, "rust", None, None);
assert!(
result.is_ok(),
"extract must succeed on source with function calls"
);
let output = result.unwrap();
assert!(
!output.calls.is_empty(),
"extract must return call entries for source with function calls"
);
}
#[test]
fn extract_calls_caps_arg_count_at_sixteen_hops() {
let src = r#"fn main() { f((((((((((((((((((((g())))))))))))))))))))); }"#;
let result = SemanticExtractor::extract(src, "rust", None, None);
assert!(
result.is_ok(),
"extract must succeed even with deeply nested parenthesized arguments"
);
let output = result.unwrap();
let g_calls: Vec<&CallInfo> = output.calls.iter().filter(|c| c.callee == "g").collect();
assert_eq!(
g_calls.len(),
1,
"expected exactly one CallInfo with callee 'g', got {}",
g_calls.len()
);
assert_eq!(
g_calls[0].arg_count,
Some(0),
"g() has 0 arguments, expected Some(0)"
);
}
}
#[cfg(test)]
mod tests_python {
use super::*;
#[test]
fn test_python_relative_import() {
let source = "from . import foo\n";
let result = SemanticExtractor::extract(source, "python", None, None).unwrap();
let relative = result.imports.iter().find(|imp| imp.module.contains("."));
assert!(
relative.is_some(),
"expected relative import in {:?}",
result.imports
);
}
#[test]
fn test_python_aliased_import() {
let source = "from os import path as p\n";
let result = SemanticExtractor::extract(source, "python", None, None).unwrap();
let path_import = result
.imports
.iter()
.find(|imp| imp.module == "os" && imp.items.iter().any(|i| i == "path"));
assert!(
path_import.is_some(),
"expected import 'path' from module 'os' in {:?}",
result.imports
);
}
#[test]
fn test_parse_no_timeout_when_none() {
let source = r#"fn hello() -> u32 { 42 }"#;
let result = SemanticExtractor::extract(source, "rust", None, None);
assert!(result.is_ok(), "extract with deadline=None should succeed");
let analysis = result.unwrap();
assert!(
analysis.functions.len() >= 1,
"should find at least one function"
);
}
#[test]
fn test_parse_timeout_triggers_error() {
let source = r#"fn hello() -> u32 { 42 }"#;
let result = SemanticExtractor::extract(source, "rust", None, Some(1u64));
assert!(
matches!(result, Err(ParserError::Timeout(_))),
"expected Timeout error, got {:?}",
result
);
}
}
#[cfg(test)]
mod tests_unsupported {
use super::*;
#[test]
fn test_element_extractor_unsupported_language() {
let result = ElementExtractor::extract_with_depth("x = 1", "cobol");
assert!(
matches!(result, Err(ParserError::UnsupportedLanguage(ref lang)) if lang == "cobol"),
"expected UnsupportedLanguage error, got {:?}",
result
);
}
#[test]
fn test_semantic_extractor_unsupported_language() {
let result = SemanticExtractor::extract("x = 1", "cobol", None, None);
assert!(
matches!(result, Err(ParserError::UnsupportedLanguage(ref lang)) if lang == "cobol"),
"expected UnsupportedLanguage error, got {:?}",
result
);
}
}