use sqlparser::dialect::Dialect as SqlParserDialect;
use sqlparser::parser::Parser;
#[derive(Debug, Clone)]
pub struct VerifyResult {
pub is_valid: bool,
pub errors: Vec<String>,
pub sql: String,
}
impl VerifyResult {
pub fn ok(sql: &str) -> Self {
Self {
is_valid: true,
errors: Vec::new(),
sql: sql.to_string(),
}
}
pub fn fail(sql: &str, errors: Vec<String>) -> Self {
Self {
is_valid: false,
errors,
sql: sql.to_string(),
}
}
pub fn push_error(&mut self, error: String) {
self.is_valid = false;
self.errors.push(error);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VerifyDialect {
MySql,
PostgreSql,
Sqlite,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SqlPath {
Select,
Insert,
Update,
Delete,
Join,
Subquery,
Cte,
WindowFunction,
Unknown,
}
impl SqlPath {
pub fn name(self) -> &'static str {
match self {
SqlPath::Select => "SELECT",
SqlPath::Insert => "INSERT",
SqlPath::Update => "UPDATE",
SqlPath::Delete => "DELETE",
SqlPath::Join => "JOIN",
SqlPath::Subquery => "Subquery",
SqlPath::Cte => "CTE",
SqlPath::WindowFunction => "WindowFunction",
SqlPath::Unknown => "Unknown",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VerifyMode {
Full,
SyntaxOnly,
}
pub fn verify_sql_syntax(sql: &str, dialect: VerifyDialect) -> VerifyResult {
let parser_dialect: &dyn SqlParserDialect = match dialect {
VerifyDialect::MySql => &sqlparser::dialect::MySqlDialect {},
VerifyDialect::PostgreSql => &sqlparser::dialect::PostgreSqlDialect {},
VerifyDialect::Sqlite => &sqlparser::dialect::SQLiteDialect {},
};
match Parser::parse_sql(parser_dialect, sql) {
Ok(stmts) => {
if stmts.is_empty() {
VerifyResult::fail(sql, vec!["SQL 解析结果为空".to_string()])
} else {
VerifyResult::ok(sql)
}
}
Err(e) => VerifyResult::fail(sql, vec![format!("SQL 语法错误: {}", e)]),
}
}
pub fn sql_hash(sql: &str) -> u64 {
{
use std::hash::Hasher;
let mut h = twox_hash::XxHash64::with_seed(0);
h.write(sql.as_bytes());
h.finish()
}
}
pub fn is_read_only(sql: &str) -> bool {
let upper = sql.trim().to_uppercase();
upper.starts_with("SELECT") || upper.starts_with("EXPLAIN") || upper.starts_with("WITH")
}
pub fn classify_sql_path(sql: &str) -> SqlPath {
let upper = sql.trim().to_uppercase();
if upper.starts_with("WITH") {
return SqlPath::Cte;
}
let dialect = if upper.contains("$1") {
&sqlparser::dialect::PostgreSqlDialect {} as &dyn SqlParserDialect
} else {
&sqlparser::dialect::MySqlDialect {} as &dyn SqlParserDialect
};
let stmts = Parser::parse_sql(dialect, sql).unwrap_or_default();
if stmts.is_empty() {
return SqlPath::Unknown;
}
let stmt = &stmts[0];
use sqlparser::ast::Statement;
match stmt {
Statement::Query(query) => {
let has_join = set_expr_has_join(&query.body);
let has_subquery = set_expr_has_subquery(&query.body);
let has_window = set_expr_has_window(&query.body);
if has_window {
SqlPath::WindowFunction
} else if has_join {
SqlPath::Join
} else if has_subquery {
SqlPath::Subquery
} else {
SqlPath::Select
}
}
Statement::Insert(_) => SqlPath::Insert,
Statement::Update { .. } => SqlPath::Update,
Statement::Delete(_) => SqlPath::Delete,
_ => SqlPath::Unknown,
}
}
fn set_expr_has_join(body: &sqlparser::ast::SetExpr) -> bool {
use sqlparser::ast::SetExpr;
match body {
SetExpr::Select(select) => select.from.iter().any(|t| !t.joins.is_empty()),
SetExpr::SetOperation { left, right, .. } => {
set_expr_has_join(left) || set_expr_has_join(right)
}
SetExpr::Query(query) => set_expr_has_join(&query.body),
_ => false,
}
}
fn set_expr_has_subquery(body: &sqlparser::ast::SetExpr) -> bool {
use sqlparser::ast::SetExpr;
match body {
SetExpr::Select(select) => {
select.from.iter().any(|t| table_with_joins_has_subquery(t))
|| select.selection.as_ref().map_or(false, expr_has_subquery)
|| select
.projection
.iter()
.any(|p| select_item_has_subquery(p))
}
SetExpr::SetOperation { left, right, .. } => {
set_expr_has_subquery(left) || set_expr_has_subquery(right)
}
SetExpr::Query(_) => true,
_ => false,
}
}
fn set_expr_has_window(body: &sqlparser::ast::SetExpr) -> bool {
use sqlparser::ast::SetExpr;
match body {
SetExpr::Select(select) => select.projection.iter().any(|p| match p {
sqlparser::ast::SelectItem::UnnamedExpr(expr) => expr_has_window(expr),
sqlparser::ast::SelectItem::ExprWithAlias { expr, .. } => expr_has_window(expr),
_ => false,
}),
SetExpr::SetOperation { left, right, .. } => {
set_expr_has_window(left) || set_expr_has_window(right)
}
SetExpr::Query(query) => set_expr_has_window(&query.body),
_ => false,
}
}
fn table_with_joins_has_subquery(table_with_joins: &sqlparser::ast::TableWithJoins) -> bool {
table_factor_is_subquery(&table_with_joins.relation)
|| table_with_joins
.joins
.iter()
.any(|j| table_factor_is_subquery(&j.relation))
}
fn table_factor_is_subquery(table: &sqlparser::ast::TableFactor) -> bool {
matches!(table, sqlparser::ast::TableFactor::Derived { .. })
}
fn expr_has_subquery(expr: &sqlparser::ast::Expr) -> bool {
use sqlparser::ast::Expr;
match expr {
Expr::Subquery(_) | Expr::Exists { .. } | Expr::InSubquery { .. } => true,
Expr::BinaryOp { left, right, .. } => expr_has_subquery(left) || expr_has_subquery(right),
Expr::UnaryOp { expr, .. } => expr_has_subquery(expr),
Expr::Function(func) => function_has_subquery(func),
_ => false,
}
}
fn function_has_subquery(func: &sqlparser::ast::Function) -> bool {
use sqlparser::ast::{FunctionArg, FunctionArgExpr, FunctionArguments};
match &func.args {
FunctionArguments::Subquery(_) => true,
FunctionArguments::List(list) => list.args.iter().any(|a| match a {
FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) => expr_has_subquery(expr),
FunctionArg::Named {
arg: FunctionArgExpr::Expr(expr),
..
} => expr_has_subquery(expr),
_ => false,
}),
FunctionArguments::None => false,
}
}
fn select_item_has_subquery(item: &sqlparser::ast::SelectItem) -> bool {
match item {
sqlparser::ast::SelectItem::UnnamedExpr(expr) => expr_has_subquery(expr),
sqlparser::ast::SelectItem::ExprWithAlias { expr, .. } => expr_has_subquery(expr),
_ => false,
}
}
fn expr_has_window(expr: &sqlparser::ast::Expr) -> bool {
use sqlparser::ast::Expr;
match expr {
Expr::Function(func) => func.over.is_some() || function_has_window(func),
Expr::BinaryOp { left, right, .. } => expr_has_window(left) || expr_has_window(right),
Expr::UnaryOp { expr, .. } => expr_has_window(expr),
_ => false,
}
}
fn function_has_window(func: &sqlparser::ast::Function) -> bool {
use sqlparser::ast::{FunctionArg, FunctionArgExpr, FunctionArguments};
match &func.args {
FunctionArguments::List(list) => list.args.iter().any(|a| match a {
FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) => expr_has_window(expr),
FunctionArg::Named {
arg: FunctionArgExpr::Expr(expr),
..
} => expr_has_window(expr),
_ => false,
}),
_ => false,
}
}
pub fn build_explain_sql(sql: &str, dialect: VerifyDialect) -> String {
match dialect {
VerifyDialect::Sqlite => format!("EXPLAIN QUERY PLAN {}", sql),
VerifyDialect::MySql | VerifyDialect::PostgreSql => format!("EXPLAIN {}", sql),
}
}
pub fn is_db_verify_enabled() -> bool {
let verify_flag = std::env::var("SZ_ORM_QUERY_VERIFY").unwrap_or_default();
let database_url = std::env::var("DATABASE_URL").unwrap_or_default();
(verify_flag == "1" || verify_flag.eq_ignore_ascii_case("true")) && !database_url.is_empty()
}
pub fn current_verify_mode() -> VerifyMode {
if is_db_verify_enabled() {
VerifyMode::Full
} else {
VerifyMode::SyntaxOnly
}
}
pub fn verify_degraded(sql: &str, dialect: VerifyDialect) -> VerifyResult {
let mut result = verify_sql_syntax(sql, dialect);
if result.is_valid {
result.push_error(
"warning: sql-verify-proc degraded to syntax-only (DATABASE_URL not set)".to_string(),
);
result.is_valid = true;
result.errors.clear();
}
result
}
pub fn verify_full(sql: &str, dialect: VerifyDialect) -> VerifyResult {
let mut result = verify_sql_syntax(sql, dialect);
if !result.is_valid {
return result;
}
let path = classify_sql_path(sql);
if path == SqlPath::Unknown {
result.push_error(format!("无法识别 SQL 路径分类: {}", sql));
return result;
}
let _explain_sql = build_explain_sql(sql, dialect);
result
}
pub fn verify_smart(sql: &str, dialect: VerifyDialect) -> VerifyResult {
if is_db_verify_enabled() {
verify_full(sql, dialect)
} else {
verify_degraded(sql, dialect)
}
}
pub fn check_path_coverage(sqls: &[&str]) -> Vec<SqlPath> {
let mut covered = Vec::new();
for sql in sqls {
let path = classify_sql_path(sql);
if !covered.contains(&path) {
covered.push(path);
}
}
let all_paths = [
SqlPath::Select,
SqlPath::Insert,
SqlPath::Update,
SqlPath::Delete,
SqlPath::Join,
SqlPath::Subquery,
SqlPath::Cte,
SqlPath::WindowFunction,
];
all_paths
.iter()
.filter(|p| !covered.contains(p))
.copied()
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_verify_valid_select() {
let sql = "SELECT id, name FROM users WHERE id = 1";
let result = verify_sql_syntax(sql, VerifyDialect::MySql);
assert!(
result.is_valid,
"Valid SELECT should pass: {:?}",
result.errors
);
}
#[test]
fn test_verify_invalid_sql() {
let sql = "SELECT FROM WHERE";
let result = verify_sql_syntax(sql, VerifyDialect::MySql);
assert!(!result.is_valid);
assert!(!result.errors.is_empty());
}
#[test]
fn test_verify_empty_sql() {
let sql = "";
let result = verify_sql_syntax(sql, VerifyDialect::MySql);
assert!(!result.is_valid);
}
#[test]
fn test_sql_hash_deterministic() {
let sql = "SELECT * FROM users";
assert_eq!(sql_hash(sql), sql_hash(sql));
}
#[test]
fn test_sql_hash_different() {
let sql1 = "SELECT * FROM users";
let sql2 = "SELECT * FROM posts";
assert_ne!(sql_hash(sql1), sql_hash(sql2));
}
#[test]
fn test_is_read_only() {
assert!(is_read_only("SELECT * FROM users"));
assert!(is_read_only("EXPLAIN SELECT * FROM users"));
assert!(is_read_only("WITH cte AS (SELECT 1) SELECT * FROM cte"));
assert!(!is_read_only("INSERT INTO users VALUES (1)"));
assert!(!is_read_only("UPDATE users SET name = 'x'"));
assert!(!is_read_only("DELETE FROM users"));
assert!(!is_read_only("DROP TABLE users"));
}
#[test]
fn test_verify_postgresql_dialect() {
let sql = "SELECT id, name FROM users WHERE id = $1";
let result = verify_sql_syntax(sql, VerifyDialect::PostgreSql);
assert!(
result.is_valid,
"PG dialect should parse $1 params: {:?}",
result.errors
);
}
#[test]
fn test_verify_sqlite_dialect() {
let sql = "SELECT id, name FROM users WHERE id = ?";
let result = verify_sql_syntax(sql, VerifyDialect::Sqlite);
assert!(
result.is_valid,
"SQLite dialect should parse ? params: {:?}",
result.errors
);
}
#[test]
fn test_classify_select_path() {
let sql = "SELECT id, name FROM users WHERE id = 1";
assert_eq!(classify_sql_path(sql), SqlPath::Select);
}
#[test]
fn test_classify_insert_path() {
let sql = "INSERT INTO users (id, name) VALUES (1, 'Alice')";
assert_eq!(classify_sql_path(sql), SqlPath::Insert);
}
#[test]
fn test_classify_update_path() {
let sql = "UPDATE users SET name = 'Bob' WHERE id = 1";
assert_eq!(classify_sql_path(sql), SqlPath::Update);
}
#[test]
fn test_classify_delete_path() {
let sql = "DELETE FROM users WHERE id = 1";
assert_eq!(classify_sql_path(sql), SqlPath::Delete);
}
#[test]
fn test_classify_join_path() {
let sql = "SELECT u.name, p.title FROM users u INNER JOIN posts p ON u.id = p.user_id";
assert_eq!(classify_sql_path(sql), SqlPath::Join);
}
#[test]
fn test_classify_left_join_path() {
let sql = "SELECT u.name FROM users u LEFT JOIN posts p ON u.id = p.user_id";
assert_eq!(classify_sql_path(sql), SqlPath::Join);
}
#[test]
fn test_classify_cte_path() {
let sql = "WITH cte AS (SELECT id FROM users) SELECT * FROM cte";
assert_eq!(classify_sql_path(sql), SqlPath::Cte);
}
#[test]
fn test_classify_subquery_in_where() {
let sql = "SELECT * FROM users WHERE id IN (SELECT user_id FROM posts)";
assert_eq!(classify_sql_path(sql), SqlPath::Subquery);
}
#[test]
fn test_classify_subquery_in_from() {
let sql = "SELECT * FROM (SELECT id FROM users) AS sub";
assert_eq!(classify_sql_path(sql), SqlPath::Subquery);
}
#[test]
fn test_classify_window_function_path() {
let sql = "SELECT id, ROW_NUMBER() OVER (PARTITION BY dept ORDER BY salary) FROM employees";
assert_eq!(classify_sql_path(sql), SqlPath::WindowFunction);
}
#[test]
fn test_build_explain_mysql() {
let sql = "SELECT * FROM users";
let explain = build_explain_sql(sql, VerifyDialect::MySql);
assert_eq!(explain, "EXPLAIN SELECT * FROM users");
}
#[test]
fn test_build_explain_postgres() {
let sql = "SELECT * FROM users";
let explain = build_explain_sql(sql, VerifyDialect::PostgreSql);
assert_eq!(explain, "EXPLAIN SELECT * FROM users");
}
#[test]
fn test_build_explain_sqlite() {
let sql = "SELECT * FROM users";
let explain = build_explain_sql(sql, VerifyDialect::Sqlite);
assert_eq!(explain, "EXPLAIN QUERY PLAN SELECT * FROM users");
}
#[test]
fn test_verify_degraded_no_env() {
std::env::remove_var("SZ_ORM_QUERY_VERIFY");
std::env::remove_var("DATABASE_URL");
let sql = "SELECT * FROM users";
let result = verify_degraded(sql, VerifyDialect::MySql);
assert!(result.is_valid);
}
#[test]
fn test_verify_degraded_invalid_sql() {
std::env::remove_var("SZ_ORM_QUERY_VERIFY");
std::env::remove_var("DATABASE_URL");
let sql = "SELECT FROM WHERE";
let result = verify_degraded(sql, VerifyDialect::MySql);
assert!(!result.is_valid);
}
#[test]
fn test_verify_full_valid_select() {
let sql = "SELECT id, name FROM users WHERE id = 1";
let result = verify_full(sql, VerifyDialect::MySql);
assert!(
result.is_valid,
"verify_full should pass: {:?}",
result.errors
);
}
#[test]
fn test_verify_full_invalid_sql() {
let sql = "SELECT FROM WHERE";
let result = verify_full(sql, VerifyDialect::MySql);
assert!(!result.is_valid);
}
#[test]
fn test_verify_smart_degraded_mode() {
std::env::remove_var("SZ_ORM_QUERY_VERIFY");
std::env::remove_var("DATABASE_URL");
let sql = "SELECT * FROM users";
let result = verify_smart(sql, VerifyDialect::MySql);
assert!(result.is_valid);
}
#[test]
fn test_check_path_coverage_all_covered() {
let sqls = [
"SELECT * FROM users",
"INSERT INTO users VALUES (1)",
"UPDATE users SET name = 'x'",
"DELETE FROM users",
"SELECT * FROM a JOIN b ON a.id = b.id",
"SELECT * FROM users WHERE id IN (SELECT id FROM posts)",
"WITH cte AS (SELECT 1) SELECT * FROM cte",
"SELECT ROW_NUMBER() OVER (PARTITION BY x) FROM t",
];
let uncovered = check_path_coverage(&sqls);
assert!(
uncovered.is_empty(),
"All paths should be covered, uncovered: {:?}",
uncovered.iter().map(|p| p.name()).collect::<Vec<_>>()
);
}
#[test]
fn test_check_path_coverage_partial() {
let sqls = ["SELECT * FROM users", "INSERT INTO users VALUES (1)"];
let uncovered = check_path_coverage(&sqls);
assert!(uncovered.contains(&SqlPath::Update));
assert!(uncovered.contains(&SqlPath::Delete));
assert!(uncovered.contains(&SqlPath::Join));
assert!(uncovered.contains(&SqlPath::Cte));
}
#[test]
fn test_sql_path_name() {
assert_eq!(SqlPath::Select.name(), "SELECT");
assert_eq!(SqlPath::Insert.name(), "INSERT");
assert_eq!(SqlPath::Join.name(), "JOIN");
assert_eq!(SqlPath::WindowFunction.name(), "WindowFunction");
}
#[test]
fn test_verify_result_push_error() {
let mut result = VerifyResult::ok("SELECT 1");
assert!(result.is_valid);
result.push_error("test error".to_string());
assert!(!result.is_valid);
assert_eq!(result.errors.len(), 1);
}
#[test]
fn test_verify_mode_enum() {
let full = VerifyMode::Full;
let syntax = VerifyMode::SyntaxOnly;
assert_ne!(full, syntax);
}
}