Skip to main content

squawk_ide/
find_references.rs

1use 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(&current_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(&current_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(&current_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        // we should see refs for `drop table` and `table`
338        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
421422            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
432433        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
461462         7 │ create function pg_catalog.bit(bit, integer, boolean) returns bit
463           │                                ─── 3. reference               ─── 4. reference
464465        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}