use crate::{detect_statement_type, tokenize, SqlStatementType, SqlToken};
use regex::Regex;
use std::fmt;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum RuleSeverity {
Info,
Warning,
Error,
Critical,
}
impl RuleSeverity {
pub fn as_str(&self) -> &'static str {
match self {
RuleSeverity::Info => "info",
RuleSeverity::Warning => "warning",
RuleSeverity::Error => "error",
RuleSeverity::Critical => "critical",
}
}
pub fn description(&self) -> &'static str {
match self {
RuleSeverity::Info => "信息",
RuleSeverity::Warning => "警告",
RuleSeverity::Error => "错误",
RuleSeverity::Critical => "严重",
}
}
pub fn is_blocking(&self) -> bool {
matches!(self, RuleSeverity::Error | RuleSeverity::Critical)
}
}
impl fmt::Display for RuleSeverity {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RuleViolation {
pub rule_name: String,
pub severity: RuleSeverity,
pub message: String,
pub position: Option<usize>,
}
impl RuleViolation {
pub fn new(
rule_name: impl Into<String>,
severity: RuleSeverity,
message: impl Into<String>,
) -> Self {
Self {
rule_name: rule_name.into(),
severity,
message: message.into(),
position: None,
}
}
pub fn with_position(
rule_name: impl Into<String>,
severity: RuleSeverity,
message: impl Into<String>,
position: usize,
) -> Self {
Self {
rule_name: rule_name.into(),
severity,
message: message.into(),
position: Some(position),
}
}
pub fn is_blocking(&self) -> bool {
self.severity.is_blocking()
}
}
impl fmt::Display for RuleViolation {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.position {
Some(pos) => write!(
f,
"[{}] {} at position {}: {}",
self.severity, self.rule_name, pos, self.message
),
None => write!(
f,
"[{}] {}: {}",
self.severity, self.rule_name, self.message
),
}
}
}
#[derive(Debug, Clone)]
pub struct RuleContext {
pub sql: String,
pub tokens: Vec<SqlToken>,
pub statement_type: SqlStatementType,
pub sql_upper: String,
}
impl RuleContext {
pub fn from_sql(sql: &str) -> Self {
Self {
sql: sql.to_string(),
tokens: tokenize(sql),
statement_type: detect_statement_type(sql),
sql_upper: sql.to_uppercase(),
}
}
pub fn keyword_count(&self, keyword: &str) -> usize {
let upper = keyword.to_uppercase();
self.tokens
.iter()
.filter(|t| matches!(t, SqlToken::Keyword(k) if *k == upper))
.count()
}
pub fn has_keyword(&self, keyword: &str) -> bool {
self.keyword_count(keyword) > 0
}
pub fn table_count(&self) -> usize {
let mut count = 0;
let mut expect_table = false;
for token in &self.tokens {
match token {
SqlToken::Keyword(k)
if matches!(k.as_str(), "FROM" | "JOIN" | "INTO" | "UPDATE") =>
{
expect_table = true;
}
SqlToken::Identifier(name) if expect_table && !name.starts_with('"') => {
let _ = name;
count += 1;
expect_table = false;
}
SqlToken::Punctuation('.') if expect_table => {
}
_ if expect_table => {
expect_table = false;
}
_ => {}
}
}
count
}
}
pub trait SqlRule: Send + Sync {
fn name(&self) -> &str;
fn severity(&self) -> RuleSeverity;
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation>;
}
#[derive(Debug, Default)]
pub struct NoSelectStarRule;
impl SqlRule for NoSelectStarRule {
fn name(&self) -> &str {
"no_select_star"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Warning
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
if ctx.statement_type != SqlStatementType::Select {
return None;
}
let mut after_select = false;
for token in &ctx.tokens {
match token {
SqlToken::Keyword(k) if k == "SELECT" => {
after_select = true;
}
SqlToken::Operator(op) if after_select && op == "*" => {
let pos = ctx.sql.find('*');
return Some(RuleViolation::with_position(
self.name(),
self.severity(),
"SELECT * 不允许,请显式列出列名",
pos.unwrap_or(0),
));
}
_ if after_select => {
after_select = false;
}
_ => {}
}
}
None
}
}
#[derive(Debug, Default)]
pub struct RequireWhereInDeleteRule;
impl SqlRule for RequireWhereInDeleteRule {
fn name(&self) -> &str {
"require_where_in_delete"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Critical
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
if ctx.statement_type != SqlStatementType::Delete {
return None;
}
if !ctx.has_keyword("WHERE") {
Some(RuleViolation::new(
self.name(),
self.severity(),
"DELETE 语句必须包含 WHERE 子句,否则将清空全表",
))
} else {
None
}
}
}
#[derive(Debug, Default)]
pub struct RequireWhereInUpdateRule;
impl SqlRule for RequireWhereInUpdateRule {
fn name(&self) -> &str {
"require_where_in_update"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Critical
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
if ctx.statement_type != SqlStatementType::Update {
return None;
}
if !ctx.has_keyword("WHERE") {
Some(RuleViolation::new(
self.name(),
self.severity(),
"UPDATE 语句必须包含 WHERE 子句,否则将更新全表",
))
} else {
None
}
}
}
#[derive(Debug)]
pub struct MaxTableCountRule {
pub max: usize,
}
impl MaxTableCountRule {
pub fn new(max: usize) -> Self {
Self { max }
}
}
impl SqlRule for MaxTableCountRule {
fn name(&self) -> &str {
"max_table_count"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Warning
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
let count = ctx.table_count();
if count > self.max {
Some(RuleViolation::new(
self.name(),
self.severity(),
format!("涉及表数量 {} 超过上限 {}", count, self.max),
))
} else {
None
}
}
}
#[derive(Debug)]
pub struct MaxJoinCountRule {
pub max: usize,
}
impl MaxJoinCountRule {
pub fn new(max: usize) -> Self {
Self { max }
}
}
impl SqlRule for MaxJoinCountRule {
fn name(&self) -> &str {
"max_join_count"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Warning
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
let count = ctx.keyword_count("JOIN");
if count > self.max {
Some(RuleViolation::new(
self.name(),
self.severity(),
format!("JOIN 数量 {} 超过上限 {}", count, self.max),
))
} else {
None
}
}
}
#[derive(Debug, Default)]
pub struct NoUnionRule;
impl SqlRule for NoUnionRule {
fn name(&self) -> &str {
"no_union"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Error
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
if ctx.has_keyword("UNION") {
let pos = ctx.sql_upper.find("UNION");
Some(RuleViolation::with_position(
self.name(),
self.severity(),
"UNION 操作不允许",
pos.unwrap_or(0),
))
} else {
None
}
}
}
#[derive(Debug)]
pub struct ForbiddenKeywordRule {
pub rule_name: String,
pub keywords: Vec<String>,
pub severity: RuleSeverity,
}
impl ForbiddenKeywordRule {
pub fn new(rule_name: impl Into<String>, keywords: &[&str]) -> Self {
Self {
rule_name: rule_name.into(),
keywords: keywords.iter().map(|k| k.to_uppercase()).collect(),
severity: RuleSeverity::Critical,
}
}
pub fn with_severity(mut self, severity: RuleSeverity) -> Self {
self.severity = severity;
self
}
}
impl SqlRule for ForbiddenKeywordRule {
fn name(&self) -> &str {
&self.rule_name
}
fn severity(&self) -> RuleSeverity {
self.severity
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
for keyword in &self.keywords {
if ctx.has_keyword(keyword) {
let pos = ctx.sql_upper.find(keyword.as_str());
return Some(RuleViolation::with_position(
self.name(),
self.severity,
format!("禁止关键字 {}", keyword),
pos.unwrap_or(0),
));
}
}
None
}
}
#[derive(Debug)]
pub struct RegexRule {
pub rule_name: String,
pub pattern: Regex,
pub severity: RuleSeverity,
pub message: String,
}
impl RegexRule {
pub fn new(
rule_name: impl Into<String>,
pattern: &str,
severity: RuleSeverity,
message: impl Into<String>,
) -> Result<Self, regex::Error> {
Ok(Self {
rule_name: rule_name.into(),
pattern: Regex::new(pattern)?,
severity,
message: message.into(),
})
}
}
impl SqlRule for RegexRule {
fn name(&self) -> &str {
&self.rule_name
}
fn severity(&self) -> RuleSeverity {
self.severity
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
self.pattern.find(&ctx.sql).map(|m| {
RuleViolation::with_position(
self.name(),
self.severity,
self.message.clone(),
m.start(),
)
})
}
}
#[derive(Debug, Default)]
pub struct RequireLimitRule;
impl SqlRule for RequireLimitRule {
fn name(&self) -> &str {
"require_limit"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Info
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
if ctx.statement_type != SqlStatementType::Select {
return None;
}
if !ctx.has_keyword("LIMIT") && !ctx.has_keyword("FETCH") {
Some(RuleViolation::new(
self.name(),
self.severity(),
"SELECT 语句建议包含 LIMIT 子句以限制结果集大小",
))
} else {
None
}
}
}
pub struct MaxColumnCountRule {
pub max_columns: usize,
}
impl Default for MaxColumnCountRule {
fn default() -> Self {
Self { max_columns: 20 }
}
}
impl SqlRule for MaxColumnCountRule {
fn name(&self) -> &str {
"max_column_count"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Warning
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
if ctx.statement_type != SqlStatementType::Select {
return None;
}
let select_idx = ctx
.tokens
.iter()
.position(|t| matches!(t, SqlToken::Keyword(k) if k == "SELECT"))?;
let from_idx = ctx
.tokens
.iter()
.position(|t| matches!(t, SqlToken::Keyword(k) if k == "FROM"))?;
if from_idx <= select_idx {
return None;
}
let comma_count = ctx.tokens[select_idx + 1..from_idx]
.iter()
.filter(|t| matches!(t, SqlToken::Punctuation(',')))
.count();
let col_count = comma_count + 1;
if col_count > self.max_columns {
Some(RuleViolation::new(
self.name(),
self.severity(),
format!("SELECT 列数 {} 超过上限 {}", col_count, self.max_columns),
))
} else {
None
}
}
}
pub struct NoSubqueryRule;
impl SqlRule for NoSubqueryRule {
fn name(&self) -> &str {
"no_subquery"
}
fn severity(&self) -> RuleSeverity {
RuleSeverity::Warning
}
fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
if ctx.statement_type != SqlStatementType::Select {
return None;
}
for i in 0..ctx.tokens.len().saturating_sub(2) {
let is_from = matches!(&ctx.tokens[i], SqlToken::Keyword(k) if k == "FROM");
let is_open = matches!(ctx.tokens[i + 1], SqlToken::Punctuation('('));
let is_select = matches!(&ctx.tokens[i + 2], SqlToken::Keyword(k) if k == "SELECT");
if is_from && is_open && is_select {
return Some(RuleViolation::new(
self.name(),
self.severity(),
"SELECT 语句中包含子查询,建议改用 JOIN",
));
}
}
None
}
}
#[derive(Debug, Clone, Default)]
pub struct RuleReport {
pub violations: Vec<RuleViolation>,
}
impl RuleReport {
pub fn is_clean(&self) -> bool {
self.violations.is_empty()
}
pub fn has_violations(&self) -> bool {
!self.violations.is_empty()
}
pub fn has_blocking(&self) -> bool {
self.violations.iter().any(|v| v.is_blocking())
}
pub fn has_errors(&self) -> bool {
self.violations
.iter()
.any(|v| v.severity == RuleSeverity::Error)
}
pub fn has_critical(&self) -> bool {
self.violations
.iter()
.any(|v| v.severity == RuleSeverity::Critical)
}
pub fn count_by_severity(&self, severity: RuleSeverity) -> usize {
self.violations
.iter()
.filter(|v| v.severity == severity)
.count()
}
pub fn violations_by_severity(&self, severity: RuleSeverity) -> Vec<&RuleViolation> {
self.violations
.iter()
.filter(|v| v.severity == severity)
.collect()
}
pub fn blocking_violations(&self) -> Vec<&RuleViolation> {
self.violations.iter().filter(|v| v.is_blocking()).collect()
}
pub fn violation_count(&self) -> usize {
self.violations.len()
}
pub fn summary(&self) -> String {
if self.is_clean() {
return "规则检查通过,无违规".to_string();
}
let mut lines = Vec::with_capacity(self.violations.len() + 2);
lines.push(format!(
"规则检查完成:共 {} 条违规({} 信息,{} 警告,{} 错误,{} 严重)",
self.violation_count(),
self.count_by_severity(RuleSeverity::Info),
self.count_by_severity(RuleSeverity::Warning),
self.count_by_severity(RuleSeverity::Error),
self.count_by_severity(RuleSeverity::Critical),
));
for v in &self.violations {
lines.push(format!(" - {}", v));
}
lines.join("\n")
}
}
impl fmt::Display for RuleReport {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.summary())
}
}
pub struct RuleEngine {
rules: Vec<Box<dyn SqlRule>>,
}
impl RuleEngine {
pub fn new() -> Self {
Self { rules: Vec::new() }
}
pub fn with_default_rules() -> Self {
let mut engine = Self::new();
engine.add_rule(Box::new(NoSelectStarRule));
engine.add_rule(Box::new(RequireWhereInDeleteRule));
engine.add_rule(Box::new(RequireWhereInUpdateRule));
engine.add_rule(Box::new(NoUnionRule));
engine.add_rule(Box::new(ForbiddenKeywordRule::new(
"forbidden_privilege_ops",
&["GRANT", "REVOKE", "EXEC", "EXECUTE"],
)));
engine
}
pub fn add_rule(&mut self, rule: Box<dyn SqlRule>) -> &mut Self {
self.rules.push(rule);
self
}
pub fn rule_count(&self) -> usize {
self.rules.len()
}
pub fn check(&self, sql: &str) -> RuleReport {
let ctx = RuleContext::from_sql(sql);
let violations = self
.rules
.iter()
.filter_map(|rule| rule.check(&ctx))
.collect();
RuleReport { violations }
}
pub fn check_batch<'a>(&self, sqls: impl IntoIterator<Item = &'a str>) -> Vec<RuleReport> {
sqls.into_iter().map(|sql| self.check(sql)).collect()
}
pub fn passes(&self, sql: &str) -> bool {
!self.check(sql).has_blocking()
}
}
impl Default for RuleEngine {
fn default() -> Self {
Self::new()
}
}
pub struct RulePresets;
impl RulePresets {
pub fn strict() -> RuleEngine {
let mut engine = RuleEngine::with_default_rules();
engine.add_rule(Box::new(MaxTableCountRule::new(5)));
engine.add_rule(Box::new(MaxJoinCountRule::new(3)));
engine.add_rule(Box::new(RequireLimitRule));
engine
}
pub fn read_only() -> RuleEngine {
let mut engine = RuleEngine::new();
engine.add_rule(Box::new(NoSelectStarRule));
engine.add_rule(Box::new(NoUnionRule));
engine.add_rule(Box::new(MaxTableCountRule::new(10)));
engine.add_rule(Box::new(MaxJoinCountRule::new(5)));
engine
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rule_severity_ordering() {
assert!(RuleSeverity::Info < RuleSeverity::Warning);
assert!(RuleSeverity::Warning < RuleSeverity::Error);
assert!(RuleSeverity::Error < RuleSeverity::Critical);
}
#[test]
fn test_rule_severity_as_str() {
assert_eq!(RuleSeverity::Info.as_str(), "info");
assert_eq!(RuleSeverity::Warning.as_str(), "warning");
assert_eq!(RuleSeverity::Error.as_str(), "error");
assert_eq!(RuleSeverity::Critical.as_str(), "critical");
}
#[test]
fn test_rule_severity_description() {
assert_eq!(RuleSeverity::Info.description(), "信息");
assert_eq!(RuleSeverity::Critical.description(), "严重");
}
#[test]
fn test_rule_severity_is_blocking() {
assert!(!RuleSeverity::Info.is_blocking());
assert!(!RuleSeverity::Warning.is_blocking());
assert!(RuleSeverity::Error.is_blocking());
assert!(RuleSeverity::Critical.is_blocking());
}
#[test]
fn test_rule_violation_new() {
let v = RuleViolation::new("test_rule", RuleSeverity::Warning, "test message");
assert_eq!(v.rule_name, "test_rule");
assert_eq!(v.severity, RuleSeverity::Warning);
assert_eq!(v.message, "test message");
assert!(v.position.is_none());
}
#[test]
fn test_rule_violation_with_position() {
let v = RuleViolation::with_position("test_rule", RuleSeverity::Error, "test message", 42);
assert_eq!(v.position, Some(42));
assert!(v.is_blocking());
}
#[test]
fn test_rule_violation_display() {
let v = RuleViolation::new("r", RuleSeverity::Warning, "msg");
let s = format!("{}", v);
assert!(s.contains("[warning]"));
assert!(s.contains("r"));
assert!(s.contains("msg"));
let v2 = RuleViolation::with_position("r", RuleSeverity::Error, "msg", 10);
let s2 = format!("{}", v2);
assert!(s2.contains("position 10"));
}
#[test]
fn test_rule_context_from_sql() {
let ctx = RuleContext::from_sql("SELECT id FROM users");
assert_eq!(ctx.statement_type, SqlStatementType::Select);
assert!(ctx.sql_upper.contains("SELECT"));
assert!(!ctx.tokens.is_empty());
}
#[test]
fn test_rule_context_keyword_count() {
let ctx = RuleContext::from_sql("SELECT id FROM users JOIN orders ON 1=1");
assert_eq!(ctx.keyword_count("JOIN"), 1);
assert_eq!(ctx.keyword_count("SELECT"), 1);
assert_eq!(ctx.keyword_count("DELETE"), 0);
}
#[test]
fn test_rule_context_has_keyword() {
let ctx = RuleContext::from_sql("SELECT id FROM users WHERE id = 1");
assert!(ctx.has_keyword("WHERE"));
assert!(!ctx.has_keyword("DELETE"));
}
#[test]
fn test_rule_context_table_count() {
let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON users.id = orders.uid");
assert!(ctx.table_count() >= 2);
}
#[test]
fn test_no_select_star_rule_detects_star() {
let rule = NoSelectStarRule;
let ctx = RuleContext::from_sql("SELECT * FROM users");
let violation = rule.check(&ctx);
assert!(violation.is_some());
let v = violation.unwrap();
assert_eq!(v.severity, RuleSeverity::Warning);
assert!(v.position.is_some());
}
#[test]
fn test_no_select_star_rule_passes_explicit_columns() {
let rule = NoSelectStarRule;
let ctx = RuleContext::from_sql("SELECT id, name FROM users");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_no_select_star_rule_ignores_non_select() {
let rule = NoSelectStarRule;
let ctx = RuleContext::from_sql("DELETE FROM users WHERE id = 1");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_require_where_in_delete_rule_detects_missing_where() {
let rule = RequireWhereInDeleteRule;
let ctx = RuleContext::from_sql("DELETE FROM users");
let violation = rule.check(&ctx);
assert!(violation.is_some());
assert_eq!(violation.unwrap().severity, RuleSeverity::Critical);
}
#[test]
fn test_require_where_in_delete_rule_passes_with_where() {
let rule = RequireWhereInDeleteRule;
let ctx = RuleContext::from_sql("DELETE FROM users WHERE id = 1");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_require_where_in_delete_rule_ignores_non_delete() {
let rule = RequireWhereInDeleteRule;
let ctx = RuleContext::from_sql("SELECT * FROM users");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_require_where_in_update_rule_detects_missing_where() {
let rule = RequireWhereInUpdateRule;
let ctx = RuleContext::from_sql("UPDATE users SET name = 'a'");
let violation = rule.check(&ctx);
assert!(violation.is_some());
assert_eq!(violation.unwrap().severity, RuleSeverity::Critical);
}
#[test]
fn test_require_where_in_update_rule_passes_with_where() {
let rule = RequireWhereInUpdateRule;
let ctx = RuleContext::from_sql("UPDATE users SET name = 'a' WHERE id = 1");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_max_table_count_rule_passes_within_limit() {
let rule = MaxTableCountRule::new(2);
let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON 1=1");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_max_table_count_rule_detects_exceed() {
let rule = MaxTableCountRule::new(1);
let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON 1=1");
let violation = rule.check(&ctx);
assert!(violation.is_some());
assert!(violation.unwrap().message.contains("超过上限"));
}
#[test]
fn test_max_join_count_rule_passes_within_limit() {
let rule = MaxJoinCountRule::new(2);
let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON 1=1 JOIN items ON 2=2");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_max_join_count_rule_detects_exceed() {
let rule = MaxJoinCountRule::new(1);
let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON 1=1 JOIN items ON 2=2");
assert!(rule.check(&ctx).is_some());
}
#[test]
fn test_no_union_rule_detects_union() {
let rule = NoUnionRule;
let ctx = RuleContext::from_sql("SELECT id FROM users UNION SELECT id FROM archived");
let violation = rule.check(&ctx);
assert!(violation.is_some());
assert_eq!(violation.unwrap().severity, RuleSeverity::Error);
}
#[test]
fn test_no_union_rule_passes_without_union() {
let rule = NoUnionRule;
let ctx = RuleContext::from_sql("SELECT id FROM users");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_forbidden_keyword_rule_detects_keyword() {
let rule = ForbiddenKeywordRule::new("no_grant", &["GRANT"]);
let ctx = RuleContext::from_sql("GRANT ALL ON users TO hacker");
let violation = rule.check(&ctx);
assert!(violation.is_some());
assert_eq!(violation.unwrap().severity, RuleSeverity::Critical);
}
#[test]
fn test_forbidden_keyword_rule_passes_without_keyword() {
let rule = ForbiddenKeywordRule::new("no_grant", &["GRANT"]);
let ctx = RuleContext::from_sql("SELECT * FROM users");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_forbidden_keyword_rule_with_severity() {
let rule =
ForbiddenKeywordRule::new("no_exec", &["EXEC"]).with_severity(RuleSeverity::Warning);
let ctx = RuleContext::from_sql("EXEC sp_executesql 'DROP TABLE users'");
let violation = rule.check(&ctx).unwrap();
assert_eq!(violation.severity, RuleSeverity::Warning);
}
#[test]
fn test_regex_rule_detects_match() {
let rule = RegexRule::new(
"no_sys_tables",
r"(?i)\bsys\.",
RuleSeverity::Error,
"禁止访问系统表",
)
.unwrap();
let ctx = RuleContext::from_sql("SELECT * FROM sys.tables");
let violation = rule.check(&ctx);
assert!(violation.is_some());
assert!(violation.unwrap().position.is_some());
}
#[test]
fn test_regex_rule_passes_no_match() {
let rule = RegexRule::new(
"no_sys_tables",
r"(?i)\bsys\.",
RuleSeverity::Error,
"禁止访问系统表",
)
.unwrap();
let ctx = RuleContext::from_sql("SELECT * FROM users");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_regex_rule_invalid_pattern() {
let result = RegexRule::new("bad", r"[invalid", RuleSeverity::Error, "msg");
assert!(result.is_err());
}
#[test]
fn test_require_limit_rule_detects_missing_limit() {
let rule = RequireLimitRule;
let ctx = RuleContext::from_sql("SELECT id FROM users");
let violation = rule.check(&ctx);
assert!(violation.is_some());
assert_eq!(violation.unwrap().severity, RuleSeverity::Info);
}
#[test]
fn test_require_limit_rule_passes_with_limit() {
let rule = RequireLimitRule;
let ctx = RuleContext::from_sql("SELECT id FROM users LIMIT 10");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_rule_report_clean() {
let report = RuleReport::default();
assert!(report.is_clean());
assert!(!report.has_violations());
assert!(!report.has_blocking());
assert_eq!(report.violation_count(), 0);
}
#[test]
fn test_rule_report_with_violations() {
let report = RuleReport {
violations: vec![
RuleViolation::new("r1", RuleSeverity::Warning, "w"),
RuleViolation::new("r2", RuleSeverity::Error, "e"),
RuleViolation::new("r3", RuleSeverity::Critical, "c"),
],
};
assert!(report.has_violations());
assert!(report.has_blocking());
assert!(report.has_errors());
assert!(report.has_critical());
assert_eq!(report.violation_count(), 3);
assert_eq!(report.count_by_severity(RuleSeverity::Warning), 1);
assert_eq!(report.count_by_severity(RuleSeverity::Error), 1);
assert_eq!(report.count_by_severity(RuleSeverity::Critical), 1);
}
#[test]
fn test_rule_report_summary() {
let report = RuleReport::default();
assert!(report.summary().contains("无违规"));
let report2 = RuleReport {
violations: vec![RuleViolation::new("r1", RuleSeverity::Error, "msg")],
};
let summary = report2.summary();
assert!(summary.contains("1 条违规"));
assert!(summary.contains("r1"));
}
#[test]
fn test_rule_report_violations_by_severity() {
let report = RuleReport {
violations: vec![
RuleViolation::new("r1", RuleSeverity::Info, "i"),
RuleViolation::new("r2", RuleSeverity::Info, "i2"),
RuleViolation::new("r3", RuleSeverity::Error, "e"),
],
};
assert_eq!(report.violations_by_severity(RuleSeverity::Info).len(), 2);
assert_eq!(report.violations_by_severity(RuleSeverity::Error).len(), 1);
assert_eq!(report.blocking_violations().len(), 1);
}
#[test]
fn test_rule_engine_new_empty() {
let engine = RuleEngine::new();
assert_eq!(engine.rule_count(), 0);
let report = engine.check("SELECT * FROM users");
assert!(report.is_clean());
}
#[test]
fn test_rule_engine_add_rule() {
let mut engine = RuleEngine::new();
engine.add_rule(Box::new(NoSelectStarRule));
assert_eq!(engine.rule_count(), 1);
}
#[test]
fn test_rule_engine_check_collects_violations() {
let mut engine = RuleEngine::new();
engine.add_rule(Box::new(NoSelectStarRule));
engine.add_rule(Box::new(RequireWhereInDeleteRule));
let report = engine.check("DELETE FROM users");
assert!(report.has_violations());
assert!(report.has_critical());
}
#[test]
fn test_rule_engine_check_clean_sql() {
let mut engine = RuleEngine::new();
engine.add_rule(Box::new(NoSelectStarRule));
engine.add_rule(Box::new(RequireWhereInDeleteRule));
let report = engine.check("SELECT id, name FROM users WHERE id = 1");
assert!(report.is_clean());
}
#[test]
fn test_rule_engine_passes() {
let mut engine = RuleEngine::new();
engine.add_rule(Box::new(RequireWhereInDeleteRule));
assert!(engine.passes("DELETE FROM users WHERE id = 1"));
assert!(!engine.passes("DELETE FROM users"));
}
#[test]
fn test_rule_engine_check_batch() {
let mut engine = RuleEngine::new();
engine.add_rule(Box::new(NoSelectStarRule));
let reports = engine.check_batch(["SELECT * FROM users", "SELECT id FROM users"]);
assert_eq!(reports.len(), 2);
assert!(reports[0].has_violations());
assert!(reports[1].is_clean());
}
#[test]
fn test_rule_engine_with_default_rules() {
let engine = RuleEngine::with_default_rules();
assert!(engine.rule_count() >= 5);
let report = engine.check("GRANT ALL ON users TO hacker");
assert!(report.has_critical());
}
#[test]
fn test_rule_engine_default_rules_select_star() {
let engine = RuleEngine::with_default_rules();
let report = engine.check("SELECT * FROM users");
assert!(report.has_violations());
}
#[test]
fn test_rule_engine_default_rules_clean_sql() {
let engine = RuleEngine::with_default_rules();
let report = engine.check("SELECT id, name FROM users WHERE id = 1 LIMIT 10");
assert!(report.is_clean());
}
#[test]
fn test_rule_presets_strict() {
let engine = RulePresets::strict();
assert!(engine.rule_count() >= 8);
let report = engine.check("SELECT id FROM users WHERE id = 1");
assert!(report.has_violations());
}
#[test]
fn test_rule_presets_read_only() {
let engine = RulePresets::read_only();
let report = engine.check("DELETE FROM users");
assert!(report.is_clean());
let report2 = engine.check("SELECT * FROM users");
assert!(report2.has_violations());
}
#[test]
fn test_rule_presets_strict_max_table_count() {
let engine = RulePresets::strict();
let sql =
"SELECT * FROM a JOIN b ON 1=1 JOIN c ON 2=2 JOIN d ON 3=3 JOIN e ON 4=4 JOIN f ON 5=5";
let report = engine.check(sql);
assert!(report.has_violations());
}
#[test]
fn test_max_column_count_pass() {
let rule = MaxColumnCountRule { max_columns: 5 };
let ctx = RuleContext::from_sql("SELECT a, b, c FROM users");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_max_column_count_fail() {
let rule = MaxColumnCountRule { max_columns: 2 };
let ctx = RuleContext::from_sql("SELECT a, b, c, d FROM users");
assert!(rule.check(&ctx).is_some());
}
#[test]
fn test_max_column_count_default() {
let rule = MaxColumnCountRule::default();
assert_eq!(rule.max_columns, 20);
}
#[test]
fn test_no_subquery_clean() {
let rule = NoSubqueryRule;
let ctx = RuleContext::from_sql("SELECT id FROM users WHERE id = 1");
assert!(rule.check(&ctx).is_none());
}
#[test]
fn test_no_subquery_detected() {
let rule = NoSubqueryRule;
let ctx = RuleContext::from_sql("SELECT * FROM (SELECT id FROM users) AS sub");
assert!(rule.check(&ctx).is_some());
}
}