Skip to main content

sequel_mcp/policy/
classifier.rs

1//! SQL classification and object extraction (`policy/classifier.ts` port).
2//!
3//! Preserves the legacy contract (categories, ast names, fast paths,
4//! multi-statement rejection, target database collection) and adds the v2
5//! object graph: distinct read vs mutated tables, locking-read and file-io
6//! flags, and executing-semantics for `EXPLAIN ANALYZE`.
7
8use crate::policy::model::SqlCategory;
9use sqlparser::ast::{Expr, ObjectName, Query, SetExpr, Statement, TableFactor};
10use sqlparser::dialect::{MySqlDialect, SQLiteDialect};
11use sqlparser::parser::Parser;
12use std::collections::BTreeSet;
13
14#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
15pub struct TableRef {
16    /// `None` for unqualified references; resolved (or denied) later.
17    pub database: Option<String>,
18    pub table: String,
19}
20
21#[derive(Debug, Clone, PartialEq)]
22pub struct ClassifiedStatement {
23    pub category: SqlCategory,
24    /// Legacy-compatible AST type name (`select`, `insert`, `update`,
25    /// `delete`, `replace`, `create`, `drop`, `alter`, `truncate`, `rename`,
26    /// `grant`, `set`, `show`, `describe`, `explain`, `pragma`,
27    /// `transaction`, `admin-keyword`).
28    pub ast_type: &'static str,
29    /// Distinct database qualifiers across all table references, sorted
30    /// (legacy `targetDatabases`).
31    pub target_databases: Vec<String>,
32    /// Tables the statement reads.
33    pub read_tables: Vec<TableRef>,
34    /// Tables the statement can mutate. For UPDATE/DELETE every
35    /// from/using table counts: MySQL may mutate any joined table.
36    pub mutated_tables: Vec<TableRef>,
37    /// `FOR UPDATE` / `LOCK IN SHARE MODE` present — not an ordinary read.
38    pub locking_read: bool,
39    /// Statement reads or writes server-side files (`INTO OUTFILE`,
40    /// `LOAD DATA [LOCAL] INFILE`) — denied by policy regardless of
41    /// category grants.
42    pub file_io: bool,
43    /// `EXPLAIN ANALYZE` — executes the wrapped statement.
44    pub executes_wrapped: bool,
45    /// `IF EXISTS` present on DROP/TRUNCATE DDL (absent-target handling:
46    /// preflight turns a missing target into an audited local no-op).
47    pub if_exists: bool,
48    /// Object type of a DROP statement (`table`, `view`, `index`,
49    /// `other`); `None` for non-DROP statements. The Mixed-IF-EXISTS
50    /// rewrite only reconstructs statements it knows (`table`/`view`)
51    /// and fails closed otherwise.
52    pub drop_object_type: Option<&'static str>,
53}
54
55#[derive(Debug, Clone, PartialEq)]
56pub enum ClassifyError {
57    Empty,
58    CommentOnly,
59    MultipleStatements,
60    Parse(String),
61    Unknown(String),
62}
63
64impl ClassifyError {
65    /// Message shapes mirror the legacy classifier error strings.
66    pub fn message(&self) -> String {
67        match self {
68            ClassifyError::Empty => "empty input".into(),
69            ClassifyError::CommentOnly => "input contains only comments".into(),
70            ClassifyError::MultipleStatements => {
71                "multiple statements not allowed (single statement only)".into()
72            }
73            ClassifyError::Parse(m) => format!("parser error: {m}"),
74            ClassifyError::Unknown(t) => format!("unknown statement type \"{t}\""),
75        }
76    }
77}
78
79#[derive(Debug, Clone, Copy, PartialEq, Eq)]
80pub enum Dialect {
81    MySql,
82    SQLite,
83}
84
85impl Dialect {
86    fn as_dyn(&self) -> &'static dyn sqlparser::dialect::Dialect {
87        match self {
88            Dialect::MySql => &MySqlDialect {},
89            Dialect::SQLite => &SQLiteDialect {},
90        }
91    }
92}
93
94/// Strip `/* */`, `-- ` and MySQL `#` comments (legacy `stripComments`).
95pub fn strip_comments(sql: &str) -> String {
96    let mut out = String::with_capacity(sql.len());
97    let bytes: Vec<char> = sql.chars().collect();
98    let mut i = 0;
99    let n = bytes.len();
100    while i < n {
101        let c = bytes[i];
102        if c == '/' && i + 1 < n && bytes[i + 1] == '*' {
103            i += 2;
104            while i + 1 < n && !(bytes[i] == '*' && bytes[i + 1] == '/') {
105                i += 1;
106            }
107            i = (i + 2).min(n);
108            out.push(' ');
109        } else if (c == '-' && i + 1 < n && bytes[i + 1] == '-') || c == '#' {
110            // `-- ` line comment or MySQL `#` comment: skip to end of line.
111            while i < n && bytes[i] != '\n' {
112                i += 1;
113            }
114            out.push(' ');
115        } else {
116            out.push(c);
117            i += 1;
118        }
119    }
120    out
121}
122
123/// Quote/paren-aware semicolon scan (legacy `looksLikeMultipleStatements`):
124/// a `;` at depth 0 outside any quoted region means multiple statements.
125pub fn looks_like_multiple_statements(sql: &str) -> bool {
126    let stripped = strip_comments(sql);
127    let trimmed = stripped.trim_end();
128    let trimmed = trimmed.strip_suffix(';').unwrap_or(trimmed).trim();
129    if trimmed.is_empty() {
130        return false;
131    }
132    let mut quote: Option<char> = None;
133    let mut depth: i32 = 0;
134    let chars: Vec<char> = trimmed.chars().collect();
135    for (i, &c) in chars.iter().enumerate() {
136        if let Some(q) = quote {
137            if c == q && (i == 0 || chars[i - 1] != '\\') {
138                quote = None;
139            }
140            continue;
141        }
142        match c {
143            '\'' | '"' | '`' => quote = Some(c),
144            '(' => depth += 1,
145            ')' => depth -= 1,
146            ';' if depth == 0 => return true,
147            _ => {}
148        }
149    }
150    false
151}
152
153fn is_tx_keyword(stripped: &str) -> bool {
154    let t = stripped.trim_start();
155    let lower = t.to_ascii_lowercase();
156    let starts = [
157        "begin",
158        "commit",
159        "rollback",
160        "start transaction",
161        "savepoint",
162        "release savepoint",
163    ];
164    let Some(first) = lower.split_whitespace().next() else {
165        return false;
166    };
167    if !starts
168        .iter()
169        .any(|s| lower.starts_with(s) && word_bounded(&lower, s))
170    {
171        return false;
172    }
173    first == "begin"
174        || first == "commit"
175        || first == "rollback"
176        || lower.starts_with("start transaction")
177        || lower.starts_with("savepoint")
178        || lower.starts_with("release savepoint")
179}
180
181fn word_bounded(lower: &str, prefix: &str) -> bool {
182    match lower.get(prefix.len()..) {
183        Some(rest) => rest.starts_with(|c: char| c.is_whitespace()) || rest.is_empty(),
184        None => false,
185    }
186}
187
188fn is_admin_keyword(stripped: &str) -> bool {
189    let lower = stripped.trim_start().to_ascii_lowercase();
190    const PREFIXES: &[&str] = &[
191        "grant ",
192        "revoke ",
193        "set global",
194        "set persist",
195        "set persist_only",
196        "set @@global",
197        "set @@persist",
198        "kill ",
199        "flush",
200        "reset master",
201        "reset slave",
202        "reset replica",
203        "lock tables",
204        "unlock tables",
205        "load data",
206        "handler ",
207        "do ",
208        "change master",
209        "change replication",
210        "start slave",
211        "stop slave",
212        "start replica",
213        "stop replica",
214        "optimize table",
215        "repair table",
216        "analyze table",
217        "check table",
218        "create user",
219        "alter user",
220        "drop user",
221        "rename user",
222        "set password",
223        "attach database",
224        "detach database",
225        "vacuum",
226        "reindex",
227    ];
228    PREFIXES.iter().any(|p| lower.starts_with(p))
229}
230
231/// The 22 read-only PRAGMA names (legacy allowlist).
232const READ_ONLY_PRAGMAS: [&str; 22] = [
233    "application_id",
234    "collation_list",
235    "compile_options",
236    "database_list",
237    "foreign_key_check",
238    "foreign_key_list",
239    "freelist_count",
240    "function_list",
241    "index_info",
242    "index_list",
243    "index_xinfo",
244    "integrity_check",
245    "module_list",
246    "page_count",
247    "page_size",
248    "quick_check",
249    "schema_version",
250    "table_info",
251    "table_list",
252    "table_xinfo",
253    "user_version",
254    "pragma_list",
255];
256
257/// SQLite PRAGMA fast path. Returns the category when the statement is a
258/// pragma: read for allowlisted pragmas without `=`, admin otherwise.
259fn classify_sqlite_pragma(stripped: &str) -> Option<SqlCategory> {
260    let lower = stripped.trim_start().to_ascii_lowercase();
261    let rest = lower.strip_prefix("pragma ")?.trim_start();
262    // Optional schema qualifier: `name.` or quoted forms.
263    let rest = match rest.find('.') {
264        Some(dot)
265            if !rest.starts_with('\'') && !rest.starts_with('"') && !rest.starts_with('[') =>
266        {
267            &rest[dot + 1..]
268        }
269        _ => rest,
270    };
271    let name: String = rest
272        .chars()
273        .take_while(|c| c.is_ascii_alphanumeric() || *c == '_')
274        .collect();
275    if name.is_empty() {
276        return None;
277    }
278    if stripped.contains('=') || !READ_ONLY_PRAGMAS.contains(&name.as_str()) {
279        Some(SqlCategory::Admin)
280    } else {
281        Some(SqlCategory::Read)
282    }
283}
284
285/// Classify one statement. See module docs for the behavioural contract.
286pub fn classify_statement(
287    sql: &str,
288    dialect: Dialect,
289) -> Result<ClassifiedStatement, ClassifyError> {
290    if sql.trim().is_empty() {
291        return Err(ClassifyError::Empty);
292    }
293    let stripped = strip_comments(sql);
294    if stripped.trim().is_empty() {
295        return Err(ClassifyError::CommentOnly);
296    }
297    if looks_like_multiple_statements(sql) {
298        return Err(ClassifyError::MultipleStatements);
299    }
300
301    if dialect == Dialect::SQLite
302        && let Some(category) = classify_sqlite_pragma(&stripped)
303    {
304        return Ok(empty_result(category, "pragma"));
305    }
306
307    if is_tx_keyword(&stripped) {
308        return Ok(empty_result(SqlCategory::TxCtrl, "transaction"));
309    }
310    if is_admin_keyword(&stripped) {
311        let file_io = {
312            let lower = stripped.to_ascii_lowercase();
313            lower.contains("infile") || lower.contains("outfile") || lower.contains("dumpfile")
314        };
315        let mut r = empty_result(SqlCategory::Admin, "admin-keyword");
316        r.file_io = file_io;
317        // Best-effort target extraction for statements the AST layer cannot
318        // express; qualified refs only.
319        r.target_databases = collect_qualified_databases(&stripped);
320        return Ok(r);
321    }
322
323    let statements = match Parser::parse_sql(dialect.as_dyn(), &stripped) {
324        Ok(stmts) => stmts,
325        Err(e) => {
326            // `SELECT … INTO OUTFILE/DUMPFILE` and MySQL `LOCK IN SHARE
327            // MODE` do not parse, but the legacy classifier accepted both;
328            // keep the read category and flag them so the policy gate
329            // applies the stricter authorization (file I/O denied; locking
330            // reads are not ordinary reads).
331            let lower = stripped.to_ascii_lowercase();
332            if lower.contains("into outfile") || lower.contains("into dumpfile") {
333                let mut r = empty_result(SqlCategory::Read, "select");
334                r.file_io = true;
335                r.target_databases = collect_qualified_databases(&stripped);
336                return Ok(r);
337            }
338            if lower.contains("lock in share mode") {
339                let mut r = empty_result(SqlCategory::Read, "select");
340                r.locking_read = true;
341                r.target_databases = collect_qualified_databases(&stripped);
342                return Ok(r);
343            }
344            // MySQL user-variable assignments (`SET @x = …`) do not parse;
345            // the legacy classifier treated every `SET` as admin.
346            if lower.trim_start().starts_with("set ") || lower.trim_start() == "set" {
347                return Ok(empty_result(SqlCategory::Admin, "set"));
348            }
349            return Err(ClassifyError::Parse(e.to_string()));
350        }
351    };
352    if statements.len() > 1 {
353        return Err(ClassifyError::MultipleStatements);
354    }
355    let stmt = statements.into_iter().next().ok_or(ClassifyError::Empty)?;
356    classify_ast(&stmt, sql)
357}
358
359fn empty_result(category: SqlCategory, ast_type: &'static str) -> ClassifiedStatement {
360    ClassifiedStatement {
361        category,
362        ast_type,
363        target_databases: Vec::new(),
364        read_tables: Vec::new(),
365        mutated_tables: Vec::new(),
366        locking_read: false,
367        file_io: false,
368        executes_wrapped: false,
369        if_exists: false,
370        drop_object_type: None,
371    }
372}
373
374/// Textual scan for `db.`-qualified identifiers used for fast-path admin
375/// statements the AST cannot represent. Quoted regions are skipped so
376/// literals like `'/tmp/x.csv'` cannot fabricate database names.
377fn collect_qualified_databases(stripped: &str) -> Vec<String> {
378    let mut out = BTreeSet::new();
379    let mut chars = stripped.chars().peekable();
380    let mut word = String::new();
381    while let Some(c) = chars.next() {
382        if c == '\'' || c == '"' || c == '`' {
383            // Skip the quoted region (backslash-escaped in single quotes).
384            while let Some(qc) = chars.next() {
385                if qc == '\\' && c == '\'' {
386                    chars.next();
387                } else if qc == c {
388                    break;
389                }
390            }
391            word.clear();
392            continue;
393        }
394        if c.is_ascii_alphanumeric() || c == '_' || c == '$' {
395            word.push(c);
396            continue;
397        }
398        if !word.is_empty()
399            && c == '.'
400            && let Some(next) = chars.peek()
401            && (next.is_ascii_alphabetic() || *next == '_' || *next == '`')
402        {
403            out.insert(word.clone());
404        }
405        word.clear();
406    }
407    out.into_iter().collect()
408}
409
410fn object_name_parts(name: &ObjectName) -> (Option<String>, String) {
411    let mut parts = Vec::new();
412    for p in &name.0 {
413        match p {
414            sqlparser::ast::ObjectNamePart::Identifier(ident) => parts.push(ident.value.clone()),
415            sqlparser::ast::ObjectNamePart::Function(_) => {}
416        }
417    }
418    match parts.len() {
419        0 => (None, String::new()),
420        1 => (None, parts.remove(0)),
421        _ => {
422            let table = parts.pop().unwrap_or_default();
423            // Multi-part (db.table for MySQL, schema.table for SQLite).
424            (Some(parts.pop().unwrap_or_default()), table)
425        }
426    }
427}
428
429/// Object-graph collector over one statement.
430#[derive(Default)]
431struct ObjectGraph {
432    read: BTreeSet<TableRef>,
433    cte_names: BTreeSet<String>,
434    locking: bool,
435}
436
437impl ObjectGraph {
438    fn add_table_factor(&mut self, factor: &TableFactor) {
439        match factor {
440            TableFactor::Table { name, .. } => {
441                let (db, table) = object_name_parts(name);
442                if self.cte_names.contains(&table) {
443                    return;
444                }
445                self.read.insert(TableRef {
446                    database: db,
447                    table,
448                });
449            }
450            TableFactor::Derived {
451                lateral, subquery, ..
452            } => {
453                let _ = lateral;
454                self.walk_query(subquery);
455            }
456            TableFactor::NestedJoin {
457                table_with_joins, ..
458            } => {
459                for j in &table_with_joins.joins {
460                    self.add_table_factor(&j.relation);
461                }
462                self.add_table_factor(&table_with_joins.relation);
463            }
464            // Table functions and UNNEST are not authorizable tables; treat
465            // conservatively elsewhere, but do not add a phantom target.
466            _ => {}
467        }
468    }
469
470    fn add_table_with_joins(&mut self, twj: &sqlparser::ast::TableWithJoins) {
471        self.add_table_factor(&twj.relation);
472        for j in &twj.joins {
473            self.add_table_factor(&j.relation);
474        }
475    }
476
477    fn walk_query(&mut self, q: &Query) {
478        if !q.locks.is_empty() {
479            self.locking = true;
480        }
481        if let Some(with) = &q.with {
482            for cte in &with.cte_tables {
483                self.cte_names.insert(cte.alias.name.value.clone());
484                self.walk_query(&cte.query);
485            }
486        }
487        // ORDER BY lives on the Query node (not Select); a scalar
488        // subquery there reads tables too.
489        if let Some(order) = &q.order_by
490            && let sqlparser::ast::OrderByKind::Expressions(exprs) = &order.kind
491        {
492            for o in exprs {
493                self.walk_expr_tables(&o.expr);
494            }
495        }
496        self.walk_set_expr(&q.body);
497    }
498
499    fn walk_set_expr(&mut self, body: &SetExpr) {
500        match body {
501            SetExpr::Select(select) => {
502                for twj in &select.from {
503                    self.add_table_with_joins(twj);
504                }
505                if let Some(expr) = &select.selection {
506                    self.walk_expr_tables(expr);
507                }
508                // Expression contexts beyond WHERE — projection, GROUP BY,
509                // HAVING, ORDER BY — also read tables through scalar
510                // subqueries (`SELECT (SELECT … FROM denied)` used to
511                // bypass table-level read denies entirely).
512                for item in &select.projection {
513                    match item {
514                        sqlparser::ast::SelectItem::UnnamedExpr(e) => {
515                            self.walk_expr_tables(e);
516                        }
517                        sqlparser::ast::SelectItem::ExprWithAlias { expr, .. } => {
518                            self.walk_expr_tables(expr);
519                        }
520                        _ => {}
521                    }
522                }
523                if let sqlparser::ast::GroupByExpr::Expressions(exprs, _) = &select.group_by {
524                    for e in exprs {
525                        self.walk_expr_tables(e);
526                    }
527                }
528                if let Some(having) = &select.having {
529                    self.walk_expr_tables(having);
530                }
531                for o in &select.sort_by {
532                    self.walk_expr_tables(&o.expr);
533                }
534            }
535            SetExpr::Query(q) => self.walk_query(q),
536            SetExpr::SetOperation { left, right, .. } => {
537                self.walk_set_expr(left);
538                self.walk_set_expr(right);
539            }
540            SetExpr::Values(_) | SetExpr::Insert(_) | SetExpr::Update(_) => {}
541            SetExpr::Delete(_) | SetExpr::Merge(_) | SetExpr::Table(_) => {}
542        }
543    }
544
545    /// Subqueries inside expressions (`IN (SELECT …)`, EXISTS, scalar
546    /// subqueries in projections/SET values, subqueries wrapped in
547    /// function calls, CASE arms, casts, …). Every expression-bearing
548    /// shape recurses so a read table can never hide from resolution.
549    fn walk_expr_tables(&mut self, expr: &Expr) {
550        match expr {
551            // Direct subquery carriers.
552            Expr::Subquery(s) => self.walk_query(s),
553            Expr::Exists { subquery, .. } => self.walk_query(subquery),
554            Expr::InSubquery { expr, subquery, .. } => {
555                self.walk_expr_tables(expr);
556                self.walk_query(subquery);
557            }
558            Expr::InUnnest {
559                expr, array_expr, ..
560            } => {
561                self.walk_expr_tables(expr);
562                self.walk_expr_tables(array_expr);
563            }
564            // Recurse through every expression-bearing shape.
565            Expr::BinaryOp { left, right, .. } => {
566                self.walk_expr_tables(left);
567                self.walk_expr_tables(right);
568            }
569            Expr::UnaryOp { expr, .. }
570            | Expr::Nested(expr)
571            | Expr::IsFalse(expr)
572            | Expr::IsNotFalse(expr)
573            | Expr::IsTrue(expr)
574            | Expr::IsNotTrue(expr)
575            | Expr::IsNull(expr)
576            | Expr::IsNotNull(expr)
577            | Expr::IsUnknown(expr)
578            | Expr::IsNotUnknown(expr)
579            | Expr::Cast { expr, .. }
580            | Expr::Convert { expr, .. }
581            | Expr::Extract { expr, .. }
582            | Expr::Ceil { expr, .. }
583            | Expr::Floor { expr, .. }
584            | Expr::Collate { expr, .. }
585            | Expr::CompoundFieldAccess { root: expr, .. }
586            | Expr::AtTimeZone {
587                timestamp: expr, ..
588            }
589            | Expr::Prefixed { value: expr, .. }
590            | Expr::IsNormalized { expr, .. }
591            | Expr::OuterJoin(expr)
592            | Expr::Prior(expr) => self.walk_expr_tables(expr),
593            Expr::IsDistinctFrom(a, b) | Expr::IsNotDistinctFrom(a, b) => {
594                self.walk_expr_tables(a);
595                self.walk_expr_tables(b);
596            }
597            Expr::InList { expr, list, .. } => {
598                self.walk_expr_tables(expr);
599                for e in list {
600                    self.walk_expr_tables(e);
601                }
602            }
603            Expr::Between {
604                expr, low, high, ..
605            } => {
606                self.walk_expr_tables(expr);
607                self.walk_expr_tables(low);
608                self.walk_expr_tables(high);
609            }
610            Expr::Like { expr, pattern, .. }
611            | Expr::ILike { expr, pattern, .. }
612            | Expr::SimilarTo { expr, pattern, .. }
613            | Expr::RLike { expr, pattern, .. } => {
614                self.walk_expr_tables(expr);
615                self.walk_expr_tables(pattern);
616            }
617            Expr::AnyOp { left, right, .. } | Expr::AllOp { left, right, .. } => {
618                self.walk_expr_tables(left);
619                self.walk_expr_tables(right);
620            }
621            Expr::Position { expr, r#in, .. } => {
622                self.walk_expr_tables(expr);
623                self.walk_expr_tables(r#in);
624            }
625            Expr::Substring {
626                expr,
627                substring_from,
628                substring_for,
629                ..
630            } => {
631                self.walk_expr_tables(expr);
632                if let Some(e) = substring_from {
633                    self.walk_expr_tables(e);
634                }
635                if let Some(e) = substring_for {
636                    self.walk_expr_tables(e);
637                }
638            }
639            Expr::Trim {
640                expr,
641                trim_what,
642                trim_characters,
643                ..
644            } => {
645                self.walk_expr_tables(expr);
646                if let Some(e) = trim_what {
647                    self.walk_expr_tables(e);
648                }
649                if let Some(list) = trim_characters {
650                    for e in list {
651                        self.walk_expr_tables(e);
652                    }
653                }
654            }
655            Expr::Overlay {
656                expr,
657                overlay_what,
658                overlay_from,
659                overlay_for,
660                ..
661            } => {
662                self.walk_expr_tables(expr);
663                self.walk_expr_tables(overlay_what);
664                self.walk_expr_tables(overlay_from);
665                if let Some(e) = overlay_for {
666                    self.walk_expr_tables(e);
667                }
668            }
669            Expr::Function(f) => {
670                self.walk_function_arguments(&f.parameters);
671                self.walk_function_arguments(&f.args);
672                if let Some(filter) = &f.filter {
673                    self.walk_expr_tables(filter);
674                }
675            }
676            Expr::Case {
677                operand,
678                conditions,
679                else_result,
680                ..
681            } => {
682                if let Some(e) = operand {
683                    self.walk_expr_tables(e);
684                }
685                for cw in conditions {
686                    self.walk_expr_tables(&cw.condition);
687                    self.walk_expr_tables(&cw.result);
688                }
689                if let Some(e) = else_result {
690                    self.walk_expr_tables(e);
691                }
692            }
693            Expr::GroupingSets(lists) | Expr::Cube(lists) | Expr::Rollup(lists) => {
694                for list in lists {
695                    for e in list {
696                        self.walk_expr_tables(e);
697                    }
698                }
699            }
700            Expr::Tuple(exprs) => {
701                for e in exprs {
702                    self.walk_expr_tables(e);
703                }
704            }
705            Expr::Struct { values, .. } => {
706                for e in values {
707                    self.walk_expr_tables(e);
708                }
709            }
710            Expr::Named { expr, .. } => self.walk_expr_tables(expr),
711            Expr::Map(m) => {
712                for entry in &m.entries {
713                    self.walk_expr_tables(&entry.key);
714                    self.walk_expr_tables(&entry.value);
715                }
716            }
717            Expr::Array(a) => {
718                for e in &a.elem {
719                    self.walk_expr_tables(e);
720                }
721            }
722            Expr::MemberOf(m) => self.walk_expr_tables(&m.value),
723            // No nested expressions (identifiers, literals, wildcards,
724            // intervals, …).
725            _ => {}
726        }
727    }
728
729    fn walk_function_arguments(&mut self, args: &sqlparser::ast::FunctionArguments) {
730        use sqlparser::ast::{FunctionArguments, OrderByKind};
731        match args {
732            FunctionArguments::None => {}
733            FunctionArguments::Subquery(q) => self.walk_query(q),
734            FunctionArguments::List(list) => {
735                for a in &list.args {
736                    match a {
737                        sqlparser::ast::FunctionArg::Named { arg, .. } => {
738                            self.walk_function_arg_expr(arg)
739                        }
740                        sqlparser::ast::FunctionArg::ExprNamed { name, arg, .. } => {
741                            self.walk_expr_tables(name);
742                            self.walk_function_arg_expr(arg);
743                        }
744                        sqlparser::ast::FunctionArg::Unnamed(arg) => {
745                            self.walk_function_arg_expr(arg)
746                        }
747                    }
748                }
749                for clause in &list.clauses {
750                    if let sqlparser::ast::FunctionArgumentClause::OrderBy(exprs) = clause {
751                        for o in exprs {
752                            self.walk_expr_tables(&o.expr);
753                        }
754                    }
755                    let _ = OrderByKind::All;
756                }
757            }
758        }
759    }
760
761    fn walk_function_arg_expr(&mut self, arg: &sqlparser::ast::FunctionArgExpr) {
762        if let sqlparser::ast::FunctionArgExpr::Expr(e) = arg {
763            self.walk_expr_tables(e);
764        }
765    }
766
767    fn finish(self) -> (Vec<TableRef>, Vec<TableRef>, bool) {
768        (self.read.into_iter().collect(), Vec::new(), self.locking)
769    }
770}
771
772fn classify_ast(
773    stmt: &Statement,
774    original_sql: &str,
775) -> Result<ClassifiedStatement, ClassifyError> {
776    let mut r = empty_result(SqlCategory::Read, "select");
777    r.if_exists = stmt_if_exists(stmt);
778    let lower = strip_comments(original_sql).to_ascii_lowercase();
779    r.file_io = lower.contains("into outfile")
780        || lower.contains("into dumpfile")
781        || lower.contains("load data");
782
783    match stmt {
784        Statement::Query(q) => {
785            r.category = SqlCategory::Read;
786            r.ast_type = "select";
787            let (read, _mutated, locking) = {
788                let mut g = ObjectGraph::default();
789                g.walk_query(q);
790                g.finish()
791            };
792            r.read_tables = read;
793            r.locking_read = locking;
794        }
795        Statement::Insert(insert) => {
796            r.category = SqlCategory::Write;
797            r.ast_type = if insert.replace_into {
798                "replace"
799            } else {
800                "insert"
801            };
802            if let sqlparser::ast::TableObject::TableName(name) = &insert.table {
803                let (db, table) = object_name_parts(name);
804                r.mutated_tables.push(TableRef {
805                    database: db,
806                    table,
807                });
808            }
809            if let Some(source) = &insert.source {
810                let mut g = ObjectGraph::default();
811                g.walk_query(source);
812                let (read, _, _) = g.finish();
813                r.read_tables = read;
814            }
815        }
816        Statement::Update(update) => {
817            r.category = SqlCategory::Write;
818            r.ast_type = "update";
819            // MySQL can mutate every table in the FROM list; authorize all
820            // of them as mutations (strictest-wins then applies per table).
821            let mut g = ObjectGraph::default();
822            g.add_table_with_joins(&update.table);
823            if let Some(from) = &update.from {
824                match from {
825                    sqlparser::ast::UpdateTableFromKind::BeforeSet(twjs)
826                    | sqlparser::ast::UpdateTableFromKind::AfterSet(twjs) => {
827                        for twj in twjs {
828                            g.add_table_with_joins(twj);
829                        }
830                    }
831                }
832            }
833            let (mutated, _, _) = g.finish();
834            // WHERE and SET values can carry subqueries that READ other
835            // tables (`SET c = (SELECT … FROM denied)`); collect them on
836            // a separate graph so they authorize as reads, not mutations.
837            let mut rg = ObjectGraph::default();
838            if let Some(sel) = &update.selection {
839                rg.walk_expr_tables(sel);
840            }
841            for a in &update.assignments {
842                rg.walk_expr_tables(&a.value);
843            }
844            let (read, _, _) = rg.finish();
845            r.mutated_tables = mutated;
846            r.read_tables = read;
847        }
848        Statement::Delete(delete) => {
849            r.category = SqlCategory::Write;
850            r.ast_type = "delete";
851            let mut g = ObjectGraph::default();
852            walk_delete_sources(delete, &mut g);
853            let (mutated, _, _) = g.finish();
854            let mut rg = ObjectGraph::default();
855            if let Some(sel) = &delete.selection {
856                rg.walk_expr_tables(sel);
857            }
858            let (read, _, _) = rg.finish();
859            r.mutated_tables = mutated;
860            r.read_tables = read;
861        }
862        Statement::Truncate(trunc) => {
863            if trunc.if_exists {
864                // Neither MySQL 8.4 nor MariaDB support TRUNCATE ... IF
865                // EXISTS (sqlparser accepts it; the servers would reject
866                // it at execution). Deny before any planning.
867                return Err(ClassifyError::Unknown(
868                    "TRUNCATE ... IF EXISTS is not valid MySQL/MariaDB syntax".to_string(),
869                ));
870            }
871            r.category = SqlCategory::Ddl;
872            r.ast_type = "truncate";
873            for target in &trunc.table_names {
874                let (db, table) = object_name_parts(&target.name);
875                r.mutated_tables.push(TableRef {
876                    database: db,
877                    table,
878                });
879            }
880        }
881        Statement::CreateTable(create) => {
882            r.category = SqlCategory::Ddl;
883            r.ast_type = "create";
884            let (db, table) = object_name_parts(&create.name);
885            r.mutated_tables.push(TableRef {
886                database: db,
887                table,
888            });
889            if let Some(q) = &create.query {
890                let mut g = ObjectGraph::default();
891                g.walk_query(q);
892                let (read, _, _) = g.finish();
893                r.read_tables = read;
894            }
895        }
896        Statement::Drop {
897            names, object_type, ..
898        } => {
899            r.category = SqlCategory::Ddl;
900            r.ast_type = "drop";
901            r.drop_object_type = Some(match object_type {
902                sqlparser::ast::ObjectType::Table => "table",
903                sqlparser::ast::ObjectType::View => "view",
904                sqlparser::ast::ObjectType::Index => "index",
905                _ => "other",
906            });
907            for obj in names {
908                let (db, table) = object_name_parts(obj);
909                r.mutated_tables.push(TableRef {
910                    database: db,
911                    table,
912                });
913            }
914        }
915        Statement::AlterTable(alter) => {
916            r.category = SqlCategory::Ddl;
917            r.ast_type = "alter";
918            let (db, table) = object_name_parts(&alter.name);
919            r.mutated_tables.push(TableRef {
920                database: db,
921                table,
922            });
923        }
924        Statement::RenameTable(renames) => {
925            r.category = SqlCategory::Ddl;
926            r.ast_type = "rename";
927            for rn in renames {
928                let (odb, otable) = object_name_parts(&rn.old_name);
929                let (ndb, ntable) = object_name_parts(&rn.new_name);
930                r.mutated_tables.push(TableRef {
931                    database: odb,
932                    table: otable,
933                });
934                r.mutated_tables.push(TableRef {
935                    database: ndb,
936                    table: ntable,
937                });
938            }
939        }
940        Statement::ShowTables { .. }
941        | Statement::ShowDatabases { .. }
942        | Statement::ShowFunctions { .. }
943        | Statement::ShowVariable { .. }
944        | Statement::ShowStatus { .. }
945        | Statement::ShowVariables { .. }
946        | Statement::ShowCreate { .. }
947        | Statement::ShowColumns { .. } => {
948            r.category = SqlCategory::Read;
949            r.ast_type = "show";
950        }
951        Statement::ExplainTable { .. } => {
952            // `DESCRIBE tbl` / `EXPLAIN tbl` / MySQL 8 `EXPLAIN ANALYZE
953            // FOR CONNECTION`-style table describes are reads.
954            r.category = SqlCategory::Read;
955            r.ast_type = "describe";
956        }
957        Statement::Explain {
958            statement: inner,
959            analyze,
960            ..
961        } => {
962            if *analyze {
963                // EXPLAIN ANALYZE executes its wrapped statement.
964                let mut inner = classify_ast(inner, original_sql)?;
965                inner.executes_wrapped = true;
966                return Ok(inner);
967            }
968            r.category = SqlCategory::Read;
969            r.ast_type = "explain";
970        }
971        Statement::Grant { .. } => {
972            r.category = SqlCategory::Admin;
973            r.ast_type = "grant";
974        }
975        Statement::Revoke { .. } => {
976            r.category = SqlCategory::Admin;
977            r.ast_type = "revoke";
978        }
979        Statement::Set(_) => {
980            r.category = SqlCategory::Admin;
981            r.ast_type = "set";
982        }
983        other => {
984            let name = variant_name(other);
985            return Err(ClassifyError::Unknown(name));
986        }
987    }
988
989    r.target_databases = {
990        let mut dbs = BTreeSet::new();
991        for t in r.read_tables.iter().chain(r.mutated_tables.iter()) {
992            if let Some(db) = &t.database {
993                dbs.insert(db.clone());
994            }
995        }
996        dbs.into_iter().collect()
997    };
998    Ok(r)
999}
1000
1001fn walk_delete_sources(delete: &sqlparser::ast::Delete, g: &mut ObjectGraph) {
1002    match &delete.from {
1003        sqlparser::ast::FromTable::WithFromKeyword(twjs)
1004        | sqlparser::ast::FromTable::WithoutKeyword(twjs) => {
1005            for twj in twjs {
1006                g.add_table_with_joins(twj);
1007            }
1008        }
1009    }
1010    if let Some(using) = &delete.using {
1011        for twj in using {
1012            g.add_table_with_joins(twj);
1013        }
1014    }
1015    for t in &delete.tables {
1016        let (db, table) = object_name_parts(t);
1017        g.mutated_extra(db, table);
1018    }
1019}
1020
1021impl ObjectGraph {
1022    fn mutated_extra(&mut self, db: Option<String>, table: String) {
1023        if self.cte_names.contains(&table) {
1024            return;
1025        }
1026        self.read.insert(TableRef {
1027            database: db,
1028            table,
1029        });
1030    }
1031}
1032
1033/// Extract `IF EXISTS` from DROP/TRUNCATE statements.
1034fn stmt_if_exists(stmt: &Statement) -> bool {
1035    match stmt {
1036        Statement::Drop { if_exists, .. } => *if_exists,
1037        Statement::Truncate(t) => t.if_exists,
1038        _ => false,
1039    }
1040}
1041
1042fn variant_name(stmt: &Statement) -> String {
1043    let debug = format!("{stmt:?}");
1044    debug
1045        .split('(')
1046        .next()
1047        .unwrap_or("unknown")
1048        .trim()
1049        .to_lowercase()
1050}
1051
1052#[cfg(test)]
1053mod tests {
1054    use super::*;
1055
1056    fn cat(sql: &str, dialect: Dialect) -> Result<SqlCategory, ClassifyError> {
1057        classify_statement(sql, dialect).map(|c| c.category)
1058    }
1059
1060    #[test]
1061    fn legacy_category_buckets() {
1062        let mysql = Dialect::MySql;
1063        assert_eq!(cat("SELECT 1", mysql).unwrap(), SqlCategory::Read);
1064        assert_eq!(
1065            cat(
1066                "SELECT u.id FROM users u JOIN orders o ON o.user_id = u.id",
1067                mysql
1068            )
1069            .unwrap(),
1070            SqlCategory::Read
1071        );
1072        assert_eq!(cat("SHOW TABLES", mysql).unwrap(), SqlCategory::Read);
1073        assert_eq!(cat("DESCRIBE users", mysql).unwrap(), SqlCategory::Read);
1074        assert_eq!(
1075            cat("EXPLAIN SELECT * FROM users", mysql).unwrap(),
1076            SqlCategory::Read
1077        );
1078        assert_eq!(
1079            cat(
1080                "WITH top AS (SELECT id FROM users ORDER BY id LIMIT 10) SELECT * FROM top",
1081                mysql
1082            )
1083            .unwrap(),
1084            SqlCategory::Read
1085        );
1086        assert_eq!(
1087            cat("INSERT INTO users (id, name) VALUES (1, 'a')", mysql).unwrap(),
1088            SqlCategory::Write
1089        );
1090        assert_eq!(
1091            cat("UPDATE users SET name = 'x' WHERE id = 1", mysql).unwrap(),
1092            SqlCategory::Write
1093        );
1094        assert_eq!(
1095            cat("DELETE FROM users WHERE id = 1", mysql).unwrap(),
1096            SqlCategory::Write
1097        );
1098        assert_eq!(
1099            cat("REPLACE INTO users (id, name) VALUES (1, 'a')", mysql).unwrap(),
1100            SqlCategory::Write
1101        );
1102        assert_eq!(
1103            cat("CREATE TABLE t1 (id INT PRIMARY KEY)", mysql).unwrap(),
1104            SqlCategory::Ddl
1105        );
1106        assert_eq!(cat("DROP TABLE users", mysql).unwrap(), SqlCategory::Ddl);
1107        assert_eq!(
1108            cat("ALTER TABLE users ADD COLUMN email TEXT", mysql).unwrap(),
1109            SqlCategory::Ddl
1110        );
1111        assert_eq!(
1112            cat("TRUNCATE TABLE users", mysql).unwrap(),
1113            SqlCategory::Ddl
1114        );
1115        assert_eq!(cat("RENAME TABLE a TO b", mysql).unwrap(), SqlCategory::Ddl);
1116        assert_eq!(cat("BEGIN", mysql).unwrap(), SqlCategory::TxCtrl);
1117        assert_eq!(cat("COMMIT", mysql).unwrap(), SqlCategory::TxCtrl);
1118        assert_eq!(cat("ROLLBACK", mysql).unwrap(), SqlCategory::TxCtrl);
1119        assert_eq!(
1120            cat("START TRANSACTION", mysql).unwrap(),
1121            SqlCategory::TxCtrl
1122        );
1123        assert_eq!(cat("SAVEPOINT sp1", mysql).unwrap(), SqlCategory::TxCtrl);
1124        assert_eq!(
1125            cat("RELEASE SAVEPOINT sp1", mysql).unwrap(),
1126            SqlCategory::TxCtrl
1127        );
1128        assert_eq!(
1129            cat("GRANT ALL ON *.* TO 'x'@'localhost'", mysql).unwrap(),
1130            SqlCategory::Admin
1131        );
1132        assert_eq!(
1133            cat("SET GLOBAL max_connections = 100", mysql).unwrap(),
1134            SqlCategory::Admin
1135        );
1136        assert_eq!(cat("KILL 42", mysql).unwrap(), SqlCategory::Admin);
1137        assert_eq!(cat("FLUSH TABLES", mysql).unwrap(), SqlCategory::Admin);
1138        assert_eq!(cat("VACUUM", mysql).unwrap(), SqlCategory::Admin);
1139        assert_eq!(
1140            cat("ATTACH DATABASE '/tmp/o.db' AS other", mysql).unwrap(),
1141            SqlCategory::Admin
1142        );
1143        assert_eq!(cat("SET @x = 1", mysql).unwrap(), SqlCategory::Admin);
1144    }
1145
1146    #[test]
1147    fn rejections() {
1148        let mysql = Dialect::MySql;
1149        assert!(matches!(
1150            classify_statement("SELECT 1; SELECT 2", mysql).unwrap_err(),
1151            ClassifyError::MultipleStatements
1152        ));
1153        assert!(matches!(
1154            classify_statement("", mysql).unwrap_err(),
1155            ClassifyError::Empty
1156        ));
1157        assert!(matches!(
1158            classify_statement("   ", mysql).unwrap_err(),
1159            ClassifyError::Empty
1160        ));
1161        assert!(matches!(
1162            classify_statement("-- just a comment", mysql).unwrap_err(),
1163            ClassifyError::CommentOnly
1164        ));
1165        assert!(matches!(
1166            classify_statement("SELECT ';' FROM t", mysql).unwrap(),
1167            classified if classified.category == SqlCategory::Read
1168        ));
1169        assert!(classify_statement("garbage not sql ((", mysql).is_err());
1170    }
1171
1172    #[test]
1173    fn sqlite_pragmas() {
1174        let sqlite = Dialect::SQLite;
1175        assert_eq!(
1176            cat("PRAGMA table_info(users)", sqlite).unwrap(),
1177            SqlCategory::Read
1178        );
1179        assert_eq!(
1180            cat("PRAGMA main.table_info(users)", sqlite).unwrap(),
1181            SqlCategory::Read
1182        );
1183        assert_eq!(
1184            cat("PRAGMA integrity_check", sqlite).unwrap(),
1185            SqlCategory::Read
1186        );
1187        assert_eq!(
1188            cat("PRAGMA user_version = 7", sqlite).unwrap(),
1189            SqlCategory::Admin
1190        );
1191        assert_eq!(
1192            cat("PRAGMA journal_mode = WAL", sqlite).unwrap(),
1193            SqlCategory::Admin
1194        );
1195        assert_eq!(
1196            cat("PRAGMA unknown_thing", sqlite).unwrap(),
1197            SqlCategory::Admin
1198        );
1199    }
1200
1201    #[test]
1202    fn objects_and_flags() {
1203        let mysql = Dialect::MySql;
1204        let c = classify_statement(
1205            "SELECT * FROM app.users WHERE id IN (SELECT uid FROM analytics.events)",
1206            mysql,
1207        )
1208        .unwrap();
1209        assert_eq!(c.read_tables.len(), 2);
1210        assert!(c.read_tables.contains(&TableRef {
1211            database: Some("app".into()),
1212            table: "users".into()
1213        }));
1214        assert!(c.read_tables.contains(&TableRef {
1215            database: Some("analytics".into()),
1216            table: "events".into()
1217        }));
1218        assert_eq!(
1219            c.target_databases,
1220            vec!["analytics".to_string(), "app".to_string()]
1221        );
1222
1223        let c = classify_statement("SELECT * FROM users FOR UPDATE", mysql).unwrap();
1224        assert!(c.locking_read);
1225
1226        let c = classify_statement("SELECT * FROM users INTO OUTFILE '/tmp/x'", mysql).unwrap();
1227        assert!(c.file_io);
1228
1229        let c = classify_statement("LOAD DATA INFILE '/tmp/x' INTO TABLE users", mysql).unwrap();
1230        assert_eq!(c.category, SqlCategory::Admin);
1231        assert!(c.file_io);
1232
1233        let c = classify_statement(
1234            "INSERT INTO app.jobs (id) SELECT id FROM staging.users",
1235            mysql,
1236        )
1237        .unwrap();
1238        assert_eq!(
1239            c.mutated_tables,
1240            vec![TableRef {
1241                database: Some("app".into()),
1242                table: "jobs".into()
1243            }]
1244        );
1245        assert_eq!(
1246            c.read_tables,
1247            vec![TableRef {
1248                database: Some("staging".into()),
1249                table: "users".into()
1250            }]
1251        );
1252
1253        // Multi-table UPDATE: every joined table is a mutation target.
1254        let c = classify_statement(
1255            "UPDATE app.jobs JOIN app.users ON app.users.id = app.jobs.user_id SET app.jobs.state = 'ok'",
1256            mysql,
1257        )
1258        .unwrap();
1259        assert_eq!(c.mutated_tables.len(), 2);
1260
1261        let c = classify_statement(
1262            "DELETE a FROM a JOIN b ON a.id = b.id WHERE b.flag = 1",
1263            mysql,
1264        )
1265        .unwrap();
1266        assert!(c.mutated_tables.iter().any(|t| t.table == "a"));
1267        assert!(c.mutated_tables.iter().any(|t| t.table == "b"));
1268
1269        // CTE names are not tables.
1270        let c = classify_statement(
1271            "WITH top AS (SELECT id FROM users ORDER BY id LIMIT 10) SELECT * FROM top",
1272            mysql,
1273        )
1274        .unwrap();
1275        assert_eq!(c.read_tables.len(), 1);
1276        assert_eq!(c.read_tables[0].table, "users");
1277
1278        let c = classify_statement("EXPLAIN ANALYZE SELECT * FROM users", mysql).unwrap();
1279        assert!(c.executes_wrapped);
1280        assert_eq!(c.category, SqlCategory::Read);
1281    }
1282
1283    /// Review finding (2026-08-23 independent review): expression-context
1284    /// subqueries must count as read tables, or table-level read denies
1285    /// are bypassable (`SELECT (SELECT … FROM denied)`).
1286    #[test]
1287    fn expression_context_subqueries_are_read_tables() {
1288        let mysql = Dialect::MySql;
1289        let secret = TableRef {
1290            database: Some("secrets".into()),
1291            table: "tokens".into(),
1292        };
1293        let cases = [
1294            // Scalar subquery in the projection.
1295            "SELECT (SELECT token FROM secrets.tokens) AS x",
1296            // Subquery wrapped in a function call.
1297            "SELECT UPPER((SELECT token FROM secrets.tokens)) AS x",
1298            // CASE with an EXISTS subquery.
1299            "SELECT CASE WHEN EXISTS (SELECT 1 FROM secrets.tokens) THEN 1 ELSE 2 END AS x",
1300            // Subquery in GROUP BY / HAVING / ORDER BY.
1301            "SELECT 1 FROM app.users GROUP BY (SELECT token FROM secrets.tokens)",
1302            "SELECT 1 FROM app.users HAVING COUNT(*) > (SELECT COUNT(*) FROM secrets.tokens)",
1303            "SELECT id FROM app.users ORDER BY (SELECT token FROM secrets.tokens)",
1304            // ANY/ALL comparisons.
1305            "SELECT id FROM app.users WHERE id = ANY (SELECT uid FROM secrets.tokens)",
1306        ];
1307        for sql in cases {
1308            let c = classify_statement(sql, mysql)
1309                .unwrap_or_else(|e| panic!("{sql}: must classify: {e:?}"));
1310            assert!(
1311                c.read_tables.contains(&secret),
1312                "{sql}: read_tables must include secrets.tokens (got {:?})",
1313                c.read_tables
1314            );
1315        }
1316    }
1317
1318    /// UPDATE SET values / WHERE and DELETE WHERE subqueries authorize as
1319    /// READS of the subquery table while the mutation targets stay exact.
1320    #[test]
1321    fn update_delete_subqueries_authorize_as_reads() {
1322        let mysql = Dialect::MySql;
1323        let secret = TableRef {
1324            database: Some("secrets".into()),
1325            table: "tokens".into(),
1326        };
1327        let jobs = TableRef {
1328            database: Some("app".into()),
1329            table: "jobs".into(),
1330        };
1331
1332        let c = classify_statement(
1333            "UPDATE app.jobs SET note = (SELECT token FROM secrets.tokens LIMIT 1)",
1334            mysql,
1335        )
1336        .unwrap();
1337        assert_eq!(c.mutated_tables, vec![jobs.clone()]);
1338        assert!(
1339            c.read_tables.contains(&secret),
1340            "SET-value subquery must be a read table (got {:?})",
1341            c.read_tables
1342        );
1343
1344        let c = classify_statement(
1345            "UPDATE app.jobs SET note = 'x' WHERE id IN (SELECT id FROM secrets.tokens)",
1346            mysql,
1347        )
1348        .unwrap();
1349        assert_eq!(c.mutated_tables, vec![jobs.clone()]);
1350        assert!(c.read_tables.contains(&secret));
1351
1352        let c = classify_statement(
1353            "DELETE FROM app.jobs WHERE id IN (SELECT id FROM secrets.tokens)",
1354            mysql,
1355        )
1356        .unwrap();
1357        assert_eq!(c.mutated_tables, vec![jobs]);
1358        assert!(c.read_tables.contains(&secret));
1359    }
1360
1361    /// Differential check against the legacy fixture corpus: categories and
1362    /// target databases must match for every case the legacy classifier
1363    /// accepted; error classes must match for structural errors.
1364    #[test]
1365    fn matches_legacy_classifier_fixtures() {
1366        let path = concat!(
1367            env!("CARGO_MANIFEST_DIR"),
1368            "/tests/fixtures/legacy/classifier.json"
1369        );
1370        // The fixture corpus is generated from the legacy checkout and is
1371        // deliberately untracked; skip where it is absent (fresh CI
1372        // checkouts) instead of failing.
1373        let Ok(data) = std::fs::read_to_string(path) else {
1374            eprintln!("skipping: legacy fixture corpus not present ({path})");
1375            return;
1376        };
1377        let cases: Vec<serde_json::Value> = serde_json::from_str(&data).unwrap();
1378        let mut checked = 0;
1379        // Legacy parser limitations the Rust parser deliberately improves on:
1380        // these statements failed to parse under node-sql-parser.
1381        let legacy_parse_failures = [
1382            "EXPLAIN ANALYZE SELECT * FROM users",
1383            "UPDATE app.jobs JOIN app.users ON app.users.id = app.jobs.user_id SET app.jobs.state = 'ok'",
1384        ];
1385        for case in cases {
1386            let sql = case["sql"].as_str().unwrap();
1387            let dialect = match case["dialect"].as_str().unwrap() {
1388                "sqlite" => Dialect::SQLite,
1389                _ => Dialect::MySql,
1390            };
1391            let legacy = &case["result"];
1392            let ours = classify_statement(sql, dialect);
1393            if legacy_parse_failures.contains(&sql) {
1394                continue;
1395            }
1396            if legacy["ok"].as_bool().unwrap_or(false) {
1397                let c = ours.unwrap_or_else(|e| panic!("rust rejected {sql:?}: {e:?}"));
1398                assert_eq!(
1399                    c.category.as_str(),
1400                    legacy["category"].as_str().unwrap(),
1401                    "category mismatch for {sql:?}"
1402                );
1403                let want_dbs: Vec<&str> = legacy["targetDatabases"]
1404                    .as_array()
1405                    .unwrap()
1406                    .iter()
1407                    .map(|v| v.as_str().unwrap())
1408                    .collect();
1409                let got_dbs: Vec<&str> = c.target_databases.iter().map(String::as_str).collect();
1410                assert_eq!(got_dbs, want_dbs, "databases mismatch for {sql:?}");
1411                checked += 1;
1412            } else {
1413                // Structural errors must match; parser message texts differ
1414                // by design (different parser).
1415                let err_msg = legacy["error"].as_str().unwrap_or_default();
1416                let structural = err_msg.starts_with("multiple statements")
1417                    || err_msg.starts_with("empty input")
1418                    || err_msg.starts_with("input contains only comments");
1419                if structural {
1420                    let e = ours.expect_err("rust accepted what legacy structurally rejected");
1421                    let msg = e.message();
1422                    assert!(
1423                        msg.starts_with(&err_msg[..err_msg.len().min(30)]),
1424                        "{sql:?}: {msg}"
1425                    );
1426                    checked += 1;
1427                }
1428            }
1429        }
1430        assert!(checked > 100, "fixture coverage collapsed: {checked}");
1431    }
1432}