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(table_with_joins_has_subquery)
236 || select.selection.as_ref().is_some_and(expr_has_subquery)
237 || select.projection.iter().any(select_item_has_subquery)
238 }
239 SetExpr::SetOperation { left, right, .. } => {
240 set_expr_has_subquery(left) || set_expr_has_subquery(right)
241 }
242 SetExpr::Query(_) => true,
243 _ => false,
244 }
245}
246
247fn set_expr_has_window(body: &sqlparser::ast::SetExpr) -> bool {
249 use sqlparser::ast::SetExpr;
250 match body {
251 SetExpr::Select(select) => select.projection.iter().any(|p| match p {
252 sqlparser::ast::SelectItem::UnnamedExpr(expr) => expr_has_window(expr),
253 sqlparser::ast::SelectItem::ExprWithAlias { expr, .. } => expr_has_window(expr),
254 _ => false,
255 }),
256 SetExpr::SetOperation { left, right, .. } => {
257 set_expr_has_window(left) || set_expr_has_window(right)
258 }
259 SetExpr::Query(query) => set_expr_has_window(&query.body),
260 _ => false,
261 }
262}
263
264fn table_with_joins_has_subquery(table_with_joins: &sqlparser::ast::TableWithJoins) -> bool {
266 table_factor_is_subquery(&table_with_joins.relation)
267 || table_with_joins
268 .joins
269 .iter()
270 .any(|j| table_factor_is_subquery(&j.relation))
271}
272
273fn table_factor_is_subquery(table: &sqlparser::ast::TableFactor) -> bool {
275 matches!(table, sqlparser::ast::TableFactor::Derived { .. })
276}
277
278fn expr_has_subquery(expr: &sqlparser::ast::Expr) -> bool {
280 use sqlparser::ast::Expr;
281 match expr {
282 Expr::Subquery(_) | Expr::Exists { .. } | Expr::InSubquery { .. } => true,
283 Expr::BinaryOp { left, right, .. } => expr_has_subquery(left) || expr_has_subquery(right),
284 Expr::UnaryOp { expr, .. } => expr_has_subquery(expr),
285 Expr::Function(func) => function_has_subquery(func),
286 _ => false,
287 }
288}
289
290fn function_has_subquery(func: &sqlparser::ast::Function) -> bool {
292 use sqlparser::ast::{FunctionArg, FunctionArgExpr, FunctionArguments};
293 match &func.args {
294 FunctionArguments::Subquery(_) => true,
295 FunctionArguments::List(list) => list.args.iter().any(|a| match a {
296 FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) => expr_has_subquery(expr),
297 FunctionArg::Named {
298 arg: FunctionArgExpr::Expr(expr),
299 ..
300 } => expr_has_subquery(expr),
301 _ => false,
302 }),
303 FunctionArguments::None => false,
304 }
305}
306
307fn select_item_has_subquery(item: &sqlparser::ast::SelectItem) -> bool {
309 match item {
310 sqlparser::ast::SelectItem::UnnamedExpr(expr) => expr_has_subquery(expr),
311 sqlparser::ast::SelectItem::ExprWithAlias { expr, .. } => expr_has_subquery(expr),
312 _ => false,
313 }
314}
315
316fn expr_has_window(expr: &sqlparser::ast::Expr) -> bool {
318 use sqlparser::ast::Expr;
319 match expr {
320 Expr::Function(func) => func.over.is_some() || function_has_window(func),
321 Expr::BinaryOp { left, right, .. } => expr_has_window(left) || expr_has_window(right),
322 Expr::UnaryOp { expr, .. } => expr_has_window(expr),
323 _ => false,
324 }
325}
326
327fn function_has_window(func: &sqlparser::ast::Function) -> bool {
329 use sqlparser::ast::{FunctionArg, FunctionArgExpr, FunctionArguments};
330 match &func.args {
331 FunctionArguments::List(list) => list.args.iter().any(|a| match a {
332 FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) => expr_has_window(expr),
333 FunctionArg::Named {
334 arg: FunctionArgExpr::Expr(expr),
335 ..
336 } => expr_has_window(expr),
337 _ => false,
338 }),
339 _ => false,
340 }
341}
342
343pub fn build_explain_sql(sql: &str, dialect: VerifyDialect) -> String {
348 match dialect {
349 VerifyDialect::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql),
350 VerifyDialect::MySql | VerifyDialect::PostgreSql => format!("EXPLAIN {}", sql),
351 }
352}
353
354pub fn is_db_verify_enabled() -> bool {
360 let verify_flag = std::env::var("SZ_ORM_QUERY_VERIFY").unwrap_or_default();
361 let database_url = std::env::var("DATABASE_URL").unwrap_or_default();
362 (verify_flag == "1" || verify_flag.eq_ignore_ascii_case("true")) && !database_url.is_empty()
363}
364
365pub fn current_verify_mode() -> VerifyMode {
369 if is_db_verify_enabled() {
370 VerifyMode::Full
371 } else {
372 VerifyMode::SyntaxOnly
373 }
374}
375
376pub fn verify_degraded(sql: &str, dialect: VerifyDialect) -> VerifyResult {
381 let mut result = verify_sql_syntax(sql, dialect);
382 if result.is_valid {
383 result.push_error(
384 "warning: sql-verify-proc degraded to syntax-only (DATABASE_URL not set)".to_string(),
385 );
386 result.is_valid = true;
387 result.errors.clear();
388 }
389 result
390}
391
392pub fn verify_full(sql: &str, dialect: VerifyDialect) -> VerifyResult {
403 let mut result = verify_sql_syntax(sql, dialect);
404 if !result.is_valid {
405 return result;
406 }
407
408 let path = classify_sql_path(sql);
409 if path == SqlPath::Unknown {
410 result.push_error(format!("无法识别 SQL 路径分类: {}", sql));
411 return result;
412 }
413
414 let _explain_sql = build_explain_sql(sql, dialect);
415 result
416}
417
418pub fn verify_smart(sql: &str, dialect: VerifyDialect) -> VerifyResult {
424 if is_db_verify_enabled() {
425 verify_full(sql, dialect)
426 } else {
427 verify_degraded(sql, dialect)
428 }
429}
430
431pub fn check_path_coverage(sqls: &[&str]) -> Vec<SqlPath> {
436 let mut covered = Vec::new();
437 for sql in sqls {
438 let path = classify_sql_path(sql);
439 if !covered.contains(&path) {
440 covered.push(path);
441 }
442 }
443
444 let all_paths = [
445 SqlPath::Select,
446 SqlPath::Insert,
447 SqlPath::Update,
448 SqlPath::Delete,
449 SqlPath::Join,
450 SqlPath::Subquery,
451 SqlPath::Cte,
452 SqlPath::WindowFunction,
453 ];
454
455 all_paths
456 .iter()
457 .filter(|p| !covered.contains(p))
458 .copied()
459 .collect()
460}
461
462#[cfg(test)]
463mod tests {
464 use super::*;
465
466 #[test]
467 fn test_verify_valid_select() {
468 let sql = "SELECT id, name FROM users WHERE id = 1";
469 let result = verify_sql_syntax(sql, VerifyDialect::MySql);
470 assert!(
471 result.is_valid,
472 "Valid SELECT should pass: {:?}",
473 result.errors
474 );
475 }
476
477 #[test]
478 fn test_verify_invalid_sql() {
479 let sql = "SELECT FROM WHERE";
480 let result = verify_sql_syntax(sql, VerifyDialect::MySql);
481 assert!(!result.is_valid);
482 assert!(!result.errors.is_empty());
483 }
484
485 #[test]
486 fn test_verify_empty_sql() {
487 let sql = "";
488 let result = verify_sql_syntax(sql, VerifyDialect::MySql);
489 assert!(!result.is_valid);
490 }
491
492 #[test]
493 fn test_sql_hash_deterministic() {
494 let sql = "SELECT * FROM users";
495 assert_eq!(sql_hash(sql), sql_hash(sql));
496 }
497
498 #[test]
499 fn test_sql_hash_different() {
500 let sql1 = "SELECT * FROM users";
501 let sql2 = "SELECT * FROM posts";
502 assert_ne!(sql_hash(sql1), sql_hash(sql2));
503 }
504
505 #[test]
506 fn test_is_read_only() {
507 assert!(is_read_only("SELECT * FROM users"));
508 assert!(is_read_only("EXPLAIN SELECT * FROM users"));
509 assert!(is_read_only("WITH cte AS (SELECT 1) SELECT * FROM cte"));
510 assert!(!is_read_only("INSERT INTO users VALUES (1)"));
511 assert!(!is_read_only("UPDATE users SET name = 'x'"));
512 assert!(!is_read_only("DELETE FROM users"));
513 assert!(!is_read_only("DROP TABLE users"));
514 }
515
516 #[test]
517 fn test_verify_postgresql_dialect() {
518 let sql = "SELECT id, name FROM users WHERE id = $1";
519 let result = verify_sql_syntax(sql, VerifyDialect::PostgreSql);
520 assert!(
521 result.is_valid,
522 "PG dialect should parse $1 params: {:?}",
523 result.errors
524 );
525 }
526
527 #[test]
528 fn test_verify_sqlite_dialect() {
529 let sql = "SELECT id, name FROM users WHERE id = ?";
530 let result = verify_sql_syntax(sql, VerifyDialect::Sqlite);
531 assert!(
532 result.is_valid,
533 "SQLite dialect should parse ? params: {:?}",
534 result.errors
535 );
536 }
537
538 #[test]
539 fn test_classify_select_path() {
540 let sql = "SELECT id, name FROM users WHERE id = 1";
541 assert_eq!(classify_sql_path(sql), SqlPath::Select);
542 }
543
544 #[test]
545 fn test_classify_insert_path() {
546 let sql = "INSERT INTO users (id, name) VALUES (1, 'Alice')";
547 assert_eq!(classify_sql_path(sql), SqlPath::Insert);
548 }
549
550 #[test]
551 fn test_classify_update_path() {
552 let sql = "UPDATE users SET name = 'Bob' WHERE id = 1";
553 assert_eq!(classify_sql_path(sql), SqlPath::Update);
554 }
555
556 #[test]
557 fn test_classify_delete_path() {
558 let sql = "DELETE FROM users WHERE id = 1";
559 assert_eq!(classify_sql_path(sql), SqlPath::Delete);
560 }
561
562 #[test]
563 fn test_classify_join_path() {
564 let sql = "SELECT u.name, p.title FROM users u INNER JOIN posts p ON u.id = p.user_id";
565 assert_eq!(classify_sql_path(sql), SqlPath::Join);
566 }
567
568 #[test]
569 fn test_classify_left_join_path() {
570 let sql = "SELECT u.name FROM users u LEFT JOIN posts p ON u.id = p.user_id";
571 assert_eq!(classify_sql_path(sql), SqlPath::Join);
572 }
573
574 #[test]
575 fn test_classify_cte_path() {
576 let sql = "WITH cte AS (SELECT id FROM users) SELECT * FROM cte";
577 assert_eq!(classify_sql_path(sql), SqlPath::Cte);
578 }
579
580 #[test]
581 fn test_classify_subquery_in_where() {
582 let sql = "SELECT * FROM users WHERE id IN (SELECT user_id FROM posts)";
583 assert_eq!(classify_sql_path(sql), SqlPath::Subquery);
584 }
585
586 #[test]
587 fn test_classify_subquery_in_from() {
588 let sql = "SELECT * FROM (SELECT id FROM users) AS sub";
589 assert_eq!(classify_sql_path(sql), SqlPath::Subquery);
590 }
591
592 #[test]
593 fn test_classify_window_function_path() {
594 let sql = "SELECT id, ROW_NUMBER() OVER (PARTITION BY dept ORDER BY salary) FROM employees";
595 assert_eq!(classify_sql_path(sql), SqlPath::WindowFunction);
596 }
597
598 #[test]
599 fn test_build_explain_mysql() {
600 let sql = "SELECT * FROM users";
601 let explain = build_explain_sql(sql, VerifyDialect::MySql);
602 assert_eq!(explain, "EXPLAIN SELECT * FROM users");
603 }
604
605 #[test]
606 fn test_build_explain_postgres() {
607 let sql = "SELECT * FROM users";
608 let explain = build_explain_sql(sql, VerifyDialect::PostgreSql);
609 assert_eq!(explain, "EXPLAIN SELECT * FROM users");
610 }
611
612 #[test]
613 fn test_build_explain_sqlite() {
614 let sql = "SELECT * FROM users";
615 let explain = build_explain_sql(sql, VerifyDialect::Sqlite);
616 assert_eq!(explain, "EXPLAIN QUERY PLAN SELECT * FROM users");
617 }
618
619 #[test]
620 fn test_verify_degraded_no_env() {
621 std::env::remove_var("SZ_ORM_QUERY_VERIFY");
622 std::env::remove_var("DATABASE_URL");
623 let sql = "SELECT * FROM users";
624 let result = verify_degraded(sql, VerifyDialect::MySql);
625 assert!(result.is_valid);
626 }
627
628 #[test]
629 fn test_verify_degraded_invalid_sql() {
630 std::env::remove_var("SZ_ORM_QUERY_VERIFY");
631 std::env::remove_var("DATABASE_URL");
632 let sql = "SELECT FROM WHERE";
633 let result = verify_degraded(sql, VerifyDialect::MySql);
634 assert!(!result.is_valid);
635 }
636
637 #[test]
638 fn test_verify_full_valid_select() {
639 let sql = "SELECT id, name FROM users WHERE id = 1";
640 let result = verify_full(sql, VerifyDialect::MySql);
641 assert!(
642 result.is_valid,
643 "verify_full should pass: {:?}",
644 result.errors
645 );
646 }
647
648 #[test]
649 fn test_verify_full_invalid_sql() {
650 let sql = "SELECT FROM WHERE";
651 let result = verify_full(sql, VerifyDialect::MySql);
652 assert!(!result.is_valid);
653 }
654
655 #[test]
656 fn test_verify_smart_degraded_mode() {
657 std::env::remove_var("SZ_ORM_QUERY_VERIFY");
658 std::env::remove_var("DATABASE_URL");
659 let sql = "SELECT * FROM users";
660 let result = verify_smart(sql, VerifyDialect::MySql);
661 assert!(result.is_valid);
662 }
663
664 #[test]
665 fn test_check_path_coverage_all_covered() {
666 let sqls = [
667 "SELECT * FROM users",
668 "INSERT INTO users VALUES (1)",
669 "UPDATE users SET name = 'x'",
670 "DELETE FROM users",
671 "SELECT * FROM a JOIN b ON a.id = b.id",
672 "SELECT * FROM users WHERE id IN (SELECT id FROM posts)",
673 "WITH cte AS (SELECT 1) SELECT * FROM cte",
674 "SELECT ROW_NUMBER() OVER (PARTITION BY x) FROM t",
675 ];
676 let uncovered = check_path_coverage(&sqls);
677 assert!(
678 uncovered.is_empty(),
679 "All paths should be covered, uncovered: {:?}",
680 uncovered.iter().map(|p| p.name()).collect::<Vec<_>>()
681 );
682 }
683
684 #[test]
685 fn test_check_path_coverage_partial() {
686 let sqls = ["SELECT * FROM users", "INSERT INTO users VALUES (1)"];
687 let uncovered = check_path_coverage(&sqls);
688 assert!(uncovered.contains(&SqlPath::Update));
689 assert!(uncovered.contains(&SqlPath::Delete));
690 assert!(uncovered.contains(&SqlPath::Join));
691 assert!(uncovered.contains(&SqlPath::Cte));
692 }
693
694 #[test]
695 fn test_sql_path_name() {
696 assert_eq!(SqlPath::Select.name(), "SELECT");
697 assert_eq!(SqlPath::Insert.name(), "INSERT");
698 assert_eq!(SqlPath::Join.name(), "JOIN");
699 assert_eq!(SqlPath::WindowFunction.name(), "WindowFunction");
700 }
701
702 #[test]
703 fn test_verify_result_push_error() {
704 let mut result = VerifyResult::ok("SELECT 1");
705 assert!(result.is_valid);
706 result.push_error("test error".to_string());
707 assert!(!result.is_valid);
708 assert_eq!(result.errors.len(), 1);
709 }
710
711 #[test]
712 fn test_verify_mode_enum() {
713 let full = VerifyMode::Full;
714 let syntax = VerifyMode::SyntaxOnly;
715 assert_ne!(full, syntax);
716 }
717}