Skip to main content

diffler_core/syntax/
scope.rs

1//! "What are we inside of": the chain of enclosing definitions (function,
2//! class, method, …) for a line, derived from a grammar's tags query. The
3//! result is plain data so the diff worker can compute it once and share it.
4
5use tree_sitter::{Node, QueryCursor, StreamingIterator};
6
7use crate::syntax::registry::LanguageRegistry;
8use crate::syntax::{MAX_PARSE_BYTES, parse};
9
10/// Definition spans for a file, queried per line for the enclosing-definition
11/// breadcrumb.
12#[derive(Debug, Clone, Default)]
13pub struct ScopeIndex {
14    defs: Vec<Def>,
15}
16
17#[derive(Debug, Clone)]
18struct Def {
19    start_row: usize,
20    end_row: usize,
21    name: String,
22}
23
24impl ScopeIndex {
25    pub fn is_empty(&self) -> bool {
26        self.defs.is_empty()
27    }
28
29    /// 0-based start rows of every definition, sorted and deduped: the jump
30    /// targets for function/definition motions.
31    pub fn def_starts(&self) -> Vec<usize> {
32        let mut rows: Vec<usize> = self.defs.iter().map(|d| d.start_row).collect();
33        rows.sort_unstable();
34        rows.dedup();
35        rows
36    }
37
38    /// 0-based row span of the definition called `name`, start and end
39    /// inclusive. What lets a reference open on a symbol and show its whole
40    /// extent rather than seating a cursor on its first line.
41    pub fn def_span(&self, name: &str) -> Option<(usize, usize)> {
42        self.defs
43            .iter()
44            .find(|def| def.name == name)
45            .map(|def| (def.start_row, def.end_row))
46    }
47
48    /// Names of the definitions enclosing `line` (0-based), outermost first. A
49    /// line inside `class A` → `method` → body returns `["A", "method"]`.
50    pub fn crumbs(&self, line: usize) -> Vec<String> {
51        let mut hits: Vec<&Def> = self
52            .defs
53            .iter()
54            .filter(|d| d.start_row <= line && line <= d.end_row)
55            .collect();
56        hits.sort_by(|a, b| {
57            a.start_row
58                .cmp(&b.start_row)
59                .then(b.end_row.cmp(&a.end_row))
60        });
61        hits.into_iter().map(|d| d.name.clone()).collect()
62    }
63}
64
65impl LanguageRegistry {
66    /// Parse `content` and index its definition spans for scope lookup. Returns
67    /// an empty index when the language is unsupported, has no tags query, the
68    /// file is too large, or parsing fails, so callers show no breadcrumb.
69    pub fn scope_index(&self, path: &str, content: &str) -> ScopeIndex {
70        if content.len() > MAX_PARSE_BYTES {
71            return ScopeIndex::default();
72        }
73        let Some(entry) = self.for_path(path) else {
74            return ScopeIndex::default();
75        };
76        let Some(query) = entry.tags() else {
77            return ScopeIndex::default();
78        };
79        let Some(tree) = parse(entry, content) else {
80            return ScopeIndex::default();
81        };
82
83        let names = query.capture_names();
84        let bytes = content.as_bytes();
85        let mut cursor = QueryCursor::new();
86        let mut matches = cursor.matches(query, tree.root_node(), bytes);
87        let mut defs = Vec::new();
88        while let Some(m) = matches.next() {
89            let mut span: Option<(usize, usize)> = None;
90            let mut name: Option<String> = None;
91            let mut qualified: Option<String> = None;
92            for cap in m.captures {
93                let cname = names.get(cap.index as usize).copied().unwrap_or("");
94                if cname.starts_with("definition.") {
95                    let node = whole_definition(cap.node);
96                    span = Some((node.start_position().row, node.end_position().row));
97                } else if cname == "name" {
98                    name = cap.node.utf8_text(bytes).ok().map(str::to_owned);
99                    qualified = cap
100                        .node
101                        .parent()
102                        .filter(|parent| parent.kind() == "qualified_identifier")
103                        .and_then(|parent| parent.utf8_text(bytes).ok())
104                        .map(str::to_owned);
105                }
106            }
107            if let (Some((start_row, end_row)), Some(name)) = (span, name) {
108                for name in std::iter::once(name).chain(qualified) {
109                    defs.push(Def {
110                        start_row,
111                        end_row,
112                        name,
113                    });
114                }
115            }
116        }
117        ScopeIndex { defs }
118    }
119}
120
121/// The node a definition spans. C and C++ tag a function by its declarator,
122/// which holds only the signature, so we climb to the function definition
123/// around it to take in the body; a prototype has none and keeps its own.
124fn whole_definition(node: Node<'_>) -> Node<'_> {
125    let mut current = node;
126    while let Some(parent) = current.parent() {
127        match parent.kind() {
128            "function_definition" => return parent,
129            kind if kind.ends_with("declarator") => current = parent,
130            _ => break,
131        }
132    }
133    node
134}
135
136#[cfg(test)]
137mod tests {
138    use super::*;
139
140    #[test]
141    fn nested_python_scope_reads_class_then_method() {
142        let reg = LanguageRegistry::build();
143        let src = "class A:\n    def method(self):\n        x = 1\n        return x\n";
144        let crumbs = reg.scope_index("a.py", src).crumbs(2);
145        let names: Vec<&str> = crumbs.iter().map(String::as_str).collect();
146        assert_eq!(names, ["A", "method"]);
147    }
148
149    #[test]
150    fn rust_function_scope() {
151        let reg = LanguageRegistry::build();
152        let src = "fn outer() {\n    let y = 2;\n}\n";
153        let crumbs = reg.scope_index("a.rs", src).crumbs(1);
154        let names: Vec<&str> = crumbs.iter().map(String::as_str).collect();
155        assert_eq!(names, ["outer"]);
156    }
157
158    #[test]
159    fn scope_index_empty_for_unsupported_language() {
160        let reg = LanguageRegistry::build();
161        assert!(reg.scope_index("a.zzz-unknown", "whatever\n").is_empty());
162    }
163
164    #[test]
165    fn line_outside_any_definition_has_no_crumbs() {
166        let reg = LanguageRegistry::build();
167        let src = "import os\n\ndef f():\n    pass\n";
168        assert!(reg.scope_index("a.py", src).crumbs(0).is_empty());
169    }
170
171    #[test]
172    fn every_tagged_language_spans_a_whole_function() {
173        let reg = LanguageRegistry::build();
174        let samples: &[(&str, &str, &str, (usize, usize))] = &[
175            ("a.rs", "fn f() {\n    let x = 1;\n    x;\n}\n", "f", (0, 3)),
176            ("a.py", "def f():\n    x = 1\n    return x\n", "f", (0, 2)),
177            (
178                "a.js",
179                "function f() {\n  const x = 1;\n  return x;\n}\n",
180                "f",
181                (0, 3),
182            ),
183            (
184                "a.ts",
185                "function f(): number {\n  const x = 1;\n  return x;\n}\n",
186                "f",
187                (0, 3),
188            ),
189            (
190                "a.tsx",
191                "function f() {\n  const x = 1;\n  return x;\n}\n",
192                "f",
193                (0, 3),
194            ),
195            (
196                "a.go",
197                "package main\n\nfunc f() int {\n\tx := 1\n\treturn x\n}\n",
198                "f",
199                (2, 5),
200            ),
201            (
202                "a.c",
203                "static int f(int a,\n             int b) {\n    int x = a;\n    return x + b;\n}\n",
204                "f",
205                (0, 4),
206            ),
207            (
208                "a.cpp",
209                "static bool f(const int* a) {\n    int x = 1;\n    return x;\n}\n",
210                "f",
211                (0, 3),
212            ),
213            (
214                "A.java",
215                "class A {\n    int f() {\n        int x = 1;\n        return x;\n    }\n}\n",
216                "f",
217                (1, 4),
218            ),
219            (
220                "a.cs",
221                "class A {\n    int F() {\n        int x = 1;\n        return x;\n    }\n}\n",
222                "F",
223                (1, 4),
224            ),
225            ("a.rb", "def f\n  x = 1\n  x\nend\n", "f", (0, 3)),
226            (
227                "a.php",
228                "<?php\nfunction f() {\n    $x = 1;\n    return $x;\n}\n",
229                "f",
230                (1, 4),
231            ),
232            (
233                "a.lua",
234                "function f()\n  local x = 1\n  return x\nend\n",
235                "f",
236                (0, 3),
237            ),
238            (
239                "a.swift",
240                "func f() -> Int {\n    let x = 1\n    return x\n}\n",
241                "f",
242                (0, 3),
243            ),
244            (
245                "a.ex",
246                "defmodule A do\n  def f do\n    x = 1\n    x\n  end\nend\n",
247                "f",
248                (1, 4),
249            ),
250            (
251                "a.dart",
252                "int f() {\n  var x = 1;\n  return x;\n}\n",
253                "f",
254                (0, 3),
255            ),
256        ];
257        let wrong: Vec<String> = samples
258            .iter()
259            .filter_map(|(path, src, name, want)| {
260                let got = reg.scope_index(path, src).def_span(name);
261                (got != Some(*want)).then(|| format!("{path}: {got:?}, want {want:?}"))
262            })
263            .collect();
264        assert!(wrong.is_empty(), "{wrong:#?}");
265    }
266
267    #[test]
268    fn a_cpp_method_defined_outside_its_class_answers_to_its_qualified_name() {
269        let reg = LanguageRegistry::build();
270        let src = "void Cache::Link(int a) {\n    int x = a;\n    use(x);\n}\n";
271        let index = reg.scope_index("a.cpp", src);
272        assert_eq!(index.def_span("Cache::Link"), Some((0, 3)));
273        assert_eq!(index.def_span("Link"), Some((0, 3)));
274    }
275
276    #[test]
277    fn a_c_prototype_keeps_its_own_line() {
278        let reg = LanguageRegistry::build();
279        let src = "int f(int a);\n\nint g(void) {\n    return f(1);\n}\n";
280        let index = reg.scope_index("a.c", src);
281        assert_eq!(index.def_span("g"), Some((2, 4)));
282        assert_eq!(index.def_span("f"), Some((0, 0)));
283    }
284}