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///
93/// That escape hatch means an unrecognised annotation can never be a hard parse
94/// error -- rejecting it outright would break every legitimate consumer-defined
95/// annotation (e.g. `@http`, `@http_auth`) alongside the typos it was meant to
96/// catch. [`suggested_keyword`](Self::suggested_keyword) is the softer signal:
97/// it flags names close enough to a known keyword to plausibly be a typo of it
98/// (`@nullible` for `@nullable`, `@optionall` for `@optional`), for a caller to
99/// surface as a warning without duplicating scythe's keyword list or edit-distance
100/// logic. See #152 -- before this field existed, a misspelled `@nullable` /
101/// `@nonnull` / `@optional` was captured here and never inspected by anything,
102/// so `scythe generate`, `scythe check` and `scythe lint` all reported success
103/// while the override it named was silently discarded.
104#[derive(Debug, Clone, PartialEq, Eq)]
105#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
106pub struct CustomAnnotation {
107    /// Annotation name, lowercased, without the leading `@` (e.g. `http`, `http_param`).
108    pub name: String,
109    /// Everything after the name on the line, trimmed. Empty if the annotation had no value.
110    pub value: String,
111    /// 1-based line number within the query SQL, for diagnostics.
112    pub line: usize,
113    /// A known annotation keyword within edit distance 2 of `name`, if any -- e.g.
114    /// `Some("nullable".to_string())` for `nullible`. `None` when `name` is not close to
115    /// any known keyword (the common case: a genuine consumer-defined annotation). An owned
116    /// `String` rather than `&'static str` so the type stays trivially `Deserialize` (a
117    /// borrowed-forever `&'static str` field cannot round-trip through an arbitrary
118    /// deserializer's input lifetime).
119    #[cfg_attr(feature = "serde", serde(default))]
120    pub suggested_keyword: Option<String>,
121}
122
123/// The annotation keywords [`parse_query_with_dialect`] recognises natively. Anything else
124/// becomes a [`CustomAnnotation`]; [`suggest_known_keyword`] checks a rejected name against
125/// this list to flag likely typos.
126const 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
138/// Maximum Levenshtein edit distance at which an unrecognised annotation name is flagged as a
139/// likely typo of a known keyword. Matches the threshold `scythe-codegen::backend_options` and
140/// `scythe-backend::manifest` already use for "did you mean" suggestions on unknown option/type
141/// keys -- close enough to catch `nullible` -> `nullable` or `optionall` -> `optional` without
142/// false-positiving on a genuinely unrelated custom annotation name.
143const ANNOTATION_TYPO_DISTANCE_THRESHOLD: usize = 2;
144
145/// Suggest a known annotation keyword for an unrecognised annotation `name`, if one is within
146/// [`ANNOTATION_TYPO_DISTANCE_THRESHOLD`] edits.
147fn 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
156/// Levenshtein edit distance between two strings, used by [`suggest_known_keyword`].
157fn 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    /// Explicit `-- @param $N name[: description]` annotations that override
191    /// the inferred or fallback `pN` name for a specific positional parameter.
192    #[cfg_attr(feature = "serde", serde(default))]
193    pub positional_param_docs: Vec<PositionalParamDoc>,
194    /// Annotations scythe does not natively recognise, preserved in source order
195    /// for crate consumers to interpret.
196    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
208/// Reject an `@name` value that cannot become an identifier in generated
209/// code.
210///
211/// Every backend uses this value verbatim (after case conversion) as the
212/// generated function's name and as the stem of its row type, so anything
213/// that is not an identifier produces a file that does not compile -- and
214/// says so nowhere. `-- @name` with no value was accepted outright until
215/// #174: it emitted `async def (conn):` in Python and collapsed the row type
216/// to a bare `Row` that collides with every other unnamed query in the file,
217/// at exit code 0.
218///
219/// The accepted set is deliberately the intersection every target language
220/// can spell: an ASCII letter or `_` followed by ASCII letters, digits and
221/// `_`. A dot, dash, space or quote in the value is a mistake in every one of
222/// them.
223fn 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
248/// Parse a single annotated SQL query into a `Query` using the PostgreSQL dialect.
249pub fn parse_query(query_sql: &str) -> Result<Query, ScytheError> {
250    parse_query_with_dialect(query_sql, &SqlDialect::PostgreSQL)
251}
252
253/// Parse a single annotated SQL query into a `Query` using the specified dialect.
254pub 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                    //   Positional: -- @param $N name[: description]
296                    //   Docs-only:  -- @param name[: description]
297                    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
482/// Strip the `WHERE …` predicate that PostgreSQL allows between
483/// `ON CONFLICT (cols)` and `DO …` (the index-inference form for partial
484/// unique indexes). sqlparser-rs through 0.61 does not parse this construct;
485/// we lift it out for the parser and let the caller keep the original SQL
486/// for codegen + runtime, where Postgres validates it.
487fn 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
542/// Build an ASCII-uppercase, fixed-byte-offset mask of `sql` where `--` line
543/// comments, `/* … */` block comments, and `'…'` / `$$…$$` string literals are
544/// replaced with spaces. Multi-byte UTF-8 is collapsed to ASCII spaces of the
545/// same byte length so positions in the mask line up with the original `sql`.
546fn 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
604/// Locate a whitespace-separated keyword in an uppercase haystack. Returns the
605/// byte offset of the keyword's start, or None if not found.
606fn 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
624/// Preprocess Oracle SQL before parsing:
625/// 1. Strip `INTO :N, :N, ...` suffix from `RETURNING ... INTO` clauses.
626/// 2. Convert `:N` positional placeholders to `$N`, preserving the declared
627///    position (#149).
628///
629/// `$N`, not bare `?`: collapsing every `:N` to the same `?` spelling
630/// discarded which N a repeated (`:1 ... :1`) or out-of-order
631/// (`:2` before `:1`) placeholder referred to, so a backend binding by SQL
632/// occurrence order could no longer recover the caller's intended argument
633/// -- a repeated placeholder produced too few binds (runtime error) and an
634/// out-of-order one silently bound the wrong argument to the wrong slot.
635/// `$N` is safe to emit here for two independent reasons: sqlparser's
636/// tokenizer accepts `$<digits>` as `Token::Placeholder` under the default
637/// `Dialect` impl that `OracleDialect` inherits (verified by inspecting
638/// `tokenize_dollar_preceded_value` in sqlparser 0.62's `tokenizer.rs`: a
639/// `$`-prefixed all-digit run with no following `$` always falls through to
640/// `Token::Placeholder(format!("${value}"))`), and the downstream SQL-text
641/// pipeline (`rewrite_placeholders_indexed` in scythe-codegen) already
642/// treats Oracle and PostgreSQL identically as "`$N`-style" dialects.
643fn 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
673/// Convert MSSQL `@pN` positional placeholders to `$N` (outside string
674/// literals), preserving the declared position (#149) instead of collapsing
675/// every occurrence to bare `?` -- see [`preprocess_oracle_sql`]'s doc
676/// comment for why `$N` and not `?`, and why it is safe to hand sqlparser's
677/// `MsSqlDialect` (which, like `OracleDialect`, inherits the default
678/// `Dialect::supports_dollar_placeholder` impl that makes `$N` tokenize as
679/// `Token::Placeholder`). MsSqlDialect treats bare `@` as an identifier
680/// start, so `@p1` would otherwise become an identifier rather than a
681/// `Placeholder` token.
682fn 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
717/// Preprocess MSSQL SQL before parsing:
718/// 1. Strip `OUTPUT INSERTED.col, ...` clauses and convert to RETURNING
719/// 2. Convert `@pN` positional placeholders to `$N` (see
720///    [`convert_mssql_placeholders`])
721fn preprocess_mssql_sql(sql: &str) -> String {
722    let sql = strip_and_convert_mssql_output(sql);
723    convert_mssql_placeholders(&sql)
724}
725
726/// Strip MSSQL `OUTPUT INSERTED.col1, INSERTED.col2, ...` from INSERT statements
727/// and convert it to a `RETURNING col1, col2, ...` clause.
728/// The OUTPUT clause appears between the column list and VALUES clause:
729///   INSERT INTO table (cols) OUTPUT INSERTED.col1, INSERTED.col2, ... VALUES (...)
730/// becomes:
731///   INSERT INTO table (cols) VALUES (...) RETURNING col1, col2, ...
732fn 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
770/// Find the position of a word (case-insensitive) in the text.
771/// The word must be a separate word, not part of another identifier.
772fn 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
799/// Parse INSERTED.col1, INSERTED.col2, ... and extract column names as "col1, col2, ..."
800fn 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
822/// Strip the `INTO :N, :N, ...` suffix from an Oracle `RETURNING ... INTO` clause.
823fn 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    /// Inverted with #174: this used to assert that `-- @name` with no value
902    /// parses to an empty name, which generated `async def (conn):` and a
903    /// bare `Row` type at exit 0.
904    #[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    /// Regression for #152: a misspelled `@nullable` / `@nonnull` / `@optional` used to be
1039    /// captured as an opaque `CustomAnnotation` that nothing ever inspected -- `scythe generate`,
1040    /// `scythe check` and `scythe lint` all reported success while the nullability override the
1041    /// user wrote never took effect. `suggested_keyword` must name the intended keyword so a
1042    /// caller can turn this into a warning instead of silence.
1043    #[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        // The override itself must NOT have silently applied under its misspelled name --
1061        // otherwise the typo would behave identically to the real keyword and there would be
1062        // nothing to warn about.
1063        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    /// A genuinely consumer-defined annotation vocabulary (`@http`, `@http_auth`, ...) must not
1069    /// be flagged -- `suggested_keyword` is a typo signal, not a closed-vocabulary rejection; see
1070    /// the `CustomAnnotation` doc comment for why an unrecognised annotation is never a hard
1071    /// error.
1072    #[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    // ---- @param positional form ----
1092
1093    #[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    /// End-to-end confirmation for #149: sqlparser must actually accept the
1176    /// `$N` spelling `preprocess_oracle_sql`/`convert_mssql_placeholders` now
1177    /// emit under `OracleDialect`/`MsSqlDialect` -- not just produce text that
1178    /// looks plausible. An out-of-order `:2 ... :1` (Oracle) / `@p2 ... @p1`
1179    /// (MSSQL) proves the position survives into `Query.sql` unrenumbered,
1180    /// which is what lets a codegen backend bind by occurrence instead of by
1181    /// declaration order.
1182    #[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    /// #149: a repeated `:1` must resolve to two occurrences of the *same*
1346    /// position, not collapse to indistinguishable bare `?`s a bind-per-
1347    /// occurrence backend can no longer tell apart.
1348    #[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    /// #149: `:2` declared before `:1` in the SQL text must keep its own
1357    /// number rather than being renumbered by text order -- renumbering is
1358    /// exactly what silently swapped a caller's arguments before this fix.
1359    #[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    /// #149: a repeated `@p1` must resolve to two occurrences of the same
1410    /// position; see the Oracle counterpart above for the full reasoning.
1411    #[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    /// #149: `@p2` declared before `@p1` must keep its own number.
1420    #[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}