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::AnyNameRef::can_cast(node.kind()) || ast::ConfigValueName::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
506    #[test]
507    fn access_method_ref() {
508        assert_snapshot!(find_refs("
509create access method heap2$0 type table handler heap_tableam_handler;
510drop access method heap2;
511"), @"
512          ╭▸ 
513        2 │ create access method heap2 type table handler heap_tableam_handler;
514          │                      ┬───┬
515          │                      │   │
516          │                      │   0. query
517          │                      1. reference
518        3 │ drop access method heap2;
519          ╰╴                   ───── 2. reference
520        ");
521    }
522
523    #[test]
524    fn channel_ref() {
525        assert_snapshot!(find_refs("
526listen updates$0;
527notify updates;
528"), @"
529          ╭▸ 
530        2 │ listen updates;
531          │        ┬─────┬
532          │        │     │
533          │        │     0. query
534          │        1. reference
535        3 │ notify updates;
536          ╰╴       ─────── 2. reference
537        ");
538    }
539
540    #[test]
541    fn column_name_ref() {
542        assert_snapshot!(find_refs("
543create table t(id$0 int);
544create table u(t_id int references t(id));
545"), @"
546          ╭▸ 
547        2 │ create table t(id int);
548          │                ┬┬
549          │                ││
550          │                │0. query
551          │                1. reference
552        3 │ create table u(t_id int references t(id));
553          ╰╴                                     ── 2. reference
554        ");
555    }
556
557    #[test]
558    fn cursor_ref() {
559        assert_snapshot!(find_refs("
560declare c$0 scroll cursor for select * from t;
561fetch forward 5 from c;
562"), @"
563          ╭▸ 
564        2 │ declare c scroll cursor for select * from t;
565          │         ┬
566          │         │
567          │         0. query
568          │         1. reference
569        3 │ fetch forward 5 from c;
570          ╰╴                     ─ 2. reference
571        ");
572    }
573
574    #[test]
575    fn database_ref() {
576        assert_snapshot!(find_refs("
577create database mydb$0;
578drop database mydb;
579"), @"
580          ╭▸ 
581        2 │ create database mydb;
582          │                 ┬──┬
583          │                 │  │
584          │                 │  0. query
585          │                 1. reference
586        3 │ drop database mydb;
587          ╰╴              ──── 2. reference
588        ");
589    }
590
591    #[test]
592    fn event_trigger_ref() {
593        assert_snapshot!(find_refs("
594create event trigger et$0 on ddl_command_start execute function f();
595drop event trigger et;
596"), @"
597          ╭▸ 
598        2 │ create event trigger et on ddl_command_start execute function f();
599          │                      ┬┬
600          │                      ││
601          │                      │0. query
602          │                      1. reference
603        3 │ drop event trigger et;
604          ╰╴                   ── 2. reference
605        ");
606    }
607
608    #[test]
609    fn extension_ref() {
610        assert_snapshot!(find_refs("
611create extension myext$0;
612drop extension myext;
613"), @"
614          ╭▸ 
615        2 │ create extension myext;
616          │                  ┬───┬
617          │                  │   │
618          │                  │   0. query
619          │                  1. reference
620        3 │ drop extension myext;
621          ╰╴               ───── 2. reference
622        ");
623    }
624
625    #[test]
626    fn foreign_data_wrapper_ref() {
627        assert_snapshot!(find_refs("
628create foreign data wrapper fdw$0;
629create server srv foreign data wrapper fdw;
630"), @"
631          ╭▸ 
632        2 │ create foreign data wrapper fdw;
633          │                             ┬─┬
634          │                             │ │
635          │                             │ 0. query
636          │                             1. reference
637        3 │ create server srv foreign data wrapper fdw;
638          ╰╴                                       ─── 2. reference
639        ");
640    }
641
642    #[test]
643    fn json_path_name_ref() {
644        assert_snapshot!(find_refs("
645select * from json_table(
646  '{}'::jsonb, '$' as root$0
647  columns (value text path '$')
648  plan (root)
649);
650"), @"
651          ╭▸ 
652        3 │   '{}'::jsonb, '$' as root
653          │                       ┬──┬
654          │                       │  │
655          │                       │  0. query
656          │                       1. reference
657        4 │   columns (value text path '$')
658        5 │   plan (root)
659          ╰╴        ──── 2. reference
660        ");
661    }
662
663    #[test]
664    fn language_ref() {
665        assert_snapshot!(find_refs("
666create language mylang$0;
667create function f() returns int language mylang as $$x$$;
668"), @"
669          ╭▸ 
670        2 │ create language mylang;
671          │                 ┬────┬
672          │                 │    │
673          │                 │    0. query
674          │                 1. reference
675        3 │ create function f() returns int language mylang as $$x$$;
676          ╰╴                                         ────── 2. reference
677        ");
678    }
679
680    #[test]
681    fn policy_ref() {
682        assert_snapshot!(find_refs("
683create table t(c int);
684create policy p$0 on t;
685drop policy p on t;
686"), @"
687          ╭▸ 
688        3 │ create policy p on t;
689          │               ┬
690          │               │
691          │               0. query
692          │               1. reference
693        4 │ drop policy p on t;
694          ╰╴            ─ 2. reference
695        ");
696    }
697
698    #[test]
699    fn prepared_statement_ref() {
700        assert_snapshot!(find_refs("
701prepare stmt$0 as select 1;
702execute stmt;
703"), @"
704          ╭▸ 
705        2 │ prepare stmt as select 1;
706          │         ┬──┬
707          │         │  │
708          │         │  0. query
709          │         1. reference
710        3 │ execute stmt;
711          ╰╴        ──── 2. reference
712        ");
713    }
714
715    #[test]
716    fn publication_ref() {
717        assert_snapshot!(find_refs("
718create table t(id int);
719create publication pub$0 for table t;
720alter publication pub add table t;
721"), @"
722          ╭▸ 
723        3 │ create publication pub for table t;
724          │                    ┬─┬
725          │                    │ │
726          │                    │ 0. query
727          │                    1. reference
728        4 │ alter publication pub add table t;
729          ╰╴                  ─── 2. reference
730        ");
731    }
732
733    #[test]
734    fn role_ref() {
735        assert_snapshot!(find_refs("
736create role reader$0;
737drop role reader;
738"), @"
739          ╭▸ 
740        2 │ create role reader;
741          │             ┬────┬
742          │             │    │
743          │             │    0. query
744          │             1. reference
745        3 │ drop role reader;
746          ╰╴          ────── 2. reference
747        ");
748    }
749
750    #[test]
751    fn rule_ref() {
752        assert_snapshot!(find_refs("
753create table t(a int);
754create rule r$0 as on select to t do instead nothing;
755drop rule r on t;
756"), @"
757          ╭▸ 
758        3 │ create rule r as on select to t do instead nothing;
759          │             ┬
760          │             │
761          │             0. query
762          │             1. reference
763        4 │ drop rule r on t;
764          ╰╴          ─ 2. reference
765        ");
766    }
767
768    #[test]
769    fn savepoint_ref() {
770        assert_snapshot!(find_refs("
771begin;
772savepoint sp$0;
773release savepoint sp;
774"), @"
775          ╭▸ 
776        3 │ savepoint sp;
777          │           ┬┬
778          │           ││
779          │           │0. query
780          │           1. reference
781        4 │ release savepoint sp;
782          ╰╴                  ── 2. reference
783        ");
784    }
785
786    #[test]
787    fn schema_ref() {
788        assert_snapshot!(find_refs("
789create schema app$0;
790drop schema app;
791"), @"
792          ╭▸ 
793        2 │ create schema app;
794          │               ┬─┬
795          │               │ │
796          │               │ 0. query
797          │               1. reference
798        3 │ drop schema app;
799          ╰╴            ─── 2. reference
800        ");
801    }
802
803    #[test]
804    fn server_ref() {
805        assert_snapshot!(find_refs("
806create server myserver$0 foreign data wrapper fdw;
807drop server myserver;
808"), @"
809          ╭▸ 
810        2 │ create server myserver foreign data wrapper fdw;
811          │               ┬──────┬
812          │               │      │
813          │               │      0. query
814          │               1. reference
815        3 │ drop server myserver;
816          ╰╴            ──────── 2. reference
817        ");
818    }
819
820    #[test]
821    fn subscription_ref() {
822        assert_snapshot!(find_refs("
823create subscription sub$0 connection $$host=localhost$$ publication pub;
824alter subscription sub refresh publication;
825"), @"
826          ╭▸ 
827        2 │ create subscription sub connection $$host=localhost$$ publication pub;
828          │                     ┬─┬
829          │                     │ │
830          │                     │ 0. query
831          │                     1. reference
832        3 │ alter subscription sub refresh publication;
833          ╰╴                   ─── 2. reference
834        ");
835    }
836
837    #[test]
838    fn tablespace_ref() {
839        assert_snapshot!(find_refs("
840create tablespace ts$0 location '/tmp/ts';
841drop tablespace ts;
842"), @"
843          ╭▸ 
844        2 │ create tablespace ts location '/tmp/ts';
845          │                   ┬┬
846          │                   ││
847          │                   │0. query
848          │                   1. reference
849        3 │ drop tablespace ts;
850          ╰╴                ── 2. reference
851        ");
852    }
853
854    #[test]
855    fn trigger_ref() {
856        assert_snapshot!(find_refs("
857create trigger tr$0 before insert on t for each row execute function f();
858drop trigger tr on t;
859"), @"
860          ╭▸ 
861        2 │ create trigger tr before insert on t for each row execute function f();
862          │                ┬┬
863          │                ││
864          │                │0. query
865          │                1. reference
866        3 │ drop trigger tr on t;
867          ╰╴             ── 2. reference
868        ");
869    }
870
871    #[test]
872    fn vertex_table_ref() {
873        assert_snapshot!(find_refs("
874create table v1(id int);
875create table e1(source_id int references v1);
876create property graph g
877  vertex tables (v1 as source_vertex)
878  edge tables (e1 source source_vertex$0 destination source_vertex);
879"), @"
880          ╭▸ 
881        5 │   vertex tables (v1 as source_vertex)
882          │                        ───────────── 1. reference
883        6 │   edge tables (e1 source source_vertex destination source_vertex);
884          │                          ┬───────────┬             ───────────── 3. reference
885          │                          │           │
886          │                          │           0. query
887          ╰╴                         2. reference
888        ");
889    }
890
891    #[test]
892    fn window_ref() {
893        assert_snapshot!(find_refs("
894select row_number() over w
895window w$0 as (partition by 1);
896"), @"
897          ╭▸ 
898        2 │ select row_number() over w
899          │                          ─ 1. reference
900        3 │ window w as (partition by 1);
901          │        ┬
902          │        │
903          │        0. query
904          ╰╴       2. reference
905        ");
906    }
907}