squawk_ide/
find_references.rs1use crate::db::parse;
2use crate::file::InFile;
3use crate::goto_definition;
4use crate::location::Location;
5use rowan::TextSize;
6use salsa::Database as Db;
7use squawk_syntax::{
8 SyntaxNode,
9 ast::{self, AstNode},
10};
11
12fn is_reference_node(node: &SyntaxNode) -> bool {
13 if ast::NameRef::can_cast(node.kind()) {
14 return true;
15 }
16
17 if let Some(literal) = ast::Literal::cast(node.clone())
18 && matches!(literal.kind(), Some(ast::LitKind::PositionalParam(_)))
19 {
20 return true;
21 }
22
23 if let Some(ty) = ast::Type::cast(node.clone()) {
24 return match ty {
25 ast::Type::BitType(_)
26 | ast::Type::CharType(_)
27 | ast::Type::DoubleType(_)
28 | ast::Type::IntervalType(_)
29 | ast::Type::TimeType(_) => true,
30 ast::Type::ArrayType(_)
31 | ast::Type::ExprType(_)
32 | ast::Type::PathType(_)
33 | ast::Type::PercentType(_) => false,
34 };
35 }
36
37 false
38}
39
40pub fn find_references(db: &dyn Db, position: InFile<TextSize>) -> Vec<Location> {
41 let file = position.file_id;
42 let targets = goto_definition::goto_definition(db, position);
43 let Some(first) = targets.first() else {
44 return vec![];
45 };
46
47 let mut refs = targets.to_vec();
48
49 for node in parse(db, file)
50 .tree()
51 .syntax()
52 .descendants()
53 .filter(is_reference_node)
54 {
55 let range = node.text_range();
56 let matches = goto_definition::goto_definition(db, InFile::new(file, range.start()))
57 .into_iter()
58 .any(|location| targets.contains(&location));
59 if matches {
60 refs.push(Location {
61 file,
62 range,
63 kind: first.kind,
64 });
65 }
66 }
67 refs.sort_by_key(|loc| (loc.file != file, loc.range.start()));
68 refs
69}
70
71#[cfg(test)]
72mod test {
73 use crate::builtins::builtins_file;
74 use crate::db::File;
75
76 use crate::find_references::find_references;
77 use crate::test_utils::Fixture;
78 use annotate_snippets::{AnnotationKind, Level, Renderer, Snippet, renderer::DecorStyle};
79 use insta::assert_snapshot;
80 use rowan::TextRange;
81 use rustc_hash::FxHashMap;
82
83 #[must_use]
84 #[track_caller]
85 fn find_refs(sql: &str) -> String {
86 let fixture = Fixture::new(sql);
87 let marker = fixture.marker();
88 let offset = marker.offset_before();
89 let query_span = marker.range();
90 let db = fixture.db();
91 let current_file = offset.file_id;
92
93 let references = find_references(db, offset);
94
95 let mut file_paths = FxHashMap::default();
96 file_paths.insert(current_file, "current.sql");
97 file_paths.insert(builtins_file(db), "builtins.sql");
98
99 let mut refs_by_file: FxHashMap<File, Vec<(usize, TextRange)>> = FxHashMap::default();
100 for (i, location) in references.iter().enumerate() {
101 refs_by_file
102 .entry(location.file)
103 .or_default()
104 .push((i + 1, location.range));
105 }
106
107 let multi_file = refs_by_file.len() > 1 || !refs_by_file.contains_key(¤t_file);
108
109 let mut snippet = Snippet::source(current_file.content(db).as_ref()).fold(true);
110 if multi_file {
111 snippet = snippet.path(*file_paths.get(¤t_file).unwrap());
112 }
113 snippet = snippet.annotation(AnnotationKind::Context.span(query_span).label("0. query"));
114 if let Some(current_refs) = refs_by_file.remove(¤t_file) {
115 snippet = annotate_refs(snippet, current_refs);
116 }
117
118 let mut groups = vec![Level::INFO.primary_title("references").element(snippet)];
119
120 for (ref_file, refs) in refs_by_file {
121 let path = file_paths.get(&ref_file).unwrap();
122 let other_snippet = Snippet::source(ref_file.content(db).as_ref())
123 .path(*path)
124 .fold(true);
125 let other_snippet = annotate_refs(other_snippet, refs);
126 groups.push(
127 Level::INFO
128 .primary_title("references")
129 .element(other_snippet),
130 );
131 }
132
133 let renderer = Renderer::plain().decor_style(DecorStyle::Unicode);
134 renderer
135 .render(&groups)
136 .to_string()
137 .replace("info: references", "")
138 }
139
140 fn annotate_refs<'a>(
141 mut snippet: Snippet<'a, annotate_snippets::Annotation<'a>>,
142 refs: Vec<(usize, TextRange)>,
143 ) -> Snippet<'a, annotate_snippets::Annotation<'a>> {
144 for (label_index, range) in refs {
145 snippet = snippet.annotation(
146 AnnotationKind::Context
147 .span(range.into())
148 .label(format!("{label_index}. reference")),
149 );
150 }
151 snippet
152 }
153
154 #[test]
155 fn simple_table_reference() {
156 assert_snapshot!(find_refs("
157create table t();
158drop table t$0;
159"), @"
160 ╭▸
161 2 │ create table t();
162 │ ─ 1. reference
163 3 │ drop table t;
164 │ ┬
165 │ │
166 │ 0. query
167 ╰╴ 2. reference
168 ");
169 }
170
171 #[test]
172 fn multiple_references() {
173 assert_snapshot!(find_refs("
174create table users();
175drop table users$0;
176table users;
177"), @r"
178 ╭▸
179 2 │ create table users();
180 │ ───── 1. reference
181 3 │ drop table users;
182 │ ┬───┬
183 │ │ │
184 │ │ 0. query
185 │ 2. reference
186 4 │ table users;
187 ╰╴ ───── 3. reference
188 ");
189 }
190
191 #[test]
192 fn join_using_column() {
193 assert_snapshot!(find_refs("
194create table t(id int);
195create table u(id int);
196select * from t join u using (id$0);
197"), @r"
198 ╭▸
199 2 │ create table t(id int);
200 │ ── 1. reference
201 3 │ create table u(id int);
202 │ ── 2. reference
203 4 │ select * from t join u using (id);
204 │ ┬┬
205 │ ││
206 │ │0. query
207 ╰╴ 3. reference
208 ");
209 }
210
211 #[test]
212 fn find_from_definition() {
213 assert_snapshot!(find_refs("
214create table t$0();
215drop table t;
216"), @r"
217 ╭▸
218 2 │ create table t();
219 │ ┬
220 │ │
221 │ 0. query
222 │ 1. reference
223 3 │ drop table t;
224 ╰╴ ─ 2. reference
225 ");
226 }
227
228 #[test]
229 fn with_schema_qualified() {
230 assert_snapshot!(find_refs("
231create table public.users();
232drop table public.users$0;
233table users;
234"), @r"
235 ╭▸
236 2 │ create table public.users();
237 │ ───── 1. reference
238 3 │ drop table public.users;
239 │ ┬───┬
240 │ │ │
241 │ │ 0. query
242 │ 2. reference
243 4 │ table users;
244 ╰╴ ───── 3. reference
245 ");
246 }
247
248 #[test]
249 fn temp_table_shadows_public() {
250 assert_snapshot!(find_refs("
251create table t();
252create temp table t$0();
253drop table t;
254"), @"
255 ╭▸
256 3 │ create temp table t();
257 │ ┬
258 │ │
259 │ 0. query
260 │ 1. reference
261 4 │ drop table t;
262 ╰╴ ─ 2. reference
263 ");
264 }
265
266 #[test]
267 fn different_schema_no_match() {
268 assert_snapshot!(find_refs("
269create table foo.t();
270create table bar.t$0();
271"), @r"
272 ╭▸
273 3 │ create table bar.t();
274 │ ┬
275 │ │
276 │ 0. query
277 ╰╴ 1. reference
278 ");
279 }
280
281 #[test]
282 fn with_search_path() {
283 assert_snapshot!(find_refs("
284set search_path to myschema;
285create table myschema.users$0();
286drop table users;
287"), @r"
288 ╭▸
289 3 │ create table myschema.users();
290 │ ┬───┬
291 │ │ │
292 │ │ 0. query
293 │ 1. reference
294 4 │ drop table users;
295 ╰╴ ───── 2. reference
296 ");
297 }
298
299 #[test]
300 fn temp_table_with_pg_temp_schema() {
301 assert_snapshot!(find_refs("
302create temp table t();
303drop table pg_temp.t$0;
304"), @r"
305 ╭▸
306 2 │ create temp table t();
307 │ ─ 1. reference
308 3 │ drop table pg_temp.t;
309 │ ┬
310 │ │
311 │ 0. query
312 ╰╴ 2. reference
313 ");
314 }
315
316 #[test]
317 fn case_insensitive() {
318 assert_snapshot!(find_refs("
319create table Users();
320drop table USERS$0;
321table users;
322"), @r"
323 ╭▸
324 2 │ create table Users();
325 │ ───── 1. reference
326 3 │ drop table USERS;
327 │ ┬───┬
328 │ │ │
329 │ │ 0. query
330 │ 2. reference
331 4 │ table users;
332 ╰╴ ───── 3. reference
333 ");
334 }
335 #[test]
336 fn case_insensitive_part_2() {
337 assert_snapshot!(find_refs(r#"
339create table actors();
340create table "Actors"();
341drop table ACTORS$0;
342table actors;
343"#), @r#"
344 ╭▸
345 2 │ create table actors();
346 │ ────── 1. reference
347 3 │ create table "Actors"();
348 4 │ drop table ACTORS;
349 │ ┬────┬
350 │ │ │
351 │ │ 0. query
352 │ 2. reference
353 5 │ table actors;
354 ╰╴ ────── 3. reference
355 "#);
356 }
357
358 #[test]
359 fn case_insensitive_with_schema() {
360 assert_snapshot!(find_refs("
361create table Public.Users();
362drop table PUBLIC.USERS$0;
363table public.users;
364"), @r"
365 ╭▸
366 2 │ create table Public.Users();
367 │ ───── 1. reference
368 3 │ drop table PUBLIC.USERS;
369 │ ┬───┬
370 │ │ │
371 │ │ 0. query
372 │ 2. reference
373 4 │ table public.users;
374 ╰╴ ───── 3. reference
375 ");
376 }
377
378 #[test]
379 fn no_partial_match() {
380 assert_snapshot!(find_refs("
381create table t$0();
382create table temp_t();
383"), @r"
384 ╭▸
385 2 │ create table t();
386 │ ┬
387 │ │
388 │ 0. query
389 ╰╴ 1. reference
390 ");
391 }
392
393 #[test]
394 fn identifier_boundaries() {
395 assert_snapshot!(find_refs("
396create table foo$0();
397drop table foo;
398drop table foo1;
399drop table barfoo;
400drop table foo_bar;
401"), @r"
402 ╭▸
403 2 │ create table foo();
404 │ ┬─┬
405 │ │ │
406 │ │ 0. query
407 │ 1. reference
408 3 │ drop table foo;
409 ╰╴ ─── 2. reference
410 ");
411 }
412
413 #[test]
414 fn builtin_function_references() {
415 assert_snapshot!(find_refs("
416-- include-builtins
417select now$0();
418select now();
419"), @"
420 ╭▸ current.sql:3:8
421 │
422 3 │ select now();
423 │ ┬─┬
424 │ │ │
425 │ │ 0. query
426 │ 1. reference
427 4 │ select now();
428 │ ─── 2. reference
429 ╰╴
430
431 ╭▸ builtins.sql:11089:28
432 │
433 11089 │ create function pg_catalog.now() returns timestamp with time zone
434 ╰╴ ─── 3. reference
435 ");
436 }
437
438 #[test]
439 fn bit() {
440 assert_snapshot!(find_refs("
441create type pg_catalog.bit$0;
442
443create function pg_catalog.bit(bigint, integer) returns bit
444 language internal;
445
446create function pg_catalog.bit(bit, integer, boolean) returns bit
447 language internal;
448
449create function pg_catalog.bit(integer, integer) returns bit
450 language internal;
451"), @"
452 ╭▸
453 2 │ create type pg_catalog.bit;
454 │ ┬─┬
455 │ │ │
456 │ │ 0. query
457 │ 1. reference
458 3 │
459 4 │ create function pg_catalog.bit(bigint, integer) returns bit
460 │ ─── 2. reference
461 ‡
462 7 │ create function pg_catalog.bit(bit, integer, boolean) returns bit
463 │ ─── 3. reference ─── 4. reference
464 ‡
465 10 │ create function pg_catalog.bit(integer, integer) returns bit
466 ╰╴ ─── 5. reference
467 ");
468 }
469
470 #[test]
471 fn positional_param_and_named_param() {
472 assert_snapshot!(find_refs("
473create function f(x$0 int) returns int language sql return x + $1;
474"), @"
475 ╭▸
476 2 │ create function f(x int) returns int language sql return x + $1;
477 │ ┬ ┬ ── 3. reference
478 │ │ │
479 │ 0. query 2. reference
480 ╰╴ 1. reference
481 ");
482 }
483
484 #[test]
485 fn char() {
486 assert_snapshot!(find_refs("
487create type pg_catalog.bpchar$0;
488
489select '1'::char;
490select '1'::bpchar;
491"), @"
492 ╭▸
493 2 │ create type pg_catalog.bpchar;
494 │ ┬────┬
495 │ │ │
496 │ │ 0. query
497 │ 1. reference
498 3 │
499 4 │ select '1'::char;
500 │ ──── 2. reference
501 5 │ select '1'::bpchar;
502 ╰╴ ────── 3. reference
503 ");
504 }
505}