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