1use sqry_core::graph::unified::build::shape::{CfBucket, ShapeMapping};
10use sqry_core::graph::unified::storage::shape::SignatureShape;
11use sqry_core::graph::{
12 GraphBuilder, GraphBuilderError, GraphResult, Language, Position, Span,
13 unified::{GraphBuildHelper, StagingGraph},
14};
15use std::path::Path;
16use std::sync::OnceLock;
17use streaming_iterator::StreamingIterator;
18use tree_sitter::{Node, Query, QueryCursor, Tree};
19
20#[derive(Debug, Clone)]
21struct SqlCallable {
22 node_id: sqry_core::graph::unified::NodeId,
23 start_byte: usize,
24 end_byte: usize,
25}
26
27#[derive(Debug, Clone)]
28struct SqlDatabaseObject {
29 node_id: sqry_core::graph::unified::NodeId,
30}
31
32#[derive(Debug, Clone)]
33enum SqlTableOpKind {
34 Read,
35 Write(sqry_core::graph::unified::TableWriteOp),
36}
37
38#[derive(Debug, Clone)]
39struct SqlTableOp {
40 op_span_bytes: (usize, usize),
41 kind: SqlTableOpKind,
42 table_name: String,
43 schema: Option<String>,
44 table_node_id: sqry_core::graph::unified::NodeId,
45 span: Span,
46}
47
48const FILE_MODULE_NAME: &str = "<file_module>";
53
54#[derive(Debug, Default, Clone, Copy)]
63pub struct SqlGraphBuilder;
64
65impl SqlGraphBuilder {
66 #[must_use]
68 pub fn new() -> Self {
69 Self
70 }
71}
72
73impl GraphBuilder for SqlGraphBuilder {
74 fn build_graph(
75 &self,
76 tree: &Tree,
77 content: &[u8],
78 file: &Path,
79 staging: &mut StagingGraph,
80 ) -> GraphResult<()> {
81 let mut helper = GraphBuildHelper::new(staging, file, Language::Sql);
83
84 let language = tree_sitter_sequel::LANGUAGE.into();
86 let queries = SqlQueries::new(&language)?;
87
88 let mut callables = extract_procedures(tree, content, &queries.procedures, &mut helper);
90
91 callables.extend(extract_triggers(
93 tree,
94 content,
95 &queries.triggers,
96 &mut helper,
97 ));
98
99 let table_reads = extract_table_reads(tree, content, &queries.table_reads, &mut helper);
101
102 let table_writes = extract_table_writes(tree, content, &queries.table_writes, &mut helper);
104
105 let function_calls = extract_function_calls(tree, content, &queries.function_calls);
107
108 let table_definitions =
110 extract_table_definitions(tree, content, &queries.table_definitions, &mut helper);
111
112 let view_definitions =
114 extract_view_definitions(tree, content, &queries.view_definitions, &mut helper);
115
116 for op in table_reads.into_iter().chain(table_writes) {
118 let Some(caller) = find_enclosing_callable(&callables, op.op_span_bytes) else {
119 continue;
120 };
121
122 match op.kind {
123 SqlTableOpKind::Read => helper.add_table_read_edge_with_span(
124 caller.node_id,
125 op.table_node_id,
126 &op.table_name,
127 op.schema.as_deref(),
128 vec![op.span],
129 ),
130 SqlTableOpKind::Write(operation) => helper.add_table_write_edge_with_span(
131 caller.node_id,
132 op.table_node_id,
133 &op.table_name,
134 op.schema.as_deref(),
135 operation,
136 vec![op.span],
137 ),
138 }
139 }
140
141 for call in function_calls {
143 if let Some(caller) = find_enclosing_callable(&callables, call.span_bytes) {
145 let callee_id =
147 helper.add_function(&call.callee_name, Some(call.span), false, false);
148 helper.add_call_edge_full_with_span(
149 caller.node_id,
150 callee_id,
151 255,
152 false,
153 vec![call.span],
154 );
155 }
156 }
159
160 extract_trigger_execute_function_calls(
163 tree,
164 content,
165 &queries.trigger_execute_function,
166 &callables,
167 &mut helper,
168 );
169
170 emit_exports(
172 &mut helper,
173 &callables,
174 &table_definitions,
175 &view_definitions,
176 );
177
178 Ok(())
179 }
180
181 fn language(&self) -> Language {
182 Language::Sql
183 }
184
185 fn shape_mapping(&self) -> Option<&dyn ShapeMapping> {
186 Some(sql_shape_mapping())
187 }
188}
189
190pub struct SqlShapeMapping {
200 cf_by_kind_id: Vec<Option<CfBucket>>,
201}
202
203impl SqlShapeMapping {
204 fn build() -> Self {
205 let lang: tree_sitter::Language = tree_sitter_sequel::LANGUAGE.into();
206 let count = lang.node_kind_count();
207 let mut cf_by_kind_id = vec![None; count];
208 for (id, slot) in cf_by_kind_id.iter_mut().enumerate() {
209 let Ok(kind_id) = u16::try_from(id) else {
210 break;
211 };
212 if !lang.node_kind_is_named(kind_id) {
213 continue;
214 }
215 if let Some(name) = lang.node_kind_for_id(kind_id) {
216 *slot = cf_bucket_for_sql_kind(name);
217 }
218 }
219 Self { cf_by_kind_id }
220 }
221}
222
223impl ShapeMapping for SqlShapeMapping {
224 fn cf_bucket(&self, ts_node_kind_id: u16) -> Option<CfBucket> {
225 self.cf_by_kind_id
226 .get(ts_node_kind_id as usize)
227 .copied()
228 .flatten()
229 }
230
231 fn signature_shape(&self, fn_node: Node, _src: &[u8]) -> SignatureShape {
232 let mut shape = SignatureShape::default();
233 let mut cursor = fn_node.walk();
236 for child in fn_node.named_children(&mut cursor) {
237 if child.kind() == "function_arguments" {
238 let mut arg_cursor = child.walk();
239 for arg in child.named_children(&mut arg_cursor) {
240 if arg.kind() == "function_argument" {
241 shape.arity_positional = shape.arity_positional.saturating_add(1);
242 let mut def_cursor = arg.walk();
245 for piece in arg.children(&mut def_cursor) {
246 if piece.kind() == "keyword_default" {
247 shape.has_defaults = true;
248 }
249 }
250 }
251 }
252 }
253 }
254 shape
255 }
256}
257
258fn cf_bucket_for_sql_kind(name: &str) -> Option<CfBucket> {
261 let bucket = match name {
262 "case" => CfBucket::Match,
264 "when_clause" => CfBucket::Branch,
265 "assignment" => CfBucket::Assign,
267 "invocation" => CfBucket::Call,
269 _ => return None,
270 };
271 Some(bucket)
272}
273
274#[must_use]
276pub fn sql_shape_mapping() -> &'static SqlShapeMapping {
277 static MAPPING: OnceLock<SqlShapeMapping> = OnceLock::new();
278 MAPPING.get_or_init(SqlShapeMapping::build)
279}
280
281struct SqlQueries {
283 procedures: Query,
284 triggers: Query,
285 trigger_execute_function: Query,
286 table_reads: Query,
287 table_writes: Query,
288 function_calls: Query,
289 table_definitions: Query,
290 view_definitions: Query,
291}
292
293impl SqlQueries {
294 #[allow(clippy::too_many_lines)]
296 fn new(language: &tree_sitter::Language) -> GraphResult<Self> {
297 let procedures = Query::new(
299 language,
300 r"
301 (create_function
302 (object_reference
303 name: (identifier) @func.name)) @func
304 ",
305 )
306 .map_err(|e| GraphBuilderError::ParseError {
307 span: Span::default(),
308 reason: format!("Failed to compile procedure query: {e}"),
309 })?;
310
311 let triggers = Query::new(
314 language,
315 r"
316 (create_trigger
317 (object_reference
318 name: (identifier) @trigger.name)
319 (keyword_on)
320 (object_reference
321 name: (identifier) @trigger.table)) @trigger
322 ",
323 )
324 .map_err(|e| GraphBuilderError::ParseError {
325 span: Span::default(),
326 reason: format!("Failed to compile trigger query: {e}"),
327 })?;
328
329 let trigger_execute_function = Query::new(
332 language,
333 r"
334 (create_trigger
335 (object_reference
336 name: (identifier) @trigger.name)
337 (keyword_execute)
338 (keyword_function)
339 (object_reference
340 name: (identifier) @func.name)) @trigger_exec
341 ",
342 )
343 .map_err(|e| GraphBuilderError::ParseError {
344 span: Span::default(),
345 reason: format!("Failed to compile trigger_execute_function query: {e}"),
346 })?;
347
348 let table_reads = Query::new(
351 language,
352 r"
353 (statement
354 (select) @select
355 (from
356 (keyword_from)
357 (relation
358 (object_reference
359 name: (identifier) @table.name))))
360 ",
361 )
362 .map_err(|e| GraphBuilderError::ParseError {
363 span: Span::default(),
364 reason: format!("Failed to compile table_reads query: {e}"),
365 })?;
366
367 let table_writes = Query::new(
372 language,
373 r"
374 [
375 (insert
376 (object_reference
377 name: (identifier) @table.name)) @write
378
379 (update
380 (relation
381 (object_reference
382 name: (identifier) @table.name))) @write
383
384 (statement
385 (delete) @write
386 (from
387 (keyword_from)
388 (object_reference
389 name: (identifier) @table.name)))
390 ]
391 ",
392 )
393 .map_err(|e| GraphBuilderError::ParseError {
394 span: Span::default(),
395 reason: format!("Failed to compile table_writes query: {e}"),
396 })?;
397
398 let function_calls = Query::new(
404 language,
405 r#"
406 [
407 (invocation
408 (object_reference
409 name: (identifier) @call.name)) @call
410
411 (ERROR
412 ":="
413 (_) @call.name
414 "(") @call
415
416 (ERROR) @call.error
417 ]
418 "#,
419 )
420 .map_err(|e| GraphBuilderError::ParseError {
421 span: Span::default(),
422 reason: format!("Failed to compile function_calls query: {e}"),
423 })?;
424
425 let table_definitions = Query::new(
427 language,
428 r"
429 (create_table
430 (object_reference
431 name: (identifier) @table.name)) @table
432 ",
433 )
434 .map_err(|e| GraphBuilderError::ParseError {
435 span: Span::default(),
436 reason: format!("Failed to compile table_definitions query: {e}"),
437 })?;
438
439 let view_definitions = Query::new(
441 language,
442 r"
443 [
444 (create_view
445 (object_reference
446 name: (identifier) @view.name)) @view
447 (create_materialized_view
448 (object_reference
449 name: (identifier) @view.name)) @view
450 ]
451 ",
452 )
453 .map_err(|e| GraphBuilderError::ParseError {
454 span: Span::default(),
455 reason: format!("Failed to compile view_definitions query: {e}"),
456 })?;
457
458 Ok(Self {
459 procedures,
460 triggers,
461 trigger_execute_function,
462 table_reads,
463 table_writes,
464 function_calls,
465 table_definitions,
466 view_definitions,
467 })
468 }
469}
470
471fn extract_procedures(
473 tree: &Tree,
474 content: &[u8],
475 query: &Query,
476 helper: &mut GraphBuildHelper,
477) -> Vec<SqlCallable> {
478 let mut callables = Vec::new();
479 let mut cursor = QueryCursor::new();
480 let capture_names = query.capture_names();
481 let mut matches = cursor.matches(query, tree.root_node(), content);
482
483 while let Some(m) = matches.next() {
484 let mut func_name = None;
485 let mut func_node = None;
486
487 for capture in m.captures {
488 let name = capture_names[capture.index as usize];
489 if name == "func.name"
490 && let Ok(text) = capture.node.utf8_text(content)
491 {
492 func_name = Some(text.to_string());
493 }
494 if name == "func" {
495 func_node = Some(capture.node);
496 }
497 }
498
499 if let (Some(name), Some(node)) = (func_name, func_node) {
500 let span = Span::from_node(&node);
501 let node_id = helper.add_function(&name, Some(span), false, false);
502 helper.mark_definition(node_id);
504 callables.push(SqlCallable {
505 node_id,
506 start_byte: node.start_byte(),
507 end_byte: node.end_byte(),
508 });
509 }
510 }
511
512 callables
513}
514
515fn extract_triggers(
517 tree: &Tree,
518 content: &[u8],
519 query: &Query,
520 helper: &mut GraphBuildHelper,
521) -> Vec<SqlCallable> {
522 let mut callables = Vec::new();
523 let mut cursor = QueryCursor::new();
524 let capture_names = query.capture_names();
525 let mut matches = cursor.matches(query, tree.root_node(), content);
526
527 while let Some(m) = matches.next() {
528 let mut trigger_name = None;
529 let mut table_name = None;
530 let mut trigger_node = None;
531
532 for capture in m.captures {
533 let name = capture_names[capture.index as usize];
534 match name {
535 "trigger.name" => {
536 if let Ok(text) = capture.node.utf8_text(content) {
537 trigger_name = Some(text.to_string());
538 }
539 }
540 "trigger.table" => {
541 if let Ok(text) = capture.node.utf8_text(content) {
542 table_name = Some(text.to_string());
543 }
544 }
545 "trigger" => {
546 trigger_node = Some(capture.node);
547 }
548 _ => {}
549 }
550 }
551
552 if let (Some(trigger), Some(table), Some(node)) = (trigger_name, table_name, trigger_node) {
553 let (schema, table_only) = split_schema_table(&table);
554 let span = Span::from_node(&node);
555
556 let trigger_id = helper.add_function(&trigger, Some(span), false, false);
557 helper.mark_definition(trigger_id);
559 callables.push(SqlCallable {
560 node_id: trigger_id,
561 start_byte: node.start_byte(),
562 end_byte: node.end_byte(),
563 });
564
565 let table_id = helper.add_variable(table_only, Some(span));
566 helper.add_triggered_by_edge_with_span(
567 trigger_id,
568 table_id,
569 &trigger,
570 schema,
571 vec![span],
572 );
573 }
574 }
575
576 callables
577}
578
579fn extract_trigger_execute_function_calls(
584 tree: &Tree,
585 content: &[u8],
586 query: &Query,
587 callables: &[SqlCallable],
588 helper: &mut GraphBuildHelper,
589) {
590 let mut cursor = QueryCursor::new();
591 let capture_names = query.capture_names();
592 let mut matches = cursor.matches(query, tree.root_node(), content);
593
594 while let Some(m) = matches.next() {
595 let mut trigger_name = None;
596 let mut func_name = None;
597 let mut trigger_node = None;
598
599 for capture in m.captures {
600 let name = capture_names[capture.index as usize];
601 match name {
602 "trigger.name" => {
603 if let Ok(text) = capture.node.utf8_text(content) {
604 trigger_name = Some(text.to_string());
605 }
606 }
607 "func.name" => {
608 if let Ok(text) = capture.node.utf8_text(content) {
609 func_name = Some(text.to_string());
610 }
611 }
612 "trigger_exec" => {
613 trigger_node = Some(capture.node);
614 }
615 _ => {}
616 }
617 }
618
619 if let (Some(_trigger), Some(func), Some(node)) = (trigger_name, func_name, trigger_node) {
620 let span = Span::from_node(&node);
621
622 if let Some(trigger_callable) = callables.iter().find(|c| {
624 c.start_byte <= node.start_byte() && node.end_byte() <= c.end_byte
626 }) {
627 let callee_id = helper.add_function(&func, Some(span), false, false);
629 helper.add_call_edge_full_with_span(
630 trigger_callable.node_id,
631 callee_id,
632 255,
633 false,
634 vec![span],
635 );
636 }
637 }
638 }
639}
640
641fn extract_table_reads(
643 tree: &Tree,
644 content: &[u8],
645 query: &Query,
646 helper: &mut GraphBuildHelper,
647) -> Vec<SqlTableOp> {
648 let mut ops = Vec::new();
649 let mut cursor = QueryCursor::new();
650 let capture_names = query.capture_names();
651 let mut matches = cursor.matches(query, tree.root_node(), content);
652
653 while let Some(m) = matches.next() {
654 let mut table_name = None;
655 let mut op_node = None;
656
657 for capture in m.captures {
658 let name = capture_names[capture.index as usize];
659 match name {
660 "table.name" => {
661 if let Ok(text) = capture.node.utf8_text(content) {
662 table_name = Some(text.to_string());
663 }
664 }
665 "select" => op_node = Some(capture.node),
666 _ => {}
667 }
668 }
669
670 if let (Some(table_name), Some(node)) = (table_name, op_node) {
671 let (schema, table_only) = split_schema_table(&table_name);
672 let span = Span::from_node(&node);
673 let table_node_id = helper.add_variable(table_only, Some(span));
674 ops.push(SqlTableOp {
675 op_span_bytes: (node.start_byte(), node.end_byte()),
676 kind: SqlTableOpKind::Read,
677 table_name: table_only.to_string(),
678 schema: schema.map(str::to_string),
679 table_node_id,
680 span,
681 });
682 }
683 }
684
685 ops
686}
687
688fn extract_table_writes(
690 tree: &Tree,
691 content: &[u8],
692 query: &Query,
693 helper: &mut GraphBuildHelper,
694) -> Vec<SqlTableOp> {
695 let mut ops = Vec::new();
696 let mut cursor = QueryCursor::new();
697 let capture_names = query.capture_names();
698 let mut matches = cursor.matches(query, tree.root_node(), content);
699
700 while let Some(m) = matches.next() {
701 let mut table_name = None;
702 let mut write_node = None;
703
704 for capture in m.captures {
705 let name = capture_names[capture.index as usize];
706 match name {
707 "table.name" => {
708 if let Ok(text) = capture.node.utf8_text(content) {
709 table_name = Some(text.to_string());
710 }
711 }
712 "write" => write_node = Some(capture.node),
713 _ => {}
714 }
715 }
716
717 let Some(table_name) = table_name else {
718 continue;
719 };
720 let Some(node) = write_node else {
721 continue;
722 };
723
724 let operation = match node.kind() {
725 "insert" => sqry_core::graph::unified::TableWriteOp::Insert,
726 "delete" => sqry_core::graph::unified::TableWriteOp::Delete,
727 _ => sqry_core::graph::unified::TableWriteOp::Update,
728 };
729
730 let (schema, table_only) = split_schema_table(&table_name);
731 let span = Span::from_node(&node);
732 let table_node_id = helper.add_variable(table_only, Some(span));
733 ops.push(SqlTableOp {
734 op_span_bytes: (node.start_byte(), node.end_byte()),
735 kind: SqlTableOpKind::Write(operation),
736 table_name: table_only.to_string(),
737 schema: schema.map(str::to_string),
738 table_node_id,
739 span,
740 });
741 }
742
743 ops
744}
745
746#[derive(Debug)]
748struct SqlFunctionCall {
749 callee_name: String,
750 span_bytes: (usize, usize),
751 span: Span,
752}
753
754fn extract_function_calls(tree: &Tree, content: &[u8], query: &Query) -> Vec<SqlFunctionCall> {
756 let mut calls = Vec::new();
757 let mut cursor = QueryCursor::new();
758 let capture_names = query.capture_names();
759 let mut matches = cursor.matches(query, tree.root_node(), content);
760
761 while let Some(m) = matches.next() {
762 let mut call_name = None;
763 let mut call_node = None;
764
765 for capture in m.captures {
766 let name = capture_names[capture.index as usize];
767 match name {
768 "call.name" => {
769 if let Ok(text) = capture.node.utf8_text(content) {
770 call_name = Some(normalize_callee_name(text));
771 }
772 }
773 "call" | "call.error" => call_node = Some(capture.node),
774 _ => {}
775 }
776 }
777
778 let Some(node) = call_node else {
779 continue;
780 };
781
782 let span_bytes = (node.start_byte(), node.end_byte());
783 let span = Span::from_node(&node);
784
785 if node.kind() == "ERROR" {
786 if let Ok(text) = node.utf8_text(content) {
787 for name in extract_error_call_names(text) {
788 calls.push(SqlFunctionCall {
789 callee_name: name,
790 span_bytes,
791 span,
792 });
793 }
794 }
795 continue;
796 }
797
798 if let Some(name) = call_name
799 && !name.is_empty()
800 {
801 calls.push(SqlFunctionCall {
802 callee_name: name,
803 span_bytes,
804 span,
805 });
806 }
807 }
808
809 calls
810}
811
812fn normalize_callee_name(name: &str) -> String {
813 name.trim()
814 .rsplit('.')
815 .next()
816 .unwrap_or_default()
817 .trim()
818 .to_string()
819}
820
821fn extract_error_call_names(text: &str) -> Vec<String> {
822 let bytes = text.as_bytes();
823 let mut offset = 0;
824 let mut call_names = Vec::new();
825
826 while offset < bytes.len() {
827 if !is_sql_identifier_start(bytes[offset]) {
828 offset += 1;
829 continue;
830 }
831
832 let start = offset;
833 offset += 1;
834 while offset < bytes.len() && is_sql_identifier_continue(bytes[offset]) {
835 offset += 1;
836 }
837
838 let token = &text[start..offset];
839 let mut lookahead = offset;
840 while lookahead < bytes.len() && bytes[lookahead].is_ascii_whitespace() {
841 lookahead += 1;
842 }
843
844 if lookahead < bytes.len() && bytes[lookahead] == b'(' {
845 let normalized = normalize_callee_name(token);
846 if !normalized.is_empty() && !call_names.iter().any(|name| name == &normalized) {
847 call_names.push(normalized);
848 }
849 }
850 }
851
852 call_names
853}
854
855const fn is_sql_identifier_start(byte: u8) -> bool {
856 byte.is_ascii_alphabetic() || byte == b'_'
857}
858
859const fn is_sql_identifier_continue(byte: u8) -> bool {
860 byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'.')
861}
862
863fn extract_table_definitions(
865 tree: &Tree,
866 content: &[u8],
867 query: &Query,
868 helper: &mut GraphBuildHelper,
869) -> Vec<SqlDatabaseObject> {
870 let mut objects = Vec::new();
871 let mut cursor = QueryCursor::new();
872 let capture_names = query.capture_names();
873 let mut matches = cursor.matches(query, tree.root_node(), content);
874
875 while let Some(m) = matches.next() {
876 let mut table_name = None;
877 let mut table_node = None;
878
879 for capture in m.captures {
880 let name = capture_names[capture.index as usize];
881 match name {
882 "table.name" => {
883 if let Ok(text) = capture.node.utf8_text(content) {
884 table_name = Some(text.to_string());
885 }
886 }
887 "table" => table_node = Some(capture.node),
888 _ => {}
889 }
890 }
891
892 if let (Some(name), Some(node)) = (table_name, table_node) {
893 let (_, table_only) = split_schema_table(&name);
895 let span = Span::from_node(&node);
896 let node_id = helper.add_variable(table_only, Some(span));
897 helper.mark_definition(node_id);
899 objects.push(SqlDatabaseObject { node_id });
900 }
901 }
902
903 objects
904}
905
906fn extract_view_definitions(
908 tree: &Tree,
909 content: &[u8],
910 query: &Query,
911 helper: &mut GraphBuildHelper,
912) -> Vec<SqlDatabaseObject> {
913 let mut objects = Vec::new();
914 let mut cursor = QueryCursor::new();
915 let capture_names = query.capture_names();
916 let mut matches = cursor.matches(query, tree.root_node(), content);
917
918 while let Some(m) = matches.next() {
919 let mut view_name = None;
920 let mut view_node = None;
921
922 for capture in m.captures {
923 let name = capture_names[capture.index as usize];
924 match name {
925 "view.name" => {
926 if let Ok(text) = capture.node.utf8_text(content) {
927 view_name = Some(text.to_string());
928 }
929 }
930 "view" => view_node = Some(capture.node),
931 _ => {}
932 }
933 }
934
935 if let (Some(name), Some(node)) = (view_name, view_node) {
936 let (_, view_only) = split_schema_table(&name);
938 let span = Span::from_node(&node);
939 let node_id = helper.add_variable(view_only, Some(span));
940 helper.mark_definition(node_id);
942 objects.push(SqlDatabaseObject { node_id });
943 }
944 }
945
946 objects
947}
948
949fn find_enclosing_callable(
950 callables: &[SqlCallable],
951 op_span_bytes: (usize, usize),
952) -> Option<&SqlCallable> {
953 let (start_byte, end_byte) = op_span_bytes;
954 callables
955 .iter()
956 .filter(|c| c.start_byte <= start_byte && end_byte <= c.end_byte)
957 .min_by_key(|c| c.end_byte.saturating_sub(c.start_byte))
958}
959
960fn split_schema_table(name: &str) -> (Option<&str>, &str) {
961 let mut parts = name.splitn(2, '.');
962 let first = parts.next().unwrap_or(name).trim();
963 let second = parts.next().map(str::trim);
964 match second {
965 Some(table) if !table.is_empty() => (Some(first), table),
966 _ => (None, first),
967 }
968}
969
970trait SpanExt {
972 fn from_node(node: &tree_sitter::Node) -> Self;
973}
974
975impl SpanExt for Span {
976 fn from_node(node: &tree_sitter::Node) -> Self {
977 Span::new(
978 Position::new(node.start_position().row, node.start_position().column),
979 Position::new(node.end_position().row, node.end_position().column),
980 )
981 }
982}
983
984fn emit_exports(
990 helper: &mut GraphBuildHelper,
991 callables: &[SqlCallable],
992 tables: &[SqlDatabaseObject],
993 views: &[SqlDatabaseObject],
994) {
995 if callables.is_empty() && tables.is_empty() && views.is_empty() {
997 return;
998 }
999
1000 let module_id = helper.add_module(FILE_MODULE_NAME, None);
1002
1003 for callable in callables {
1005 helper.add_export_edge(module_id, callable.node_id);
1006 }
1007
1008 for table in tables {
1010 helper.add_export_edge(module_id, table.node_id);
1011 }
1012
1013 for view in views {
1015 helper.add_export_edge(module_id, view.node_id);
1016 }
1017}
1018
1019#[cfg(test)]
1020mod tests {
1021 use super::*;
1022 use sqry_core::graph::unified::StagingOp;
1023 use sqry_core::graph::unified::TableWriteOp;
1024 use sqry_core::graph::unified::edge::EdgeKind;
1025 use std::path::PathBuf;
1026
1027 fn parse_sql(sql: &str) -> Tree {
1028 let mut parser = tree_sitter::Parser::new();
1029 parser
1030 .set_language(&tree_sitter_sequel::LANGUAGE.into())
1031 .expect("Failed to set SQL language");
1032 parser
1033 .parse(sql.as_bytes(), None)
1034 .expect("Failed to parse SQL")
1035 }
1036
1037 #[allow(dead_code)]
1039 fn get_table_read_edges(staging: &StagingGraph) -> Vec<String> {
1040 staging
1041 .operations()
1042 .iter()
1043 .filter_map(|op| {
1044 if let StagingOp::AddEdge {
1045 kind: EdgeKind::TableRead { table_name, .. },
1046 ..
1047 } = op
1048 {
1049 Some(format!("TableRead({table_name:?})"))
1052 } else {
1053 None
1054 }
1055 })
1056 .collect()
1057 }
1058
1059 #[allow(dead_code)]
1061 fn get_table_write_edges(staging: &StagingGraph) -> Vec<(String, TableWriteOp)> {
1062 staging
1063 .operations()
1064 .iter()
1065 .filter_map(|op| {
1066 if let StagingOp::AddEdge {
1067 kind:
1068 EdgeKind::TableWrite {
1069 table_name,
1070 operation,
1071 ..
1072 },
1073 ..
1074 } = op
1075 {
1076 Some((format!("TableWrite({table_name:?})"), *operation))
1077 } else {
1078 None
1079 }
1080 })
1081 .collect()
1082 }
1083
1084 fn count_table_read_edges(staging: &StagingGraph) -> usize {
1086 staging
1087 .operations()
1088 .iter()
1089 .filter(|op| {
1090 matches!(
1091 op,
1092 StagingOp::AddEdge {
1093 kind: EdgeKind::TableRead { .. },
1094 ..
1095 }
1096 )
1097 })
1098 .count()
1099 }
1100
1101 fn count_table_write_edges(staging: &StagingGraph) -> usize {
1102 staging
1103 .operations()
1104 .iter()
1105 .filter(|op| {
1106 matches!(
1107 op,
1108 StagingOp::AddEdge {
1109 kind: EdgeKind::TableWrite { .. },
1110 ..
1111 }
1112 )
1113 })
1114 .count()
1115 }
1116
1117 fn count_table_write_edges_by_op(staging: &StagingGraph, expected_op: TableWriteOp) -> usize {
1118 staging
1119 .operations()
1120 .iter()
1121 .filter(|op| {
1122 matches!(
1123 op,
1124 StagingOp::AddEdge { kind: EdgeKind::TableWrite { operation, .. }, .. }
1125 if *operation == expected_op
1126 )
1127 })
1128 .count()
1129 }
1130
1131 fn count_call_edges(staging: &StagingGraph) -> usize {
1132 staging
1133 .operations()
1134 .iter()
1135 .filter(|op| {
1136 matches!(
1137 op,
1138 StagingOp::AddEdge {
1139 kind: EdgeKind::Calls { .. },
1140 ..
1141 }
1142 )
1143 })
1144 .count()
1145 }
1146
1147 fn count_export_edges(staging: &StagingGraph) -> usize {
1149 staging
1150 .operations()
1151 .iter()
1152 .filter(|op| {
1153 matches!(
1154 op,
1155 StagingOp::AddEdge {
1156 kind: EdgeKind::Exports { .. },
1157 ..
1158 }
1159 )
1160 })
1161 .count()
1162 }
1163
1164 #[test]
1165 fn test_sql_graph_builder_new() {
1166 let builder = SqlGraphBuilder::new();
1167 assert_eq!(builder.language(), Language::Sql);
1168 }
1169
1170 #[test]
1171 fn test_select_creates_table_read_edge() {
1172 let sql = r"
1173 CREATE FUNCTION get_users()
1174 RETURNS TABLE (id INT, name TEXT) AS $$
1175 SELECT * FROM users;
1176 $$ LANGUAGE sql;
1177 ";
1178
1179 let tree = parse_sql(sql);
1180 let mut staging = StagingGraph::new();
1181 let builder = SqlGraphBuilder::new();
1182 let file = PathBuf::from("test.sql");
1183
1184 builder
1185 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1186 .expect("Graph building should succeed");
1187
1188 let read_count = count_table_read_edges(&staging);
1189 assert!(
1190 read_count >= 1,
1191 "Expected at least 1 TableRead edge, got {read_count}"
1192 );
1193 }
1194
1195 #[test]
1196 fn test_insert_creates_table_write_edge() {
1197 let sql = r"
1198 CREATE FUNCTION create_user(user_name TEXT)
1199 RETURNS VOID AS $$
1200 INSERT INTO users (name) VALUES (user_name);
1201 $$ LANGUAGE sql;
1202 ";
1203
1204 let tree = parse_sql(sql);
1205 let mut staging = StagingGraph::new();
1206 let builder = SqlGraphBuilder::new();
1207 let file = PathBuf::from("test.sql");
1208
1209 builder
1210 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1211 .expect("Graph building should succeed");
1212
1213 let insert_count = count_table_write_edges_by_op(&staging, TableWriteOp::Insert);
1214 assert!(
1215 insert_count >= 1,
1216 "Expected at least 1 TableWrite(Insert) edge, got {insert_count}"
1217 );
1218 }
1219
1220 #[test]
1221 fn test_update_creates_table_write_edge() {
1222 let sql = r"
1223 CREATE FUNCTION update_user(user_id INT, new_name TEXT)
1224 RETURNS VOID AS $$
1225 UPDATE users SET name = new_name WHERE id = user_id;
1226 $$ LANGUAGE sql;
1227 ";
1228
1229 let tree = parse_sql(sql);
1230 let mut staging = StagingGraph::new();
1231 let builder = SqlGraphBuilder::new();
1232 let file = PathBuf::from("test.sql");
1233
1234 builder
1235 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1236 .expect("Graph building should succeed");
1237
1238 let update_count = count_table_write_edges_by_op(&staging, TableWriteOp::Update);
1239 assert!(
1240 update_count >= 1,
1241 "Expected at least 1 TableWrite(Update) edge, got {update_count}"
1242 );
1243 }
1244
1245 #[test]
1246 fn test_delete_creates_table_write_edge() {
1247 let sql = r"
1248 CREATE FUNCTION delete_user(user_id INT)
1249 RETURNS VOID AS $$
1250 DELETE FROM users WHERE id = user_id;
1251 $$ LANGUAGE sql;
1252 ";
1253
1254 let tree = parse_sql(sql);
1255 let mut staging = StagingGraph::new();
1256 let builder = SqlGraphBuilder::new();
1257 let file = PathBuf::from("test.sql");
1258
1259 builder
1260 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1261 .expect("Graph building should succeed");
1262
1263 let delete_count = count_table_write_edges_by_op(&staging, TableWriteOp::Delete);
1264 assert!(
1265 delete_count >= 1,
1266 "Expected at least 1 TableWrite(Delete) edge, got {delete_count}"
1267 );
1268 }
1269
1270 #[test]
1271 fn test_join_creates_table_read_edge_for_primary_table() {
1272 let sql = r"
1275 CREATE FUNCTION get_user_orders()
1276 RETURNS TABLE (user_name TEXT, order_id INT) AS $$
1277 SELECT u.name, o.id FROM users u JOIN orders o ON u.id = o.user_id;
1278 $$ LANGUAGE sql;
1279 ";
1280
1281 let tree = parse_sql(sql);
1282 let mut staging = StagingGraph::new();
1283 let builder = SqlGraphBuilder::new();
1284 let file = PathBuf::from("test.sql");
1285
1286 builder
1287 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1288 .expect("Graph building should succeed");
1289
1290 let read_count = count_table_read_edges(&staging);
1292 assert!(
1293 read_count >= 1,
1294 "Expected at least 1 TableRead edge for FROM clause, got {read_count}"
1295 );
1296 }
1297
1298 #[test]
1299 fn test_multiple_joins_creates_table_read_edge_for_primary_table() {
1300 let sql = r"
1302 CREATE FUNCTION get_order_details()
1303 RETURNS TABLE (user_name TEXT, product_name TEXT, quantity INT) AS $$
1304 SELECT u.name, p.name, o.quantity
1305 FROM users u
1306 JOIN orders o ON u.id = o.user_id
1307 LEFT JOIN products p ON o.product_id = p.id;
1308 $$ LANGUAGE sql;
1309 ";
1310
1311 let tree = parse_sql(sql);
1312 let mut staging = StagingGraph::new();
1313 let builder = SqlGraphBuilder::new();
1314 let file = PathBuf::from("test.sql");
1315
1316 builder
1317 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1318 .expect("Graph building should succeed");
1319
1320 let read_count = count_table_read_edges(&staging);
1322 assert!(
1323 read_count >= 1,
1324 "Expected at least 1 TableRead edge for FROM clause, got {read_count}"
1325 );
1326 }
1327
1328 #[test]
1329 fn test_mixed_read_write_operations() {
1330 let sql = r"
1332 CREATE FUNCTION transfer_funds(from_id INT, to_id INT, amount DECIMAL)
1333 RETURNS VOID AS $$
1334 BEGIN
1335 SELECT balance FROM accounts WHERE id = from_id;
1336 UPDATE accounts SET balance = balance - amount WHERE id = from_id;
1337 UPDATE accounts SET balance = balance + amount WHERE id = to_id;
1338 INSERT INTO transactions (from_account, to_account, amount) VALUES (from_id, to_id, amount);
1339 END;
1340 $$ LANGUAGE plpgsql;
1341 ";
1342
1343 let tree = parse_sql(sql);
1344 let mut staging = StagingGraph::new();
1345 let builder = SqlGraphBuilder::new();
1346 let file = PathBuf::from("test.sql");
1347
1348 builder
1349 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1350 .expect("Graph building should succeed");
1351
1352 let read_count = count_table_read_edges(&staging);
1353 let write_count = count_table_write_edges(&staging);
1354
1355 assert!(
1357 read_count >= 1,
1358 "Expected at least 1 TableRead edge, got {read_count}"
1359 );
1360 assert!(
1361 write_count >= 1,
1362 "Expected at least 1 TableWrite edge, got {write_count}"
1363 );
1364 }
1365
1366 #[test]
1367 fn test_plpgsql_assignment_function_calls_create_call_edges() {
1368 let sql = r"
1369 CREATE FUNCTION add(a INT, b INT) RETURNS INT AS $$
1370 BEGIN
1371 RETURN a + b;
1372 END;
1373 $$ LANGUAGE plpgsql;
1374
1375 CREATE FUNCTION multiply(a INT, b INT) RETURNS INT AS $$
1376 BEGIN
1377 RETURN a * b;
1378 END;
1379 $$ LANGUAGE plpgsql;
1380
1381 CREATE FUNCTION compute(x INT, y INT, z INT) RETURNS INT AS $$
1382 DECLARE
1383 sum_val INT;
1384 BEGIN
1385 sum_val := add(x, y);
1386 RETURN multiply(sum_val, z);
1387 END;
1388 $$ LANGUAGE plpgsql;
1389 ";
1390
1391 let tree = parse_sql(sql);
1392 let mut staging = StagingGraph::new();
1393 let builder = SqlGraphBuilder::new();
1394 let file = PathBuf::from("nested_calls.sql");
1395
1396 builder
1397 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1398 .expect("Graph building should succeed");
1399
1400 let call_count = count_call_edges(&staging);
1401 assert!(
1402 call_count >= 2,
1403 "Expected at least 2 call edges for add() and multiply(), got {call_count}"
1404 );
1405 }
1406
1407 #[test]
1408 fn test_plpgsql_multiple_assignment_calls_create_call_edges() {
1409 let sql = r"
1410 CREATE FUNCTION helper_one() RETURNS INT AS $$
1411 BEGIN
1412 RETURN 42;
1413 END;
1414 $$ LANGUAGE plpgsql;
1415
1416 CREATE FUNCTION helper_two() RETURNS INT AS $$
1417 BEGIN
1418 RETURN 100;
1419 END;
1420 $$ LANGUAGE plpgsql;
1421
1422 CREATE FUNCTION orchestrator() RETURNS INT AS $$
1423 DECLARE
1424 val1 INT;
1425 val2 INT;
1426 BEGIN
1427 val1 := helper_one();
1428 val2 := helper_two();
1429 RETURN val1 + val2;
1430 END;
1431 $$ LANGUAGE plpgsql;
1432 ";
1433
1434 let tree = parse_sql(sql);
1435 let mut staging = StagingGraph::new();
1436 let builder = SqlGraphBuilder::new();
1437 let file = PathBuf::from("multiple_assignment_calls.sql");
1438
1439 builder
1440 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1441 .expect("Graph building should succeed");
1442
1443 let call_count = count_call_edges(&staging);
1444 assert!(
1445 call_count >= 2,
1446 "Expected at least 2 call edges for helper_one() and helper_two(), got {call_count}"
1447 );
1448 }
1449
1450 #[test]
1451 fn test_schema_qualified_table_name() {
1452 let sql = r"
1453 CREATE FUNCTION get_public_users()
1454 RETURNS TABLE (id INT, name TEXT) AS $$
1455 SELECT * FROM public.users;
1456 $$ LANGUAGE sql;
1457 ";
1458
1459 let tree = parse_sql(sql);
1460 let mut staging = StagingGraph::new();
1461 let builder = SqlGraphBuilder::new();
1462 let file = PathBuf::from("test.sql");
1463
1464 let result = builder.build_graph(&tree, sql.as_bytes(), &file, &mut staging);
1466 assert!(result.is_ok(), "Should handle schema-qualified table names");
1467 }
1468
1469 #[test]
1470 fn test_split_schema_table_with_schema() {
1471 let (schema, table) = split_schema_table("public.users");
1472 assert_eq!(schema, Some("public"));
1473 assert_eq!(table, "users");
1474 }
1475
1476 #[test]
1477 fn test_split_schema_table_without_schema() {
1478 let (schema, table) = split_schema_table("users");
1479 assert_eq!(schema, None);
1480 assert_eq!(table, "users");
1481 }
1482
1483 #[test]
1484 fn test_split_schema_table_with_whitespace() {
1485 let (schema, table) = split_schema_table(" public . users ");
1486 assert_eq!(schema, Some("public"));
1487 assert_eq!(table, "users");
1488 }
1489
1490 #[test]
1491 fn test_empty_sql_file() {
1492 let sql = "";
1493 let tree = parse_sql(sql);
1494 let mut staging = StagingGraph::new();
1495 let builder = SqlGraphBuilder::new();
1496 let file = PathBuf::from("empty.sql");
1497
1498 let result = builder.build_graph(&tree, sql.as_bytes(), &file, &mut staging);
1499 assert!(result.is_ok(), "Should handle empty SQL files");
1500 }
1501
1502 #[test]
1503 fn test_standalone_select_without_function() {
1504 let sql = "SELECT * FROM users;";
1506
1507 let tree = parse_sql(sql);
1508 let mut staging = StagingGraph::new();
1509 let builder = SqlGraphBuilder::new();
1510 let file = PathBuf::from("query.sql");
1511
1512 builder
1513 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1514 .expect("Graph building should succeed");
1515
1516 let read_count = count_table_read_edges(&staging);
1519 assert_eq!(
1520 read_count, 0,
1521 "Standalone SELECT should not create edges without enclosing function"
1522 );
1523 }
1524
1525 #[test]
1526 fn test_export_edges_for_table_definitions() {
1527 let sql = r"
1528 CREATE TABLE users (
1529 id SERIAL PRIMARY KEY,
1530 name TEXT NOT NULL
1531 );
1532
1533 CREATE TABLE orders (
1534 id SERIAL PRIMARY KEY,
1535 user_id INTEGER REFERENCES users(id)
1536 );
1537 ";
1538
1539 let tree = parse_sql(sql);
1540 let mut staging = StagingGraph::new();
1541 let builder = SqlGraphBuilder::new();
1542 let file = PathBuf::from("schema.sql");
1543
1544 builder
1545 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1546 .expect("Graph building should succeed");
1547
1548 let export_count = count_export_edges(&staging);
1549 assert_eq!(
1550 export_count, 2,
1551 "Expected 2 Export edges (users and orders), got {export_count}"
1552 );
1553 }
1554
1555 #[test]
1556 fn test_export_edges_for_view_definitions() {
1557 let sql = r"
1558 CREATE TABLE users (id INT, created_at TIMESTAMP);
1559
1560 CREATE VIEW active_users AS
1561 SELECT * FROM users WHERE created_at > NOW() - INTERVAL '30 days';
1562
1563 CREATE MATERIALIZED VIEW user_stats AS
1564 SELECT COUNT(*) as total FROM users;
1565 ";
1566
1567 let tree = parse_sql(sql);
1568 let mut staging = StagingGraph::new();
1569 let builder = SqlGraphBuilder::new();
1570 let file = PathBuf::from("views.sql");
1571
1572 builder
1573 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1574 .expect("Graph building should succeed");
1575
1576 let export_count = count_export_edges(&staging);
1577 assert_eq!(
1579 export_count, 3,
1580 "Expected 3 Export edges (1 table + 2 views), got {export_count}"
1581 );
1582 }
1583
1584 #[test]
1585 fn test_export_edges_for_functions_and_triggers() {
1586 let sql = r"
1587 CREATE FUNCTION get_balance(account_id INT) RETURNS BIGINT AS $$
1588 BEGIN
1589 RETURN 42;
1590 END;
1591 $$ LANGUAGE plpgsql;
1592
1593 CREATE FUNCTION update_balance() RETURNS TRIGGER AS $$
1594 BEGIN
1595 RETURN NEW;
1596 END;
1597 $$ LANGUAGE plpgsql;
1598
1599 CREATE TRIGGER balance_updated
1600 BEFORE INSERT ON accounts
1601 FOR EACH ROW
1602 EXECUTE FUNCTION update_balance();
1603 ";
1604
1605 let tree = parse_sql(sql);
1606 let mut staging = StagingGraph::new();
1607 let builder = SqlGraphBuilder::new();
1608 let file = PathBuf::from("banking.sql");
1609
1610 builder
1611 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1612 .expect("Graph building should succeed");
1613
1614 let export_count = count_export_edges(&staging);
1615 assert!(
1619 export_count >= 3,
1620 "Expected at least 3 Export edges (2 functions + 1 trigger), got {export_count}"
1621 );
1622 }
1623
1624 #[test]
1625 fn test_export_edges_with_schema_qualified_names() {
1626 let sql = r"
1627 CREATE TABLE public.customers (
1628 id SERIAL PRIMARY KEY,
1629 name TEXT NOT NULL
1630 );
1631
1632 CREATE FUNCTION public.get_customer_name(cust_id INT) RETURNS TEXT AS $$
1633 BEGIN
1634 RETURN 'test';
1635 END;
1636 $$ LANGUAGE plpgsql;
1637 ";
1638
1639 let tree = parse_sql(sql);
1640 let mut staging = StagingGraph::new();
1641 let builder = SqlGraphBuilder::new();
1642 let file = PathBuf::from("public_schema.sql");
1643
1644 builder
1645 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1646 .expect("Graph building should succeed");
1647
1648 let export_count = count_export_edges(&staging);
1649 assert_eq!(
1651 export_count, 2,
1652 "Expected 2 Export edges (table + function), got {export_count}"
1653 );
1654 }
1655
1656 #[test]
1657 fn test_mixed_database_objects_exports() {
1658 let sql = r"
1659 CREATE TABLE accounts (
1660 id SERIAL PRIMARY KEY,
1661 balance_cents BIGINT NOT NULL
1662 );
1663
1664 CREATE VIEW positive_balances AS
1665 SELECT * FROM accounts WHERE balance_cents > 0;
1666
1667 CREATE FUNCTION get_balance(account_id INT) RETURNS BIGINT AS $$
1668 BEGIN
1669 RETURN (SELECT balance_cents FROM accounts WHERE id = account_id);
1670 END;
1671 $$ LANGUAGE plpgsql;
1672 ";
1673
1674 let tree = parse_sql(sql);
1675 let mut staging = StagingGraph::new();
1676 let builder = SqlGraphBuilder::new();
1677 let file = PathBuf::from("mixed.sql");
1678
1679 builder
1680 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1681 .expect("Graph building should succeed");
1682
1683 let export_count = count_export_edges(&staging);
1684 assert_eq!(
1686 export_count, 3,
1687 "Expected 3 Export edges (table + view + function), got {export_count}"
1688 );
1689 }
1690
1691 #[test]
1692 fn test_no_exports_for_empty_file() {
1693 let sql = "";
1694 let tree = parse_sql(sql);
1695 let mut staging = StagingGraph::new();
1696 let builder = SqlGraphBuilder::new();
1697 let file = PathBuf::from("empty.sql");
1698
1699 builder
1700 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1701 .expect("Graph building should succeed");
1702
1703 let export_count = count_export_edges(&staging);
1704 assert_eq!(
1705 export_count, 0,
1706 "Expected 0 Export edges for empty file, got {export_count}"
1707 );
1708 }
1709}
1710
1711#[cfg(test)]
1712mod shape_tests {
1713 use super::*;
1714 use sqry_core::graph::unified::build::shape::{ShapeBudget, compute_shape_descriptor};
1715
1716 const SAMPLE: &str = include_str!(concat!(
1717 env!("CARGO_MANIFEST_DIR"),
1718 "/../test-fixtures/shape/data/sample.sql"
1719 ));
1720
1721 fn parse(src: &str) -> Tree {
1722 let mut parser = tree_sitter::Parser::new();
1723 parser
1724 .set_language(&tree_sitter_sequel::LANGUAGE.into())
1725 .expect("load sql grammar");
1726 parser.parse(src, None).expect("parse")
1727 }
1728
1729 fn first_create_function(node: Node<'_>) -> Option<Node<'_>> {
1731 if node.kind() == "create_function" {
1732 return Some(node);
1733 }
1734 let mut cursor = node.walk();
1735 for child in node.children(&mut cursor) {
1736 if let Some(found) = first_create_function(child) {
1737 return Some(found);
1738 }
1739 }
1740 None
1741 }
1742
1743 #[test]
1744 fn cf_map_is_non_empty_and_covers_real_kinds() {
1745 let mapping = sql_shape_mapping();
1746 let populated = mapping.cf_by_kind_id.iter().filter(|s| s.is_some()).count();
1747 assert!(
1748 populated > 0,
1749 "SQL cf map must map at least one real grammar kind"
1750 );
1751
1752 let lang: tree_sitter::Language = tree_sitter_sequel::LANGUAGE.into();
1755 let case_id = lang.id_for_node_kind("case", true);
1756 let when_id = lang.id_for_node_kind("when_clause", true);
1757 assert_eq!(mapping.cf_bucket(case_id), Some(CfBucket::Match));
1758 assert_eq!(mapping.cf_bucket(when_id), Some(CfBucket::Branch));
1759 }
1760
1761 #[test]
1762 fn descriptor_counts_case_control_flow_in_sql_body() {
1763 let tree = parse(SAMPLE);
1764 let func = first_create_function(tree.root_node()).expect("create_function in fixture");
1765 let descriptor = compute_shape_descriptor(
1766 func,
1767 SAMPLE.as_bytes(),
1768 sql_shape_mapping(),
1769 &ShapeBudget::default(),
1770 );
1771 assert!(
1772 !descriptor.is_unhashable(),
1773 "a function with a parsed CASE body must be hashable"
1774 );
1775 assert!(
1776 descriptor.cf_histogram[CfBucket::Match.index()] >= 1,
1777 "the CASE expression must be counted in the Match bucket"
1778 );
1779 }
1780
1781 #[test]
1782 fn signature_shape_reads_arguments_and_defaults() {
1783 let tree = parse(SAMPLE);
1784 let func = first_create_function(tree.root_node()).expect("create_function in fixture");
1785 let shape = sql_shape_mapping().signature_shape(func, SAMPLE.as_bytes());
1786 assert_eq!(
1787 shape.arity_positional, 2,
1788 "grade(score, bonus) has two arguments"
1789 );
1790 assert!(
1791 shape.has_defaults,
1792 "the DEFAULT 0 argument must set has_defaults"
1793 );
1794 }
1795}