Skip to main content

scythe_core/parser/
mod.rs

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/// An explicit `-- @param $N name[: description]` annotation that binds a
61/// position number to a human-chosen parameter name, overriding the `pN`
62/// fallback name that the analyzer would otherwise infer.
63///
64/// The positional form (`$N` as first token) is distinct from the existing
65/// docs-only `-- @param name: description` form, which is preserved unchanged
66/// for backward compatibility.
67#[derive(Debug, Clone, PartialEq, Eq)]
68#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
69pub struct PositionalParamDoc {
70    /// 1-based parameter position (`$N`).
71    pub position: i64,
72    /// Caller-chosen name that replaces the `pN` fallback.
73    pub name: String,
74    /// Optional description (everything after the `:` separator).
75    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/// A custom (non-native) annotation captured verbatim from the SQL source.
86///
87/// Scythe parses its known annotations (`@name`, `@returns`, `@param`, `@nullable`,
88/// `@nonnull`, `@json`, `@optional`, `@group_by`, `@deprecated`) into typed fields.
89/// Any other `-- @<name> <value>` line is captured here as an opaque triple and
90/// exposed to crate consumers, who can layer their own annotation vocabulary on
91/// top of scythe without coupling scythe to their domain.
92#[derive(Debug, Clone, PartialEq, Eq)]
93#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
94pub struct CustomAnnotation {
95    /// Annotation name, lowercased, without the leading `@` (e.g. `http`, `http_param`).
96    pub name: String,
97    /// Everything after the name on the line, trimmed. Empty if the annotation had no value.
98    pub value: String,
99    /// 1-based line number within the query SQL, for diagnostics.
100    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    /// Explicit `-- @param $N name[: description]` annotations that override
116    /// the inferred or fallback `pN` name for a specific positional parameter.
117    #[cfg_attr(feature = "serde", serde(default))]
118    pub positional_param_docs: Vec<PositionalParamDoc>,
119    /// Annotations scythe does not natively recognise, preserved in source order
120    /// for crate consumers to interpret.
121    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
133/// Parse a single annotated SQL query into a `Query` using the PostgreSQL dialect.
134pub fn parse_query(query_sql: &str) -> Result<Query, ScytheError> {
135    parse_query_with_dialect(query_sql, &SqlDialect::PostgreSQL)
136}
137
138/// Parse a single annotated SQL query into a `Query` using the specified dialect.
139pub 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        // Check for annotation: "-- @..." or "--@..."
159        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            // Parse the annotation keyword and value
168            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                    // Two forms:
183                    //   Positional: -- @param $N name[: description]
184                    //     → overrides the inferred name for parameter at position N.
185                    //   Docs-only:  -- @param name[: description]
186                    //     → preserved in param_docs for consumers; does not affect codegen.
187                    //
188                    // Detect the positional form by checking whether the first whitespace-
189                    // separated token matches `$<digits>`.
190                    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                        // Positional form: extract name and optional description from the rest.
198                        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                        // Docs-only form: "<name>: <description>" or "<name>:<description>"
218                        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                    // format: "<col> = <Type>"
251                    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                    // Unknown annotation — capture verbatim for crate consumers.
273                    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    // Preprocess dialect-specific syntax before parsing:
301    //   * Oracle: strip `RETURNING ... INTO` output binds, convert `:N` → `?`.
302    //   * MSSQL: convert `OUTPUT INSERTED.*` → `RETURNING` for parsing,
303    //     convert `@pN` → `?` for parsing; keep original SQL for codegen.
304    //   * PostgreSQL: strip `WHERE …` between `ON CONFLICT (cols)` and `DO …`
305    //     for parsing (sqlparser-rs <= 0.61 doesn't recognise the
306    //     partial-index inference form); keep original SQL for codegen.
307    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        // For codegen: only convert @pN → ? placeholders (keep OUTPUT syntax)
312        let codegen_sql = convert_mssql_placeholders(&sql);
313        // For parsing: also convert OUTPUT INSERTED → RETURNING
314        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        // sqlparser may produce an extra empty statement from a trailing semicolon —
329        // filter those out by checking for exactly one non-empty statement.
330        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
388/// Strip the `WHERE …` predicate that PostgreSQL allows between
389/// `ON CONFLICT (cols)` and `DO …` (the index-inference form for partial
390/// unique indexes). sqlparser-rs through 0.61 does not parse this construct;
391/// we lift it out for the parser and let the caller keep the original SQL
392/// for codegen + runtime, where Postgres validates it.
393fn preprocess_postgres_sql(sql: &str) -> String {
394    // Strip line comments + string literals first so we only scan structural SQL.
395    // (We still emit the original `sql` slice byte-for-byte; the upper-mask is
396    //  only used to decide *where* to cut.)
397    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            // Slice from the original SQL (preserves casing + UTF-8) up to
441            // the byte before WHERE; skip ahead to DO.
442            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
453/// Build an ASCII-uppercase, fixed-byte-offset mask of `sql` where `--` line
454/// comments, `/* … */` block comments, and `'…'` / `$$…$$` string literals are
455/// replaced with spaces. Multi-byte UTF-8 is collapsed to ASCII spaces of the
456/// same byte length so positions in the mask line up with the original `sql`.
457fn 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            // Line comment — replace through end-of-line with spaces.
465            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            // Block comment — replace through `*/`.
473            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        // ASCII goes through as-uppercase; non-ASCII bytes become spaces so the
508        // mask stays single-byte-per-position and positions line up.
509        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
519/// Locate a whitespace-separated keyword in an uppercase haystack. Returns the
520/// byte offset of the keyword's start, or None if not found.
521fn 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
539/// Preprocess Oracle SQL before parsing:
540/// 1. Strip `INTO :N, :N, ...` suffix from `RETURNING ... INTO` clauses.
541/// 2. Convert `:N` positional placeholders to `?` (universally supported).
542fn preprocess_oracle_sql(sql: &str) -> String {
543    // Strip Oracle RETURNING ... INTO clause (output bind variables)
544    // e.g. "INSERT ... RETURNING id, name INTO :4, :5" → "INSERT ... RETURNING id, name"
545    let sql = strip_returning_into(sql);
546
547    // Convert :N → ? (outside string literals)
548    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            // Skip string literals
553            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            // Convert :N → ?
566            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
577/// Convert MSSQL `@pN` positional placeholders to `?` (outside string literals).
578/// MsSqlDialect treats `@` as an identifier start, so `@p1` becomes an identifier
579/// rather than a `Placeholder` token — preprocessing normalises it to `?`.
580fn 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            // Skip string literals verbatim
586            result.push(ch);
587            while let Some(inner) = chars.next() {
588                result.push(inner);
589                if inner == '\'' {
590                    if chars.peek() == Some(&'\'') {
591                        // Escaped quote inside string literal
592                        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            // Peek ahead: must be `@p` followed by at least one digit
600            let mut lookahead = chars.clone();
601            lookahead.next(); // consume the 'p'/'P'
602            if lookahead.peek().is_some_and(|c| c.is_ascii_digit()) {
603                // It is an `@pN` placeholder — consume `p` and all digits
604                chars.next(); // consume 'p'/'P'
605                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
619/// Preprocess MSSQL SQL before parsing:
620/// 1. Strip `OUTPUT INSERTED.col, ...` clauses and convert to RETURNING
621/// 2. Convert `@pN` positional placeholders to `?`
622fn preprocess_mssql_sql(sql: &str) -> String {
623    // First pass: convert OUTPUT INSERTED.col to RETURNING col
624    let sql = strip_and_convert_mssql_output(sql);
625    // Second pass: convert @pN to ?
626    convert_mssql_placeholders(&sql)
627}
628
629/// Strip MSSQL `OUTPUT INSERTED.col1, INSERTED.col2, ...` from INSERT statements
630/// and convert it to a `RETURNING col1, col2, ...` clause.
631/// The OUTPUT clause appears between the column list and VALUES clause:
632///   INSERT INTO table (cols) OUTPUT INSERTED.col1, INSERTED.col2, ... VALUES (...)
633/// becomes:
634///   INSERT INTO table (cols) VALUES (...) RETURNING col1, col2, ...
635fn strip_and_convert_mssql_output(sql: &str) -> String {
636    // Case-insensitive search for OUTPUT keyword in INSERT statements
637    let upper = sql.to_uppercase();
638
639    // Only process INSERT statements with OUTPUT
640    if !upper.contains("INSERT") || !upper.contains("OUTPUT") {
641        return sql.to_string();
642    }
643
644    // Find the OUTPUT keyword
645    if let Some(output_pos) = find_word_position(&upper, "OUTPUT") {
646        // Check if this is actually part of an INSERT statement by finding INSERT before it
647        let before_output = &upper[..output_pos];
648        if !before_output.contains("INSERT") {
649            return sql.to_string();
650        }
651
652        // Look for the VALUES keyword after OUTPUT
653        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            // Extract the OUTPUT column list (between OUTPUT and VALUES)
658            let output_cols_str = &sql[output_pos + "OUTPUT".len()..values_pos];
659
660            // Parse column names: strip "INSERTED." prefix from each column name
661            let cols = parse_inserted_columns(output_cols_str);
662
663            if !cols.is_empty() {
664                // Build result: keep everything before OUTPUT, then VALUES clause,
665                // then RETURNING clause (before any trailing semicolon)
666                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
682/// Find the position of a word (case-insensitive) in the text.
683/// The word must be a separate word, not part of another identifier.
684fn 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        // Check character before
691        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        // Check character after
698        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
713/// Parse INSERTED.col1, INSERTED.col2, ... and extract column names as "col1, col2, ..."
714fn 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        // Try to extract column name after INSERTED.
721        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
737/// Strip the `INTO :N, :N, ...` suffix from an Oracle `RETURNING ... INTO` clause.
738fn strip_returning_into(sql: &str) -> String {
739    // Case-insensitive search for "INTO" after "RETURNING" at the end of the statement
740    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            // Keep everything before INTO, trim trailing whitespace/semicolons
746            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        // An empty name is accepted by the parser (it stores "")
821        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        // Unknown @xxx lines are captured verbatim as CustomAnnotation triples;
870        // native annotations remain in their typed fields.
871        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    // ---- @param positional form ----
928
929    #[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        // docs-only list must be empty
938        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        // The docs-only form still works alongside the positional form.
957        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 must accept the cleaned form.
1067        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        // The DELETE's WHERE is its own clause, not an ON-CONFLICT predicate;
1095        // it must survive untouched.
1096        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        // Earlier scans treated this as a real `ON CONFLICT (col) WHERE … DO`
1103        // and excised the entire comment + INSERT body up to the next `DO`.
1104        // Comments must be opaque to the predicate-stripping pass.
1105        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        // @variable (not @pN pattern) must not be touched
1203        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        // Should convert OUTPUT INSERTED.col to RETURNING col and @pN to ?
1216        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        // The original lowercase "values" is preserved, then @p1 becomes ?
1239        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        // @p1 inside a string should be preserved by placeholder conversion
1256        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}