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)]
105#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
106pub struct CustomAnnotation {
107 pub name: String,
109 pub value: String,
111 pub line: usize,
113 #[cfg_attr(feature = "serde", serde(default))]
120 pub suggested_keyword: Option<String>,
121}
122
123const KNOWN_ANNOTATION_KEYWORDS: &[&str] = &[
127 "name",
128 "returns",
129 "param",
130 "nullable",
131 "nonnull",
132 "json",
133 "deprecated",
134 "group_by",
135 "optional",
136];
137
138const ANNOTATION_TYPO_DISTANCE_THRESHOLD: usize = 2;
144
145fn suggest_known_keyword(name: &str) -> Option<&'static str> {
148 KNOWN_ANNOTATION_KEYWORDS
149 .iter()
150 .map(|&candidate| (candidate, annotation_levenshtein_distance(name, candidate)))
151 .filter(|&(_, distance)| distance <= ANNOTATION_TYPO_DISTANCE_THRESHOLD)
152 .min_by_key(|&(_, distance)| distance)
153 .map(|(candidate, _)| candidate)
154}
155
156fn annotation_levenshtein_distance(a: &str, b: &str) -> usize {
158 let a: Vec<char> = a.chars().collect();
159 let b: Vec<char> = b.chars().collect();
160
161 let mut prev_row: Vec<usize> = (0..=b.len()).collect();
162 let mut curr_row = vec![0usize; b.len() + 1];
163
164 for (i, &char_a) in a.iter().enumerate() {
165 curr_row[0] = i + 1;
166 for (j, &char_b) in b.iter().enumerate() {
167 let substitution_cost = usize::from(char_a != char_b);
168 curr_row[j + 1] = (prev_row[j + 1] + 1)
169 .min(curr_row[j] + 1)
170 .min(prev_row[j] + substitution_cost);
171 }
172 std::mem::swap(&mut prev_row, &mut curr_row);
173 }
174
175 prev_row[b.len()]
176}
177
178#[derive(Debug, Clone, Default, PartialEq, Eq)]
179#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
180pub struct Annotations {
181 pub name: String,
182 pub command: QueryCommand,
183 pub param_docs: Vec<ParamDoc>,
184 pub nullable_overrides: Vec<String>,
185 pub nonnull_overrides: Vec<String>,
186 pub json_mappings: Vec<JsonMapping>,
187 pub deprecated: Option<String>,
188 pub optional_params: Vec<String>,
189 pub group_by: Option<String>,
190 #[cfg_attr(feature = "serde", serde(default))]
193 pub positional_param_docs: Vec<PositionalParamDoc>,
194 pub custom: Vec<CustomAnnotation>,
197}
198
199#[derive(Debug)]
200pub struct Query {
201 pub name: String,
202 pub command: QueryCommand,
203 pub sql: String,
204 pub stmt: sqlparser::ast::Statement,
205 pub annotations: Annotations,
206}
207
208fn validate_query_name(name: &str) -> Result<(), ScytheError> {
224 if name.is_empty() {
225 return Err(ScytheError::invalid_annotation(
226 "@name requires a value (e.g. `-- @name GetUser`)",
227 ));
228 }
229
230 let mut chars = name.chars();
231 let first = chars.next().expect("name is non-empty");
232 if !first.is_ascii_alphabetic() && first != '_' {
233 return Err(ScytheError::invalid_annotation(format!(
234 "@name \"{name}\" must start with an ASCII letter or underscore to be a valid \
235 identifier in generated code"
236 )));
237 }
238 if let Some(invalid) = chars.find(|c| !c.is_ascii_alphanumeric() && *c != '_') {
239 return Err(ScytheError::invalid_annotation(format!(
240 "@name \"{name}\" contains '{invalid}'; only ASCII letters, digits and underscores \
241 are valid in generated code"
242 )));
243 }
244
245 Ok(())
246}
247
248pub fn parse_query(query_sql: &str) -> Result<Query, ScytheError> {
250 parse_query_with_dialect(query_sql, &SqlDialect::PostgreSQL)
251}
252
253pub fn parse_query_with_dialect(query_sql: &str, dialect: &SqlDialect) -> Result<Query, ScytheError> {
255 let mut name: Option<String> = None;
256 let mut command: Option<QueryCommand> = None;
257 let mut param_docs = Vec::new();
258 let mut positional_param_docs: Vec<PositionalParamDoc> = Vec::new();
259 let mut nullable_overrides = Vec::new();
260 let mut nonnull_overrides = Vec::new();
261 let mut json_mappings = Vec::new();
262 let mut deprecated: Option<String> = None;
263 let mut optional_params = Vec::new();
264 let mut group_by: Option<String> = None;
265 let mut custom: Vec<CustomAnnotation> = Vec::new();
266
267 let mut sql_lines = Vec::new();
268
269 for (line_idx, line) in query_sql.lines().enumerate() {
270 let line_no = line_idx + 1;
271 let trimmed = line.trim();
272
273 let annotation_body = if let Some(rest) = trimmed.strip_prefix("--") {
274 let rest = rest.trim_start();
275 rest.strip_prefix('@')
276 } else {
277 None
278 };
279
280 if let Some(body) = annotation_body {
281 let (keyword, value) = match body.find(|c: char| c.is_whitespace()) {
282 Some(pos) => (&body[..pos], body[pos..].trim()),
283 None => (body, ""),
284 };
285
286 match keyword.to_ascii_lowercase().as_str() {
287 "name" => {
288 name = Some(value.to_string());
289 }
290 "returns" => {
291 let cmd_str = value.strip_prefix(':').unwrap_or(value);
292 command = Some(QueryCommand::from_str(cmd_str)?);
293 }
294 "param" => {
295 let first_end = value.find(|c: char| c.is_whitespace());
298 let first_token = first_end.map(|p| &value[..p]).unwrap_or(value);
299
300 if let Some(digits) = first_token.strip_prefix('$')
301 && let Ok(pos) = digits.parse::<i64>()
302 && pos > 0
303 {
304 let rest = first_end.map(|p| value[p..].trim()).unwrap_or("").trim();
305 if !rest.is_empty() {
306 let (param_name, description) = if let Some(colon_pos) = rest.find(':') {
307 (
308 rest[..colon_pos].trim().to_string(),
309 rest[colon_pos + 1..].trim().to_string(),
310 )
311 } else {
312 (rest.to_string(), String::new())
313 };
314 if !param_name.is_empty() {
315 positional_param_docs.push(PositionalParamDoc {
316 position: pos,
317 name: param_name,
318 description,
319 });
320 }
321 }
322 } else {
323 if let Some(colon_pos) = value.find(':') {
324 let param_name = value[..colon_pos].trim().to_string();
325 let description = value[colon_pos + 1..].trim().to_string();
326 param_docs.push(ParamDoc {
327 name: param_name,
328 description,
329 });
330 } else {
331 param_docs.push(ParamDoc {
332 name: value.to_string(),
333 description: String::new(),
334 });
335 }
336 }
337 }
338 "nullable" => {
339 for col in value.split(',') {
340 let col = col.trim();
341 if !col.is_empty() {
342 nullable_overrides.push(col.to_string());
343 }
344 }
345 }
346 "nonnull" => {
347 for col in value.split(',') {
348 let col = col.trim();
349 if !col.is_empty() {
350 nonnull_overrides.push(col.to_string());
351 }
352 }
353 }
354 "json" => {
355 if let Some(eq_pos) = value.find('=') {
356 let column = value[..eq_pos].trim().to_string();
357 let rust_type = value[eq_pos + 1..].trim().to_string();
358 json_mappings.push(JsonMapping { column, rust_type });
359 }
360 }
361 "deprecated" => {
362 deprecated = Some(value.to_string());
363 }
364 "group_by" => {
365 group_by = Some(value.to_string());
366 }
367 "optional" => {
368 for param in value.split(',') {
369 let param = param.trim();
370 if !param.is_empty() {
371 optional_params.push(param.to_string());
372 }
373 }
374 }
375 other => {
376 custom.push(CustomAnnotation {
377 name: other.to_string(),
378 value: value.to_string(),
379 line: line_no,
380 suggested_keyword: suggest_known_keyword(other).map(str::to_string),
381 });
382 }
383 }
384 } else {
385 sql_lines.push(line);
386 }
387 }
388
389 let name = name.ok_or_else(|| ScytheError::missing_annotation("name"))?;
390 validate_query_name(&name)?;
391 let command = command.ok_or_else(|| ScytheError::missing_annotation("returns"))?;
392
393 if command == QueryCommand::Grouped && group_by.is_none() {
394 return Err(ScytheError::invalid_annotation(
395 "@returns :grouped requires a @group_by annotation (e.g. @group_by users.id)",
396 ));
397 }
398
399 let sql = sql_lines.join("\n").trim().to_string();
400
401 if sql.is_empty() {
402 return Err(ScytheError::syntax("empty SQL body"));
403 }
404
405 let (sql, parse_sql) = if *dialect == SqlDialect::Oracle {
406 let processed = preprocess_oracle_sql(&sql);
407 (processed.clone(), processed)
408 } else if *dialect == SqlDialect::MsSql {
409 let codegen_sql = convert_mssql_placeholders(&sql);
410 let parse_sql = preprocess_mssql_sql(&sql);
411 (codegen_sql, parse_sql)
412 } else if *dialect == SqlDialect::PostgreSQL {
413 let parse_sql = preprocess_postgres_sql(&sql);
414 (sql.clone(), parse_sql)
415 } else {
416 (sql.clone(), sql)
417 };
418
419 let parser_dialect = dialect.to_sqlparser_dialect();
420 let statements = Parser::parse_sql(parser_dialect.as_ref(), &parse_sql)
421 .map_err(|e| ScytheError::syntax(format!("syntax error: {}", e)))?;
422
423 if statements.len() != 1 {
424 let non_empty: Vec<_> = statements
425 .into_iter()
426 .filter(|s| !matches!(s, sqlparser::ast::Statement::Flush { .. }) && format!("{s}") != "")
427 .collect();
428 if non_empty.len() != 1 {
429 return Err(ScytheError::syntax("expected exactly one SQL statement"));
430 }
431 let stmt = non_empty.into_iter().next().expect("filtered to exactly one statement");
432 let annotations = Annotations {
433 name: name.clone(),
434 command: command.clone(),
435 param_docs,
436 positional_param_docs: positional_param_docs.clone(),
437 nullable_overrides,
438 nonnull_overrides,
439 json_mappings,
440 deprecated,
441 optional_params,
442 group_by: group_by.clone(),
443 custom,
444 };
445 return Ok(Query {
446 name,
447 command,
448 sql,
449 stmt,
450 annotations,
451 });
452 }
453
454 let stmt = statements
455 .into_iter()
456 .next()
457 .expect("filtered to exactly one statement");
458
459 let annotations = Annotations {
460 name: name.clone(),
461 command: command.clone(),
462 param_docs,
463 positional_param_docs,
464 nullable_overrides,
465 nonnull_overrides,
466 json_mappings,
467 deprecated,
468 optional_params,
469 group_by,
470 custom,
471 };
472
473 Ok(Query {
474 name,
475 command,
476 sql,
477 stmt,
478 annotations,
479 })
480}
481
482fn preprocess_postgres_sql(sql: &str) -> String {
488 let mask = mask_postgres_for_scan(sql);
489 let mask_bytes = mask.as_bytes();
490 let bytes = sql.as_bytes();
491 let mut search_from = 0;
492 let mut result = String::with_capacity(sql.len());
493 let mut last = 0;
494 while let Some(rel) = find_keyword(&mask[search_from..], "ON CONFLICT") {
495 let on_conflict_pos = search_from + rel;
496 let after_on_conflict = on_conflict_pos + "ON CONFLICT".len();
497 let mut idx = after_on_conflict;
498 while idx < mask_bytes.len() && mask_bytes[idx].is_ascii_whitespace() {
499 idx += 1;
500 }
501 if idx >= mask_bytes.len() || mask_bytes[idx] != b'(' {
502 search_from = after_on_conflict;
503 continue;
504 }
505 let mut depth = 0i32;
506 let mut close = idx;
507 while close < mask_bytes.len() {
508 match mask_bytes[close] {
509 b'(' => depth += 1,
510 b')' => {
511 depth -= 1;
512 if depth == 0 {
513 break;
514 }
515 }
516 _ => {}
517 }
518 close += 1;
519 }
520 if depth != 0 {
521 return sql.to_string();
522 }
523 let mut after_cols = close + 1;
524 while after_cols < mask_bytes.len() && mask_bytes[after_cols].is_ascii_whitespace() {
525 after_cols += 1;
526 }
527 if mask[after_cols..].starts_with("WHERE")
528 && let Some(do_rel) = find_keyword(&mask[after_cols + "WHERE".len()..], "DO")
529 {
530 let do_abs = after_cols + "WHERE".len() + do_rel;
531 result.push_str(std::str::from_utf8(&bytes[last..after_cols]).unwrap_or(""));
532 last = do_abs;
533 search_from = do_abs;
534 continue;
535 }
536 search_from = close + 1;
537 }
538 result.push_str(std::str::from_utf8(&bytes[last..]).unwrap_or(""));
539 result
540}
541
542fn mask_postgres_for_scan(sql: &str) -> String {
547 let bytes = sql.as_bytes();
548 let mut out = vec![b' '; bytes.len()];
549 let mut i = 0;
550 while i < bytes.len() {
551 let b = bytes[i];
552 if b == b'-' && i + 1 < bytes.len() && bytes[i + 1] == b'-' {
553 while i < bytes.len() && bytes[i] != b'\n' {
554 out[i] = b' ';
555 i += 1;
556 }
557 continue;
558 }
559 if b == b'/' && i + 1 < bytes.len() && bytes[i + 1] == b'*' {
560 out[i] = b' ';
561 out[i + 1] = b' ';
562 i += 2;
563 while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') {
564 out[i] = b' ';
565 i += 1;
566 }
567 if i + 1 < bytes.len() {
568 out[i] = b' ';
569 out[i + 1] = b' ';
570 i += 2;
571 }
572 continue;
573 }
574 if b == b'\'' {
575 out[i] = b' ';
576 i += 1;
577 while i < bytes.len() {
578 if bytes[i] == b'\'' {
579 if i + 1 < bytes.len() && bytes[i + 1] == b'\'' {
580 out[i] = b' ';
581 out[i + 1] = b' ';
582 i += 2;
583 continue;
584 }
585 out[i] = b' ';
586 i += 1;
587 break;
588 }
589 out[i] = b' ';
590 i += 1;
591 }
592 continue;
593 }
594 if b.is_ascii() {
595 out[i] = b.to_ascii_uppercase();
596 } else {
597 out[i] = b' ';
598 }
599 i += 1;
600 }
601 String::from_utf8(out).expect("mask is ASCII by construction")
602}
603
604fn find_keyword(haystack: &str, keyword: &str) -> Option<usize> {
607 let bytes = haystack.as_bytes();
608 let key = keyword.as_bytes();
609 let mut i = 0;
610 while i + key.len() <= bytes.len() {
611 if &bytes[i..i + key.len()] == key {
612 let prev_ok = i == 0 || !bytes[i - 1].is_ascii_alphanumeric();
613 let next = i + key.len();
614 let next_ok = next >= bytes.len() || !bytes[next].is_ascii_alphanumeric();
615 if prev_ok && next_ok {
616 return Some(i);
617 }
618 }
619 i += 1;
620 }
621 None
622}
623
624fn preprocess_oracle_sql(sql: &str) -> String {
644 let sql = strip_returning_into(sql);
645
646 let mut result = String::with_capacity(sql.len());
647 let mut chars = sql.chars().peekable();
648 while let Some(ch) = chars.next() {
649 if ch == '\'' {
650 result.push(ch);
651 while let Some(inner) = chars.next() {
652 result.push(inner);
653 if inner == '\'' {
654 if chars.peek() == Some(&'\'') {
655 result.push(chars.next().unwrap());
656 } else {
657 break;
658 }
659 }
660 }
661 } else if ch == ':' && chars.peek().is_some_and(|c| c.is_ascii_digit()) {
662 result.push('$');
663 while chars.peek().is_some_and(|c| c.is_ascii_digit()) {
664 result.push(chars.next().unwrap());
665 }
666 } else {
667 result.push(ch);
668 }
669 }
670 result
671}
672
673fn convert_mssql_placeholders(sql: &str) -> String {
683 let mut result = String::with_capacity(sql.len());
684 let mut chars = sql.chars().peekable();
685 while let Some(ch) = chars.next() {
686 if ch == '\'' {
687 result.push(ch);
688 while let Some(inner) = chars.next() {
689 result.push(inner);
690 if inner == '\'' {
691 if chars.peek() == Some(&'\'') {
692 result.push(chars.next().unwrap());
693 } else {
694 break;
695 }
696 }
697 }
698 } else if ch == '@' && chars.peek().is_some_and(|c| *c == 'p' || *c == 'P') {
699 let mut lookahead = chars.clone();
700 lookahead.next();
701 if lookahead.peek().is_some_and(|c| c.is_ascii_digit()) {
702 chars.next();
703 result.push('$');
704 while chars.peek().is_some_and(|c| c.is_ascii_digit()) {
705 result.push(chars.next().unwrap());
706 }
707 } else {
708 result.push(ch);
709 }
710 } else {
711 result.push(ch);
712 }
713 }
714 result
715}
716
717fn preprocess_mssql_sql(sql: &str) -> String {
722 let sql = strip_and_convert_mssql_output(sql);
723 convert_mssql_placeholders(&sql)
724}
725
726fn strip_and_convert_mssql_output(sql: &str) -> String {
733 let upper = sql.to_uppercase();
734
735 if !upper.contains("INSERT") || !upper.contains("OUTPUT") {
736 return sql.to_string();
737 }
738
739 if let Some(output_pos) = find_word_position(&upper, "OUTPUT") {
740 let before_output = &upper[..output_pos];
741 if !before_output.contains("INSERT") {
742 return sql.to_string();
743 }
744
745 let after_output = &upper[output_pos + "OUTPUT".len()..];
746 if let Some(values_offset) = find_word_position(after_output, "VALUES") {
747 let values_pos = output_pos + "OUTPUT".len() + values_offset;
748
749 let output_cols_str = &sql[output_pos + "OUTPUT".len()..values_pos];
750
751 let cols = parse_inserted_columns(output_cols_str);
752
753 if !cols.is_empty() {
754 let before_output_sql = sql[..output_pos].trim_end();
755 let after_values = sql[values_pos..].trim_end();
756 let (values_body, trailing) = if let Some(stripped) = after_values.strip_suffix(';') {
757 (stripped, ";")
758 } else {
759 (after_values, "")
760 };
761
762 return format!("{}\n{} RETURNING {}{}", before_output_sql, values_body, cols, trailing);
763 }
764 }
765 }
766
767 sql.to_string()
768}
769
770fn find_word_position(text: &str, word: &str) -> Option<usize> {
773 let mut pos = 0;
774 let word_len = word.len();
775 while let Some(idx) = text[pos..].find(word) {
776 let abs_idx = pos + idx;
777
778 let before_ok = abs_idx == 0
779 || !text
780 .as_bytes()
781 .get(abs_idx - 1)
782 .is_some_and(|&b| b.is_ascii_alphanumeric() || b == b'_');
783
784 let after_idx = abs_idx + word_len;
785 let after_ok = after_idx >= text.len()
786 || !text
787 .as_bytes()
788 .get(after_idx)
789 .is_some_and(|&b| b.is_ascii_alphanumeric() || b == b'_');
790
791 if before_ok && after_ok {
792 return Some(abs_idx);
793 }
794 pos = abs_idx + 1;
795 }
796 None
797}
798
799fn parse_inserted_columns(output_str: &str) -> String {
801 let mut cols = Vec::new();
802
803 for part in output_str.split(',') {
804 let trimmed = part.trim();
805
806 if let Some(after_inserted) = trimmed
807 .strip_prefix("INSERTED.")
808 .or_else(|| trimmed.strip_prefix("inserted."))
809 .or_else(|| trimmed.strip_prefix("INSERTED"))
810 .or_else(|| trimmed.strip_prefix("inserted"))
811 {
812 let col_name = after_inserted.trim().to_string();
813 if !col_name.is_empty() {
814 cols.push(col_name);
815 }
816 }
817 }
818
819 cols.join(", ")
820}
821
822fn strip_returning_into(sql: &str) -> String {
824 let upper = sql.to_uppercase();
825 if let Some(ret_pos) = upper.rfind("RETURNING") {
826 let after_returning = &upper[ret_pos + "RETURNING".len()..];
827 if let Some(into_offset) = after_returning.find("INTO") {
828 let into_pos = ret_pos + "RETURNING".len() + into_offset;
829 let trimmed = sql[..into_pos].trim_end();
830 return trimmed.to_string();
831 }
832 }
833 sql.to_string()
834}
835
836#[cfg(test)]
837mod tests {
838 use super::*;
839 use crate::errors::ErrorCode;
840
841 fn parse(sql: &str) -> Result<Query, ScytheError> {
842 parse_query(sql)
843 }
844
845 #[test]
846 fn test_basic_parse() {
847 let input = "-- @name GetUsers\n-- @returns :many\nSELECT * FROM users;";
848 let q = parse(input).unwrap();
849 assert_eq!(q.name, "GetUsers");
850 assert_eq!(q.command, QueryCommand::Many);
851 assert!(q.sql.contains("SELECT"));
852 }
853
854 #[test]
855 fn test_all_command_types() {
856 let cases = vec![
857 (":one", QueryCommand::One),
858 (":many", QueryCommand::Many),
859 (":exec", QueryCommand::Exec),
860 (":exec_result", QueryCommand::ExecResult),
861 (":exec_rows", QueryCommand::ExecRows),
862 ];
863 for (tag, expected) in cases {
864 let input = format!("-- @name Q\n-- @returns {}\nSELECT 1", tag);
865 let q = parse(&input).unwrap();
866 assert_eq!(q.command, expected, "failed for {}", tag);
867 }
868 }
869
870 #[test]
871 fn test_case_insensitive_keywords() {
872 let input = "-- @Name GetUsers\n-- @RETURNS :many\nSELECT 1";
873 let q = parse(input).unwrap();
874 assert_eq!(q.name, "GetUsers");
875 assert_eq!(q.command, QueryCommand::Many);
876 }
877
878 #[test]
879 fn test_missing_name_errors() {
880 let input = "-- @returns :many\nSELECT 1";
881 let err = parse(input).unwrap_err();
882 assert_eq!(err.code, ErrorCode::MissingAnnotation);
883 assert!(err.message.contains("name"));
884 }
885
886 #[test]
887 fn test_missing_returns_errors() {
888 let input = "-- @name Foo\nSELECT 1";
889 let err = parse(input).unwrap_err();
890 assert_eq!(err.code, ErrorCode::MissingAnnotation);
891 assert!(err.message.contains("returns"));
892 }
893
894 #[test]
895 fn test_invalid_returns_value() {
896 let input = "-- @name Foo\n-- @returns :invalid\nSELECT 1";
897 let err = parse(input).unwrap_err();
898 assert_eq!(err.code, ErrorCode::InvalidAnnotation);
899 }
900
901 #[test]
905 fn test_empty_name_value_is_rejected() {
906 let input = "-- @name\n-- @returns :one\nSELECT 1";
907 let err = parse(input).unwrap_err();
908 assert_eq!(err.code, ErrorCode::InvalidAnnotation);
909 assert!(
910 err.message.contains("@name requires a value"),
911 "message must say what is missing; got: {}",
912 err.message
913 );
914 }
915
916 #[test]
917 fn test_name_value_that_is_not_an_identifier_is_rejected() {
918 for bad in ["Get User", "get-user", "users.get", "2fast", "\"GetUser\""] {
919 let input = format!("-- @name {bad}\n-- @returns :one\nSELECT 1");
920 let err = parse(&input).unwrap_err();
921 assert_eq!(
922 err.code,
923 ErrorCode::InvalidAnnotation,
924 "@name \"{bad}\" must be rejected"
925 );
926 }
927 }
928
929 #[test]
930 fn test_identifier_name_values_are_accepted() {
931 for good in ["GetUser", "get_user", "_private", "Query2"] {
932 let input = format!("-- @name {good}\n-- @returns :one\nSELECT 1");
933 let query = parse(&input).unwrap_or_else(|e| panic!("@name \"{good}\" must parse; got: {e}"));
934 assert_eq!(query.name, good);
935 }
936 }
937
938 #[test]
939 fn test_param_annotation() {
940 let input = "-- @name Foo\n-- @returns :one\n-- @param id: the user ID\nSELECT 1";
941 let q = parse(input).unwrap();
942 assert_eq!(q.annotations.param_docs.len(), 1);
943 assert_eq!(q.annotations.param_docs[0].name, "id");
944 assert_eq!(q.annotations.param_docs[0].description, "the user ID");
945 }
946
947 #[test]
948 fn test_param_no_description() {
949 let input = "-- @name Foo\n-- @returns :one\n-- @param id\nSELECT 1";
950 let q = parse(input).unwrap();
951 assert_eq!(q.annotations.param_docs.len(), 1);
952 assert_eq!(q.annotations.param_docs[0].name, "id");
953 assert_eq!(q.annotations.param_docs[0].description, "");
954 }
955
956 #[test]
957 fn test_nullable_annotation() {
958 let input = "-- @name Foo\n-- @returns :one\n-- @nullable col1, col2\nSELECT 1";
959 let q = parse(input).unwrap();
960 assert_eq!(q.annotations.nullable_overrides, vec!["col1", "col2"]);
961 }
962
963 #[test]
964 fn test_nonnull_annotation() {
965 let input = "-- @name Foo\n-- @returns :one\n-- @nonnull col1\nSELECT 1";
966 let q = parse(input).unwrap();
967 assert_eq!(q.annotations.nonnull_overrides, vec!["col1"]);
968 }
969
970 #[test]
971 fn test_json_annotation() {
972 let input = "-- @name Foo\n-- @returns :one\n-- @json data = EventData\nSELECT 1";
973 let q = parse(input).unwrap();
974 assert_eq!(q.annotations.json_mappings.len(), 1);
975 assert_eq!(q.annotations.json_mappings[0].column, "data");
976 assert_eq!(q.annotations.json_mappings[0].rust_type, "EventData");
977 }
978
979 #[test]
980 fn test_custom_annotations_captured() {
981 let input = "-- @name GetUser
982-- @returns :one
983-- @http GET /users/{id}
984-- @http_auth bearer:jwt
985-- @http_status 200,404
986SELECT id FROM users WHERE id = $1";
987 let q = parse(input).unwrap();
988 assert_eq!(q.annotations.custom.len(), 3);
989 assert_eq!(q.annotations.custom[0].name, "http");
990 assert_eq!(q.annotations.custom[0].value, "GET /users/{id}");
991 assert_eq!(q.annotations.custom[0].line, 3);
992 assert_eq!(q.annotations.custom[1].name, "http_auth");
993 assert_eq!(q.annotations.custom[1].value, "bearer:jwt");
994 assert_eq!(q.annotations.custom[1].line, 4);
995 assert_eq!(q.annotations.custom[2].name, "http_status");
996 assert_eq!(q.annotations.custom[2].value, "200,404");
997 assert_eq!(q.annotations.custom[2].line, 5);
998 }
999
1000 #[test]
1001 fn test_custom_annotation_without_value() {
1002 let input = "-- @name GetUser
1003-- @returns :one
1004-- @http_internal
1005SELECT 1";
1006 let q = parse(input).unwrap();
1007 assert_eq!(q.annotations.custom.len(), 1);
1008 assert_eq!(q.annotations.custom[0].name, "http_internal");
1009 assert_eq!(q.annotations.custom[0].value, "");
1010 }
1011
1012 #[cfg(feature = "serde")]
1013 #[test]
1014 fn test_custom_annotation_serde_round_trip() {
1015 let original = CustomAnnotation {
1016 name: "http".to_string(),
1017 value: "GET /users/{id}".to_string(),
1018 line: 7,
1019 suggested_keyword: None,
1020 };
1021 let json = serde_json::to_string(&original).unwrap();
1022 let back: CustomAnnotation = serde_json::from_str(&json).unwrap();
1023 assert_eq!(back, original);
1024 }
1025
1026 #[test]
1027 fn test_custom_annotation_name_lowercased() {
1028 let input = "-- @name GetUser
1029-- @returns :one
1030-- @HTTP_Auth Bearer
1031SELECT 1";
1032 let q = parse(input).unwrap();
1033 assert_eq!(q.annotations.custom.len(), 1);
1034 assert_eq!(q.annotations.custom[0].name, "http_auth");
1035 assert_eq!(q.annotations.custom[0].value, "Bearer");
1036 }
1037
1038 #[test]
1044 fn test_typo_annotation_suggests_known_keyword() {
1045 let input = "-- @name GetUser
1046-- @returns :one
1047-- @nullible email
1048-- @optionall email
1049-- @nonull name
1050SELECT id, name, email FROM users WHERE id = $1";
1051 let q = parse(input).unwrap();
1052 assert_eq!(q.annotations.custom.len(), 3);
1053 assert_eq!(q.annotations.custom[0].name, "nullible");
1054 assert_eq!(q.annotations.custom[0].suggested_keyword.as_deref(), Some("nullable"));
1055 assert_eq!(q.annotations.custom[1].name, "optionall");
1056 assert_eq!(q.annotations.custom[1].suggested_keyword.as_deref(), Some("optional"));
1057 assert_eq!(q.annotations.custom[2].name, "nonull");
1058 assert_eq!(q.annotations.custom[2].suggested_keyword.as_deref(), Some("nonnull"));
1059
1060 assert!(q.annotations.nullable_overrides.is_empty());
1064 assert!(q.annotations.optional_params.is_empty());
1065 assert!(q.annotations.nonnull_overrides.is_empty());
1066 }
1067
1068 #[test]
1073 fn test_custom_annotation_far_from_any_keyword_has_no_suggestion() {
1074 let input = "-- @name GetUser
1075-- @returns :one
1076-- @http GET /users/{id}
1077-- @http_auth bearer:jwt
1078-- @http_status 200,404
1079SELECT id FROM users WHERE id = $1";
1080 let q = parse(input).unwrap();
1081 assert_eq!(q.annotations.custom.len(), 3);
1082 for annotation in &q.annotations.custom {
1083 assert_eq!(
1084 annotation.suggested_keyword, None,
1085 "{:?} must not be flagged as a typo of a known keyword",
1086 annotation.name
1087 );
1088 }
1089 }
1090
1091 #[test]
1094 fn test_positional_param_basic() {
1095 let input = "-- @name Foo\n-- @returns :one\n-- @param $1 user_id\nSELECT 1";
1096 let q = parse(input).unwrap();
1097 assert_eq!(q.annotations.positional_param_docs.len(), 1);
1098 assert_eq!(q.annotations.positional_param_docs[0].position, 1);
1099 assert_eq!(q.annotations.positional_param_docs[0].name, "user_id");
1100 assert_eq!(q.annotations.positional_param_docs[0].description, "");
1101 assert_eq!(q.annotations.param_docs.len(), 0);
1102 }
1103
1104 #[test]
1105 fn test_positional_param_with_description() {
1106 let input = "-- @name Foo\n-- @returns :one\n-- @param $4 bucket: time bucket as text\nSELECT 1";
1107 let q = parse(input).unwrap();
1108 assert_eq!(q.annotations.positional_param_docs.len(), 1);
1109 assert_eq!(q.annotations.positional_param_docs[0].position, 4);
1110 assert_eq!(q.annotations.positional_param_docs[0].name, "bucket");
1111 assert_eq!(
1112 q.annotations.positional_param_docs[0].description,
1113 "time bucket as text"
1114 );
1115 }
1116
1117 #[test]
1118 fn test_positional_param_does_not_affect_docs_only_param() {
1119 let input = "-- @name Foo\n-- @returns :one\n-- @param id: the user ID\n-- @param $2 name\nSELECT 1";
1120 let q = parse(input).unwrap();
1121 assert_eq!(q.annotations.param_docs.len(), 1);
1122 assert_eq!(q.annotations.param_docs[0].name, "id");
1123 assert_eq!(q.annotations.positional_param_docs.len(), 1);
1124 assert_eq!(q.annotations.positional_param_docs[0].position, 2);
1125 assert_eq!(q.annotations.positional_param_docs[0].name, "name");
1126 }
1127
1128 #[test]
1129 fn test_positional_param_multiple() {
1130 let input =
1131 "-- @name Foo\n-- @returns :one\n-- @param $1 start_date: start\n-- @param $2 end_date: end\nSELECT 1";
1132 let q = parse(input).unwrap();
1133 assert_eq!(q.annotations.positional_param_docs.len(), 2);
1134 assert_eq!(q.annotations.positional_param_docs[0].position, 1);
1135 assert_eq!(q.annotations.positional_param_docs[0].name, "start_date");
1136 assert_eq!(q.annotations.positional_param_docs[1].position, 2);
1137 assert_eq!(q.annotations.positional_param_docs[1].name, "end_date");
1138 }
1139
1140 #[test]
1141 fn test_deprecated_annotation() {
1142 let input = "-- @name Foo\n-- @returns :one\n-- @deprecated Use V2\nSELECT 1";
1143 let q = parse(input).unwrap();
1144 assert_eq!(q.annotations.deprecated, Some("Use V2".to_string()));
1145 }
1146
1147 #[test]
1148 fn test_sql_syntax_error() {
1149 let input = "-- @name Foo\n-- @returns :one\nSELCT * FROM users";
1150 let err = parse(input).unwrap_err();
1151 assert_eq!(err.code, ErrorCode::SyntaxError);
1152 }
1153
1154 #[test]
1155 fn test_trailing_semicolon() {
1156 let input = "-- @name Foo\n-- @returns :one\nSELECT 1;";
1157 let q = parse(input).unwrap();
1158 assert_eq!(q.name, "Foo");
1159 }
1160
1161 #[test]
1162 fn test_multiple_statements_error() {
1163 let input = "-- @name Foo\n-- @returns :one\nSELECT 1; SELECT 2;";
1164 let err = parse(input).unwrap_err();
1165 assert_eq!(err.code, ErrorCode::SyntaxError);
1166 }
1167
1168 #[test]
1169 fn test_sql_preserved_without_annotations() {
1170 let input = "-- @name Foo\n-- @returns :one\nSELECT id, name FROM users WHERE id = $1";
1171 let q = parse(input).unwrap();
1172 assert_eq!(q.sql, "SELECT id, name FROM users WHERE id = $1");
1173 }
1174
1175 #[test]
1183 fn test_oracle_out_of_order_placeholders_parse_and_preserve_position_in_sql() {
1184 let input = "-- @name Foo\n-- @returns :one\nSELECT * FROM users WHERE b = :2 AND a = :1";
1185 let q = parse_query_with_dialect(input, &SqlDialect::Oracle).unwrap();
1186 assert_eq!(q.sql, "SELECT * FROM users WHERE b = $2 AND a = $1");
1187 }
1188
1189 #[test]
1190 fn test_mssql_out_of_order_placeholders_parse_and_preserve_position_in_sql() {
1191 let input = "-- @name Foo\n-- @returns :one\nSELECT * FROM users WHERE b = @p2 AND a = @p1";
1192 let q = parse_query_with_dialect(input, &SqlDialect::MsSql).unwrap();
1193 assert_eq!(q.sql, "SELECT * FROM users WHERE b = $2 AND a = $1");
1194 }
1195
1196 #[test]
1197 fn test_returns_without_colon_prefix() {
1198 let input = "-- @name Foo\n-- @returns many\nSELECT 1";
1199 let q = parse(input).unwrap();
1200 assert_eq!(q.command, QueryCommand::Many);
1201 }
1202
1203 #[test]
1204 fn test_batch_command() {
1205 let input = "-- @name Foo\n-- @returns :batch\nSELECT 1";
1206 let q = parse(input).unwrap();
1207 assert_eq!(q.command, QueryCommand::Batch);
1208 }
1209
1210 #[test]
1211 fn test_grouped_command_with_group_by() {
1212 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";
1213 let q = parse(input).unwrap();
1214 assert_eq!(q.command, QueryCommand::Grouped);
1215 assert_eq!(q.annotations.group_by, Some("users.id".to_string()));
1216 }
1217
1218 #[test]
1219 fn test_grouped_command_without_group_by_errors() {
1220 let input = "-- @name Foo\n-- @returns :grouped\nSELECT 1";
1221 let err = parse(input).unwrap_err();
1222 assert_eq!(err.code, ErrorCode::InvalidAnnotation);
1223 assert!(err.message.contains("@group_by"));
1224 }
1225
1226 #[test]
1227 fn test_group_by_without_grouped_is_ignored() {
1228 let input = "-- @name Foo\n-- @returns :many\n-- @group_by users.id\nSELECT 1";
1229 let q = parse(input).unwrap();
1230 assert_eq!(q.command, QueryCommand::Many);
1231 assert_eq!(q.annotations.group_by, Some("users.id".to_string()));
1232 }
1233
1234 #[test]
1235 fn test_preprocess_postgres_strips_partial_index_where() {
1236 let sql = "INSERT INTO billing_events (project_id, stripe_event_id) \
1237 VALUES ($1, $2) \
1238 ON CONFLICT (stripe_event_id) WHERE stripe_event_id IS NOT NULL DO NOTHING";
1239 let cleaned = preprocess_postgres_sql(sql);
1240 assert!(
1241 !cleaned.to_uppercase().contains("WHERE STRIPE_EVENT_ID IS NOT NULL"),
1242 "WHERE clause must be stripped between ON CONFLICT cols and DO; got: {cleaned}"
1243 );
1244 assert!(
1245 cleaned
1246 .to_uppercase()
1247 .contains("ON CONFLICT (STRIPE_EVENT_ID) DO NOTHING")
1248 );
1249 sqlparser::parser::Parser::parse_sql(&sqlparser::dialect::PostgreSqlDialect {}, &cleaned)
1250 .expect("cleaned SQL should parse");
1251 }
1252
1253 #[test]
1254 fn test_preprocess_postgres_no_op_when_no_partial_clause() {
1255 let sql = "INSERT INTO t (a) VALUES ($1) ON CONFLICT (a) DO UPDATE SET a = EXCLUDED.a";
1256 assert_eq!(preprocess_postgres_sql(sql), sql);
1257 }
1258
1259 #[test]
1260 fn test_preprocess_postgres_leaves_on_conflict_on_constraint_alone() {
1261 let sql = "INSERT INTO t (a) VALUES ($1) ON CONFLICT ON CONSTRAINT t_a_uidx DO NOTHING";
1262 assert_eq!(preprocess_postgres_sql(sql), sql);
1263 }
1264
1265 #[test]
1266 fn test_preprocess_postgres_handles_compound_index_cols() {
1267 let sql = "INSERT INTO t (a, b) VALUES ($1, $2) \
1268 ON CONFLICT (a, b) WHERE a IS NOT NULL AND b > 0 DO UPDATE SET b = EXCLUDED.b";
1269 let cleaned = preprocess_postgres_sql(sql);
1270 assert!(cleaned.to_uppercase().contains("ON CONFLICT (A, B) DO UPDATE"));
1271 assert!(!cleaned.to_uppercase().contains("WHERE A IS NOT NULL"));
1272 }
1273
1274 #[test]
1275 fn test_preprocess_postgres_preserves_unrelated_where() {
1276 let sql = "DELETE FROM t WHERE id = $1";
1277 assert_eq!(preprocess_postgres_sql(sql), sql);
1278 }
1279
1280 #[test]
1281 fn test_preprocess_postgres_ignores_text_inside_line_comments() {
1282 let sql = "-- inline doc: `ON CONFLICT (col) WHERE …` is the partial form\n\
1283 INSERT INTO t (a) VALUES ($1) \
1284 ON CONFLICT (a) WHERE a IS NOT NULL DO NOTHING";
1285 let cleaned = preprocess_postgres_sql(sql);
1286 assert!(
1287 cleaned.contains("-- inline doc"),
1288 "comment must survive the pass; got: {cleaned}"
1289 );
1290 assert!(cleaned.contains("ON CONFLICT (a) DO NOTHING"));
1291 }
1292
1293 #[test]
1294 fn test_preprocess_postgres_ignores_text_inside_string_literals() {
1295 let sql = "SELECT 'ON CONFLICT (a) WHERE a IS NOT NULL DO NOTHING' AS s";
1296 assert_eq!(preprocess_postgres_sql(sql), sql);
1297 }
1298
1299 #[test]
1300 fn test_preprocess_oracle_colon_placeholders() {
1301 assert_eq!(
1302 preprocess_oracle_sql("SELECT * FROM users WHERE id = :1"),
1303 "SELECT * FROM users WHERE id = $1"
1304 );
1305 assert_eq!(
1306 preprocess_oracle_sql("INSERT INTO users (name, email) VALUES (:1, :2)"),
1307 "INSERT INTO users (name, email) VALUES ($1, $2)"
1308 );
1309 }
1310
1311 #[test]
1312 fn test_preprocess_oracle_preserves_string_literals() {
1313 assert_eq!(
1314 preprocess_oracle_sql("SELECT * FROM users WHERE name = ':1' AND id = :1"),
1315 "SELECT * FROM users WHERE name = ':1' AND id = $1"
1316 );
1317 }
1318
1319 #[test]
1320 fn test_preprocess_oracle_strips_returning_into() {
1321 assert_eq!(
1322 preprocess_oracle_sql("INSERT INTO users (name) VALUES (:1) RETURNING id, name INTO :2, :3"),
1323 "INSERT INTO users (name) VALUES ($1) RETURNING id, name"
1324 );
1325 }
1326
1327 #[test]
1328 fn test_preprocess_oracle_full_insert_returning_into() {
1329 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";
1330 let result = preprocess_oracle_sql(sql);
1331 assert_eq!(
1332 result,
1333 "INSERT INTO users (name, email, active) VALUES ($1, $2, $3) RETURNING id, name, email, active, created_at"
1334 );
1335 }
1336
1337 #[test]
1338 fn test_preprocess_oracle_no_returning_into_unchanged() {
1339 assert_eq!(
1340 preprocess_oracle_sql("DELETE FROM users WHERE id = :1"),
1341 "DELETE FROM users WHERE id = $1"
1342 );
1343 }
1344
1345 #[test]
1349 fn test_preprocess_oracle_repeated_placeholder_keeps_its_own_position() {
1350 assert_eq!(
1351 preprocess_oracle_sql("SELECT * FROM users WHERE id = :1 OR parent_id = :1"),
1352 "SELECT * FROM users WHERE id = $1 OR parent_id = $1"
1353 );
1354 }
1355
1356 #[test]
1360 fn test_preprocess_oracle_out_of_order_placeholders_keep_their_own_numbers() {
1361 assert_eq!(
1362 preprocess_oracle_sql("SELECT * FROM users WHERE b = :2 AND a = :1"),
1363 "SELECT * FROM users WHERE b = $2 AND a = $1"
1364 );
1365 }
1366
1367 #[test]
1368 fn test_preprocess_mssql_single_placeholder() {
1369 assert_eq!(
1370 preprocess_mssql_sql("SELECT * FROM users WHERE id = @p1"),
1371 "SELECT * FROM users WHERE id = $1"
1372 );
1373 }
1374
1375 #[test]
1376 fn test_preprocess_mssql_multiple_placeholders() {
1377 assert_eq!(
1378 preprocess_mssql_sql("INSERT INTO users (name, email) VALUES (@p1, @p2)"),
1379 "INSERT INTO users (name, email) VALUES ($1, $2)"
1380 );
1381 }
1382
1383 #[test]
1384 fn test_preprocess_mssql_preserves_string_literals() {
1385 assert_eq!(
1386 preprocess_mssql_sql("SELECT * FROM users WHERE name = '@p1' AND id = @p1"),
1387 "SELECT * FROM users WHERE name = '@p1' AND id = $1"
1388 );
1389 }
1390
1391 #[test]
1392 fn test_preprocess_mssql_case_insensitive_p() {
1393 assert_eq!(
1394 preprocess_mssql_sql("SELECT * FROM users WHERE id = @P1"),
1395 "SELECT * FROM users WHERE id = $1"
1396 );
1397 }
1398
1399 #[test]
1400 fn test_preprocess_mssql_non_placeholder_at_variable_unchanged() {
1401 assert_eq!(preprocess_mssql_sql("SELECT @myvar"), "SELECT @myvar");
1402 }
1403
1404 #[test]
1405 fn test_preprocess_mssql_multi_digit_placeholder() {
1406 assert_eq!(preprocess_mssql_sql("SELECT @p10, @p2"), "SELECT $10, $2");
1407 }
1408
1409 #[test]
1412 fn test_preprocess_mssql_repeated_placeholder_keeps_its_own_position() {
1413 assert_eq!(
1414 preprocess_mssql_sql("SELECT * FROM users WHERE id = @p1 OR parent_id = @p1"),
1415 "SELECT * FROM users WHERE id = $1 OR parent_id = $1"
1416 );
1417 }
1418
1419 #[test]
1421 fn test_preprocess_mssql_out_of_order_placeholders_keep_their_own_numbers() {
1422 assert_eq!(
1423 preprocess_mssql_sql("SELECT * FROM users WHERE b = @p2 AND a = @p1"),
1424 "SELECT * FROM users WHERE b = $2 AND a = $1"
1425 );
1426 }
1427
1428 #[test]
1429 fn test_preprocess_mssql_output_inserted_simple() {
1430 let sql = "INSERT INTO users (id, name) OUTPUT INSERTED.id, INSERTED.name VALUES (@p1, @p2)";
1431 let result = preprocess_mssql_sql(sql);
1432 assert!(result.contains("RETURNING id, name"), "got: {}", result);
1433 assert!(result.contains("VALUES ($1, $2)"), "got: {}", result);
1434 assert!(!result.contains("OUTPUT"), "got: {}", result);
1435 }
1436
1437 #[test]
1438 fn test_preprocess_mssql_output_inserted_full_example() {
1439 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)";
1440 let result = preprocess_mssql_sql(sql);
1441 assert!(
1442 result.contains("RETURNING id, name, email, active, created_at"),
1443 "got: {}",
1444 result
1445 );
1446 assert!(result.contains("VALUES ($1, $2, $3, $4)"), "got: {}", result);
1447 }
1448
1449 #[test]
1450 fn test_preprocess_mssql_output_case_insensitive() {
1451 let sql = "INSERT INTO users (id) output inserted.id values (@p1)";
1452 let result = preprocess_mssql_sql(sql);
1453 assert!(result.contains("RETURNING id"), "got: {}", result);
1454 assert!(
1455 result.contains("values ($1)") || result.contains("VALUES ($1)"),
1456 "got: {}",
1457 result
1458 );
1459 }
1460
1461 #[test]
1462 fn test_preprocess_mssql_no_output_unchanged() {
1463 let sql = "INSERT INTO users (id, name) VALUES (@p1, @p2)";
1464 let result = preprocess_mssql_sql(sql);
1465 assert_eq!(result, "INSERT INTO users (id, name) VALUES ($1, $2)");
1466 }
1467
1468 #[test]
1469 fn test_preprocess_mssql_output_with_string_literal() {
1470 let sql = "INSERT INTO users (id, name) OUTPUT INSERTED.id, INSERTED.name VALUES (@p1, '@p2')";
1471 let result = preprocess_mssql_sql(sql);
1472 assert!(result.contains("RETURNING id, name"), "got: {}", result);
1473 assert!(result.contains("($1, '@p2')"), "got: {}", result);
1474 }
1475
1476 #[test]
1477 fn test_preprocess_mssql_output_with_whitespace() {
1478 let sql = "INSERT INTO users (id, name)\nOUTPUT INSERTED.id,\n INSERTED.name\nVALUES (@p1, @p2)";
1479 let result = preprocess_mssql_sql(sql);
1480 assert!(result.contains("RETURNING id, name"), "got: {}", result);
1481 assert!(result.contains("VALUES ($1, $2)"), "got: {}", result);
1482 }
1483}