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