1use std::path::Path;
2
3use tree_sitter::StreamingIterator;
4
5use crate::language::LanguageRegistry;
6use crate::selector::{ParsedSelector, SelectorTarget, SymbolKind, SymbolSelector};
7
8#[derive(Debug, Clone, PartialEq, Eq)]
9pub struct SymbolMatch {
10 pub name: String,
11 pub kind: SymbolKind,
12 pub kind_prefix: &'static str,
13 pub start_line: usize,
14 pub end_line: usize,
15}
16
17impl SymbolMatch {
18 pub fn canonical_selector(&self, file_path: &Path, project_root: &Path) -> String {
19 let rel_path = file_path
20 .strip_prefix(project_root)
21 .ok()
22 .map(|p| p.to_string_lossy().to_string())
23 .unwrap_or_else(|| file_path.to_string_lossy().to_string())
24 .replace('\\', "/");
25
26 format!(
27 "{}::{}{} #L{}-L{}",
28 rel_path, self.kind_prefix, self.name, self.start_line, self.end_line
29 )
30 }
31
32 pub fn source_from(&self, content: &str) -> String {
33 let lines: Vec<&str> = content.lines().collect();
34 if self.start_line == 0 || self.end_line < self.start_line || self.start_line > lines.len()
35 {
36 return String::new();
37 }
38
39 let start_idx = self.start_line - 1;
40 let end_idx = self.end_line.min(lines.len());
41 let mut snippet = lines[start_idx..end_idx].join("\n");
42 if content.ends_with('\n') || self.end_line < lines.len() {
43 snippet.push('\n');
44 }
45 snippet
46 }
47}
48
49pub struct TreeSitterAnalyzer {
50 registry: LanguageRegistry,
51}
52
53impl Default for TreeSitterAnalyzer {
54 fn default() -> Self {
55 Self::new()
56 }
57}
58
59impl TreeSitterAnalyzer {
60 pub fn new() -> Self {
61 Self {
62 registry: LanguageRegistry::new(),
63 }
64 }
65
66 pub fn find_containing_symbol(
71 &self,
72 file_path: &Path,
73 line_number: usize,
74 project_root: &Path,
75 ) -> Option<String> {
76 self.find_containing_symbol_match(file_path, line_number)
77 .map(|m| m.canonical_selector(file_path, project_root))
78 }
79
80 pub fn find_containing_symbol_match(
81 &self,
82 file_path: &Path,
83 line_number: usize,
84 ) -> Option<SymbolMatch> {
85 let symbols = self.symbols_in_file(file_path).ok()?;
86 symbols
87 .into_iter()
88 .filter(|m| line_number >= m.start_line && line_number <= m.end_line)
89 .max_by_key(|m| (m.start_line, usize::MAX - m.end_line))
90 }
91
92 pub fn resolve_selector(
93 &self,
94 file_path: &Path,
95 parsed: &ParsedSelector,
96 ) -> Result<SymbolMatch, String> {
97 let symbols = self.symbols_in_file(file_path)?;
98 let mut matches: Vec<SymbolMatch> = symbols
99 .into_iter()
100 .filter(|m| symbol_matches_selector(m, parsed))
101 .collect();
102
103 let symbol = match &parsed.target {
104 SelectorTarget::Symbol(symbol) => symbol,
105 _ => {
106 return Err(format!(
107 "selector target is not a symbol and cannot be resolved as a symbol: {}",
108 file_path.display()
109 ));
110 }
111 };
112
113 if let Some((start, end)) = symbol.line_range {
114 matches.retain(|m| m.start_line == start && m.end_line == end);
115 }
116
117 match matches.len() {
118 0 => Err(format!(
119 "symbol '{}' not found in {}",
120 symbol.name,
121 file_path.display()
122 )),
123 1 => Ok(matches.remove(0)),
124 _ => {
125 let candidates = matches
126 .iter()
127 .map(|m| {
128 format!(
129 "{}{} #L{}-L{}",
130 m.kind_prefix, m.name, m.start_line, m.end_line
131 )
132 })
133 .collect::<Vec<_>>()
134 .join(", ");
135 Err(format!(
136 "ambiguous selector for '{}' in {}; candidates: {}",
137 symbol.name,
138 file_path.display(),
139 candidates
140 ))
141 }
142 }
143 }
144
145 pub fn symbols_in_file(&self, file_path: &Path) -> Result<Vec<SymbolMatch>, String> {
146 let ext = file_path
147 .extension()
148 .and_then(|e| e.to_str())
149 .ok_or_else(|| {
150 format!(
151 "cannot determine language from file: {}",
152 file_path.display()
153 )
154 })?;
155 let adapter = self
156 .registry
157 .get(ext)
158 .ok_or_else(|| format!("unsupported language extension: {ext}"))?;
159
160 let content = std::fs::read_to_string(file_path)
161 .map_err(|e| format!("failed to read {}: {e}", file_path.display()))?;
162 let mut parser = adapter.parser();
163 let tree = parser
164 .parse(&content, None)
165 .ok_or_else(|| format!("failed to parse {}", file_path.display()))?;
166 let mut symbols = Vec::new();
167 self.collect_symbols(tree.root_node(), &content, &mut symbols);
168 Ok(symbols)
169 }
170
171 pub fn is_import_only_reference(&self, file_path: &Path, line_number: usize) -> bool {
172 let ext = match file_path.extension().and_then(|e| e.to_str()) {
173 Some(ext) => ext,
174 None => return false,
175 };
176 let Some(adapter) = self.registry.get(ext) else {
177 return false;
178 };
179 if adapter.language_name() != "rust" {
180 return false;
181 }
182
183 let content = match std::fs::read_to_string(file_path) {
184 Ok(content) => content,
185 Err(_) => return false,
186 };
187 let mut parser = adapter.parser();
188 let Some(tree) = parser.parse(&content, None) else {
189 return false;
190 };
191
192 let query = match tree_sitter::Query::new(&adapter.language(), RUST_USE_IMPORT_QUERY) {
193 Ok(query) => query,
194 Err(_) => return false,
195 };
196 let mut cursor = tree_sitter::QueryCursor::new();
197 let mut matches = cursor.matches(&query, tree.root_node(), content.as_bytes());
198 while let Some(query_match) = matches.next() {
199 for capture in query_match.captures {
200 let node = capture.node;
201 let start = node.start_position().row + 1;
202 let end = node.end_position().row + 1;
203 if line_number >= start && line_number <= end {
204 return true;
205 }
206 }
207 }
208
209 false
210 }
211
212 pub fn can_parse(&self, ext: &str, content: &str) -> bool {
216 let adapter = match self.registry.get(ext) {
217 Some(a) => a,
218 None => return false,
219 };
220 let mut parser = adapter.parser();
221 parser
222 .parse(content, None)
223 .is_some_and(|tree| !node_has_parse_error(tree.root_node()))
224 }
225
226 pub fn responsible_language_for_extension(&self, ext: &str) -> Option<&'static str> {
229 self.registry
230 .get(ext)
231 .map(|adapter| adapter.language_name())
232 }
233
234 pub fn is_responsible_source_path(&self, file_path: &Path) -> bool {
236 file_path
237 .extension()
238 .and_then(|ext| ext.to_str())
239 .and_then(|ext| self.responsible_language_for_extension(ext))
240 .is_some()
241 }
242
243 fn collect_symbols(
244 &self,
245 node: tree_sitter::Node,
246 source: &str,
247 symbols: &mut Vec<SymbolMatch>,
248 ) {
249 let kind = node.kind();
250 if is_definition_kind(kind)
251 && let Some(name) = self.extract_def_name(node, source)
252 {
253 let start_line = node.start_position().row + 1;
254 let end_line = node.end_position().row + 1;
255 symbols.push(SymbolMatch {
256 name,
257 kind: SymbolKind::from_ts_node_kind(kind),
258 kind_prefix: kind_prefix(kind),
259 start_line,
260 end_line,
261 });
262 }
263
264 for i in 0..node.child_count() {
265 if let Some(child) = node.child(i) {
266 self.collect_symbols(child, source, symbols);
267 }
268 }
269 }
270
271 fn extract_def_name(&self, node: tree_sitter::Node, source: &str) -> Option<String> {
272 for i in 0..node.child_count() {
273 let child = node.child(i)?;
274 let kind = child.kind();
275 if kind == "identifier" || kind == "type_identifier" {
276 return child
277 .utf8_text(source.as_bytes())
278 .ok()
279 .map(|s| s.to_string());
280 }
281 }
282 None
283 }
284}
285
286fn symbol_matches_selector(symbol: &SymbolMatch, parsed: &ParsedSelector) -> bool {
287 let Some(selector) = parsed.as_symbol() else {
288 return false;
289 };
290 symbol_matches_symbol_selector(symbol, selector)
291}
292
293fn symbol_matches_symbol_selector(symbol: &SymbolMatch, selector: &SymbolSelector) -> bool {
294 symbol.name == selector.name
295 && (selector.kind == SymbolKind::Unknown || symbol.kind == selector.kind)
296}
297
298fn is_definition_kind(kind: &str) -> bool {
299 matches!(
300 kind,
301 "function_item"
302 | "struct_item"
303 | "enum_item"
304 | "trait_item"
305 | "impl_item"
306 | "function_definition"
307 | "class_definition"
308 | "decorated_definition"
309 | "function_declaration"
310 | "class_declaration"
311 | "interface_declaration"
312 | "enum_declaration"
313 | "method_definition"
314 | "type_alias_declaration"
315 )
316}
317
318fn kind_prefix(kind: &str) -> &'static str {
319 match kind {
320 "function_item" => "fn ",
321 "struct_item" => "struct ",
322 "enum_item" => "enum ",
323 "trait_item" => "trait ",
324 "impl_item" => "impl ",
325 "function_definition" => "fn ",
326 "class_definition" => "class ",
327 "function_declaration" => "fn ",
328 "class_declaration" => "class ",
329 "interface_declaration" => "trait ",
330 "enum_declaration" => "enum ",
331 "method_definition" => "fn ",
332 "type_alias_declaration" => "type ",
333 _ => "",
334 }
335}
336
337fn node_has_parse_error(node: tree_sitter::Node<'_>) -> bool {
338 if node.has_error() || node.is_error() || node.is_missing() {
339 return true;
340 }
341
342 let mut cursor = node.walk();
343 node.children(&mut cursor).any(node_has_parse_error)
344}
345
346const RUST_USE_IMPORT_QUERY: &str = "(use_declaration) @import";
347
348#[cfg(test)]
349mod tests {
350 use super::*;
351 use std::io::Write;
352 use std::path::PathBuf;
353
354 fn write_temp_rust_file(dir: &Path, name: &str, content: &str) -> PathBuf {
355 let path = dir.join(name);
356 let mut f = std::fs::File::create(&path).unwrap();
357 f.write_all(content.as_bytes()).unwrap();
358 path
359 }
360
361 const RUST_CODE: &str = "// line 1\n fn startup() {\n inner_call();\n }\n }\n ";
362
363 #[test]
364 fn test_find_containing_symbol_fn() {
365 let dir = tempfile::tempdir().unwrap();
366 let path = write_temp_rust_file(dir.path(), "test.rs", RUST_CODE);
367 let analyzer = TreeSitterAnalyzer::new();
368 let result = analyzer.find_containing_symbol(&path, 3, dir.path());
370 println!("find_containing_symbol result: {:?}", result);
372 }
373
374 #[test]
375 fn symbol_match_source_from_returns_exact_line_range() {
376 let symbol = SymbolMatch {
377 name: "target".to_string(),
378 kind: SymbolKind::Function,
379 kind_prefix: "fn ",
380 start_line: 3,
381 end_line: 5,
382 };
383 let content = "line 1\nline 2\nfn target() {\n body();\n}\nfn other() {}\n";
384 assert_eq!(
385 symbol.source_from(content),
386 "fn target() {\n body();\n}\n"
387 );
388 }
389
390 #[test]
391 fn tsx_files_use_tsx_parser_and_expose_top_level_functions() {
392 let dir = tempfile::tempdir().unwrap();
393 let path = dir.path().join("status-page.tsx");
394 std::fs::write(
395 &path,
396 "function AgentChatActivityHeader() {\n return <div />;\n}\n\nfunction agentChatActivityGlyph(bubble: { kind: string }) {\n return bubble.kind;\n}\n",
397 )
398 .unwrap();
399 let analyzer = TreeSitterAnalyzer::new();
400
401 assert!(analyzer.can_parse("tsx", &std::fs::read_to_string(&path).unwrap()));
402 let symbol = analyzer
403 .resolve_selector(
404 &path,
405 &crate::selector::parse_selector("status-page.tsx::fn agentChatActivityGlyph")
406 .unwrap(),
407 )
408 .expect("TSX top-level function should resolve");
409 assert_eq!(symbol.name, "agentChatActivityGlyph");
410 assert_eq!(symbol.start_line, 5);
411 }
412
413 #[test]
414 fn canonical_selector_disambiguates_duplicate_method_names() {
415 let dir = tempfile::tempdir().unwrap();
416 let code = r#"trait Hints {
417 fn setup_hints(&self);
418}
419
420struct Alpha;
421struct Beta;
422
423impl Hints for Alpha {
424 fn setup_hints(&self) {
425 println!("alpha");
426 }
427}
428
429impl Hints for Beta {
430 fn setup_hints(&self) {
431 println!("beta");
432 }
433}
434"#;
435 let path = write_temp_rust_file(dir.path(), "dup.rs", code);
436 let analyzer = TreeSitterAnalyzer::new();
437
438 let canonical = analyzer
439 .find_containing_symbol(&path, 16, dir.path())
440 .expect("line inside Beta::setup_hints should resolve");
441 assert!(canonical.starts_with("dup.rs::fn setup_hints #L"));
442 assert!(canonical.contains("-L"));
443
444 let parsed = crate::selector::parse_selector(&canonical).unwrap();
445 let resolved = analyzer.resolve_selector(&path, &parsed).unwrap();
446 assert_eq!(resolved.name, "setup_hints");
447 assert_eq!(resolved.start_line, 15);
448 }
449
450 #[test]
451 fn legacy_duplicate_method_selector_is_rejected_as_ambiguous() {
452 let dir = tempfile::tempdir().unwrap();
453 let code = r#"trait Hints {
454 fn setup_hints(&self);
455}
456
457struct Alpha;
458struct Beta;
459
460impl Hints for Alpha {
461 fn setup_hints(&self) {}
462}
463
464impl Hints for Beta {
465 fn setup_hints(&self) {}
466}
467"#;
468 let path = write_temp_rust_file(dir.path(), "dup.rs", code);
469 let analyzer = TreeSitterAnalyzer::new();
470 let parsed = crate::selector::parse_selector("dup.rs::fn setup_hints").unwrap();
471 let err = analyzer.resolve_selector(&path, &parsed).unwrap_err();
472 assert!(err.contains("ambiguous selector"));
473 assert!(err.contains("#L"));
474 }
475
476 #[test]
477 fn test_can_parse_valid_rust() {
478 let analyzer = TreeSitterAnalyzer::new();
479 let valid = "fn main() { println!(\"hello\"); }";
480 assert!(analyzer.can_parse("rs", valid));
481 }
482
483 #[test]
484 fn rust_use_declaration_is_import_only_reference() {
485 let dir = tempfile::tempdir().unwrap();
486 let code = r#"use crate::parser::Parser;
487use crate::{engine::Engine, runtime};
488
489fn run() {
490 Parser::new();
491}
492"#;
493 let path = write_temp_rust_file(dir.path(), "imports.rs", code);
494 let analyzer = TreeSitterAnalyzer::new();
495
496 assert!(analyzer.is_import_only_reference(&path, 1));
497 assert!(analyzer.is_import_only_reference(&path, 2));
498 assert!(!analyzer.is_import_only_reference(&path, 5));
499 }
500
501 #[test]
502 fn test_can_parse_rejects_rust_error_nodes() {
503 let analyzer = TreeSitterAnalyzer::new();
504 let invalid = "fn main( {\n";
505 assert!(!analyzer.can_parse("rs", invalid));
506 }
507
508 #[test]
509 fn test_can_parse_empty_string() {
510 let analyzer = TreeSitterAnalyzer::new();
511 assert!(analyzer.can_parse("rs", ""));
512 }
513
514 #[test]
515 fn test_can_parse_unknown_language_returns_false() {
516 let analyzer = TreeSitterAnalyzer::new();
517 assert!(!analyzer.can_parse("unknown_ext", "fn main() {}"));
518 }
519
520 #[test]
521 fn test_can_parse_valid_python() {
522 let analyzer = TreeSitterAnalyzer::new();
523 let py_code = "def greet(name):\n return f\"Hello, {name}!\"\n";
524 assert!(analyzer.can_parse("py", py_code));
525 }
526
527 #[test]
528 fn test_can_parse_valid_go() {
529 let analyzer = TreeSitterAnalyzer::new();
530 let go_code = "package main\nfunc greet(name string) string { return \"Hello\" }\n";
531 assert!(analyzer.can_parse("go", go_code));
532 }
533
534 #[test]
535 fn test_can_parse_valid_java() {
536 let analyzer = TreeSitterAnalyzer::new();
537 let java_code = "public class Hello { public static void main(String[] args) {} }\n";
538 assert!(analyzer.can_parse("java", java_code));
539 }
540
541 #[test]
542 fn test_can_parse_valid_typescript() {
543 let analyzer = TreeSitterAnalyzer::new();
544 let ts_code = "function greet(name: string): string { return \"Hello\"; }\n";
545 assert!(analyzer.can_parse("ts", ts_code));
546 }
547
548 #[test]
549 fn test_can_parse_valid_javascript() {
550 let analyzer = TreeSitterAnalyzer::new();
551 let js_code = "function greet(name) { return \"Hello\"; }\n";
552 assert!(analyzer.can_parse("js", js_code));
553 }
554
555 #[test]
556 fn test_can_parse_valid_c() {
557 let analyzer = TreeSitterAnalyzer::new();
558 let c_code = "int main() { return 0; }\n";
559 assert!(analyzer.can_parse("c", c_code));
560 }
561
562 #[test]
563 fn test_can_parse_valid_cpp() {
564 let analyzer = TreeSitterAnalyzer::new();
565 let cpp_code = "class Hello { public: void greet() {} };\n";
566 assert!(analyzer.can_parse("cpp", cpp_code));
567 }
568
569 #[test]
570 fn test_can_parse_valid_ruby() {
571 let analyzer = TreeSitterAnalyzer::new();
572 let ruby_code = "def greet(name)\n \"Hello, #{name}!\"\nend\n";
573 assert!(analyzer.can_parse("rb", ruby_code));
574 }
575
576 #[test]
577 fn test_can_parse_valid_php() {
578 let analyzer = TreeSitterAnalyzer::new();
579 let php_code = "<?php\nfunction greet($name) { return \"Hello\"; }\n";
580 assert!(analyzer.can_parse("php", php_code));
581 }
582
583 #[test]
584 fn responsible_source_path_matches_registered_languages() {
585 let analyzer = TreeSitterAnalyzer::new();
586 assert!(analyzer.is_responsible_source_path(Path::new("src/lib.rs")));
587 assert!(analyzer.is_responsible_source_path(Path::new("script.py")));
588 assert!(!analyzer.is_responsible_source_path(Path::new("README.md")));
589 assert!(!analyzer.is_responsible_source_path(Path::new("Makefile")));
590 assert_eq!(
591 analyzer.responsible_language_for_extension("rs"),
592 Some("rust")
593 );
594 assert_eq!(analyzer.responsible_language_for_extension("md"), None);
595 }
596 #[test]
597 fn test_language_registry_has_all_languages() {
598 let registry = LanguageRegistry::new();
599 assert!(registry.get("rs").is_some(), "Rust should be registered");
600 assert!(registry.get("py").is_some(), "Python should be registered");
601 assert!(registry.get("go").is_some(), "Go should be registered");
602 assert!(registry.get("java").is_some(), "Java should be registered");
603 assert!(
604 registry.get("ts").is_some(),
605 "TypeScript should be registered"
606 );
607 assert!(
608 registry.get("js").is_some(),
609 "JavaScript should be registered"
610 );
611 assert!(registry.get("c").is_some(), "C should be registered");
612 assert!(registry.get("cpp").is_some(), "C++ should be registered");
613 assert!(registry.get("rb").is_some(), "Ruby should be registered");
614 assert!(registry.get("php").is_some(), "PHP should be registered");
615 }
616
617 #[test]
618 fn test_language_registry_all_names() {
619 let registry = LanguageRegistry::new();
620 let langs = registry.list_languages();
621 let names: Vec<&str> = langs.iter().map(|(n, _)| *n).collect();
622 assert!(names.contains(&"rust"), "rust in {:?}", names);
623 assert!(names.contains(&"python"), "python in {:?}", names);
624 assert!(names.contains(&"go"), "go in {:?}", names);
625 assert!(names.contains(&"java"), "java in {:?}", names);
626 assert!(names.contains(&"typescript"), "typescript in {:?}", names);
627 assert!(names.contains(&"javascript"), "javascript in {:?}", names);
628 assert!(names.contains(&"c"), "c in {:?}", names);
629 assert!(names.contains(&"cpp"), "cpp in {:?}", names);
630 assert!(names.contains(&"ruby"), "ruby in {:?}", names);
631 assert!(names.contains(&"php"), "php in {:?}", names);
632 }
633}