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(¤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
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}