1use 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 pub database: Option<String>,
18 pub table: String,
19}
20
21#[derive(Debug, Clone, PartialEq)]
22pub struct ClassifiedStatement {
23 pub category: SqlCategory,
24 pub ast_type: &'static str,
29 pub target_databases: Vec<String>,
32 pub read_tables: Vec<TableRef>,
34 pub mutated_tables: Vec<TableRef>,
37 pub locking_read: bool,
39 pub file_io: bool,
43 pub executes_wrapped: bool,
45 pub if_exists: bool,
48 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 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
94pub 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 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
123pub 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
231const 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
257fn 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 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
285pub 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 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 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 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
374fn 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 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 (Some(parts.pop().unwrap_or_default()), table)
425 }
426 }
427}
428
429#[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 _ => {}
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 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 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 fn walk_expr_tables(&mut self, expr: &Expr) {
550 match expr {
551 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 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 _ => {}
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 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 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 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 r.category = SqlCategory::Read;
955 r.ast_type = "describe";
956 }
957 Statement::Explain {
958 statement: inner,
959 analyze,
960 ..
961 } => {
962 if *analyze {
963 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
1033fn 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 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 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 #[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 "SELECT (SELECT token FROM secrets.tokens) AS x",
1296 "SELECT UPPER((SELECT token FROM secrets.tokens)) AS x",
1298 "SELECT CASE WHEN EXISTS (SELECT 1 FROM secrets.tokens) THEN 1 ELSE 2 END AS x",
1300 "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 "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 #[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 #[test]
1365 fn matches_legacy_classifier_fixtures() {
1366 let path = concat!(
1367 env!("CARGO_MANIFEST_DIR"),
1368 "/tests/fixtures/legacy/classifier.json"
1369 );
1370 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 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 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}