1use 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#[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 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 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 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 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 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
98pub(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
141fn 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}