1use sqlparser::dialect::Dialect as SqlParserDialect;
32use sqlparser::parser::Parser;
33
34#[derive(Debug, Clone)]
36pub struct VerifyResult {
37 pub is_valid: bool,
39 pub errors: Vec<String>,
41 pub sql: String,
43}
44
45impl VerifyResult {
46 pub fn ok(sql: &str) -> Self {
48 Self {
49 is_valid: true,
50 errors: Vec::new(),
51 sql: sql.to_string(),
52 }
53 }
54
55 pub fn fail(sql: &str, errors: Vec<String>) -> Self {
57 Self {
58 is_valid: false,
59 errors,
60 sql: sql.to_string(),
61 }
62 }
63
64 pub fn push_error(&mut self, error: String) {
66 self.is_valid = false;
67 self.errors.push(error);
68 }
69}
70
71#[derive(Debug, Clone, Copy, PartialEq, Eq)]
73pub enum VerifyDialect {
74 MySql,
76 PostgreSql,
78 Sqlite,
80}
81
82#[derive(Debug, Clone, Copy, PartialEq, Eq)]
84pub enum SqlPath {
85 Select,
87 Insert,
89 Update,
91 Delete,
93 Join,
95 Subquery,
97 Cte,
99 WindowFunction,
101 Unknown,
103}
104
105impl SqlPath {
106 pub fn name(self) -> &'static str {
108 match self {
109 SqlPath::Select => "SELECT",
110 SqlPath::Insert => "INSERT",
111 SqlPath::Update => "UPDATE",
112 SqlPath::Delete => "DELETE",
113 SqlPath::Join => "JOIN",
114 SqlPath::Subquery => "Subquery",
115 SqlPath::Cte => "CTE",
116 SqlPath::WindowFunction => "WindowFunction",
117 SqlPath::Unknown => "Unknown",
118 }
119 }
120}
121
122#[derive(Debug, Clone, Copy, PartialEq, Eq)]
124pub enum VerifyMode {
125 Full,
127 SyntaxOnly,
129}
130
131pub fn verify_sql_syntax(sql: &str, dialect: VerifyDialect) -> VerifyResult {
135 let parser_dialect: &dyn SqlParserDialect = match dialect {
136 VerifyDialect::MySql => &sqlparser::dialect::MySqlDialect {},
137 VerifyDialect::PostgreSql => &sqlparser::dialect::PostgreSqlDialect {},
138 VerifyDialect::Sqlite => &sqlparser::dialect::SQLiteDialect {},
139 };
140
141 match Parser::parse_sql(parser_dialect, sql) {
142 Ok(stmts) => {
143 if stmts.is_empty() {
144 VerifyResult::fail(sql, vec!["SQL 解析结果为空".to_string()])
145 } else {
146 VerifyResult::ok(sql)
147 }
148 }
149 Err(e) => VerifyResult::fail(sql, vec![format!("SQL 语法错误: {}", e)]),
150 }
151}
152
153pub fn sql_hash(sql: &str) -> u64 {
155 {
156 use std::hash::Hasher;
157 let mut h = twox_hash::XxHash64::with_seed(0);
158 h.write(sql.as_bytes());
159 h.finish()
160 }
161}
162
163pub fn is_read_only(sql: &str) -> bool {
165 let upper = sql.trim().to_uppercase();
166 upper.starts_with("SELECT") || upper.starts_with("EXPLAIN") || upper.starts_with("WITH")
167}
168
169pub fn classify_sql_path(sql: &str) -> SqlPath {
174 let upper = sql.trim().to_uppercase();
175
176 if upper.starts_with("WITH") {
177 return SqlPath::Cte;
178 }
179
180 let dialect = if upper.contains("$1") {
181 &sqlparser::dialect::PostgreSqlDialect {} as &dyn SqlParserDialect
182 } else {
183 &sqlparser::dialect::MySqlDialect {} as &dyn SqlParserDialect
184 };
185
186 let stmts = Parser::parse_sql(dialect, sql).unwrap_or_default();
187 if stmts.is_empty() {
188 return SqlPath::Unknown;
189 }
190
191 let stmt = &stmts[0];
192 use sqlparser::ast::Statement;
193
194 match stmt {
195 Statement::Query(query) => {
196 let has_join = set_expr_has_join(&query.body);
197 let has_subquery = set_expr_has_subquery(&query.body);
198 let has_window = set_expr_has_window(&query.body);
199
200 if has_window {
201 SqlPath::WindowFunction
202 } else if has_join {
203 SqlPath::Join
204 } else if has_subquery {
205 SqlPath::Subquery
206 } else {
207 SqlPath::Select
208 }
209 }
210 Statement::Insert(_) => SqlPath::Insert,
211 Statement::Update { .. } => SqlPath::Update,
212 Statement::Delete(_) => SqlPath::Delete,
213 _ => SqlPath::Unknown,
214 }
215}
216
217fn set_expr_has_join(body: &sqlparser::ast::SetExpr) -> bool {
219 use sqlparser::ast::SetExpr;
220 match body {
221 SetExpr::Select(select) => select.from.iter().any(|t| !t.joins.is_empty()),
222 SetExpr::SetOperation { left, right, .. } => {
223 set_expr_has_join(left) || set_expr_has_join(right)
224 }
225 SetExpr::Query(query) => set_expr_has_join(&query.body),
226 _ => false,
227 }
228}
229
230fn set_expr_has_subquery(body: &sqlparser::ast::SetExpr) -> bool {
232 use sqlparser::ast::SetExpr;
233 match body {
234 SetExpr::Select(select) => {
235 select.from.iter().any(|t| table_with_joins_has_subquery(t))
236 || select.selection.as_ref().map_or(false, expr_has_subquery)
237 || select
238 .projection
239 .iter()
240 .any(|p| select_item_has_subquery(p))
241 }
242 SetExpr::SetOperation { left, right, .. } => {
243 set_expr_has_subquery(left) || set_expr_has_subquery(right)
244 }
245 SetExpr::Query(_) => true,
246 _ => false,
247 }
248}
249
250fn set_expr_has_window(body: &sqlparser::ast::SetExpr) -> bool {
252 use sqlparser::ast::SetExpr;
253 match body {
254 SetExpr::Select(select) => select.projection.iter().any(|p| match p {
255 sqlparser::ast::SelectItem::UnnamedExpr(expr) => expr_has_window(expr),
256 sqlparser::ast::SelectItem::ExprWithAlias { expr, .. } => expr_has_window(expr),
257 _ => false,
258 }),
259 SetExpr::SetOperation { left, right, .. } => {
260 set_expr_has_window(left) || set_expr_has_window(right)
261 }
262 SetExpr::Query(query) => set_expr_has_window(&query.body),
263 _ => false,
264 }
265}
266
267fn table_with_joins_has_subquery(table_with_joins: &sqlparser::ast::TableWithJoins) -> bool {
269 table_factor_is_subquery(&table_with_joins.relation)
270 || table_with_joins
271 .joins
272 .iter()
273 .any(|j| table_factor_is_subquery(&j.relation))
274}
275
276fn table_factor_is_subquery(table: &sqlparser::ast::TableFactor) -> bool {
278 matches!(table, sqlparser::ast::TableFactor::Derived { .. })
279}
280
281fn expr_has_subquery(expr: &sqlparser::ast::Expr) -> bool {
283 use sqlparser::ast::Expr;
284 match expr {
285 Expr::Subquery(_) | Expr::Exists { .. } | Expr::InSubquery { .. } => true,
286 Expr::BinaryOp { left, right, .. } => expr_has_subquery(left) || expr_has_subquery(right),
287 Expr::UnaryOp { expr, .. } => expr_has_subquery(expr),
288 Expr::Function(func) => function_has_subquery(func),
289 _ => false,
290 }
291}
292
293fn function_has_subquery(func: &sqlparser::ast::Function) -> bool {
295 use sqlparser::ast::{FunctionArg, FunctionArgExpr, FunctionArguments};
296 match &func.args {
297 FunctionArguments::Subquery(_) => true,
298 FunctionArguments::List(list) => list.args.iter().any(|a| match a {
299 FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) => expr_has_subquery(expr),
300 FunctionArg::Named {
301 arg: FunctionArgExpr::Expr(expr),
302 ..
303 } => expr_has_subquery(expr),
304 _ => false,
305 }),
306 FunctionArguments::None => false,
307 }
308}
309
310fn select_item_has_subquery(item: &sqlparser::ast::SelectItem) -> bool {
312 match item {
313 sqlparser::ast::SelectItem::UnnamedExpr(expr) => expr_has_subquery(expr),
314 sqlparser::ast::SelectItem::ExprWithAlias { expr, .. } => expr_has_subquery(expr),
315 _ => false,
316 }
317}
318
319fn expr_has_window(expr: &sqlparser::ast::Expr) -> bool {
321 use sqlparser::ast::Expr;
322 match expr {
323 Expr::Function(func) => func.over.is_some() || function_has_window(func),
324 Expr::BinaryOp { left, right, .. } => expr_has_window(left) || expr_has_window(right),
325 Expr::UnaryOp { expr, .. } => expr_has_window(expr),
326 _ => false,
327 }
328}
329
330fn function_has_window(func: &sqlparser::ast::Function) -> bool {
332 use sqlparser::ast::{FunctionArg, FunctionArgExpr, FunctionArguments};
333 match &func.args {
334 FunctionArguments::List(list) => list.args.iter().any(|a| match a {
335 FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) => expr_has_window(expr),
336 FunctionArg::Named {
337 arg: FunctionArgExpr::Expr(expr),
338 ..
339 } => expr_has_window(expr),
340 _ => false,
341 }),
342 _ => false,
343 }
344}
345
346pub fn build_explain_sql(sql: &str, dialect: VerifyDialect) -> String {
351 match dialect {
352 VerifyDialect::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql),
353 VerifyDialect::MySql | VerifyDialect::PostgreSql => format!("EXPLAIN {}", sql),
354 }
355}
356
357pub fn is_db_verify_enabled() -> bool {
363 let verify_flag = std::env::var("SZ_ORM_QUERY_VERIFY").unwrap_or_default();
364 let database_url = std::env::var("DATABASE_URL").unwrap_or_default();
365 (verify_flag == "1" || verify_flag.eq_ignore_ascii_case("true")) && !database_url.is_empty()
366}
367
368pub fn current_verify_mode() -> VerifyMode {
372 if is_db_verify_enabled() {
373 VerifyMode::Full
374 } else {
375 VerifyMode::SyntaxOnly
376 }
377}
378
379pub fn verify_degraded(sql: &str, dialect: VerifyDialect) -> VerifyResult {
384 let mut result = verify_sql_syntax(sql, dialect);
385 if result.is_valid {
386 result.push_error(
387 "warning: sql-verify-proc degraded to syntax-only (DATABASE_URL not set)".to_string(),
388 );
389 result.is_valid = true;
390 result.errors.clear();
391 }
392 result
393}
394
395pub fn verify_full(sql: &str, dialect: VerifyDialect) -> VerifyResult {
406 let mut result = verify_sql_syntax(sql, dialect);
407 if !result.is_valid {
408 return result;
409 }
410
411 let path = classify_sql_path(sql);
412 if path == SqlPath::Unknown {
413 result.push_error(format!("无法识别 SQL 路径分类: {}", sql));
414 return result;
415 }
416
417 let _explain_sql = build_explain_sql(sql, dialect);
418 result
419}
420
421pub fn verify_smart(sql: &str, dialect: VerifyDialect) -> VerifyResult {
427 if is_db_verify_enabled() {
428 verify_full(sql, dialect)
429 } else {
430 verify_degraded(sql, dialect)
431 }
432}
433
434pub fn check_path_coverage(sqls: &[&str]) -> Vec<SqlPath> {
439 let mut covered = Vec::new();
440 for sql in sqls {
441 let path = classify_sql_path(sql);
442 if !covered.contains(&path) {
443 covered.push(path);
444 }
445 }
446
447 let all_paths = [
448 SqlPath::Select,
449 SqlPath::Insert,
450 SqlPath::Update,
451 SqlPath::Delete,
452 SqlPath::Join,
453 SqlPath::Subquery,
454 SqlPath::Cte,
455 SqlPath::WindowFunction,
456 ];
457
458 all_paths
459 .iter()
460 .filter(|p| !covered.contains(p))
461 .copied()
462 .collect()
463}
464
465#[cfg(test)]
466mod tests {
467 use super::*;
468
469 #[test]
470 fn test_verify_valid_select() {
471 let sql = "SELECT id, name FROM users WHERE id = 1";
472 let result = verify_sql_syntax(sql, VerifyDialect::MySql);
473 assert!(
474 result.is_valid,
475 "Valid SELECT should pass: {:?}",
476 result.errors
477 );
478 }
479
480 #[test]
481 fn test_verify_invalid_sql() {
482 let sql = "SELECT FROM WHERE";
483 let result = verify_sql_syntax(sql, VerifyDialect::MySql);
484 assert!(!result.is_valid);
485 assert!(!result.errors.is_empty());
486 }
487
488 #[test]
489 fn test_verify_empty_sql() {
490 let sql = "";
491 let result = verify_sql_syntax(sql, VerifyDialect::MySql);
492 assert!(!result.is_valid);
493 }
494
495 #[test]
496 fn test_sql_hash_deterministic() {
497 let sql = "SELECT * FROM users";
498 assert_eq!(sql_hash(sql), sql_hash(sql));
499 }
500
501 #[test]
502 fn test_sql_hash_different() {
503 let sql1 = "SELECT * FROM users";
504 let sql2 = "SELECT * FROM posts";
505 assert_ne!(sql_hash(sql1), sql_hash(sql2));
506 }
507
508 #[test]
509 fn test_is_read_only() {
510 assert!(is_read_only("SELECT * FROM users"));
511 assert!(is_read_only("EXPLAIN SELECT * FROM users"));
512 assert!(is_read_only("WITH cte AS (SELECT 1) SELECT * FROM cte"));
513 assert!(!is_read_only("INSERT INTO users VALUES (1)"));
514 assert!(!is_read_only("UPDATE users SET name = 'x'"));
515 assert!(!is_read_only("DELETE FROM users"));
516 assert!(!is_read_only("DROP TABLE users"));
517 }
518
519 #[test]
520 fn test_verify_postgresql_dialect() {
521 let sql = "SELECT id, name FROM users WHERE id = $1";
522 let result = verify_sql_syntax(sql, VerifyDialect::PostgreSql);
523 assert!(
524 result.is_valid,
525 "PG dialect should parse $1 params: {:?}",
526 result.errors
527 );
528 }
529
530 #[test]
531 fn test_verify_sqlite_dialect() {
532 let sql = "SELECT id, name FROM users WHERE id = ?";
533 let result = verify_sql_syntax(sql, VerifyDialect::Sqlite);
534 assert!(
535 result.is_valid,
536 "SQLite dialect should parse ? params: {:?}",
537 result.errors
538 );
539 }
540
541 #[test]
542 fn test_classify_select_path() {
543 let sql = "SELECT id, name FROM users WHERE id = 1";
544 assert_eq!(classify_sql_path(sql), SqlPath::Select);
545 }
546
547 #[test]
548 fn test_classify_insert_path() {
549 let sql = "INSERT INTO users (id, name) VALUES (1, 'Alice')";
550 assert_eq!(classify_sql_path(sql), SqlPath::Insert);
551 }
552
553 #[test]
554 fn test_classify_update_path() {
555 let sql = "UPDATE users SET name = 'Bob' WHERE id = 1";
556 assert_eq!(classify_sql_path(sql), SqlPath::Update);
557 }
558
559 #[test]
560 fn test_classify_delete_path() {
561 let sql = "DELETE FROM users WHERE id = 1";
562 assert_eq!(classify_sql_path(sql), SqlPath::Delete);
563 }
564
565 #[test]
566 fn test_classify_join_path() {
567 let sql = "SELECT u.name, p.title FROM users u INNER JOIN posts p ON u.id = p.user_id";
568 assert_eq!(classify_sql_path(sql), SqlPath::Join);
569 }
570
571 #[test]
572 fn test_classify_left_join_path() {
573 let sql = "SELECT u.name FROM users u LEFT JOIN posts p ON u.id = p.user_id";
574 assert_eq!(classify_sql_path(sql), SqlPath::Join);
575 }
576
577 #[test]
578 fn test_classify_cte_path() {
579 let sql = "WITH cte AS (SELECT id FROM users) SELECT * FROM cte";
580 assert_eq!(classify_sql_path(sql), SqlPath::Cte);
581 }
582
583 #[test]
584 fn test_classify_subquery_in_where() {
585 let sql = "SELECT * FROM users WHERE id IN (SELECT user_id FROM posts)";
586 assert_eq!(classify_sql_path(sql), SqlPath::Subquery);
587 }
588
589 #[test]
590 fn test_classify_subquery_in_from() {
591 let sql = "SELECT * FROM (SELECT id FROM users) AS sub";
592 assert_eq!(classify_sql_path(sql), SqlPath::Subquery);
593 }
594
595 #[test]
596 fn test_classify_window_function_path() {
597 let sql = "SELECT id, ROW_NUMBER() OVER (PARTITION BY dept ORDER BY salary) FROM employees";
598 assert_eq!(classify_sql_path(sql), SqlPath::WindowFunction);
599 }
600
601 #[test]
602 fn test_build_explain_mysql() {
603 let sql = "SELECT * FROM users";
604 let explain = build_explain_sql(sql, VerifyDialect::MySql);
605 assert_eq!(explain, "EXPLAIN SELECT * FROM users");
606 }
607
608 #[test]
609 fn test_build_explain_postgres() {
610 let sql = "SELECT * FROM users";
611 let explain = build_explain_sql(sql, VerifyDialect::PostgreSql);
612 assert_eq!(explain, "EXPLAIN SELECT * FROM users");
613 }
614
615 #[test]
616 fn test_build_explain_sqlite() {
617 let sql = "SELECT * FROM users";
618 let explain = build_explain_sql(sql, VerifyDialect::Sqlite);
619 assert_eq!(explain, "EXPLAIN QUERY PLAN SELECT * FROM users");
620 }
621
622 #[test]
623 fn test_verify_degraded_no_env() {
624 std::env::remove_var("SZ_ORM_QUERY_VERIFY");
625 std::env::remove_var("DATABASE_URL");
626 let sql = "SELECT * FROM users";
627 let result = verify_degraded(sql, VerifyDialect::MySql);
628 assert!(result.is_valid);
629 }
630
631 #[test]
632 fn test_verify_degraded_invalid_sql() {
633 std::env::remove_var("SZ_ORM_QUERY_VERIFY");
634 std::env::remove_var("DATABASE_URL");
635 let sql = "SELECT FROM WHERE";
636 let result = verify_degraded(sql, VerifyDialect::MySql);
637 assert!(!result.is_valid);
638 }
639
640 #[test]
641 fn test_verify_full_valid_select() {
642 let sql = "SELECT id, name FROM users WHERE id = 1";
643 let result = verify_full(sql, VerifyDialect::MySql);
644 assert!(
645 result.is_valid,
646 "verify_full should pass: {:?}",
647 result.errors
648 );
649 }
650
651 #[test]
652 fn test_verify_full_invalid_sql() {
653 let sql = "SELECT FROM WHERE";
654 let result = verify_full(sql, VerifyDialect::MySql);
655 assert!(!result.is_valid);
656 }
657
658 #[test]
659 fn test_verify_smart_degraded_mode() {
660 std::env::remove_var("SZ_ORM_QUERY_VERIFY");
661 std::env::remove_var("DATABASE_URL");
662 let sql = "SELECT * FROM users";
663 let result = verify_smart(sql, VerifyDialect::MySql);
664 assert!(result.is_valid);
665 }
666
667 #[test]
668 fn test_check_path_coverage_all_covered() {
669 let sqls = [
670 "SELECT * FROM users",
671 "INSERT INTO users VALUES (1)",
672 "UPDATE users SET name = 'x'",
673 "DELETE FROM users",
674 "SELECT * FROM a JOIN b ON a.id = b.id",
675 "SELECT * FROM users WHERE id IN (SELECT id FROM posts)",
676 "WITH cte AS (SELECT 1) SELECT * FROM cte",
677 "SELECT ROW_NUMBER() OVER (PARTITION BY x) FROM t",
678 ];
679 let uncovered = check_path_coverage(&sqls);
680 assert!(
681 uncovered.is_empty(),
682 "All paths should be covered, uncovered: {:?}",
683 uncovered.iter().map(|p| p.name()).collect::<Vec<_>>()
684 );
685 }
686
687 #[test]
688 fn test_check_path_coverage_partial() {
689 let sqls = ["SELECT * FROM users", "INSERT INTO users VALUES (1)"];
690 let uncovered = check_path_coverage(&sqls);
691 assert!(uncovered.contains(&SqlPath::Update));
692 assert!(uncovered.contains(&SqlPath::Delete));
693 assert!(uncovered.contains(&SqlPath::Join));
694 assert!(uncovered.contains(&SqlPath::Cte));
695 }
696
697 #[test]
698 fn test_sql_path_name() {
699 assert_eq!(SqlPath::Select.name(), "SELECT");
700 assert_eq!(SqlPath::Insert.name(), "INSERT");
701 assert_eq!(SqlPath::Join.name(), "JOIN");
702 assert_eq!(SqlPath::WindowFunction.name(), "WindowFunction");
703 }
704
705 #[test]
706 fn test_verify_result_push_error() {
707 let mut result = VerifyResult::ok("SELECT 1");
708 assert!(result.is_valid);
709 result.push_error("test error".to_string());
710 assert!(!result.is_valid);
711 assert_eq!(result.errors.len(), 1);
712 }
713
714 #[test]
715 fn test_verify_mode_enum() {
716 let full = VerifyMode::Full;
717 let syntax = VerifyMode::SyntaxOnly;
718 assert_ne!(full, syntax);
719 }
720}