1use sqlparser::parser::Parser;
2
3use crate::dialect::SqlDialect;
4use crate::errors::ScytheError;
5
6#[derive(Debug, Clone, Default, PartialEq, Eq)]
7#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
8pub enum QueryCommand {
9 One,
10 Opt,
11 Many,
12 #[default]
13 Exec,
14 ExecResult,
15 ExecRows,
16 Batch,
17 Grouped,
18}
19
20impl std::fmt::Display for QueryCommand {
21 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
22 match self {
23 QueryCommand::One => write!(f, "one"),
24 QueryCommand::Opt => write!(f, "opt"),
25 QueryCommand::Many => write!(f, "many"),
26 QueryCommand::Exec => write!(f, "exec"),
27 QueryCommand::ExecResult => write!(f, "exec_result"),
28 QueryCommand::ExecRows => write!(f, "exec_rows"),
29 QueryCommand::Batch => write!(f, "batch"),
30 QueryCommand::Grouped => write!(f, "grouped"),
31 }
32 }
33}
34
35impl QueryCommand {
36 fn from_str(s: &str) -> Result<Self, ScytheError> {
37 match s {
38 "one" => Ok(QueryCommand::One),
39 "opt" => Ok(QueryCommand::Opt),
40 "many" => Ok(QueryCommand::Many),
41 "exec" => Ok(QueryCommand::Exec),
42 "exec_result" => Ok(QueryCommand::ExecResult),
43 "exec_rows" => Ok(QueryCommand::ExecRows),
44 "batch" => Ok(QueryCommand::Batch),
45 "grouped" => Ok(QueryCommand::Grouped),
46 other => Err(ScytheError::invalid_annotation(format!(
47 "invalid @returns value: {other}"
48 ))),
49 }
50 }
51}
52
53#[derive(Debug, Clone, PartialEq, Eq)]
54#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
55pub struct ParamDoc {
56 pub name: String,
57 pub description: String,
58}
59
60#[derive(Debug, Clone, PartialEq, Eq)]
68#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
69pub struct PositionalParamDoc {
70 pub position: i64,
72 pub name: String,
74 pub description: String,
76}
77
78#[derive(Debug, Clone, PartialEq, Eq)]
79#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
80pub struct JsonMapping {
81 pub column: String,
82 pub rust_type: String,
83}
84
85#[derive(Debug, Clone, PartialEq, Eq)]
93#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
94pub struct CustomAnnotation {
95 pub name: String,
97 pub value: String,
99 pub line: usize,
101}
102
103#[derive(Debug, Clone, Default, PartialEq, Eq)]
104#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
105pub struct Annotations {
106 pub name: String,
107 pub command: QueryCommand,
108 pub param_docs: Vec<ParamDoc>,
109 pub nullable_overrides: Vec<String>,
110 pub nonnull_overrides: Vec<String>,
111 pub json_mappings: Vec<JsonMapping>,
112 pub deprecated: Option<String>,
113 pub optional_params: Vec<String>,
114 pub group_by: Option<String>,
115 #[cfg_attr(feature = "serde", serde(default))]
118 pub positional_param_docs: Vec<PositionalParamDoc>,
119 pub custom: Vec<CustomAnnotation>,
122}
123
124#[derive(Debug)]
125pub struct Query {
126 pub name: String,
127 pub command: QueryCommand,
128 pub sql: String,
129 pub stmt: sqlparser::ast::Statement,
130 pub annotations: Annotations,
131}
132
133pub fn parse_query(query_sql: &str) -> Result<Query, ScytheError> {
135 parse_query_with_dialect(query_sql, &SqlDialect::PostgreSQL)
136}
137
138pub fn parse_query_with_dialect(query_sql: &str, dialect: &SqlDialect) -> Result<Query, ScytheError> {
140 let mut name: Option<String> = None;
141 let mut command: Option<QueryCommand> = None;
142 let mut param_docs = Vec::new();
143 let mut positional_param_docs: Vec<PositionalParamDoc> = Vec::new();
144 let mut nullable_overrides = Vec::new();
145 let mut nonnull_overrides = Vec::new();
146 let mut json_mappings = Vec::new();
147 let mut deprecated: Option<String> = None;
148 let mut optional_params = Vec::new();
149 let mut group_by: Option<String> = None;
150 let mut custom: Vec<CustomAnnotation> = Vec::new();
151
152 let mut sql_lines = Vec::new();
153
154 for (line_idx, line) in query_sql.lines().enumerate() {
155 let line_no = line_idx + 1;
156 let trimmed = line.trim();
157
158 let annotation_body = if let Some(rest) = trimmed.strip_prefix("--") {
160 let rest = rest.trim_start();
161 rest.strip_prefix('@')
162 } else {
163 None
164 };
165
166 if let Some(body) = annotation_body {
167 let (keyword, value) = match body.find(|c: char| c.is_whitespace()) {
169 Some(pos) => (&body[..pos], body[pos..].trim()),
170 None => (body, ""),
171 };
172
173 match keyword.to_ascii_lowercase().as_str() {
174 "name" => {
175 name = Some(value.to_string());
176 }
177 "returns" => {
178 let cmd_str = value.strip_prefix(':').unwrap_or(value);
179 command = Some(QueryCommand::from_str(cmd_str)?);
180 }
181 "param" => {
182 let first_end = value.find(|c: char| c.is_whitespace());
191 let first_token = first_end.map(|p| &value[..p]).unwrap_or(value);
192
193 if let Some(digits) = first_token.strip_prefix('$')
194 && let Ok(pos) = digits.parse::<i64>()
195 && pos > 0
196 {
197 let rest = first_end.map(|p| value[p..].trim()).unwrap_or("").trim();
199 if !rest.is_empty() {
200 let (param_name, description) = if let Some(colon_pos) = rest.find(':') {
201 (
202 rest[..colon_pos].trim().to_string(),
203 rest[colon_pos + 1..].trim().to_string(),
204 )
205 } else {
206 (rest.to_string(), String::new())
207 };
208 if !param_name.is_empty() {
209 positional_param_docs.push(PositionalParamDoc {
210 position: pos,
211 name: param_name,
212 description,
213 });
214 }
215 }
216 } else {
217 if let Some(colon_pos) = value.find(':') {
219 let param_name = value[..colon_pos].trim().to_string();
220 let description = value[colon_pos + 1..].trim().to_string();
221 param_docs.push(ParamDoc {
222 name: param_name,
223 description,
224 });
225 } else {
226 param_docs.push(ParamDoc {
227 name: value.to_string(),
228 description: String::new(),
229 });
230 }
231 }
232 }
233 "nullable" => {
234 for col in value.split(',') {
235 let col = col.trim();
236 if !col.is_empty() {
237 nullable_overrides.push(col.to_string());
238 }
239 }
240 }
241 "nonnull" => {
242 for col in value.split(',') {
243 let col = col.trim();
244 if !col.is_empty() {
245 nonnull_overrides.push(col.to_string());
246 }
247 }
248 }
249 "json" => {
250 if let Some(eq_pos) = value.find('=') {
252 let column = value[..eq_pos].trim().to_string();
253 let rust_type = value[eq_pos + 1..].trim().to_string();
254 json_mappings.push(JsonMapping { column, rust_type });
255 }
256 }
257 "deprecated" => {
258 deprecated = Some(value.to_string());
259 }
260 "group_by" => {
261 group_by = Some(value.to_string());
262 }
263 "optional" => {
264 for param in value.split(',') {
265 let param = param.trim();
266 if !param.is_empty() {
267 optional_params.push(param.to_string());
268 }
269 }
270 }
271 other => {
272 custom.push(CustomAnnotation {
274 name: other.to_string(),
275 value: value.to_string(),
276 line: line_no,
277 });
278 }
279 }
280 } else {
281 sql_lines.push(line);
282 }
283 }
284
285 let name = name.ok_or_else(|| ScytheError::missing_annotation("name"))?;
286 let command = command.ok_or_else(|| ScytheError::missing_annotation("returns"))?;
287
288 if command == QueryCommand::Grouped && group_by.is_none() {
289 return Err(ScytheError::invalid_annotation(
290 "@returns :grouped requires a @group_by annotation (e.g. @group_by users.id)",
291 ));
292 }
293
294 let sql = sql_lines.join("\n").trim().to_string();
295
296 if sql.is_empty() {
297 return Err(ScytheError::syntax("empty SQL body"));
298 }
299
300 let (sql, parse_sql) = if *dialect == SqlDialect::Oracle {
308 let processed = preprocess_oracle_sql(&sql);
309 (processed.clone(), processed)
310 } else if *dialect == SqlDialect::MsSql {
311 let codegen_sql = convert_mssql_placeholders(&sql);
313 let parse_sql = preprocess_mssql_sql(&sql);
315 (codegen_sql, parse_sql)
316 } else if *dialect == SqlDialect::PostgreSQL {
317 let parse_sql = preprocess_postgres_sql(&sql);
318 (sql.clone(), parse_sql)
319 } else {
320 (sql.clone(), sql)
321 };
322
323 let parser_dialect = dialect.to_sqlparser_dialect();
324 let statements = Parser::parse_sql(parser_dialect.as_ref(), &parse_sql)
325 .map_err(|e| ScytheError::syntax(format!("syntax error: {}", e)))?;
326
327 if statements.len() != 1 {
328 let non_empty: Vec<_> = statements
331 .into_iter()
332 .filter(|s| !matches!(s, sqlparser::ast::Statement::Flush { .. }) && format!("{s}") != "")
333 .collect();
334 if non_empty.len() != 1 {
335 return Err(ScytheError::syntax("expected exactly one SQL statement"));
336 }
337 let stmt = non_empty.into_iter().next().expect("filtered to exactly one statement");
338 let annotations = Annotations {
339 name: name.clone(),
340 command: command.clone(),
341 param_docs,
342 positional_param_docs: positional_param_docs.clone(),
343 nullable_overrides,
344 nonnull_overrides,
345 json_mappings,
346 deprecated,
347 optional_params,
348 group_by: group_by.clone(),
349 custom,
350 };
351 return Ok(Query {
352 name,
353 command,
354 sql,
355 stmt,
356 annotations,
357 });
358 }
359
360 let stmt = statements
361 .into_iter()
362 .next()
363 .expect("filtered to exactly one statement");
364
365 let annotations = Annotations {
366 name: name.clone(),
367 command: command.clone(),
368 param_docs,
369 positional_param_docs,
370 nullable_overrides,
371 nonnull_overrides,
372 json_mappings,
373 deprecated,
374 optional_params,
375 group_by,
376 custom,
377 };
378
379 Ok(Query {
380 name,
381 command,
382 sql,
383 stmt,
384 annotations,
385 })
386}
387
388fn preprocess_postgres_sql(sql: &str) -> String {
394 let mask = mask_postgres_for_scan(sql);
398 let mask_bytes = mask.as_bytes();
399 let bytes = sql.as_bytes();
400 let mut search_from = 0;
401 let mut result = String::with_capacity(sql.len());
402 let mut last = 0;
403 while let Some(rel) = find_keyword(&mask[search_from..], "ON CONFLICT") {
404 let on_conflict_pos = search_from + rel;
405 let after_on_conflict = on_conflict_pos + "ON CONFLICT".len();
406 let mut idx = after_on_conflict;
407 while idx < mask_bytes.len() && mask_bytes[idx].is_ascii_whitespace() {
408 idx += 1;
409 }
410 if idx >= mask_bytes.len() || mask_bytes[idx] != b'(' {
411 search_from = after_on_conflict;
412 continue;
413 }
414 let mut depth = 0i32;
415 let mut close = idx;
416 while close < mask_bytes.len() {
417 match mask_bytes[close] {
418 b'(' => depth += 1,
419 b')' => {
420 depth -= 1;
421 if depth == 0 {
422 break;
423 }
424 }
425 _ => {}
426 }
427 close += 1;
428 }
429 if depth != 0 {
430 return sql.to_string();
431 }
432 let mut after_cols = close + 1;
433 while after_cols < mask_bytes.len() && mask_bytes[after_cols].is_ascii_whitespace() {
434 after_cols += 1;
435 }
436 if mask[after_cols..].starts_with("WHERE")
437 && let Some(do_rel) = find_keyword(&mask[after_cols + "WHERE".len()..], "DO")
438 {
439 let do_abs = after_cols + "WHERE".len() + do_rel;
440 result.push_str(std::str::from_utf8(&bytes[last..after_cols]).unwrap_or(""));
443 last = do_abs;
444 search_from = do_abs;
445 continue;
446 }
447 search_from = close + 1;
448 }
449 result.push_str(std::str::from_utf8(&bytes[last..]).unwrap_or(""));
450 result
451}
452
453fn mask_postgres_for_scan(sql: &str) -> String {
458 let bytes = sql.as_bytes();
459 let mut out = vec![b' '; bytes.len()];
460 let mut i = 0;
461 while i < bytes.len() {
462 let b = bytes[i];
463 if b == b'-' && i + 1 < bytes.len() && bytes[i + 1] == b'-' {
464 while i < bytes.len() && bytes[i] != b'\n' {
466 out[i] = b' ';
467 i += 1;
468 }
469 continue;
470 }
471 if b == b'/' && i + 1 < bytes.len() && bytes[i + 1] == b'*' {
472 out[i] = b' ';
474 out[i + 1] = b' ';
475 i += 2;
476 while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') {
477 out[i] = b' ';
478 i += 1;
479 }
480 if i + 1 < bytes.len() {
481 out[i] = b' ';
482 out[i + 1] = b' ';
483 i += 2;
484 }
485 continue;
486 }
487 if b == b'\'' {
488 out[i] = b' ';
489 i += 1;
490 while i < bytes.len() {
491 if bytes[i] == b'\'' {
492 if i + 1 < bytes.len() && bytes[i + 1] == b'\'' {
493 out[i] = b' ';
494 out[i + 1] = b' ';
495 i += 2;
496 continue;
497 }
498 out[i] = b' ';
499 i += 1;
500 break;
501 }
502 out[i] = b' ';
503 i += 1;
504 }
505 continue;
506 }
507 if b.is_ascii() {
510 out[i] = b.to_ascii_uppercase();
511 } else {
512 out[i] = b' ';
513 }
514 i += 1;
515 }
516 String::from_utf8(out).expect("mask is ASCII by construction")
517}
518
519fn find_keyword(haystack: &str, keyword: &str) -> Option<usize> {
522 let bytes = haystack.as_bytes();
523 let key = keyword.as_bytes();
524 let mut i = 0;
525 while i + key.len() <= bytes.len() {
526 if &bytes[i..i + key.len()] == key {
527 let prev_ok = i == 0 || !bytes[i - 1].is_ascii_alphanumeric();
528 let next = i + key.len();
529 let next_ok = next >= bytes.len() || !bytes[next].is_ascii_alphanumeric();
530 if prev_ok && next_ok {
531 return Some(i);
532 }
533 }
534 i += 1;
535 }
536 None
537}
538
539fn preprocess_oracle_sql(sql: &str) -> String {
543 let sql = strip_returning_into(sql);
546
547 let mut result = String::with_capacity(sql.len());
549 let mut chars = sql.chars().peekable();
550 while let Some(ch) = chars.next() {
551 if ch == '\'' {
552 result.push(ch);
554 while let Some(inner) = chars.next() {
555 result.push(inner);
556 if inner == '\'' {
557 if chars.peek() == Some(&'\'') {
558 result.push(chars.next().unwrap());
559 } else {
560 break;
561 }
562 }
563 }
564 } else if ch == ':' && chars.peek().is_some_and(|c| c.is_ascii_digit()) {
565 result.push('?');
567 while chars.peek().is_some_and(|c| c.is_ascii_digit()) {
568 chars.next();
569 }
570 } else {
571 result.push(ch);
572 }
573 }
574 result
575}
576
577fn convert_mssql_placeholders(sql: &str) -> String {
581 let mut result = String::with_capacity(sql.len());
582 let mut chars = sql.chars().peekable();
583 while let Some(ch) = chars.next() {
584 if ch == '\'' {
585 result.push(ch);
587 while let Some(inner) = chars.next() {
588 result.push(inner);
589 if inner == '\'' {
590 if chars.peek() == Some(&'\'') {
591 result.push(chars.next().unwrap());
593 } else {
594 break;
595 }
596 }
597 }
598 } else if ch == '@' && chars.peek().is_some_and(|c| *c == 'p' || *c == 'P') {
599 let mut lookahead = chars.clone();
601 lookahead.next(); if lookahead.peek().is_some_and(|c| c.is_ascii_digit()) {
603 chars.next(); while chars.peek().is_some_and(|c| c.is_ascii_digit()) {
606 chars.next();
607 }
608 result.push('?');
609 } else {
610 result.push(ch);
611 }
612 } else {
613 result.push(ch);
614 }
615 }
616 result
617}
618
619fn preprocess_mssql_sql(sql: &str) -> String {
623 let sql = strip_and_convert_mssql_output(sql);
625 convert_mssql_placeholders(&sql)
627}
628
629fn strip_and_convert_mssql_output(sql: &str) -> String {
636 let upper = sql.to_uppercase();
638
639 if !upper.contains("INSERT") || !upper.contains("OUTPUT") {
641 return sql.to_string();
642 }
643
644 if let Some(output_pos) = find_word_position(&upper, "OUTPUT") {
646 let before_output = &upper[..output_pos];
648 if !before_output.contains("INSERT") {
649 return sql.to_string();
650 }
651
652 let after_output = &upper[output_pos + "OUTPUT".len()..];
654 if let Some(values_offset) = find_word_position(after_output, "VALUES") {
655 let values_pos = output_pos + "OUTPUT".len() + values_offset;
656
657 let output_cols_str = &sql[output_pos + "OUTPUT".len()..values_pos];
659
660 let cols = parse_inserted_columns(output_cols_str);
662
663 if !cols.is_empty() {
664 let before_output_sql = sql[..output_pos].trim_end();
667 let after_values = sql[values_pos..].trim_end();
668 let (values_body, trailing) = if let Some(stripped) = after_values.strip_suffix(';') {
669 (stripped, ";")
670 } else {
671 (after_values, "")
672 };
673
674 return format!("{}\n{} RETURNING {}{}", before_output_sql, values_body, cols, trailing);
675 }
676 }
677 }
678
679 sql.to_string()
680}
681
682fn find_word_position(text: &str, word: &str) -> Option<usize> {
685 let mut pos = 0;
686 let word_len = word.len();
687 while let Some(idx) = text[pos..].find(word) {
688 let abs_idx = pos + idx;
689
690 let before_ok = abs_idx == 0
692 || !text
693 .as_bytes()
694 .get(abs_idx - 1)
695 .is_some_and(|&b| b.is_ascii_alphanumeric() || b == b'_');
696
697 let after_idx = abs_idx + word_len;
699 let after_ok = after_idx >= text.len()
700 || !text
701 .as_bytes()
702 .get(after_idx)
703 .is_some_and(|&b| b.is_ascii_alphanumeric() || b == b'_');
704
705 if before_ok && after_ok {
706 return Some(abs_idx);
707 }
708 pos = abs_idx + 1;
709 }
710 None
711}
712
713fn parse_inserted_columns(output_str: &str) -> String {
715 let mut cols = Vec::new();
716
717 for part in output_str.split(',') {
718 let trimmed = part.trim();
719
720 if let Some(after_inserted) = trimmed
722 .strip_prefix("INSERTED.")
723 .or_else(|| trimmed.strip_prefix("inserted."))
724 .or_else(|| trimmed.strip_prefix("INSERTED"))
725 .or_else(|| trimmed.strip_prefix("inserted"))
726 {
727 let col_name = after_inserted.trim().to_string();
728 if !col_name.is_empty() {
729 cols.push(col_name);
730 }
731 }
732 }
733
734 cols.join(", ")
735}
736
737fn strip_returning_into(sql: &str) -> String {
739 let upper = sql.to_uppercase();
741 if let Some(ret_pos) = upper.rfind("RETURNING") {
742 let after_returning = &upper[ret_pos + "RETURNING".len()..];
743 if let Some(into_offset) = after_returning.find("INTO") {
744 let into_pos = ret_pos + "RETURNING".len() + into_offset;
745 let trimmed = sql[..into_pos].trim_end();
747 return trimmed.to_string();
748 }
749 }
750 sql.to_string()
751}
752
753#[cfg(test)]
754mod tests {
755 use super::*;
756 use crate::errors::ErrorCode;
757
758 fn parse(sql: &str) -> Result<Query, ScytheError> {
759 parse_query(sql)
760 }
761
762 #[test]
763 fn test_basic_parse() {
764 let input = "-- @name GetUsers\n-- @returns :many\nSELECT * FROM users;";
765 let q = parse(input).unwrap();
766 assert_eq!(q.name, "GetUsers");
767 assert_eq!(q.command, QueryCommand::Many);
768 assert!(q.sql.contains("SELECT"));
769 }
770
771 #[test]
772 fn test_all_command_types() {
773 let cases = vec![
774 (":one", QueryCommand::One),
775 (":many", QueryCommand::Many),
776 (":exec", QueryCommand::Exec),
777 (":exec_result", QueryCommand::ExecResult),
778 (":exec_rows", QueryCommand::ExecRows),
779 ];
780 for (tag, expected) in cases {
781 let input = format!("-- @name Q\n-- @returns {}\nSELECT 1", tag);
782 let q = parse(&input).unwrap();
783 assert_eq!(q.command, expected, "failed for {}", tag);
784 }
785 }
786
787 #[test]
788 fn test_case_insensitive_keywords() {
789 let input = "-- @Name GetUsers\n-- @RETURNS :many\nSELECT 1";
790 let q = parse(input).unwrap();
791 assert_eq!(q.name, "GetUsers");
792 assert_eq!(q.command, QueryCommand::Many);
793 }
794
795 #[test]
796 fn test_missing_name_errors() {
797 let input = "-- @returns :many\nSELECT 1";
798 let err = parse(input).unwrap_err();
799 assert_eq!(err.code, ErrorCode::MissingAnnotation);
800 assert!(err.message.contains("name"));
801 }
802
803 #[test]
804 fn test_missing_returns_errors() {
805 let input = "-- @name Foo\nSELECT 1";
806 let err = parse(input).unwrap_err();
807 assert_eq!(err.code, ErrorCode::MissingAnnotation);
808 assert!(err.message.contains("returns"));
809 }
810
811 #[test]
812 fn test_invalid_returns_value() {
813 let input = "-- @name Foo\n-- @returns :invalid\nSELECT 1";
814 let err = parse(input).unwrap_err();
815 assert_eq!(err.code, ErrorCode::InvalidAnnotation);
816 }
817
818 #[test]
819 fn test_empty_name_value() {
820 let input = "-- @name\n-- @returns :one\nSELECT 1";
822 let q = parse(input).unwrap();
823 assert_eq!(q.name, "");
824 }
825
826 #[test]
827 fn test_param_annotation() {
828 let input = "-- @name Foo\n-- @returns :one\n-- @param id: the user ID\nSELECT 1";
829 let q = parse(input).unwrap();
830 assert_eq!(q.annotations.param_docs.len(), 1);
831 assert_eq!(q.annotations.param_docs[0].name, "id");
832 assert_eq!(q.annotations.param_docs[0].description, "the user ID");
833 }
834
835 #[test]
836 fn test_param_no_description() {
837 let input = "-- @name Foo\n-- @returns :one\n-- @param id\nSELECT 1";
838 let q = parse(input).unwrap();
839 assert_eq!(q.annotations.param_docs.len(), 1);
840 assert_eq!(q.annotations.param_docs[0].name, "id");
841 assert_eq!(q.annotations.param_docs[0].description, "");
842 }
843
844 #[test]
845 fn test_nullable_annotation() {
846 let input = "-- @name Foo\n-- @returns :one\n-- @nullable col1, col2\nSELECT 1";
847 let q = parse(input).unwrap();
848 assert_eq!(q.annotations.nullable_overrides, vec!["col1", "col2"]);
849 }
850
851 #[test]
852 fn test_nonnull_annotation() {
853 let input = "-- @name Foo\n-- @returns :one\n-- @nonnull col1\nSELECT 1";
854 let q = parse(input).unwrap();
855 assert_eq!(q.annotations.nonnull_overrides, vec!["col1"]);
856 }
857
858 #[test]
859 fn test_json_annotation() {
860 let input = "-- @name Foo\n-- @returns :one\n-- @json data = EventData\nSELECT 1";
861 let q = parse(input).unwrap();
862 assert_eq!(q.annotations.json_mappings.len(), 1);
863 assert_eq!(q.annotations.json_mappings[0].column, "data");
864 assert_eq!(q.annotations.json_mappings[0].rust_type, "EventData");
865 }
866
867 #[test]
868 fn test_custom_annotations_captured() {
869 let input = "-- @name GetUser
872-- @returns :one
873-- @http GET /users/{id}
874-- @http_auth bearer:jwt
875-- @http_status 200,404
876SELECT id FROM users WHERE id = $1";
877 let q = parse(input).unwrap();
878 assert_eq!(q.annotations.custom.len(), 3);
879 assert_eq!(q.annotations.custom[0].name, "http");
880 assert_eq!(q.annotations.custom[0].value, "GET /users/{id}");
881 assert_eq!(q.annotations.custom[0].line, 3);
882 assert_eq!(q.annotations.custom[1].name, "http_auth");
883 assert_eq!(q.annotations.custom[1].value, "bearer:jwt");
884 assert_eq!(q.annotations.custom[1].line, 4);
885 assert_eq!(q.annotations.custom[2].name, "http_status");
886 assert_eq!(q.annotations.custom[2].value, "200,404");
887 assert_eq!(q.annotations.custom[2].line, 5);
888 }
889
890 #[test]
891 fn test_custom_annotation_without_value() {
892 let input = "-- @name GetUser
893-- @returns :one
894-- @http_internal
895SELECT 1";
896 let q = parse(input).unwrap();
897 assert_eq!(q.annotations.custom.len(), 1);
898 assert_eq!(q.annotations.custom[0].name, "http_internal");
899 assert_eq!(q.annotations.custom[0].value, "");
900 }
901
902 #[cfg(feature = "serde")]
903 #[test]
904 fn test_custom_annotation_serde_round_trip() {
905 let original = CustomAnnotation {
906 name: "http".to_string(),
907 value: "GET /users/{id}".to_string(),
908 line: 7,
909 };
910 let json = serde_json::to_string(&original).unwrap();
911 let back: CustomAnnotation = serde_json::from_str(&json).unwrap();
912 assert_eq!(back, original);
913 }
914
915 #[test]
916 fn test_custom_annotation_name_lowercased() {
917 let input = "-- @name GetUser
918-- @returns :one
919-- @HTTP_Auth Bearer
920SELECT 1";
921 let q = parse(input).unwrap();
922 assert_eq!(q.annotations.custom.len(), 1);
923 assert_eq!(q.annotations.custom[0].name, "http_auth");
924 assert_eq!(q.annotations.custom[0].value, "Bearer");
925 }
926
927 #[test]
930 fn test_positional_param_basic() {
931 let input = "-- @name Foo\n-- @returns :one\n-- @param $1 user_id\nSELECT 1";
932 let q = parse(input).unwrap();
933 assert_eq!(q.annotations.positional_param_docs.len(), 1);
934 assert_eq!(q.annotations.positional_param_docs[0].position, 1);
935 assert_eq!(q.annotations.positional_param_docs[0].name, "user_id");
936 assert_eq!(q.annotations.positional_param_docs[0].description, "");
937 assert_eq!(q.annotations.param_docs.len(), 0);
939 }
940
941 #[test]
942 fn test_positional_param_with_description() {
943 let input = "-- @name Foo\n-- @returns :one\n-- @param $4 bucket: time bucket as text\nSELECT 1";
944 let q = parse(input).unwrap();
945 assert_eq!(q.annotations.positional_param_docs.len(), 1);
946 assert_eq!(q.annotations.positional_param_docs[0].position, 4);
947 assert_eq!(q.annotations.positional_param_docs[0].name, "bucket");
948 assert_eq!(
949 q.annotations.positional_param_docs[0].description,
950 "time bucket as text"
951 );
952 }
953
954 #[test]
955 fn test_positional_param_does_not_affect_docs_only_param() {
956 let input = "-- @name Foo\n-- @returns :one\n-- @param id: the user ID\n-- @param $2 name\nSELECT 1";
958 let q = parse(input).unwrap();
959 assert_eq!(q.annotations.param_docs.len(), 1);
960 assert_eq!(q.annotations.param_docs[0].name, "id");
961 assert_eq!(q.annotations.positional_param_docs.len(), 1);
962 assert_eq!(q.annotations.positional_param_docs[0].position, 2);
963 assert_eq!(q.annotations.positional_param_docs[0].name, "name");
964 }
965
966 #[test]
967 fn test_positional_param_multiple() {
968 let input =
969 "-- @name Foo\n-- @returns :one\n-- @param $1 start_date: start\n-- @param $2 end_date: end\nSELECT 1";
970 let q = parse(input).unwrap();
971 assert_eq!(q.annotations.positional_param_docs.len(), 2);
972 assert_eq!(q.annotations.positional_param_docs[0].position, 1);
973 assert_eq!(q.annotations.positional_param_docs[0].name, "start_date");
974 assert_eq!(q.annotations.positional_param_docs[1].position, 2);
975 assert_eq!(q.annotations.positional_param_docs[1].name, "end_date");
976 }
977
978 #[test]
979 fn test_deprecated_annotation() {
980 let input = "-- @name Foo\n-- @returns :one\n-- @deprecated Use V2\nSELECT 1";
981 let q = parse(input).unwrap();
982 assert_eq!(q.annotations.deprecated, Some("Use V2".to_string()));
983 }
984
985 #[test]
986 fn test_sql_syntax_error() {
987 let input = "-- @name Foo\n-- @returns :one\nSELCT * FROM users";
988 let err = parse(input).unwrap_err();
989 assert_eq!(err.code, ErrorCode::SyntaxError);
990 }
991
992 #[test]
993 fn test_trailing_semicolon() {
994 let input = "-- @name Foo\n-- @returns :one\nSELECT 1;";
995 let q = parse(input).unwrap();
996 assert_eq!(q.name, "Foo");
997 }
998
999 #[test]
1000 fn test_multiple_statements_error() {
1001 let input = "-- @name Foo\n-- @returns :one\nSELECT 1; SELECT 2;";
1002 let err = parse(input).unwrap_err();
1003 assert_eq!(err.code, ErrorCode::SyntaxError);
1004 }
1005
1006 #[test]
1007 fn test_sql_preserved_without_annotations() {
1008 let input = "-- @name Foo\n-- @returns :one\nSELECT id, name FROM users WHERE id = $1";
1009 let q = parse(input).unwrap();
1010 assert_eq!(q.sql, "SELECT id, name FROM users WHERE id = $1");
1011 }
1012
1013 #[test]
1014 fn test_returns_without_colon_prefix() {
1015 let input = "-- @name Foo\n-- @returns many\nSELECT 1";
1016 let q = parse(input).unwrap();
1017 assert_eq!(q.command, QueryCommand::Many);
1018 }
1019
1020 #[test]
1021 fn test_batch_command() {
1022 let input = "-- @name Foo\n-- @returns :batch\nSELECT 1";
1023 let q = parse(input).unwrap();
1024 assert_eq!(q.command, QueryCommand::Batch);
1025 }
1026
1027 #[test]
1028 fn test_grouped_command_with_group_by() {
1029 let input = "-- @name GetUsersWithOrders\n-- @returns :grouped\n-- @group_by users.id\nSELECT u.id, u.name FROM users u JOIN orders o ON o.user_id = u.id";
1030 let q = parse(input).unwrap();
1031 assert_eq!(q.command, QueryCommand::Grouped);
1032 assert_eq!(q.annotations.group_by, Some("users.id".to_string()));
1033 }
1034
1035 #[test]
1036 fn test_grouped_command_without_group_by_errors() {
1037 let input = "-- @name Foo\n-- @returns :grouped\nSELECT 1";
1038 let err = parse(input).unwrap_err();
1039 assert_eq!(err.code, ErrorCode::InvalidAnnotation);
1040 assert!(err.message.contains("@group_by"));
1041 }
1042
1043 #[test]
1044 fn test_group_by_without_grouped_is_ignored() {
1045 let input = "-- @name Foo\n-- @returns :many\n-- @group_by users.id\nSELECT 1";
1046 let q = parse(input).unwrap();
1047 assert_eq!(q.command, QueryCommand::Many);
1048 assert_eq!(q.annotations.group_by, Some("users.id".to_string()));
1049 }
1050
1051 #[test]
1052 fn test_preprocess_postgres_strips_partial_index_where() {
1053 let sql = "INSERT INTO billing_events (project_id, stripe_event_id) \
1054 VALUES ($1, $2) \
1055 ON CONFLICT (stripe_event_id) WHERE stripe_event_id IS NOT NULL DO NOTHING";
1056 let cleaned = preprocess_postgres_sql(sql);
1057 assert!(
1058 !cleaned.to_uppercase().contains("WHERE STRIPE_EVENT_ID IS NOT NULL"),
1059 "WHERE clause must be stripped between ON CONFLICT cols and DO; got: {cleaned}"
1060 );
1061 assert!(
1062 cleaned
1063 .to_uppercase()
1064 .contains("ON CONFLICT (STRIPE_EVENT_ID) DO NOTHING")
1065 );
1066 sqlparser::parser::Parser::parse_sql(&sqlparser::dialect::PostgreSqlDialect {}, &cleaned)
1068 .expect("cleaned SQL should parse");
1069 }
1070
1071 #[test]
1072 fn test_preprocess_postgres_no_op_when_no_partial_clause() {
1073 let sql = "INSERT INTO t (a) VALUES ($1) ON CONFLICT (a) DO UPDATE SET a = EXCLUDED.a";
1074 assert_eq!(preprocess_postgres_sql(sql), sql);
1075 }
1076
1077 #[test]
1078 fn test_preprocess_postgres_leaves_on_conflict_on_constraint_alone() {
1079 let sql = "INSERT INTO t (a) VALUES ($1) ON CONFLICT ON CONSTRAINT t_a_uidx DO NOTHING";
1080 assert_eq!(preprocess_postgres_sql(sql), sql);
1081 }
1082
1083 #[test]
1084 fn test_preprocess_postgres_handles_compound_index_cols() {
1085 let sql = "INSERT INTO t (a, b) VALUES ($1, $2) \
1086 ON CONFLICT (a, b) WHERE a IS NOT NULL AND b > 0 DO UPDATE SET b = EXCLUDED.b";
1087 let cleaned = preprocess_postgres_sql(sql);
1088 assert!(cleaned.to_uppercase().contains("ON CONFLICT (A, B) DO UPDATE"));
1089 assert!(!cleaned.to_uppercase().contains("WHERE A IS NOT NULL"));
1090 }
1091
1092 #[test]
1093 fn test_preprocess_postgres_preserves_unrelated_where() {
1094 let sql = "DELETE FROM t WHERE id = $1";
1097 assert_eq!(preprocess_postgres_sql(sql), sql);
1098 }
1099
1100 #[test]
1101 fn test_preprocess_postgres_ignores_text_inside_line_comments() {
1102 let sql = "-- inline doc: `ON CONFLICT (col) WHERE …` is the partial form\n\
1106 INSERT INTO t (a) VALUES ($1) \
1107 ON CONFLICT (a) WHERE a IS NOT NULL DO NOTHING";
1108 let cleaned = preprocess_postgres_sql(sql);
1109 assert!(
1110 cleaned.contains("-- inline doc"),
1111 "comment must survive the pass; got: {cleaned}"
1112 );
1113 assert!(cleaned.contains("ON CONFLICT (a) DO NOTHING"));
1114 }
1115
1116 #[test]
1117 fn test_preprocess_postgres_ignores_text_inside_string_literals() {
1118 let sql = "SELECT 'ON CONFLICT (a) WHERE a IS NOT NULL DO NOTHING' AS s";
1119 assert_eq!(preprocess_postgres_sql(sql), sql);
1120 }
1121
1122 #[test]
1123 fn test_preprocess_oracle_colon_placeholders() {
1124 assert_eq!(
1125 preprocess_oracle_sql("SELECT * FROM users WHERE id = :1"),
1126 "SELECT * FROM users WHERE id = ?"
1127 );
1128 assert_eq!(
1129 preprocess_oracle_sql("INSERT INTO users (name, email) VALUES (:1, :2)"),
1130 "INSERT INTO users (name, email) VALUES (?, ?)"
1131 );
1132 }
1133
1134 #[test]
1135 fn test_preprocess_oracle_preserves_string_literals() {
1136 assert_eq!(
1137 preprocess_oracle_sql("SELECT * FROM users WHERE name = ':1' AND id = :1"),
1138 "SELECT * FROM users WHERE name = ':1' AND id = ?"
1139 );
1140 }
1141
1142 #[test]
1143 fn test_preprocess_oracle_strips_returning_into() {
1144 assert_eq!(
1145 preprocess_oracle_sql("INSERT INTO users (name) VALUES (:1) RETURNING id, name INTO :2, :3"),
1146 "INSERT INTO users (name) VALUES (?) RETURNING id, name"
1147 );
1148 }
1149
1150 #[test]
1151 fn test_preprocess_oracle_full_insert_returning_into() {
1152 let sql = "INSERT INTO users (name, email, active) VALUES (:1, :2, :3) RETURNING id, name, email, active, created_at INTO :4, :5, :6, :7, :8";
1153 let result = preprocess_oracle_sql(sql);
1154 assert_eq!(
1155 result,
1156 "INSERT INTO users (name, email, active) VALUES (?, ?, ?) RETURNING id, name, email, active, created_at"
1157 );
1158 }
1159
1160 #[test]
1161 fn test_preprocess_oracle_no_returning_into_unchanged() {
1162 assert_eq!(
1163 preprocess_oracle_sql("DELETE FROM users WHERE id = :1"),
1164 "DELETE FROM users WHERE id = ?"
1165 );
1166 }
1167
1168 #[test]
1169 fn test_preprocess_mssql_single_placeholder() {
1170 assert_eq!(
1171 preprocess_mssql_sql("SELECT * FROM users WHERE id = @p1"),
1172 "SELECT * FROM users WHERE id = ?"
1173 );
1174 }
1175
1176 #[test]
1177 fn test_preprocess_mssql_multiple_placeholders() {
1178 assert_eq!(
1179 preprocess_mssql_sql("INSERT INTO users (name, email) VALUES (@p1, @p2)"),
1180 "INSERT INTO users (name, email) VALUES (?, ?)"
1181 );
1182 }
1183
1184 #[test]
1185 fn test_preprocess_mssql_preserves_string_literals() {
1186 assert_eq!(
1187 preprocess_mssql_sql("SELECT * FROM users WHERE name = '@p1' AND id = @p1"),
1188 "SELECT * FROM users WHERE name = '@p1' AND id = ?"
1189 );
1190 }
1191
1192 #[test]
1193 fn test_preprocess_mssql_case_insensitive_p() {
1194 assert_eq!(
1195 preprocess_mssql_sql("SELECT * FROM users WHERE id = @P1"),
1196 "SELECT * FROM users WHERE id = ?"
1197 );
1198 }
1199
1200 #[test]
1201 fn test_preprocess_mssql_non_placeholder_at_variable_unchanged() {
1202 assert_eq!(preprocess_mssql_sql("SELECT @myvar"), "SELECT @myvar");
1204 }
1205
1206 #[test]
1207 fn test_preprocess_mssql_multi_digit_placeholder() {
1208 assert_eq!(preprocess_mssql_sql("SELECT @p10, @p2"), "SELECT ?, ?");
1209 }
1210
1211 #[test]
1212 fn test_preprocess_mssql_output_inserted_simple() {
1213 let sql = "INSERT INTO users (id, name) OUTPUT INSERTED.id, INSERTED.name VALUES (@p1, @p2)";
1214 let result = preprocess_mssql_sql(sql);
1215 assert!(result.contains("RETURNING id, name"), "got: {}", result);
1217 assert!(result.contains("VALUES (?, ?)"), "got: {}", result);
1218 assert!(!result.contains("OUTPUT"), "got: {}", result);
1219 }
1220
1221 #[test]
1222 fn test_preprocess_mssql_output_inserted_full_example() {
1223 let sql = "INSERT INTO users (id, name, email, active) OUTPUT INSERTED.id, INSERTED.name, INSERTED.email, INSERTED.active, INSERTED.created_at VALUES (@p1, @p2, @p3, @p4)";
1224 let result = preprocess_mssql_sql(sql);
1225 assert!(
1226 result.contains("RETURNING id, name, email, active, created_at"),
1227 "got: {}",
1228 result
1229 );
1230 assert!(result.contains("VALUES (?, ?, ?, ?)"), "got: {}", result);
1231 }
1232
1233 #[test]
1234 fn test_preprocess_mssql_output_case_insensitive() {
1235 let sql = "INSERT INTO users (id) output inserted.id values (@p1)";
1236 let result = preprocess_mssql_sql(sql);
1237 assert!(result.contains("RETURNING id"), "got: {}", result);
1238 assert!(
1240 result.contains("values (?)") || result.contains("VALUES (?)"),
1241 "got: {}",
1242 result
1243 );
1244 }
1245
1246 #[test]
1247 fn test_preprocess_mssql_no_output_unchanged() {
1248 let sql = "INSERT INTO users (id, name) VALUES (@p1, @p2)";
1249 let result = preprocess_mssql_sql(sql);
1250 assert_eq!(result, "INSERT INTO users (id, name) VALUES (?, ?)");
1251 }
1252
1253 #[test]
1254 fn test_preprocess_mssql_output_with_string_literal() {
1255 let sql = "INSERT INTO users (id, name) OUTPUT INSERTED.id, INSERTED.name VALUES (@p1, '@p2')";
1257 let result = preprocess_mssql_sql(sql);
1258 assert!(result.contains("RETURNING id, name"), "got: {}", result);
1259 assert!(result.contains("(?, '@p2')"), "got: {}", result);
1260 }
1261
1262 #[test]
1263 fn test_preprocess_mssql_output_with_whitespace() {
1264 let sql = "INSERT INTO users (id, name)\nOUTPUT INSERTED.id,\n INSERTED.name\nVALUES (@p1, @p2)";
1265 let result = preprocess_mssql_sql(sql);
1266 assert!(result.contains("RETURNING id, name"), "got: {}", result);
1267 assert!(result.contains("VALUES (?, ?)"), "got: {}", result);
1268 }
1269}