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 callables.push(SqlCallable {
503 node_id,
504 start_byte: node.start_byte(),
505 end_byte: node.end_byte(),
506 });
507 }
508 }
509
510 callables
511}
512
513fn extract_triggers(
515 tree: &Tree,
516 content: &[u8],
517 query: &Query,
518 helper: &mut GraphBuildHelper,
519) -> Vec<SqlCallable> {
520 let mut callables = Vec::new();
521 let mut cursor = QueryCursor::new();
522 let capture_names = query.capture_names();
523 let mut matches = cursor.matches(query, tree.root_node(), content);
524
525 while let Some(m) = matches.next() {
526 let mut trigger_name = None;
527 let mut table_name = None;
528 let mut trigger_node = None;
529
530 for capture in m.captures {
531 let name = capture_names[capture.index as usize];
532 match name {
533 "trigger.name" => {
534 if let Ok(text) = capture.node.utf8_text(content) {
535 trigger_name = Some(text.to_string());
536 }
537 }
538 "trigger.table" => {
539 if let Ok(text) = capture.node.utf8_text(content) {
540 table_name = Some(text.to_string());
541 }
542 }
543 "trigger" => {
544 trigger_node = Some(capture.node);
545 }
546 _ => {}
547 }
548 }
549
550 if let (Some(trigger), Some(table), Some(node)) = (trigger_name, table_name, trigger_node) {
551 let (schema, table_only) = split_schema_table(&table);
552 let span = Span::from_node(&node);
553
554 let trigger_id = helper.add_function(&trigger, Some(span), false, false);
555 callables.push(SqlCallable {
556 node_id: trigger_id,
557 start_byte: node.start_byte(),
558 end_byte: node.end_byte(),
559 });
560
561 let table_id = helper.add_variable(table_only, Some(span));
562 helper.add_triggered_by_edge_with_span(
563 trigger_id,
564 table_id,
565 &trigger,
566 schema,
567 vec![span],
568 );
569 }
570 }
571
572 callables
573}
574
575fn extract_trigger_execute_function_calls(
580 tree: &Tree,
581 content: &[u8],
582 query: &Query,
583 callables: &[SqlCallable],
584 helper: &mut GraphBuildHelper,
585) {
586 let mut cursor = QueryCursor::new();
587 let capture_names = query.capture_names();
588 let mut matches = cursor.matches(query, tree.root_node(), content);
589
590 while let Some(m) = matches.next() {
591 let mut trigger_name = None;
592 let mut func_name = None;
593 let mut trigger_node = None;
594
595 for capture in m.captures {
596 let name = capture_names[capture.index as usize];
597 match name {
598 "trigger.name" => {
599 if let Ok(text) = capture.node.utf8_text(content) {
600 trigger_name = Some(text.to_string());
601 }
602 }
603 "func.name" => {
604 if let Ok(text) = capture.node.utf8_text(content) {
605 func_name = Some(text.to_string());
606 }
607 }
608 "trigger_exec" => {
609 trigger_node = Some(capture.node);
610 }
611 _ => {}
612 }
613 }
614
615 if let (Some(_trigger), Some(func), Some(node)) = (trigger_name, func_name, trigger_node) {
616 let span = Span::from_node(&node);
617
618 if let Some(trigger_callable) = callables.iter().find(|c| {
620 c.start_byte <= node.start_byte() && node.end_byte() <= c.end_byte
622 }) {
623 let callee_id = helper.add_function(&func, Some(span), false, false);
625 helper.add_call_edge_full_with_span(
626 trigger_callable.node_id,
627 callee_id,
628 255,
629 false,
630 vec![span],
631 );
632 }
633 }
634 }
635}
636
637fn extract_table_reads(
639 tree: &Tree,
640 content: &[u8],
641 query: &Query,
642 helper: &mut GraphBuildHelper,
643) -> Vec<SqlTableOp> {
644 let mut ops = Vec::new();
645 let mut cursor = QueryCursor::new();
646 let capture_names = query.capture_names();
647 let mut matches = cursor.matches(query, tree.root_node(), content);
648
649 while let Some(m) = matches.next() {
650 let mut table_name = None;
651 let mut op_node = None;
652
653 for capture in m.captures {
654 let name = capture_names[capture.index as usize];
655 match name {
656 "table.name" => {
657 if let Ok(text) = capture.node.utf8_text(content) {
658 table_name = Some(text.to_string());
659 }
660 }
661 "select" => op_node = Some(capture.node),
662 _ => {}
663 }
664 }
665
666 if let (Some(table_name), Some(node)) = (table_name, op_node) {
667 let (schema, table_only) = split_schema_table(&table_name);
668 let span = Span::from_node(&node);
669 let table_node_id = helper.add_variable(table_only, Some(span));
670 ops.push(SqlTableOp {
671 op_span_bytes: (node.start_byte(), node.end_byte()),
672 kind: SqlTableOpKind::Read,
673 table_name: table_only.to_string(),
674 schema: schema.map(str::to_string),
675 table_node_id,
676 span,
677 });
678 }
679 }
680
681 ops
682}
683
684fn extract_table_writes(
686 tree: &Tree,
687 content: &[u8],
688 query: &Query,
689 helper: &mut GraphBuildHelper,
690) -> Vec<SqlTableOp> {
691 let mut ops = Vec::new();
692 let mut cursor = QueryCursor::new();
693 let capture_names = query.capture_names();
694 let mut matches = cursor.matches(query, tree.root_node(), content);
695
696 while let Some(m) = matches.next() {
697 let mut table_name = None;
698 let mut write_node = None;
699
700 for capture in m.captures {
701 let name = capture_names[capture.index as usize];
702 match name {
703 "table.name" => {
704 if let Ok(text) = capture.node.utf8_text(content) {
705 table_name = Some(text.to_string());
706 }
707 }
708 "write" => write_node = Some(capture.node),
709 _ => {}
710 }
711 }
712
713 let Some(table_name) = table_name else {
714 continue;
715 };
716 let Some(node) = write_node else {
717 continue;
718 };
719
720 let operation = match node.kind() {
721 "insert" => sqry_core::graph::unified::TableWriteOp::Insert,
722 "delete" => sqry_core::graph::unified::TableWriteOp::Delete,
723 _ => sqry_core::graph::unified::TableWriteOp::Update,
724 };
725
726 let (schema, table_only) = split_schema_table(&table_name);
727 let span = Span::from_node(&node);
728 let table_node_id = helper.add_variable(table_only, Some(span));
729 ops.push(SqlTableOp {
730 op_span_bytes: (node.start_byte(), node.end_byte()),
731 kind: SqlTableOpKind::Write(operation),
732 table_name: table_only.to_string(),
733 schema: schema.map(str::to_string),
734 table_node_id,
735 span,
736 });
737 }
738
739 ops
740}
741
742#[derive(Debug)]
744struct SqlFunctionCall {
745 callee_name: String,
746 span_bytes: (usize, usize),
747 span: Span,
748}
749
750fn extract_function_calls(tree: &Tree, content: &[u8], query: &Query) -> Vec<SqlFunctionCall> {
752 let mut calls = Vec::new();
753 let mut cursor = QueryCursor::new();
754 let capture_names = query.capture_names();
755 let mut matches = cursor.matches(query, tree.root_node(), content);
756
757 while let Some(m) = matches.next() {
758 let mut call_name = None;
759 let mut call_node = None;
760
761 for capture in m.captures {
762 let name = capture_names[capture.index as usize];
763 match name {
764 "call.name" => {
765 if let Ok(text) = capture.node.utf8_text(content) {
766 call_name = Some(normalize_callee_name(text));
767 }
768 }
769 "call" | "call.error" => call_node = Some(capture.node),
770 _ => {}
771 }
772 }
773
774 let Some(node) = call_node else {
775 continue;
776 };
777
778 let span_bytes = (node.start_byte(), node.end_byte());
779 let span = Span::from_node(&node);
780
781 if node.kind() == "ERROR" {
782 if let Ok(text) = node.utf8_text(content) {
783 for name in extract_error_call_names(text) {
784 calls.push(SqlFunctionCall {
785 callee_name: name,
786 span_bytes,
787 span,
788 });
789 }
790 }
791 continue;
792 }
793
794 if let Some(name) = call_name
795 && !name.is_empty()
796 {
797 calls.push(SqlFunctionCall {
798 callee_name: name,
799 span_bytes,
800 span,
801 });
802 }
803 }
804
805 calls
806}
807
808fn normalize_callee_name(name: &str) -> String {
809 name.trim()
810 .rsplit('.')
811 .next()
812 .unwrap_or_default()
813 .trim()
814 .to_string()
815}
816
817fn extract_error_call_names(text: &str) -> Vec<String> {
818 let bytes = text.as_bytes();
819 let mut offset = 0;
820 let mut call_names = Vec::new();
821
822 while offset < bytes.len() {
823 if !is_sql_identifier_start(bytes[offset]) {
824 offset += 1;
825 continue;
826 }
827
828 let start = offset;
829 offset += 1;
830 while offset < bytes.len() && is_sql_identifier_continue(bytes[offset]) {
831 offset += 1;
832 }
833
834 let token = &text[start..offset];
835 let mut lookahead = offset;
836 while lookahead < bytes.len() && bytes[lookahead].is_ascii_whitespace() {
837 lookahead += 1;
838 }
839
840 if lookahead < bytes.len() && bytes[lookahead] == b'(' {
841 let normalized = normalize_callee_name(token);
842 if !normalized.is_empty() && !call_names.iter().any(|name| name == &normalized) {
843 call_names.push(normalized);
844 }
845 }
846 }
847
848 call_names
849}
850
851const fn is_sql_identifier_start(byte: u8) -> bool {
852 byte.is_ascii_alphabetic() || byte == b'_'
853}
854
855const fn is_sql_identifier_continue(byte: u8) -> bool {
856 byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'.')
857}
858
859fn extract_table_definitions(
861 tree: &Tree,
862 content: &[u8],
863 query: &Query,
864 helper: &mut GraphBuildHelper,
865) -> Vec<SqlDatabaseObject> {
866 let mut objects = Vec::new();
867 let mut cursor = QueryCursor::new();
868 let capture_names = query.capture_names();
869 let mut matches = cursor.matches(query, tree.root_node(), content);
870
871 while let Some(m) = matches.next() {
872 let mut table_name = None;
873 let mut table_node = None;
874
875 for capture in m.captures {
876 let name = capture_names[capture.index as usize];
877 match name {
878 "table.name" => {
879 if let Ok(text) = capture.node.utf8_text(content) {
880 table_name = Some(text.to_string());
881 }
882 }
883 "table" => table_node = Some(capture.node),
884 _ => {}
885 }
886 }
887
888 if let (Some(name), Some(node)) = (table_name, table_node) {
889 let (_, table_only) = split_schema_table(&name);
891 let span = Span::from_node(&node);
892 let node_id = helper.add_variable(table_only, Some(span));
893 objects.push(SqlDatabaseObject { node_id });
894 }
895 }
896
897 objects
898}
899
900fn extract_view_definitions(
902 tree: &Tree,
903 content: &[u8],
904 query: &Query,
905 helper: &mut GraphBuildHelper,
906) -> Vec<SqlDatabaseObject> {
907 let mut objects = Vec::new();
908 let mut cursor = QueryCursor::new();
909 let capture_names = query.capture_names();
910 let mut matches = cursor.matches(query, tree.root_node(), content);
911
912 while let Some(m) = matches.next() {
913 let mut view_name = None;
914 let mut view_node = None;
915
916 for capture in m.captures {
917 let name = capture_names[capture.index as usize];
918 match name {
919 "view.name" => {
920 if let Ok(text) = capture.node.utf8_text(content) {
921 view_name = Some(text.to_string());
922 }
923 }
924 "view" => view_node = Some(capture.node),
925 _ => {}
926 }
927 }
928
929 if let (Some(name), Some(node)) = (view_name, view_node) {
930 let (_, view_only) = split_schema_table(&name);
932 let span = Span::from_node(&node);
933 let node_id = helper.add_variable(view_only, Some(span));
934 objects.push(SqlDatabaseObject { node_id });
935 }
936 }
937
938 objects
939}
940
941fn find_enclosing_callable(
942 callables: &[SqlCallable],
943 op_span_bytes: (usize, usize),
944) -> Option<&SqlCallable> {
945 let (start_byte, end_byte) = op_span_bytes;
946 callables
947 .iter()
948 .filter(|c| c.start_byte <= start_byte && end_byte <= c.end_byte)
949 .min_by_key(|c| c.end_byte.saturating_sub(c.start_byte))
950}
951
952fn split_schema_table(name: &str) -> (Option<&str>, &str) {
953 let mut parts = name.splitn(2, '.');
954 let first = parts.next().unwrap_or(name).trim();
955 let second = parts.next().map(str::trim);
956 match second {
957 Some(table) if !table.is_empty() => (Some(first), table),
958 _ => (None, first),
959 }
960}
961
962trait SpanExt {
964 fn from_node(node: &tree_sitter::Node) -> Self;
965}
966
967impl SpanExt for Span {
968 fn from_node(node: &tree_sitter::Node) -> Self {
969 Span::new(
970 Position::new(node.start_position().row, node.start_position().column),
971 Position::new(node.end_position().row, node.end_position().column),
972 )
973 }
974}
975
976fn emit_exports(
982 helper: &mut GraphBuildHelper,
983 callables: &[SqlCallable],
984 tables: &[SqlDatabaseObject],
985 views: &[SqlDatabaseObject],
986) {
987 if callables.is_empty() && tables.is_empty() && views.is_empty() {
989 return;
990 }
991
992 let module_id = helper.add_module(FILE_MODULE_NAME, None);
994
995 for callable in callables {
997 helper.add_export_edge(module_id, callable.node_id);
998 }
999
1000 for table in tables {
1002 helper.add_export_edge(module_id, table.node_id);
1003 }
1004
1005 for view in views {
1007 helper.add_export_edge(module_id, view.node_id);
1008 }
1009}
1010
1011#[cfg(test)]
1012mod tests {
1013 use super::*;
1014 use sqry_core::graph::unified::StagingOp;
1015 use sqry_core::graph::unified::TableWriteOp;
1016 use sqry_core::graph::unified::edge::EdgeKind;
1017 use std::path::PathBuf;
1018
1019 fn parse_sql(sql: &str) -> Tree {
1020 let mut parser = tree_sitter::Parser::new();
1021 parser
1022 .set_language(&tree_sitter_sequel::LANGUAGE.into())
1023 .expect("Failed to set SQL language");
1024 parser
1025 .parse(sql.as_bytes(), None)
1026 .expect("Failed to parse SQL")
1027 }
1028
1029 #[allow(dead_code)]
1031 fn get_table_read_edges(staging: &StagingGraph) -> Vec<String> {
1032 staging
1033 .operations()
1034 .iter()
1035 .filter_map(|op| {
1036 if let StagingOp::AddEdge {
1037 kind: EdgeKind::TableRead { table_name, .. },
1038 ..
1039 } = op
1040 {
1041 Some(format!("TableRead({table_name:?})"))
1044 } else {
1045 None
1046 }
1047 })
1048 .collect()
1049 }
1050
1051 #[allow(dead_code)]
1053 fn get_table_write_edges(staging: &StagingGraph) -> Vec<(String, TableWriteOp)> {
1054 staging
1055 .operations()
1056 .iter()
1057 .filter_map(|op| {
1058 if let StagingOp::AddEdge {
1059 kind:
1060 EdgeKind::TableWrite {
1061 table_name,
1062 operation,
1063 ..
1064 },
1065 ..
1066 } = op
1067 {
1068 Some((format!("TableWrite({table_name:?})"), *operation))
1069 } else {
1070 None
1071 }
1072 })
1073 .collect()
1074 }
1075
1076 fn count_table_read_edges(staging: &StagingGraph) -> usize {
1078 staging
1079 .operations()
1080 .iter()
1081 .filter(|op| {
1082 matches!(
1083 op,
1084 StagingOp::AddEdge {
1085 kind: EdgeKind::TableRead { .. },
1086 ..
1087 }
1088 )
1089 })
1090 .count()
1091 }
1092
1093 fn count_table_write_edges(staging: &StagingGraph) -> usize {
1094 staging
1095 .operations()
1096 .iter()
1097 .filter(|op| {
1098 matches!(
1099 op,
1100 StagingOp::AddEdge {
1101 kind: EdgeKind::TableWrite { .. },
1102 ..
1103 }
1104 )
1105 })
1106 .count()
1107 }
1108
1109 fn count_table_write_edges_by_op(staging: &StagingGraph, expected_op: TableWriteOp) -> usize {
1110 staging
1111 .operations()
1112 .iter()
1113 .filter(|op| {
1114 matches!(
1115 op,
1116 StagingOp::AddEdge { kind: EdgeKind::TableWrite { operation, .. }, .. }
1117 if *operation == expected_op
1118 )
1119 })
1120 .count()
1121 }
1122
1123 fn count_call_edges(staging: &StagingGraph) -> usize {
1124 staging
1125 .operations()
1126 .iter()
1127 .filter(|op| {
1128 matches!(
1129 op,
1130 StagingOp::AddEdge {
1131 kind: EdgeKind::Calls { .. },
1132 ..
1133 }
1134 )
1135 })
1136 .count()
1137 }
1138
1139 fn count_export_edges(staging: &StagingGraph) -> usize {
1141 staging
1142 .operations()
1143 .iter()
1144 .filter(|op| {
1145 matches!(
1146 op,
1147 StagingOp::AddEdge {
1148 kind: EdgeKind::Exports { .. },
1149 ..
1150 }
1151 )
1152 })
1153 .count()
1154 }
1155
1156 #[test]
1157 fn test_sql_graph_builder_new() {
1158 let builder = SqlGraphBuilder::new();
1159 assert_eq!(builder.language(), Language::Sql);
1160 }
1161
1162 #[test]
1163 fn test_select_creates_table_read_edge() {
1164 let sql = r"
1165 CREATE FUNCTION get_users()
1166 RETURNS TABLE (id INT, name TEXT) AS $$
1167 SELECT * FROM users;
1168 $$ LANGUAGE sql;
1169 ";
1170
1171 let tree = parse_sql(sql);
1172 let mut staging = StagingGraph::new();
1173 let builder = SqlGraphBuilder::new();
1174 let file = PathBuf::from("test.sql");
1175
1176 builder
1177 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1178 .expect("Graph building should succeed");
1179
1180 let read_count = count_table_read_edges(&staging);
1181 assert!(
1182 read_count >= 1,
1183 "Expected at least 1 TableRead edge, got {read_count}"
1184 );
1185 }
1186
1187 #[test]
1188 fn test_insert_creates_table_write_edge() {
1189 let sql = r"
1190 CREATE FUNCTION create_user(user_name TEXT)
1191 RETURNS VOID AS $$
1192 INSERT INTO users (name) VALUES (user_name);
1193 $$ LANGUAGE sql;
1194 ";
1195
1196 let tree = parse_sql(sql);
1197 let mut staging = StagingGraph::new();
1198 let builder = SqlGraphBuilder::new();
1199 let file = PathBuf::from("test.sql");
1200
1201 builder
1202 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1203 .expect("Graph building should succeed");
1204
1205 let insert_count = count_table_write_edges_by_op(&staging, TableWriteOp::Insert);
1206 assert!(
1207 insert_count >= 1,
1208 "Expected at least 1 TableWrite(Insert) edge, got {insert_count}"
1209 );
1210 }
1211
1212 #[test]
1213 fn test_update_creates_table_write_edge() {
1214 let sql = r"
1215 CREATE FUNCTION update_user(user_id INT, new_name TEXT)
1216 RETURNS VOID AS $$
1217 UPDATE users SET name = new_name WHERE id = user_id;
1218 $$ LANGUAGE sql;
1219 ";
1220
1221 let tree = parse_sql(sql);
1222 let mut staging = StagingGraph::new();
1223 let builder = SqlGraphBuilder::new();
1224 let file = PathBuf::from("test.sql");
1225
1226 builder
1227 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1228 .expect("Graph building should succeed");
1229
1230 let update_count = count_table_write_edges_by_op(&staging, TableWriteOp::Update);
1231 assert!(
1232 update_count >= 1,
1233 "Expected at least 1 TableWrite(Update) edge, got {update_count}"
1234 );
1235 }
1236
1237 #[test]
1238 fn test_delete_creates_table_write_edge() {
1239 let sql = r"
1240 CREATE FUNCTION delete_user(user_id INT)
1241 RETURNS VOID AS $$
1242 DELETE FROM users WHERE id = user_id;
1243 $$ LANGUAGE sql;
1244 ";
1245
1246 let tree = parse_sql(sql);
1247 let mut staging = StagingGraph::new();
1248 let builder = SqlGraphBuilder::new();
1249 let file = PathBuf::from("test.sql");
1250
1251 builder
1252 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1253 .expect("Graph building should succeed");
1254
1255 let delete_count = count_table_write_edges_by_op(&staging, TableWriteOp::Delete);
1256 assert!(
1257 delete_count >= 1,
1258 "Expected at least 1 TableWrite(Delete) edge, got {delete_count}"
1259 );
1260 }
1261
1262 #[test]
1263 fn test_join_creates_table_read_edge_for_primary_table() {
1264 let sql = r"
1267 CREATE FUNCTION get_user_orders()
1268 RETURNS TABLE (user_name TEXT, order_id INT) AS $$
1269 SELECT u.name, o.id FROM users u JOIN orders o ON u.id = o.user_id;
1270 $$ LANGUAGE sql;
1271 ";
1272
1273 let tree = parse_sql(sql);
1274 let mut staging = StagingGraph::new();
1275 let builder = SqlGraphBuilder::new();
1276 let file = PathBuf::from("test.sql");
1277
1278 builder
1279 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1280 .expect("Graph building should succeed");
1281
1282 let read_count = count_table_read_edges(&staging);
1284 assert!(
1285 read_count >= 1,
1286 "Expected at least 1 TableRead edge for FROM clause, got {read_count}"
1287 );
1288 }
1289
1290 #[test]
1291 fn test_multiple_joins_creates_table_read_edge_for_primary_table() {
1292 let sql = r"
1294 CREATE FUNCTION get_order_details()
1295 RETURNS TABLE (user_name TEXT, product_name TEXT, quantity INT) AS $$
1296 SELECT u.name, p.name, o.quantity
1297 FROM users u
1298 JOIN orders o ON u.id = o.user_id
1299 LEFT JOIN products p ON o.product_id = p.id;
1300 $$ LANGUAGE sql;
1301 ";
1302
1303 let tree = parse_sql(sql);
1304 let mut staging = StagingGraph::new();
1305 let builder = SqlGraphBuilder::new();
1306 let file = PathBuf::from("test.sql");
1307
1308 builder
1309 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1310 .expect("Graph building should succeed");
1311
1312 let read_count = count_table_read_edges(&staging);
1314 assert!(
1315 read_count >= 1,
1316 "Expected at least 1 TableRead edge for FROM clause, got {read_count}"
1317 );
1318 }
1319
1320 #[test]
1321 fn test_mixed_read_write_operations() {
1322 let sql = r"
1324 CREATE FUNCTION transfer_funds(from_id INT, to_id INT, amount DECIMAL)
1325 RETURNS VOID AS $$
1326 BEGIN
1327 SELECT balance FROM accounts WHERE id = from_id;
1328 UPDATE accounts SET balance = balance - amount WHERE id = from_id;
1329 UPDATE accounts SET balance = balance + amount WHERE id = to_id;
1330 INSERT INTO transactions (from_account, to_account, amount) VALUES (from_id, to_id, amount);
1331 END;
1332 $$ LANGUAGE plpgsql;
1333 ";
1334
1335 let tree = parse_sql(sql);
1336 let mut staging = StagingGraph::new();
1337 let builder = SqlGraphBuilder::new();
1338 let file = PathBuf::from("test.sql");
1339
1340 builder
1341 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1342 .expect("Graph building should succeed");
1343
1344 let read_count = count_table_read_edges(&staging);
1345 let write_count = count_table_write_edges(&staging);
1346
1347 assert!(
1349 read_count >= 1,
1350 "Expected at least 1 TableRead edge, got {read_count}"
1351 );
1352 assert!(
1353 write_count >= 1,
1354 "Expected at least 1 TableWrite edge, got {write_count}"
1355 );
1356 }
1357
1358 #[test]
1359 fn test_plpgsql_assignment_function_calls_create_call_edges() {
1360 let sql = r"
1361 CREATE FUNCTION add(a INT, b INT) RETURNS INT AS $$
1362 BEGIN
1363 RETURN a + b;
1364 END;
1365 $$ LANGUAGE plpgsql;
1366
1367 CREATE FUNCTION multiply(a INT, b INT) RETURNS INT AS $$
1368 BEGIN
1369 RETURN a * b;
1370 END;
1371 $$ LANGUAGE plpgsql;
1372
1373 CREATE FUNCTION compute(x INT, y INT, z INT) RETURNS INT AS $$
1374 DECLARE
1375 sum_val INT;
1376 BEGIN
1377 sum_val := add(x, y);
1378 RETURN multiply(sum_val, z);
1379 END;
1380 $$ LANGUAGE plpgsql;
1381 ";
1382
1383 let tree = parse_sql(sql);
1384 let mut staging = StagingGraph::new();
1385 let builder = SqlGraphBuilder::new();
1386 let file = PathBuf::from("nested_calls.sql");
1387
1388 builder
1389 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1390 .expect("Graph building should succeed");
1391
1392 let call_count = count_call_edges(&staging);
1393 assert!(
1394 call_count >= 2,
1395 "Expected at least 2 call edges for add() and multiply(), got {call_count}"
1396 );
1397 }
1398
1399 #[test]
1400 fn test_plpgsql_multiple_assignment_calls_create_call_edges() {
1401 let sql = r"
1402 CREATE FUNCTION helper_one() RETURNS INT AS $$
1403 BEGIN
1404 RETURN 42;
1405 END;
1406 $$ LANGUAGE plpgsql;
1407
1408 CREATE FUNCTION helper_two() RETURNS INT AS $$
1409 BEGIN
1410 RETURN 100;
1411 END;
1412 $$ LANGUAGE plpgsql;
1413
1414 CREATE FUNCTION orchestrator() RETURNS INT AS $$
1415 DECLARE
1416 val1 INT;
1417 val2 INT;
1418 BEGIN
1419 val1 := helper_one();
1420 val2 := helper_two();
1421 RETURN val1 + val2;
1422 END;
1423 $$ LANGUAGE plpgsql;
1424 ";
1425
1426 let tree = parse_sql(sql);
1427 let mut staging = StagingGraph::new();
1428 let builder = SqlGraphBuilder::new();
1429 let file = PathBuf::from("multiple_assignment_calls.sql");
1430
1431 builder
1432 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1433 .expect("Graph building should succeed");
1434
1435 let call_count = count_call_edges(&staging);
1436 assert!(
1437 call_count >= 2,
1438 "Expected at least 2 call edges for helper_one() and helper_two(), got {call_count}"
1439 );
1440 }
1441
1442 #[test]
1443 fn test_schema_qualified_table_name() {
1444 let sql = r"
1445 CREATE FUNCTION get_public_users()
1446 RETURNS TABLE (id INT, name TEXT) AS $$
1447 SELECT * FROM public.users;
1448 $$ LANGUAGE sql;
1449 ";
1450
1451 let tree = parse_sql(sql);
1452 let mut staging = StagingGraph::new();
1453 let builder = SqlGraphBuilder::new();
1454 let file = PathBuf::from("test.sql");
1455
1456 let result = builder.build_graph(&tree, sql.as_bytes(), &file, &mut staging);
1458 assert!(result.is_ok(), "Should handle schema-qualified table names");
1459 }
1460
1461 #[test]
1462 fn test_split_schema_table_with_schema() {
1463 let (schema, table) = split_schema_table("public.users");
1464 assert_eq!(schema, Some("public"));
1465 assert_eq!(table, "users");
1466 }
1467
1468 #[test]
1469 fn test_split_schema_table_without_schema() {
1470 let (schema, table) = split_schema_table("users");
1471 assert_eq!(schema, None);
1472 assert_eq!(table, "users");
1473 }
1474
1475 #[test]
1476 fn test_split_schema_table_with_whitespace() {
1477 let (schema, table) = split_schema_table(" public . users ");
1478 assert_eq!(schema, Some("public"));
1479 assert_eq!(table, "users");
1480 }
1481
1482 #[test]
1483 fn test_empty_sql_file() {
1484 let sql = "";
1485 let tree = parse_sql(sql);
1486 let mut staging = StagingGraph::new();
1487 let builder = SqlGraphBuilder::new();
1488 let file = PathBuf::from("empty.sql");
1489
1490 let result = builder.build_graph(&tree, sql.as_bytes(), &file, &mut staging);
1491 assert!(result.is_ok(), "Should handle empty SQL files");
1492 }
1493
1494 #[test]
1495 fn test_standalone_select_without_function() {
1496 let sql = "SELECT * FROM users;";
1498
1499 let tree = parse_sql(sql);
1500 let mut staging = StagingGraph::new();
1501 let builder = SqlGraphBuilder::new();
1502 let file = PathBuf::from("query.sql");
1503
1504 builder
1505 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1506 .expect("Graph building should succeed");
1507
1508 let read_count = count_table_read_edges(&staging);
1511 assert_eq!(
1512 read_count, 0,
1513 "Standalone SELECT should not create edges without enclosing function"
1514 );
1515 }
1516
1517 #[test]
1518 fn test_export_edges_for_table_definitions() {
1519 let sql = r"
1520 CREATE TABLE users (
1521 id SERIAL PRIMARY KEY,
1522 name TEXT NOT NULL
1523 );
1524
1525 CREATE TABLE orders (
1526 id SERIAL PRIMARY KEY,
1527 user_id INTEGER REFERENCES users(id)
1528 );
1529 ";
1530
1531 let tree = parse_sql(sql);
1532 let mut staging = StagingGraph::new();
1533 let builder = SqlGraphBuilder::new();
1534 let file = PathBuf::from("schema.sql");
1535
1536 builder
1537 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1538 .expect("Graph building should succeed");
1539
1540 let export_count = count_export_edges(&staging);
1541 assert_eq!(
1542 export_count, 2,
1543 "Expected 2 Export edges (users and orders), got {export_count}"
1544 );
1545 }
1546
1547 #[test]
1548 fn test_export_edges_for_view_definitions() {
1549 let sql = r"
1550 CREATE TABLE users (id INT, created_at TIMESTAMP);
1551
1552 CREATE VIEW active_users AS
1553 SELECT * FROM users WHERE created_at > NOW() - INTERVAL '30 days';
1554
1555 CREATE MATERIALIZED VIEW user_stats AS
1556 SELECT COUNT(*) as total FROM users;
1557 ";
1558
1559 let tree = parse_sql(sql);
1560 let mut staging = StagingGraph::new();
1561 let builder = SqlGraphBuilder::new();
1562 let file = PathBuf::from("views.sql");
1563
1564 builder
1565 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1566 .expect("Graph building should succeed");
1567
1568 let export_count = count_export_edges(&staging);
1569 assert_eq!(
1571 export_count, 3,
1572 "Expected 3 Export edges (1 table + 2 views), got {export_count}"
1573 );
1574 }
1575
1576 #[test]
1577 fn test_export_edges_for_functions_and_triggers() {
1578 let sql = r"
1579 CREATE FUNCTION get_balance(account_id INT) RETURNS BIGINT AS $$
1580 BEGIN
1581 RETURN 42;
1582 END;
1583 $$ LANGUAGE plpgsql;
1584
1585 CREATE FUNCTION update_balance() RETURNS TRIGGER AS $$
1586 BEGIN
1587 RETURN NEW;
1588 END;
1589 $$ LANGUAGE plpgsql;
1590
1591 CREATE TRIGGER balance_updated
1592 BEFORE INSERT ON accounts
1593 FOR EACH ROW
1594 EXECUTE FUNCTION update_balance();
1595 ";
1596
1597 let tree = parse_sql(sql);
1598 let mut staging = StagingGraph::new();
1599 let builder = SqlGraphBuilder::new();
1600 let file = PathBuf::from("banking.sql");
1601
1602 builder
1603 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1604 .expect("Graph building should succeed");
1605
1606 let export_count = count_export_edges(&staging);
1607 assert!(
1611 export_count >= 3,
1612 "Expected at least 3 Export edges (2 functions + 1 trigger), got {export_count}"
1613 );
1614 }
1615
1616 #[test]
1617 fn test_export_edges_with_schema_qualified_names() {
1618 let sql = r"
1619 CREATE TABLE public.customers (
1620 id SERIAL PRIMARY KEY,
1621 name TEXT NOT NULL
1622 );
1623
1624 CREATE FUNCTION public.get_customer_name(cust_id INT) RETURNS TEXT AS $$
1625 BEGIN
1626 RETURN 'test';
1627 END;
1628 $$ LANGUAGE plpgsql;
1629 ";
1630
1631 let tree = parse_sql(sql);
1632 let mut staging = StagingGraph::new();
1633 let builder = SqlGraphBuilder::new();
1634 let file = PathBuf::from("public_schema.sql");
1635
1636 builder
1637 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1638 .expect("Graph building should succeed");
1639
1640 let export_count = count_export_edges(&staging);
1641 assert_eq!(
1643 export_count, 2,
1644 "Expected 2 Export edges (table + function), got {export_count}"
1645 );
1646 }
1647
1648 #[test]
1649 fn test_mixed_database_objects_exports() {
1650 let sql = r"
1651 CREATE TABLE accounts (
1652 id SERIAL PRIMARY KEY,
1653 balance_cents BIGINT NOT NULL
1654 );
1655
1656 CREATE VIEW positive_balances AS
1657 SELECT * FROM accounts WHERE balance_cents > 0;
1658
1659 CREATE FUNCTION get_balance(account_id INT) RETURNS BIGINT AS $$
1660 BEGIN
1661 RETURN (SELECT balance_cents FROM accounts WHERE id = account_id);
1662 END;
1663 $$ LANGUAGE plpgsql;
1664 ";
1665
1666 let tree = parse_sql(sql);
1667 let mut staging = StagingGraph::new();
1668 let builder = SqlGraphBuilder::new();
1669 let file = PathBuf::from("mixed.sql");
1670
1671 builder
1672 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1673 .expect("Graph building should succeed");
1674
1675 let export_count = count_export_edges(&staging);
1676 assert_eq!(
1678 export_count, 3,
1679 "Expected 3 Export edges (table + view + function), got {export_count}"
1680 );
1681 }
1682
1683 #[test]
1684 fn test_no_exports_for_empty_file() {
1685 let sql = "";
1686 let tree = parse_sql(sql);
1687 let mut staging = StagingGraph::new();
1688 let builder = SqlGraphBuilder::new();
1689 let file = PathBuf::from("empty.sql");
1690
1691 builder
1692 .build_graph(&tree, sql.as_bytes(), &file, &mut staging)
1693 .expect("Graph building should succeed");
1694
1695 let export_count = count_export_edges(&staging);
1696 assert_eq!(
1697 export_count, 0,
1698 "Expected 0 Export edges for empty file, got {export_count}"
1699 );
1700 }
1701}
1702
1703#[cfg(test)]
1704mod shape_tests {
1705 use super::*;
1706 use sqry_core::graph::unified::build::shape::{ShapeBudget, compute_shape_descriptor};
1707
1708 const SAMPLE: &str = include_str!(concat!(
1709 env!("CARGO_MANIFEST_DIR"),
1710 "/../test-fixtures/shape/data/sample.sql"
1711 ));
1712
1713 fn parse(src: &str) -> Tree {
1714 let mut parser = tree_sitter::Parser::new();
1715 parser
1716 .set_language(&tree_sitter_sequel::LANGUAGE.into())
1717 .expect("load sql grammar");
1718 parser.parse(src, None).expect("parse")
1719 }
1720
1721 fn first_create_function(node: Node<'_>) -> Option<Node<'_>> {
1723 if node.kind() == "create_function" {
1724 return Some(node);
1725 }
1726 let mut cursor = node.walk();
1727 for child in node.children(&mut cursor) {
1728 if let Some(found) = first_create_function(child) {
1729 return Some(found);
1730 }
1731 }
1732 None
1733 }
1734
1735 #[test]
1736 fn cf_map_is_non_empty_and_covers_real_kinds() {
1737 let mapping = sql_shape_mapping();
1738 let populated = mapping.cf_by_kind_id.iter().filter(|s| s.is_some()).count();
1739 assert!(
1740 populated > 0,
1741 "SQL cf map must map at least one real grammar kind"
1742 );
1743
1744 let lang: tree_sitter::Language = tree_sitter_sequel::LANGUAGE.into();
1747 let case_id = lang.id_for_node_kind("case", true);
1748 let when_id = lang.id_for_node_kind("when_clause", true);
1749 assert_eq!(mapping.cf_bucket(case_id), Some(CfBucket::Match));
1750 assert_eq!(mapping.cf_bucket(when_id), Some(CfBucket::Branch));
1751 }
1752
1753 #[test]
1754 fn descriptor_counts_case_control_flow_in_sql_body() {
1755 let tree = parse(SAMPLE);
1756 let func = first_create_function(tree.root_node()).expect("create_function in fixture");
1757 let descriptor = compute_shape_descriptor(
1758 func,
1759 SAMPLE.as_bytes(),
1760 sql_shape_mapping(),
1761 &ShapeBudget::default(),
1762 );
1763 assert!(
1764 !descriptor.is_unhashable(),
1765 "a function with a parsed CASE body must be hashable"
1766 );
1767 assert!(
1768 descriptor.cf_histogram[CfBucket::Match.index()] >= 1,
1769 "the CASE expression must be counted in the Match bucket"
1770 );
1771 }
1772
1773 #[test]
1774 fn signature_shape_reads_arguments_and_defaults() {
1775 let tree = parse(SAMPLE);
1776 let func = first_create_function(tree.root_node()).expect("create_function in fixture");
1777 let shape = sql_shape_mapping().signature_shape(func, SAMPLE.as_bytes());
1778 assert_eq!(
1779 shape.arity_positional, 2,
1780 "grade(score, bonus) has two arguments"
1781 );
1782 assert!(
1783 shape.has_defaults,
1784 "the DEFAULT 0 argument must set has_defaults"
1785 );
1786 }
1787}