use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tokio::sync::Mutex;
use tree_sitter::{Node, Parser};
use super::engine::{ResolvedDefinition, StackGraphEngine};
use super::fallback::ScmFallbackResolver;
use crate::parsers::{crate_root_for, module_path_for};
use crate::payload::{RawCodeGraphEdge, RawScope, RawSymbol, RawSymbolReference};
use stack_graphs::NoCancellation;
#[derive(Debug, Clone, Default)]
pub struct ExtractedReferences {
pub file_path: PathBuf,
pub references: Vec<RawSymbolReference>,
pub edges: Vec<RawCodeGraphEdge>,
}
fn tree_sitter_language_for(rel_path: &Path) -> Option<tree_sitter::Language> {
let ext = rel_path.extension()?.to_str()?.to_ascii_lowercase();
match ext.as_str() {
"rs" => Some(tree_sitter_rust::LANGUAGE.into()),
"py" => Some(tree_sitter_python::LANGUAGE.into()),
"ts" | "js" => Some(tree_sitter_typescript::LANGUAGE_TYPESCRIPT.into()),
"tsx" | "jsx" => Some(tree_sitter_typescript::LANGUAGE_TSX.into()),
_ => None,
}
}
pub struct CrossFileResolver;
impl CrossFileResolver {
pub async fn resolve_project(
engine: Arc<Mutex<StackGraphEngine>>,
files: &[(PathBuf, String)],
all_symbols: &HashMap<PathBuf, Vec<RawSymbol>>,
all_scopes: &HashMap<PathBuf, Vec<RawScope>>,
) -> anyhow::Result<HashMap<PathBuf, ExtractedReferences>> {
{
let mut eng = engine.lock().await;
for (rel_path, content) in files {
let _ = eng.add_file(rel_path, content);
}
let _ = eng.precompute_all_files(&NoCancellation);
}
let mut global_symbols = HashMap::with_capacity(all_symbols.len() * 4);
for syms in all_symbols.values() {
for sym in syms {
global_symbols
.entry(sym.name.as_str())
.or_insert(sym.qualified_name.as_str());
}
}
let mut results = HashMap::new();
for (rel_path, content) in files {
let extracted = Self::resolve_file_references_indexed(
Arc::clone(&engine),
rel_path,
content,
all_symbols,
all_scopes,
&global_symbols,
)
.await?;
results.insert(rel_path.clone(), extracted);
}
Ok(results)
}
pub async fn resolve_file_references(
engine: Arc<Mutex<StackGraphEngine>>,
rel_path: &Path,
content: &str,
all_symbols: &HashMap<PathBuf, Vec<RawSymbol>>,
all_scopes: &HashMap<PathBuf, Vec<RawScope>>,
) -> anyhow::Result<ExtractedReferences> {
let mut global_symbols = HashMap::with_capacity(all_symbols.len() * 4);
for syms in all_symbols.values() {
for sym in syms {
global_symbols
.entry(sym.name.as_str())
.or_insert(sym.qualified_name.as_str());
}
}
Self::resolve_file_references_indexed(
engine,
rel_path,
content,
all_symbols,
all_scopes,
&global_symbols,
)
.await
}
async fn resolve_file_references_indexed(
engine: Arc<Mutex<StackGraphEngine>>,
rel_path: &Path,
content: &str,
all_symbols: &HashMap<PathBuf, Vec<RawSymbol>>,
all_scopes: &HashMap<PathBuf, Vec<RawScope>>,
global_symbols: &HashMap<&str, &str>,
) -> anyhow::Result<ExtractedReferences> {
let mut references = Vec::new();
let mut edges = Vec::new();
let Some(language) = tree_sitter_language_for(rel_path) else {
return Ok(ExtractedReferences {
file_path: rel_path.to_path_buf(),
..Default::default()
});
};
let mut parser = Parser::new();
parser.set_language(&language).map_err(|e| {
anyhow::anyhow!(
"Failed to set tree-sitter language for {}: {e}",
rel_path.display()
)
})?;
let tree = match parser.parse(content, None) {
Some(t) => t,
None => {
return Ok(ExtractedReferences {
file_path: rel_path.to_path_buf(),
references,
edges,
});
}
};
let file_symbols = all_symbols.get(rel_path).cloned().unwrap_or_default();
let file_scopes = all_scopes.get(rel_path).cloned().unwrap_or_default();
let mut visitor = ReferenceVisitor {
source: content,
candidates: Vec::new(),
};
visitor.traverse(tree.root_node(), None);
let is_rust = !matches!(
rel_path.extension().and_then(|e| e.to_str()),
Some("py" | "ts" | "tsx" | "js" | "jsx")
);
for cand in visitor.candidates {
if is_rust
&& (cand.role == "import" || cand.role == "reexport")
&& let Some(path) = cand.path.as_deref()
{
let mut eng = engine.lock().await;
if let Some(&file_handle) = eng.files.get(rel_path) {
if cand.role == "reexport" {
let _ = eng.add_reexport_alias(file_handle, None, &cand.identifier, path);
} else {
let _ = eng.add_import_stitching(file_handle, &cand.identifier, path);
}
}
}
let mut resolved_target: Option<String> = None;
if is_rust && let Some(path) = cand.path.as_deref() {
resolved_target = normalize_rust_path(path, rel_path);
}
if resolved_target.is_none() {
let mut eng = engine.lock().await;
if let Ok(Some(def)) =
eng.resolve_at_location(rel_path, cand.line_number, cand.col_number)
{
resolved_target =
Some(qualify_definition(&def, all_symbols).unwrap_or(def.symbol_name));
}
}
if resolved_target.is_none() {
if let Some(res) = ScmFallbackResolver::resolve_identifier(
&cand.identifier,
rel_path,
cand.line_number,
cand.col_number,
&file_scopes,
&file_symbols,
) {
resolved_target = Some(res.qualified_name);
} else if let Some(&target_qualified) = global_symbols.get(cand.identifier.as_str())
{
resolved_target = Some(target_qualified.to_string());
} else {
resolved_target = Some(cand.identifier.clone());
}
}
if let Some(target_qualified) = resolved_target {
references.push(RawSymbolReference {
target_symbol_name: target_qualified.clone(),
role: cand.role.clone(),
start_byte: cand.start_byte,
end_byte: cand.end_byte,
line_number: cand.line_number,
});
if cand.role == "call" {
let enclosing = enclosing_function(&file_symbols, cand.start_byte)
.or(cand.enclosing_symbol);
if let Some(enclosing) = enclosing {
edges.push(RawCodeGraphEdge {
source_symbol_name: enclosing,
target_symbol_name: target_qualified.clone(),
edge_type: "calls".to_string(),
});
}
} else if cand.role == "reexport"
&& let Some(reexport_sym) = file_symbols
.iter()
.find(|s| s.name == cand.identifier && s.kind == "reexport")
{
edges.push(RawCodeGraphEdge {
source_symbol_name: reexport_sym.qualified_name.clone(),
target_symbol_name: target_qualified,
edge_type: "reexports".to_string(),
});
}
}
}
references.dedup_by(|a, b| {
a.start_byte == b.start_byte
&& a.role == b.role
&& a.target_symbol_name == b.target_symbol_name
});
edges.dedup_by(|a, b| {
a.source_symbol_name == b.source_symbol_name
&& a.target_symbol_name == b.target_symbol_name
&& a.edge_type == b.edge_type
});
Ok(ExtractedReferences {
file_path: rel_path.to_path_buf(),
references,
edges,
})
}
}
struct CandidateReference {
identifier: String,
path: Option<String>,
role: String,
start_byte: usize,
end_byte: usize,
line_number: u32,
col_number: u32,
enclosing_symbol: Option<String>,
}
struct ReferenceVisitor<'a> {
source: &'a str,
candidates: Vec<CandidateReference>,
}
impl<'a> ReferenceVisitor<'a> {
fn text(&self, node: Node) -> &'a str {
&self.source[node.start_byte()..node.end_byte()]
}
fn traverse_callee(&mut self, func: Node, current_enclosing: Option<&str>) {
match func.kind() {
"identifier" | "scoped_identifier" | "field_identifier" | "type_identifier" => {}
"field_expression" | "attribute" | "member_expression" => {
if let Some(receiver) = func
.child_by_field_name("value")
.or_else(|| func.child_by_field_name("object"))
{
self.traverse(receiver, current_enclosing);
}
}
"generic_function" => {
if let Some(inner) = func.child_by_field_name("function") {
self.traverse_callee(inner, current_enclosing);
}
}
_ => self.traverse(func, current_enclosing),
}
}
fn traverse(&mut self, node: Node, current_enclosing: Option<&str>) {
match node.kind() {
"function_item"
| "function_definition"
| "function_declaration"
| "method_definition" => {
let name = node
.child_by_field_name("name")
.map(|n| self.text(n))
.unwrap_or("anonymous_fn");
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
if child.kind() != "visibility_modifier"
&& child.kind() != "parameters"
&& child.kind() != "formal_parameters"
{
self.traverse(child, Some(name));
}
}
}
"call_expression" | "call" => {
if let Some(func) = node.child_by_field_name("function") {
let func = if func.kind() == "generic_function" {
func.child_by_field_name("function").unwrap_or(func)
} else {
func
};
let path = (func.kind() == "scoped_identifier")
.then(|| strip_generics(self.text(func)));
let ident_info = match func.kind() {
"identifier" => Some((
self.text(func).to_string(),
func.start_byte(),
func.end_byte(),
(func.start_position().row + 1) as u32,
func.start_position().column as u32,
)),
"scoped_identifier" => func.child_by_field_name("name").map(|n| {
(
self.text(n).to_string(),
n.start_byte(),
n.end_byte(),
(n.start_position().row + 1) as u32,
n.start_position().column as u32,
)
}),
"field_expression" => func.child_by_field_name("field").map(|n| {
(
self.text(n).to_string(),
n.start_byte(),
n.end_byte(),
(n.start_position().row + 1) as u32,
n.start_position().column as u32,
)
}),
"attribute" => func.child_by_field_name("attribute").map(|n| {
(
self.text(n).to_string(),
n.start_byte(),
n.end_byte(),
(n.start_position().row + 1) as u32,
n.start_position().column as u32,
)
}),
"member_expression" => func.child_by_field_name("property").map(|n| {
(
self.text(n).to_string(),
n.start_byte(),
n.end_byte(),
(n.start_position().row + 1) as u32,
n.start_position().column as u32,
)
}),
_ => None,
};
if let Some((id, sb, eb, ln, col)) = ident_info {
self.candidates.push(CandidateReference {
identifier: id,
path,
role: "call".to_string(),
start_byte: sb,
end_byte: eb,
line_number: ln,
col_number: col,
enclosing_symbol: current_enclosing.map(|s| s.to_string()),
});
}
}
if let Some(args) = node
.child_by_field_name("arguments")
.or_else(|| node.child_by_field_name("argument_list"))
{
let mut cursor = args.walk();
for child in args.children(&mut cursor) {
self.traverse(child, current_enclosing);
}
}
if let Some(func) = node.child_by_field_name("function") {
self.traverse_callee(func, current_enclosing);
}
}
"type_identifier" => {
let is_decl = node.parent().is_some_and(|p| {
(p.kind() == "struct_item"
|| p.kind() == "enum_item"
|| p.kind() == "trait_item"
|| p.kind() == "type_item")
&& p.child_by_field_name("name") == Some(node)
});
if !is_decl {
let path = node
.parent()
.filter(|p| p.kind() == "scoped_type_identifier")
.map(|p| strip_generics(self.text(p)));
self.candidates.push(CandidateReference {
identifier: self.text(node).to_string(),
path,
role: "type_annotation".to_string(),
start_byte: node.start_byte(),
end_byte: node.end_byte(),
line_number: (node.start_position().row + 1) as u32,
col_number: node.start_position().column as u32,
enclosing_symbol: current_enclosing.map(|s| s.to_string()),
});
}
}
"use_declaration" => {
let is_pub = node
.children(&mut node.walk())
.any(|c| c.kind() == "visibility_modifier" && self.text(c).trim() == "pub");
let role = if is_pub { "reexport" } else { "import" };
let items = crate::parsers::extract_use_items(self.source, node);
for item in items {
self.candidates.push(CandidateReference {
identifier: item.local_name,
path: Some(item.target_path),
role: role.to_string(),
start_byte: item.start_byte,
end_byte: item.end_byte,
line_number: item.start_line,
col_number: item.col_number,
enclosing_symbol: current_enclosing.map(|s| s.to_string()),
});
}
}
_ => {
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
self.traverse(child, current_enclosing);
}
}
}
}
}
fn strip_generics(path: &str) -> String {
let mut out = String::with_capacity(path.len());
let mut depth = 0usize;
for c in path.chars() {
match c {
'<' => depth += 1,
'>' => depth = depth.saturating_sub(1),
c if depth == 0 && !c.is_whitespace() => out.push(c),
_ => {}
}
}
out.replace("::::", "::").trim_end_matches("::").to_string()
}
pub fn normalize_rust_path(path: &str, rel_path: &Path) -> Option<String> {
let segments: Vec<&str> = path
.trim_start_matches("::")
.split("::")
.filter(|s| !s.is_empty())
.collect();
let (&first, rest) = segments.split_first()?;
if rest.is_empty() {
return Some(first.to_string());
}
let mut base: Vec<String> = match first {
"crate" => vec![crate_root_for(rel_path)],
"self" | "super" => module_path_for(rel_path)
.split("::")
.map(str::to_string)
.collect(),
"Self" => return rest.last().map(|s| s.to_string()),
_ => return Some(segments.join("::")),
};
let mut rest = rest;
if first == "super" {
base.pop();
}
while let Some((&"super", tail)) = rest.split_first() {
base.pop();
rest = tail;
}
base.extend(rest.iter().map(|s| s.to_string()));
Some(base.join("::"))
}
fn qualify_definition(
def: &ResolvedDefinition,
all_symbols: &HashMap<PathBuf, Vec<RawSymbol>>,
) -> Option<String> {
all_symbols
.get(&def.file_path)?
.iter()
.find(|s| {
s.name == def.symbol_name
&& s.start_line <= def.start_line
&& def.start_line <= s.end_line
})
.map(|s| s.qualified_name.clone())
}
fn enclosing_function(symbols: &[RawSymbol], byte: usize) -> Option<String> {
symbols
.iter()
.filter(|s| {
matches!(s.kind.as_str(), "fn" | "method") && s.start_byte <= byte && byte < s.end_byte
})
.min_by_key(|s| s.end_byte - s.start_byte)
.map(|s| s.qualified_name.clone())
}
pub fn reference_target_at(rel_path: &Path, content: &str, line: u32, col: u32) -> Option<String> {
let language = tree_sitter_language_for(rel_path)?;
let mut parser = Parser::new();
parser.set_language(&language).ok()?;
let tree = parser.parse(content, None)?;
let point = tree_sitter::Point::new(line.checked_sub(1)? as usize, col as usize);
let node = tree
.root_node()
.named_descendant_for_point_range(point, point)?;
if !node.kind().ends_with("identifier") {
return None;
}
let text = |n: Node| &content[n.start_byte()..n.end_byte()];
let is_rust = !matches!(
rel_path.extension().and_then(|e| e.to_str()),
Some("py" | "ts" | "tsx" | "js" | "jsx")
);
if is_rust {
let scoped_parent = node.parent().filter(|p| {
matches!(p.kind(), "scoped_identifier" | "scoped_type_identifier")
&& p.child_by_field_name("name") == Some(node)
});
if let Some(parent) = scoped_parent {
return normalize_rust_path(&strip_generics(text(parent)), rel_path);
}
}
Some(text(node).to_string())
}