Skip to main content

scope_engine/
treesitter.rs

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    /// Given a file path and a 1-based line number, find the innermost
67    /// named definition (function, struct, enum, trait, impl) that contains
68    /// that line. Returns a canonical CodeStruct-style selector like
69    /// `src/foo.rs::fn authenticate #L10-L20`.
70    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    /// Validate that a file's content can be parsed by tree-sitter.
213    /// Returns true if parsing succeeds (i.e. the file is syntactically valid
214    /// for the given language), false otherwise.
215    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    /// Return the SCOPE language adapter that owns semantic source operations
227    /// for the given extension.
228    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    /// Return true when SCOPE owns semantic source operations for this path.
235    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        // Line 4 should be inside startup() (adjusted for the actual structure)
369        let result = analyzer.find_containing_symbol(&path, 3, dir.path());
370        // Just check it doesn't crash; exact line numbers depend on the test string
371        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}