Skip to main content

sz_orm_sql_validator/
sql_rules.rs

1//! SQL 规则引擎:可组合的规则集合、规则匹配与违规报告
2//!
3//! 本模块提供基于规则的 SQL 质量检查框架,支持:
4//!
5//! - **规则定义**([`SqlRule`] trait):自定义检查逻辑,返回违规报告
6//! - **内置规则**:禁止 `SELECT *`、DELETE/UPDATE 必须带 WHERE、表数量上限、
7//!   JOIN 数量上限、禁止 UNION、禁止关键字、正则匹配规则
8//! - **规则引擎**([`RuleEngine`]):组合多条规则,批量检查 SQL,生成汇总报告
9//! - **违规报告**([`RuleReport`] / [`RuleViolation`]):按严重级别分类,
10//!   支持位置定位与人类可读摘要
11//!
12//! ## 示例
13//!
14//! ```rust
15//! use sz_orm_sql_validator::sql_rules::{RuleEngine, NoSelectStarRule, RequireWhereInDeleteRule};
16//!
17//! let mut engine = RuleEngine::new();
18//! engine.add_rule(Box::new(NoSelectStarRule));
19//! engine.add_rule(Box::new(RequireWhereInDeleteRule));
20//!
21//! let report = engine.check("DELETE FROM users");
22//! assert!(report.has_violations());
23//! ```
24
25use crate::{detect_statement_type, tokenize, SqlStatementType, SqlToken};
26use regex::Regex;
27use std::fmt;
28
29// ============================================================================
30// 规则严重级别
31// ============================================================================
32
33/// 规则违规的严重级别
34///
35/// 从低到高依次为 `Info` < `Warning` < `Error` < `Critical`。
36/// [`RuleReport::has_errors`] 将 `Error` 与 `Critical` 均视为阻断级。
37#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
38pub enum RuleSeverity {
39    /// 信息级:仅提示,不阻断
40    Info,
41    /// 警告级:潜在问题,建议修正
42    Warning,
43    /// 错误级:明确违规,应阻断
44    Error,
45    /// 严重级:高危违规,必须阻断
46    Critical,
47}
48
49impl RuleSeverity {
50    /// 返回级别的英文小写标识,便于日志聚合与序列化
51    pub fn as_str(&self) -> &'static str {
52        match self {
53            RuleSeverity::Info => "info",
54            RuleSeverity::Warning => "warning",
55            RuleSeverity::Error => "error",
56            RuleSeverity::Critical => "critical",
57        }
58    }
59
60    /// 返回级别的中文描述
61    pub fn description(&self) -> &'static str {
62        match self {
63            RuleSeverity::Info => "信息",
64            RuleSeverity::Warning => "警告",
65            RuleSeverity::Error => "错误",
66            RuleSeverity::Critical => "严重",
67        }
68    }
69
70    /// 是否为阻断级(Error 或 Critical)
71    pub fn is_blocking(&self) -> bool {
72        matches!(self, RuleSeverity::Error | RuleSeverity::Critical)
73    }
74}
75
76impl fmt::Display for RuleSeverity {
77    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
78        f.write_str(self.as_str())
79    }
80}
81
82// ============================================================================
83// 规则违规报告
84// ============================================================================
85
86/// 单条规则违规报告
87///
88/// 由 [`SqlRule::check`] 返回,描述一条 SQL 在某规则下的违规详情。
89#[derive(Debug, Clone, PartialEq, Eq)]
90pub struct RuleViolation {
91    /// 触发违规的规则名称
92    pub rule_name: String,
93    /// 违规严重级别
94    pub severity: RuleSeverity,
95    /// 人类可读的违规说明
96    pub message: String,
97    /// 违规在 SQL 文本中的字节偏移(0-based),`None` 表示无法定位
98    pub position: Option<usize>,
99}
100
101impl RuleViolation {
102    /// 创建一条无位置信息的违规
103    pub fn new(
104        rule_name: impl Into<String>,
105        severity: RuleSeverity,
106        message: impl Into<String>,
107    ) -> Self {
108        Self {
109            rule_name: rule_name.into(),
110            severity,
111            message: message.into(),
112            position: None,
113        }
114    }
115
116    /// 创建一条带位置信息的违规
117    pub fn with_position(
118        rule_name: impl Into<String>,
119        severity: RuleSeverity,
120        message: impl Into<String>,
121        position: usize,
122    ) -> Self {
123        Self {
124            rule_name: rule_name.into(),
125            severity,
126            message: message.into(),
127            position: Some(position),
128        }
129    }
130
131    /// 违规是否为阻断级
132    pub fn is_blocking(&self) -> bool {
133        self.severity.is_blocking()
134    }
135}
136
137impl fmt::Display for RuleViolation {
138    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
139        match self.position {
140            Some(pos) => write!(
141                f,
142                "[{}] {} at position {}: {}",
143                self.severity, self.rule_name, pos, self.message
144            ),
145            None => write!(
146                f,
147                "[{}] {}: {}",
148                self.severity, self.rule_name, self.message
149            ),
150        }
151    }
152}
153
154// ============================================================================
155// 规则匹配上下文
156// ============================================================================
157
158/// 规则匹配上下文,封装 SQL 文本及其预计算信息
159///
160/// 由 [`RuleEngine::check`] 在执行规则前构建,避免每条规则重复词法分析。
161#[derive(Debug, Clone)]
162pub struct RuleContext {
163    /// 原始 SQL 文本
164    pub sql: String,
165    /// 词法分析后的令牌序列
166    pub tokens: Vec<SqlToken>,
167    /// 语句类型
168    pub statement_type: SqlStatementType,
169    /// SQL 全大写形式(便于关键字匹配,避免重复计算)
170    pub sql_upper: String,
171}
172
173impl RuleContext {
174    /// 从 SQL 文本构建上下文
175    pub fn from_sql(sql: &str) -> Self {
176        Self {
177            sql: sql.to_string(),
178            tokens: tokenize(sql),
179            statement_type: detect_statement_type(sql),
180            sql_upper: sql.to_uppercase(),
181        }
182    }
183
184    /// 统计指定关键字的个数(基于令牌序列,精确匹配关键字令牌)
185    pub fn keyword_count(&self, keyword: &str) -> usize {
186        let upper = keyword.to_uppercase();
187        self.tokens
188            .iter()
189            .filter(|t| matches!(t, SqlToken::Keyword(k) if *k == upper))
190            .count()
191    }
192
193    /// 判断是否包含指定关键字
194    pub fn has_keyword(&self, keyword: &str) -> bool {
195        self.keyword_count(keyword) > 0
196    }
197
198    /// 统计 FROM/JOIN/INTO/UPDATE 后的表标识符个数
199    pub fn table_count(&self) -> usize {
200        let mut count = 0;
201        let mut expect_table = false;
202        for token in &self.tokens {
203            match token {
204                SqlToken::Keyword(k)
205                    if matches!(k.as_str(), "FROM" | "JOIN" | "INTO" | "UPDATE") =>
206                {
207                    expect_table = true;
208                }
209                SqlToken::Identifier(name) if expect_table && !name.starts_with('"') => {
210                    // 跳过别名(单个字母后跟 ON 或 WHERE)
211                    let _ = name;
212                    count += 1;
213                    expect_table = false;
214                }
215                SqlToken::Punctuation('.') if expect_table => {
216                    // 保留点号用于限定名
217                }
218                _ if expect_table => {
219                    // 遇到非标识符则取消等待
220                    expect_table = false;
221                }
222                _ => {}
223            }
224        }
225        count
226    }
227}
228
229// ============================================================================
230// 规则 trait
231// ============================================================================
232
233/// SQL 规则接口
234///
235/// 实现者定义一条具体的检查逻辑。规则应为无状态且线程安全(`Send + Sync`),
236/// 以便在规则引擎中共享。
237pub trait SqlRule: Send + Sync {
238    /// 规则名称,需在引擎内唯一
239    fn name(&self) -> &str;
240
241    /// 规则的默认严重级别
242    fn severity(&self) -> RuleSeverity;
243
244    /// 检查 SQL 上下文,返回 `Some(violation)` 表示违规,`None` 表示通过
245    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation>;
246}
247
248// ============================================================================
249// 内置规则
250// ============================================================================
251
252/// 禁止 `SELECT *` 规则
253///
254/// 检测 `SELECT *` 模式,建议显式列出列名以避免列顺序依赖与全表扫描。
255#[derive(Debug, Default)]
256pub struct NoSelectStarRule;
257
258impl SqlRule for NoSelectStarRule {
259    fn name(&self) -> &str {
260        "no_select_star"
261    }
262
263    fn severity(&self) -> RuleSeverity {
264        RuleSeverity::Warning
265    }
266
267    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
268        if ctx.statement_type != SqlStatementType::Select {
269            return None;
270        }
271        // 检测 SELECT 后紧跟 *(tokenize 将 * 归为 Operator)
272        let mut after_select = false;
273        for token in &ctx.tokens {
274            match token {
275                SqlToken::Keyword(k) if k == "SELECT" => {
276                    after_select = true;
277                }
278                SqlToken::Operator(op) if after_select && op == "*" => {
279                    let pos = ctx.sql.find('*');
280                    return Some(RuleViolation::with_position(
281                        self.name(),
282                        self.severity(),
283                        "SELECT * 不允许,请显式列出列名",
284                        pos.unwrap_or(0),
285                    ));
286                }
287                _ if after_select => {
288                    // 遇到非 * 的令牌则不再跟踪
289                    after_select = false;
290                }
291                _ => {}
292            }
293        }
294        None
295    }
296}
297
298/// DELETE 必须包含 WHERE 子句规则
299///
300/// 无 WHERE 的 DELETE 会清空全表,属高危操作。
301#[derive(Debug, Default)]
302pub struct RequireWhereInDeleteRule;
303
304impl SqlRule for RequireWhereInDeleteRule {
305    fn name(&self) -> &str {
306        "require_where_in_delete"
307    }
308
309    fn severity(&self) -> RuleSeverity {
310        RuleSeverity::Critical
311    }
312
313    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
314        if ctx.statement_type != SqlStatementType::Delete {
315            return None;
316        }
317        if !ctx.has_keyword("WHERE") {
318            Some(RuleViolation::new(
319                self.name(),
320                self.severity(),
321                "DELETE 语句必须包含 WHERE 子句,否则将清空全表",
322            ))
323        } else {
324            None
325        }
326    }
327}
328
329/// UPDATE 必须包含 WHERE 子句规则
330///
331/// 无 WHERE 的 UPDATE 会更新全表,属高危操作。
332#[derive(Debug, Default)]
333pub struct RequireWhereInUpdateRule;
334
335impl SqlRule for RequireWhereInUpdateRule {
336    fn name(&self) -> &str {
337        "require_where_in_update"
338    }
339
340    fn severity(&self) -> RuleSeverity {
341        RuleSeverity::Critical
342    }
343
344    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
345        if ctx.statement_type != SqlStatementType::Update {
346            return None;
347        }
348        if !ctx.has_keyword("WHERE") {
349            Some(RuleViolation::new(
350                self.name(),
351                self.severity(),
352                "UPDATE 语句必须包含 WHERE 子句,否则将更新全表",
353            ))
354        } else {
355            None
356        }
357    }
358}
359
360/// 表数量上限规则
361///
362/// 限制单条 SQL 涉及的表数量(FROM/JOIN/INTO/UPDATE 后的标识符),
363/// 防止过度复杂的跨表查询。
364#[derive(Debug)]
365pub struct MaxTableCountRule {
366    /// 允许的最大表数量
367    pub max: usize,
368}
369
370impl MaxTableCountRule {
371    /// 创建规则,指定最大表数量
372    pub fn new(max: usize) -> Self {
373        Self { max }
374    }
375}
376
377impl SqlRule for MaxTableCountRule {
378    fn name(&self) -> &str {
379        "max_table_count"
380    }
381
382    fn severity(&self) -> RuleSeverity {
383        RuleSeverity::Warning
384    }
385
386    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
387        let count = ctx.table_count();
388        if count > self.max {
389            Some(RuleViolation::new(
390                self.name(),
391                self.severity(),
392                format!("涉及表数量 {} 超过上限 {}", count, self.max),
393            ))
394        } else {
395            None
396        }
397    }
398}
399
400/// JOIN 数量上限规则
401#[derive(Debug)]
402pub struct MaxJoinCountRule {
403    /// 允许的最大 JOIN 数量
404    pub max: usize,
405}
406
407impl MaxJoinCountRule {
408    /// 创建规则,指定最大 JOIN 数量
409    pub fn new(max: usize) -> Self {
410        Self { max }
411    }
412}
413
414impl SqlRule for MaxJoinCountRule {
415    fn name(&self) -> &str {
416        "max_join_count"
417    }
418
419    fn severity(&self) -> RuleSeverity {
420        RuleSeverity::Warning
421    }
422
423    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
424        let count = ctx.keyword_count("JOIN");
425        if count > self.max {
426            Some(RuleViolation::new(
427                self.name(),
428                self.severity(),
429                format!("JOIN 数量 {} 超过上限 {}", count, self.max),
430            ))
431        } else {
432            None
433        }
434    }
435}
436
437/// 禁止 UNION 规则
438///
439/// UNION 可能导致结果集不可预测与性能问题,部分场景需禁用。
440#[derive(Debug, Default)]
441pub struct NoUnionRule;
442
443impl SqlRule for NoUnionRule {
444    fn name(&self) -> &str {
445        "no_union"
446    }
447
448    fn severity(&self) -> RuleSeverity {
449        RuleSeverity::Error
450    }
451
452    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
453        if ctx.has_keyword("UNION") {
454            let pos = ctx.sql_upper.find("UNION");
455            Some(RuleViolation::with_position(
456                self.name(),
457                self.severity(),
458                "UNION 操作不允许",
459                pos.unwrap_or(0),
460            ))
461        } else {
462            None
463        }
464    }
465}
466
467/// 禁止关键字规则
468///
469/// 检测 SQL 中是否包含指定的禁止关键字(如 `GRANT`、`REVOKE`、`EXEC`)。
470#[derive(Debug)]
471pub struct ForbiddenKeywordRule {
472    /// 规则名称
473    pub rule_name: String,
474    /// 禁止的关键字列表(大写)
475    pub keywords: Vec<String>,
476    /// 违规严重级别
477    pub severity: RuleSeverity,
478}
479
480impl ForbiddenKeywordRule {
481    /// 创建规则,指定名称与关键字列表
482    pub fn new(rule_name: impl Into<String>, keywords: &[&str]) -> Self {
483        Self {
484            rule_name: rule_name.into(),
485            keywords: keywords.iter().map(|k| k.to_uppercase()).collect(),
486            severity: RuleSeverity::Critical,
487        }
488    }
489
490    /// 设置违规严重级别
491    pub fn with_severity(mut self, severity: RuleSeverity) -> Self {
492        self.severity = severity;
493        self
494    }
495}
496
497impl SqlRule for ForbiddenKeywordRule {
498    fn name(&self) -> &str {
499        &self.rule_name
500    }
501
502    fn severity(&self) -> RuleSeverity {
503        self.severity
504    }
505
506    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
507        for keyword in &self.keywords {
508            if ctx.has_keyword(keyword) {
509                let pos = ctx.sql_upper.find(keyword.as_str());
510                return Some(RuleViolation::with_position(
511                    self.name(),
512                    self.severity,
513                    format!("禁止关键字 {}", keyword),
514                    pos.unwrap_or(0),
515                ));
516            }
517        }
518        None
519    }
520}
521
522/// 正则匹配规则
523///
524/// 使用正则表达式检测 SQL 文本,匹配则视为违规。
525/// 适用于自定义模式检测,如敏感表访问、特定函数调用等。
526#[derive(Debug)]
527pub struct RegexRule {
528    /// 规则名称
529    pub rule_name: String,
530    /// 编译后的正则表达式
531    pub pattern: Regex,
532    /// 违规严重级别
533    pub severity: RuleSeverity,
534    /// 违规说明模板
535    pub message: String,
536}
537
538impl RegexRule {
539    /// 创建正则规则,指定名称与正则字符串
540    pub fn new(
541        rule_name: impl Into<String>,
542        pattern: &str,
543        severity: RuleSeverity,
544        message: impl Into<String>,
545    ) -> Result<Self, regex::Error> {
546        Ok(Self {
547            rule_name: rule_name.into(),
548            pattern: Regex::new(pattern)?,
549            severity,
550            message: message.into(),
551        })
552    }
553}
554
555impl SqlRule for RegexRule {
556    fn name(&self) -> &str {
557        &self.rule_name
558    }
559
560    fn severity(&self) -> RuleSeverity {
561        self.severity
562    }
563
564    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
565        self.pattern.find(&ctx.sql).map(|m| {
566            RuleViolation::with_position(
567                self.name(),
568                self.severity,
569                self.message.clone(),
570                m.start(),
571            )
572        })
573    }
574}
575
576/// 限制结果集行数规则
577///
578/// 检测 SELECT 语句是否包含 LIMIT 子句,未包含则报告违规。
579/// 适用于防止全表扫描返回过大结果集的场景。
580#[derive(Debug, Default)]
581pub struct RequireLimitRule;
582
583impl SqlRule for RequireLimitRule {
584    fn name(&self) -> &str {
585        "require_limit"
586    }
587
588    fn severity(&self) -> RuleSeverity {
589        RuleSeverity::Info
590    }
591
592    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
593        if ctx.statement_type != SqlStatementType::Select {
594            return None;
595        }
596        // 子查询中的 SELECT 不要求 LIMIT,仅检查顶层
597        if !ctx.has_keyword("LIMIT") && !ctx.has_keyword("FETCH") {
598            Some(RuleViolation::new(
599                self.name(),
600                self.severity(),
601                "SELECT 语句建议包含 LIMIT 子句以限制结果集大小",
602            ))
603        } else {
604            None
605        }
606    }
607}
608
609/// 限制 SELECT 列数上限(防止超宽结果集)
610pub struct MaxColumnCountRule {
611    /// 最大允许列数
612    pub max_columns: usize,
613}
614
615impl Default for MaxColumnCountRule {
616    fn default() -> Self {
617        Self { max_columns: 20 }
618    }
619}
620
621impl SqlRule for MaxColumnCountRule {
622    fn name(&self) -> &str {
623        "max_column_count"
624    }
625
626    fn severity(&self) -> RuleSeverity {
627        RuleSeverity::Warning
628    }
629
630    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
631        if ctx.statement_type != SqlStatementType::Select {
632            return None;
633        }
634        let select_idx = ctx
635            .tokens
636            .iter()
637            .position(|t| matches!(t, SqlToken::Keyword(k) if k == "SELECT"))?;
638        let from_idx = ctx
639            .tokens
640            .iter()
641            .position(|t| matches!(t, SqlToken::Keyword(k) if k == "FROM"))?;
642        if from_idx <= select_idx {
643            return None;
644        }
645        let comma_count = ctx.tokens[select_idx + 1..from_idx]
646            .iter()
647            .filter(|t| matches!(t, SqlToken::Punctuation(',')))
648            .count();
649        let col_count = comma_count + 1;
650        if col_count > self.max_columns {
651            Some(RuleViolation::new(
652                self.name(),
653                self.severity(),
654                format!("SELECT 列数 {} 超过上限 {}", col_count, self.max_columns),
655            ))
656        } else {
657            None
658        }
659    }
660}
661
662/// 禁止在 SELECT 中使用子查询(SELECT ... FROM (SELECT ...))
663pub struct NoSubqueryRule;
664
665impl SqlRule for NoSubqueryRule {
666    fn name(&self) -> &str {
667        "no_subquery"
668    }
669
670    fn severity(&self) -> RuleSeverity {
671        RuleSeverity::Warning
672    }
673
674    fn check(&self, ctx: &RuleContext) -> Option<RuleViolation> {
675        if ctx.statement_type != SqlStatementType::Select {
676            return None;
677        }
678        for i in 0..ctx.tokens.len().saturating_sub(2) {
679            let is_from = matches!(&ctx.tokens[i], SqlToken::Keyword(k) if k == "FROM");
680            let is_open = matches!(ctx.tokens[i + 1], SqlToken::Punctuation('('));
681            let is_select = matches!(&ctx.tokens[i + 2], SqlToken::Keyword(k) if k == "SELECT");
682            if is_from && is_open && is_select {
683                return Some(RuleViolation::new(
684                    self.name(),
685                    self.severity(),
686                    "SELECT 语句中包含子查询,建议改用 JOIN",
687                ));
688            }
689        }
690        None
691    }
692}
693
694// ============================================================================
695// 规则报告
696// ============================================================================
697
698/// 规则检查汇总报告
699#[derive(Debug, Clone, Default)]
700pub struct RuleReport {
701    /// 所有违规项
702    pub violations: Vec<RuleViolation>,
703}
704
705impl RuleReport {
706    /// 是否无违规
707    pub fn is_clean(&self) -> bool {
708        self.violations.is_empty()
709    }
710
711    /// 是否存在违规
712    pub fn has_violations(&self) -> bool {
713        !self.violations.is_empty()
714    }
715
716    /// 是否存在阻断级违规(Error 或 Critical)
717    pub fn has_blocking(&self) -> bool {
718        self.violations.iter().any(|v| v.is_blocking())
719    }
720
721    /// 是否存在 Error 级违规
722    pub fn has_errors(&self) -> bool {
723        self.violations
724            .iter()
725            .any(|v| v.severity == RuleSeverity::Error)
726    }
727
728    /// 是否存在 Critical 级违规
729    pub fn has_critical(&self) -> bool {
730        self.violations
731            .iter()
732            .any(|v| v.severity == RuleSeverity::Critical)
733    }
734
735    /// 返回指定级别的违规数量
736    pub fn count_by_severity(&self, severity: RuleSeverity) -> usize {
737        self.violations
738            .iter()
739            .filter(|v| v.severity == severity)
740            .count()
741    }
742
743    /// 返回指定级别的违规引用列表
744    pub fn violations_by_severity(&self, severity: RuleSeverity) -> Vec<&RuleViolation> {
745        self.violations
746            .iter()
747            .filter(|v| v.severity == severity)
748            .collect()
749    }
750
751    /// 返回阻断级违规引用列表
752    pub fn blocking_violations(&self) -> Vec<&RuleViolation> {
753        self.violations.iter().filter(|v| v.is_blocking()).collect()
754    }
755
756    /// 违规总数
757    pub fn violation_count(&self) -> usize {
758        self.violations.len()
759    }
760
761    /// 生成人类可读的摘要报告
762    pub fn summary(&self) -> String {
763        if self.is_clean() {
764            return "规则检查通过,无违规".to_string();
765        }
766        let mut lines = Vec::with_capacity(self.violations.len() + 2);
767        lines.push(format!(
768            "规则检查完成:共 {} 条违规({} 信息,{} 警告,{} 错误,{} 严重)",
769            self.violation_count(),
770            self.count_by_severity(RuleSeverity::Info),
771            self.count_by_severity(RuleSeverity::Warning),
772            self.count_by_severity(RuleSeverity::Error),
773            self.count_by_severity(RuleSeverity::Critical),
774        ));
775        for v in &self.violations {
776            lines.push(format!("  - {}", v));
777        }
778        lines.join("\n")
779    }
780}
781
782impl fmt::Display for RuleReport {
783    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
784        f.write_str(&self.summary())
785    }
786}
787
788// ============================================================================
789// 规则引擎
790// ============================================================================
791
792/// 规则引擎,管理规则集合并批量执行检查
793///
794/// 规则按添加顺序执行,所有规则均会运行(不短路),
795/// 以便一次性收集全部违规。
796pub struct RuleEngine {
797    rules: Vec<Box<dyn SqlRule>>,
798}
799
800impl RuleEngine {
801    /// 创建空规则引擎
802    pub fn new() -> Self {
803        Self { rules: Vec::new() }
804    }
805
806    /// 创建包含默认规则集的引擎
807    ///
808    /// 默认规则集包含:
809    /// - [`NoSelectStarRule`]
810    /// - [`RequireWhereInDeleteRule`]
811    /// - [`RequireWhereInUpdateRule`]
812    /// - [`NoUnionRule`]
813    /// - [`ForbiddenKeywordRule`](禁止 GRANT/REVOKE/EXEC/EXECUTE)
814    pub fn with_default_rules() -> Self {
815        let mut engine = Self::new();
816        engine.add_rule(Box::new(NoSelectStarRule));
817        engine.add_rule(Box::new(RequireWhereInDeleteRule));
818        engine.add_rule(Box::new(RequireWhereInUpdateRule));
819        engine.add_rule(Box::new(NoUnionRule));
820        engine.add_rule(Box::new(ForbiddenKeywordRule::new(
821            "forbidden_privilege_ops",
822            &["GRANT", "REVOKE", "EXEC", "EXECUTE"],
823        )));
824        engine
825    }
826
827    /// 添加一条规则
828    pub fn add_rule(&mut self, rule: Box<dyn SqlRule>) -> &mut Self {
829        self.rules.push(rule);
830        self
831    }
832
833    /// 返回规则数量
834    pub fn rule_count(&self) -> usize {
835        self.rules.len()
836    }
837
838    /// 对单条 SQL 执行所有规则,返回汇总报告
839    pub fn check(&self, sql: &str) -> RuleReport {
840        let ctx = RuleContext::from_sql(sql);
841        let violations = self
842            .rules
843            .iter()
844            .filter_map(|rule| rule.check(&ctx))
845            .collect();
846        RuleReport { violations }
847    }
848
849    /// 对多条 SQL 批量执行检查,返回每条 SQL 的报告
850    pub fn check_batch<'a>(&self, sqls: impl IntoIterator<Item = &'a str>) -> Vec<RuleReport> {
851        sqls.into_iter().map(|sql| self.check(sql)).collect()
852    }
853
854    /// 检查 SQL 是否通过所有规则(无阻断级违规)
855    pub fn passes(&self, sql: &str) -> bool {
856        !self.check(sql).has_blocking()
857    }
858}
859
860impl Default for RuleEngine {
861    fn default() -> Self {
862        Self::new()
863    }
864}
865
866// ============================================================================
867// 规则集预设
868// ============================================================================
869
870/// 规则集预设,提供常见场景的规则组合
871pub struct RulePresets;
872
873impl RulePresets {
874    /// 生产环境严格规则集
875    ///
876    /// 包含默认规则外加:
877    /// - 表数量上限 5
878    /// - JOIN 数量上限 3
879    /// - 要求 LIMIT
880    pub fn strict() -> RuleEngine {
881        let mut engine = RuleEngine::with_default_rules();
882        engine.add_rule(Box::new(MaxTableCountRule::new(5)));
883        engine.add_rule(Box::new(MaxJoinCountRule::new(3)));
884        engine.add_rule(Box::new(RequireLimitRule));
885        engine
886    }
887
888    /// 只读查询规则集
889    ///
890    /// 仅检查 SELECT 相关规则,不检查 DML。
891    pub fn read_only() -> RuleEngine {
892        let mut engine = RuleEngine::new();
893        engine.add_rule(Box::new(NoSelectStarRule));
894        engine.add_rule(Box::new(NoUnionRule));
895        engine.add_rule(Box::new(MaxTableCountRule::new(10)));
896        engine.add_rule(Box::new(MaxJoinCountRule::new(5)));
897        engine
898    }
899}
900
901// ============================================================================
902// 测试
903// ============================================================================
904
905#[cfg(test)]
906mod tests {
907    use super::*;
908
909    // ---- RuleSeverity 测试 ----
910
911    #[test]
912    fn test_rule_severity_ordering() {
913        assert!(RuleSeverity::Info < RuleSeverity::Warning);
914        assert!(RuleSeverity::Warning < RuleSeverity::Error);
915        assert!(RuleSeverity::Error < RuleSeverity::Critical);
916    }
917
918    #[test]
919    fn test_rule_severity_as_str() {
920        assert_eq!(RuleSeverity::Info.as_str(), "info");
921        assert_eq!(RuleSeverity::Warning.as_str(), "warning");
922        assert_eq!(RuleSeverity::Error.as_str(), "error");
923        assert_eq!(RuleSeverity::Critical.as_str(), "critical");
924    }
925
926    #[test]
927    fn test_rule_severity_description() {
928        assert_eq!(RuleSeverity::Info.description(), "信息");
929        assert_eq!(RuleSeverity::Critical.description(), "严重");
930    }
931
932    #[test]
933    fn test_rule_severity_is_blocking() {
934        assert!(!RuleSeverity::Info.is_blocking());
935        assert!(!RuleSeverity::Warning.is_blocking());
936        assert!(RuleSeverity::Error.is_blocking());
937        assert!(RuleSeverity::Critical.is_blocking());
938    }
939
940    // ---- RuleViolation 测试 ----
941
942    #[test]
943    fn test_rule_violation_new() {
944        let v = RuleViolation::new("test_rule", RuleSeverity::Warning, "test message");
945        assert_eq!(v.rule_name, "test_rule");
946        assert_eq!(v.severity, RuleSeverity::Warning);
947        assert_eq!(v.message, "test message");
948        assert!(v.position.is_none());
949    }
950
951    #[test]
952    fn test_rule_violation_with_position() {
953        let v = RuleViolation::with_position("test_rule", RuleSeverity::Error, "test message", 42);
954        assert_eq!(v.position, Some(42));
955        assert!(v.is_blocking());
956    }
957
958    #[test]
959    fn test_rule_violation_display() {
960        let v = RuleViolation::new("r", RuleSeverity::Warning, "msg");
961        let s = format!("{}", v);
962        assert!(s.contains("[warning]"));
963        assert!(s.contains("r"));
964        assert!(s.contains("msg"));
965
966        let v2 = RuleViolation::with_position("r", RuleSeverity::Error, "msg", 10);
967        let s2 = format!("{}", v2);
968        assert!(s2.contains("position 10"));
969    }
970
971    // ---- RuleContext 测试 ----
972
973    #[test]
974    fn test_rule_context_from_sql() {
975        let ctx = RuleContext::from_sql("SELECT id FROM users");
976        assert_eq!(ctx.statement_type, SqlStatementType::Select);
977        assert!(ctx.sql_upper.contains("SELECT"));
978        assert!(!ctx.tokens.is_empty());
979    }
980
981    #[test]
982    fn test_rule_context_keyword_count() {
983        let ctx = RuleContext::from_sql("SELECT id FROM users JOIN orders ON 1=1");
984        assert_eq!(ctx.keyword_count("JOIN"), 1);
985        assert_eq!(ctx.keyword_count("SELECT"), 1);
986        assert_eq!(ctx.keyword_count("DELETE"), 0);
987    }
988
989    #[test]
990    fn test_rule_context_has_keyword() {
991        let ctx = RuleContext::from_sql("SELECT id FROM users WHERE id = 1");
992        assert!(ctx.has_keyword("WHERE"));
993        assert!(!ctx.has_keyword("DELETE"));
994    }
995
996    #[test]
997    fn test_rule_context_table_count() {
998        let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON users.id = orders.uid");
999        assert!(ctx.table_count() >= 2);
1000    }
1001
1002    // ---- NoSelectStarRule 测试 ----
1003
1004    #[test]
1005    fn test_no_select_star_rule_detects_star() {
1006        let rule = NoSelectStarRule;
1007        let ctx = RuleContext::from_sql("SELECT * FROM users");
1008        let violation = rule.check(&ctx);
1009        assert!(violation.is_some());
1010        let v = violation.unwrap();
1011        assert_eq!(v.severity, RuleSeverity::Warning);
1012        assert!(v.position.is_some());
1013    }
1014
1015    #[test]
1016    fn test_no_select_star_rule_passes_explicit_columns() {
1017        let rule = NoSelectStarRule;
1018        let ctx = RuleContext::from_sql("SELECT id, name FROM users");
1019        assert!(rule.check(&ctx).is_none());
1020    }
1021
1022    #[test]
1023    fn test_no_select_star_rule_ignores_non_select() {
1024        let rule = NoSelectStarRule;
1025        let ctx = RuleContext::from_sql("DELETE FROM users WHERE id = 1");
1026        assert!(rule.check(&ctx).is_none());
1027    }
1028
1029    // ---- RequireWhereInDeleteRule 测试 ----
1030
1031    #[test]
1032    fn test_require_where_in_delete_rule_detects_missing_where() {
1033        let rule = RequireWhereInDeleteRule;
1034        let ctx = RuleContext::from_sql("DELETE FROM users");
1035        let violation = rule.check(&ctx);
1036        assert!(violation.is_some());
1037        assert_eq!(violation.unwrap().severity, RuleSeverity::Critical);
1038    }
1039
1040    #[test]
1041    fn test_require_where_in_delete_rule_passes_with_where() {
1042        let rule = RequireWhereInDeleteRule;
1043        let ctx = RuleContext::from_sql("DELETE FROM users WHERE id = 1");
1044        assert!(rule.check(&ctx).is_none());
1045    }
1046
1047    #[test]
1048    fn test_require_where_in_delete_rule_ignores_non_delete() {
1049        let rule = RequireWhereInDeleteRule;
1050        let ctx = RuleContext::from_sql("SELECT * FROM users");
1051        assert!(rule.check(&ctx).is_none());
1052    }
1053
1054    // ---- RequireWhereInUpdateRule 测试 ----
1055
1056    #[test]
1057    fn test_require_where_in_update_rule_detects_missing_where() {
1058        let rule = RequireWhereInUpdateRule;
1059        let ctx = RuleContext::from_sql("UPDATE users SET name = 'a'");
1060        let violation = rule.check(&ctx);
1061        assert!(violation.is_some());
1062        assert_eq!(violation.unwrap().severity, RuleSeverity::Critical);
1063    }
1064
1065    #[test]
1066    fn test_require_where_in_update_rule_passes_with_where() {
1067        let rule = RequireWhereInUpdateRule;
1068        let ctx = RuleContext::from_sql("UPDATE users SET name = 'a' WHERE id = 1");
1069        assert!(rule.check(&ctx).is_none());
1070    }
1071
1072    // ---- MaxTableCountRule 测试 ----
1073
1074    #[test]
1075    fn test_max_table_count_rule_passes_within_limit() {
1076        let rule = MaxTableCountRule::new(2);
1077        let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON 1=1");
1078        assert!(rule.check(&ctx).is_none());
1079    }
1080
1081    #[test]
1082    fn test_max_table_count_rule_detects_exceed() {
1083        let rule = MaxTableCountRule::new(1);
1084        let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON 1=1");
1085        let violation = rule.check(&ctx);
1086        assert!(violation.is_some());
1087        assert!(violation.unwrap().message.contains("超过上限"));
1088    }
1089
1090    // ---- MaxJoinCountRule 测试 ----
1091
1092    #[test]
1093    fn test_max_join_count_rule_passes_within_limit() {
1094        let rule = MaxJoinCountRule::new(2);
1095        let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON 1=1 JOIN items ON 2=2");
1096        assert!(rule.check(&ctx).is_none());
1097    }
1098
1099    #[test]
1100    fn test_max_join_count_rule_detects_exceed() {
1101        let rule = MaxJoinCountRule::new(1);
1102        let ctx = RuleContext::from_sql("SELECT * FROM users JOIN orders ON 1=1 JOIN items ON 2=2");
1103        assert!(rule.check(&ctx).is_some());
1104    }
1105
1106    // ---- NoUnionRule 测试 ----
1107
1108    #[test]
1109    fn test_no_union_rule_detects_union() {
1110        let rule = NoUnionRule;
1111        let ctx = RuleContext::from_sql("SELECT id FROM users UNION SELECT id FROM archived");
1112        let violation = rule.check(&ctx);
1113        assert!(violation.is_some());
1114        assert_eq!(violation.unwrap().severity, RuleSeverity::Error);
1115    }
1116
1117    #[test]
1118    fn test_no_union_rule_passes_without_union() {
1119        let rule = NoUnionRule;
1120        let ctx = RuleContext::from_sql("SELECT id FROM users");
1121        assert!(rule.check(&ctx).is_none());
1122    }
1123
1124    // ---- ForbiddenKeywordRule 测试 ----
1125
1126    #[test]
1127    fn test_forbidden_keyword_rule_detects_keyword() {
1128        let rule = ForbiddenKeywordRule::new("no_grant", &["GRANT"]);
1129        let ctx = RuleContext::from_sql("GRANT ALL ON users TO hacker");
1130        let violation = rule.check(&ctx);
1131        assert!(violation.is_some());
1132        assert_eq!(violation.unwrap().severity, RuleSeverity::Critical);
1133    }
1134
1135    #[test]
1136    fn test_forbidden_keyword_rule_passes_without_keyword() {
1137        let rule = ForbiddenKeywordRule::new("no_grant", &["GRANT"]);
1138        let ctx = RuleContext::from_sql("SELECT * FROM users");
1139        assert!(rule.check(&ctx).is_none());
1140    }
1141
1142    #[test]
1143    fn test_forbidden_keyword_rule_with_severity() {
1144        let rule =
1145            ForbiddenKeywordRule::new("no_exec", &["EXEC"]).with_severity(RuleSeverity::Warning);
1146        let ctx = RuleContext::from_sql("EXEC sp_executesql 'DROP TABLE users'");
1147        let violation = rule.check(&ctx).unwrap();
1148        assert_eq!(violation.severity, RuleSeverity::Warning);
1149    }
1150
1151    // ---- RegexRule 测试 ----
1152
1153    #[test]
1154    fn test_regex_rule_detects_match() {
1155        let rule = RegexRule::new(
1156            "no_sys_tables",
1157            r"(?i)\bsys\.",
1158            RuleSeverity::Error,
1159            "禁止访问系统表",
1160        )
1161        .unwrap();
1162        let ctx = RuleContext::from_sql("SELECT * FROM sys.tables");
1163        let violation = rule.check(&ctx);
1164        assert!(violation.is_some());
1165        assert!(violation.unwrap().position.is_some());
1166    }
1167
1168    #[test]
1169    fn test_regex_rule_passes_no_match() {
1170        let rule = RegexRule::new(
1171            "no_sys_tables",
1172            r"(?i)\bsys\.",
1173            RuleSeverity::Error,
1174            "禁止访问系统表",
1175        )
1176        .unwrap();
1177        let ctx = RuleContext::from_sql("SELECT * FROM users");
1178        assert!(rule.check(&ctx).is_none());
1179    }
1180
1181    #[test]
1182    fn test_regex_rule_invalid_pattern() {
1183        let result = RegexRule::new("bad", r"[invalid", RuleSeverity::Error, "msg");
1184        assert!(result.is_err());
1185    }
1186
1187    // ---- RequireLimitRule 测试 ----
1188
1189    #[test]
1190    fn test_require_limit_rule_detects_missing_limit() {
1191        let rule = RequireLimitRule;
1192        let ctx = RuleContext::from_sql("SELECT id FROM users");
1193        let violation = rule.check(&ctx);
1194        assert!(violation.is_some());
1195        assert_eq!(violation.unwrap().severity, RuleSeverity::Info);
1196    }
1197
1198    #[test]
1199    fn test_require_limit_rule_passes_with_limit() {
1200        let rule = RequireLimitRule;
1201        let ctx = RuleContext::from_sql("SELECT id FROM users LIMIT 10");
1202        assert!(rule.check(&ctx).is_none());
1203    }
1204
1205    // ---- RuleReport 测试 ----
1206
1207    #[test]
1208    fn test_rule_report_clean() {
1209        let report = RuleReport::default();
1210        assert!(report.is_clean());
1211        assert!(!report.has_violations());
1212        assert!(!report.has_blocking());
1213        assert_eq!(report.violation_count(), 0);
1214    }
1215
1216    #[test]
1217    fn test_rule_report_with_violations() {
1218        let report = RuleReport {
1219            violations: vec![
1220                RuleViolation::new("r1", RuleSeverity::Warning, "w"),
1221                RuleViolation::new("r2", RuleSeverity::Error, "e"),
1222                RuleViolation::new("r3", RuleSeverity::Critical, "c"),
1223            ],
1224        };
1225        assert!(report.has_violations());
1226        assert!(report.has_blocking());
1227        assert!(report.has_errors());
1228        assert!(report.has_critical());
1229        assert_eq!(report.violation_count(), 3);
1230        assert_eq!(report.count_by_severity(RuleSeverity::Warning), 1);
1231        assert_eq!(report.count_by_severity(RuleSeverity::Error), 1);
1232        assert_eq!(report.count_by_severity(RuleSeverity::Critical), 1);
1233    }
1234
1235    #[test]
1236    fn test_rule_report_summary() {
1237        let report = RuleReport::default();
1238        assert!(report.summary().contains("无违规"));
1239
1240        let report2 = RuleReport {
1241            violations: vec![RuleViolation::new("r1", RuleSeverity::Error, "msg")],
1242        };
1243        let summary = report2.summary();
1244        assert!(summary.contains("1 条违规"));
1245        assert!(summary.contains("r1"));
1246    }
1247
1248    #[test]
1249    fn test_rule_report_violations_by_severity() {
1250        let report = RuleReport {
1251            violations: vec![
1252                RuleViolation::new("r1", RuleSeverity::Info, "i"),
1253                RuleViolation::new("r2", RuleSeverity::Info, "i2"),
1254                RuleViolation::new("r3", RuleSeverity::Error, "e"),
1255            ],
1256        };
1257        assert_eq!(report.violations_by_severity(RuleSeverity::Info).len(), 2);
1258        assert_eq!(report.violations_by_severity(RuleSeverity::Error).len(), 1);
1259        assert_eq!(report.blocking_violations().len(), 1);
1260    }
1261
1262    // ---- RuleEngine 测试 ----
1263
1264    #[test]
1265    fn test_rule_engine_new_empty() {
1266        let engine = RuleEngine::new();
1267        assert_eq!(engine.rule_count(), 0);
1268        let report = engine.check("SELECT * FROM users");
1269        assert!(report.is_clean());
1270    }
1271
1272    #[test]
1273    fn test_rule_engine_add_rule() {
1274        let mut engine = RuleEngine::new();
1275        engine.add_rule(Box::new(NoSelectStarRule));
1276        assert_eq!(engine.rule_count(), 1);
1277    }
1278
1279    #[test]
1280    fn test_rule_engine_check_collects_violations() {
1281        let mut engine = RuleEngine::new();
1282        engine.add_rule(Box::new(NoSelectStarRule));
1283        engine.add_rule(Box::new(RequireWhereInDeleteRule));
1284        let report = engine.check("DELETE FROM users");
1285        // RequireWhereInDeleteRule 应触发
1286        assert!(report.has_violations());
1287        assert!(report.has_critical());
1288    }
1289
1290    #[test]
1291    fn test_rule_engine_check_clean_sql() {
1292        let mut engine = RuleEngine::new();
1293        engine.add_rule(Box::new(NoSelectStarRule));
1294        engine.add_rule(Box::new(RequireWhereInDeleteRule));
1295        let report = engine.check("SELECT id, name FROM users WHERE id = 1");
1296        assert!(report.is_clean());
1297    }
1298
1299    #[test]
1300    fn test_rule_engine_passes() {
1301        let mut engine = RuleEngine::new();
1302        engine.add_rule(Box::new(RequireWhereInDeleteRule));
1303        assert!(engine.passes("DELETE FROM users WHERE id = 1"));
1304        assert!(!engine.passes("DELETE FROM users"));
1305    }
1306
1307    #[test]
1308    fn test_rule_engine_check_batch() {
1309        let mut engine = RuleEngine::new();
1310        engine.add_rule(Box::new(NoSelectStarRule));
1311        let reports = engine.check_batch(["SELECT * FROM users", "SELECT id FROM users"]);
1312        assert_eq!(reports.len(), 2);
1313        assert!(reports[0].has_violations());
1314        assert!(reports[1].is_clean());
1315    }
1316
1317    #[test]
1318    fn test_rule_engine_with_default_rules() {
1319        let engine = RuleEngine::with_default_rules();
1320        assert!(engine.rule_count() >= 5);
1321        // DROP 不在默认规则中(由 firewall 处理),但 GRANT 应被禁止
1322        let report = engine.check("GRANT ALL ON users TO hacker");
1323        assert!(report.has_critical());
1324    }
1325
1326    #[test]
1327    fn test_rule_engine_default_rules_select_star() {
1328        let engine = RuleEngine::with_default_rules();
1329        let report = engine.check("SELECT * FROM users");
1330        assert!(report.has_violations());
1331    }
1332
1333    #[test]
1334    fn test_rule_engine_default_rules_clean_sql() {
1335        let engine = RuleEngine::with_default_rules();
1336        let report = engine.check("SELECT id, name FROM users WHERE id = 1 LIMIT 10");
1337        assert!(report.is_clean());
1338    }
1339
1340    // ---- RulePresets 测试 ----
1341
1342    #[test]
1343    fn test_rule_presets_strict() {
1344        let engine = RulePresets::strict();
1345        assert!(engine.rule_count() >= 8);
1346        // 缺少 LIMIT 应触发 Info 级违规
1347        let report = engine.check("SELECT id FROM users WHERE id = 1");
1348        assert!(report.has_violations());
1349    }
1350
1351    #[test]
1352    fn test_rule_presets_read_only() {
1353        let engine = RulePresets::read_only();
1354        // DELETE 不受只读规则集限制(无 RequireWhereInDeleteRule)
1355        let report = engine.check("DELETE FROM users");
1356        assert!(report.is_clean());
1357        // SELECT * 应触发
1358        let report2 = engine.check("SELECT * FROM users");
1359        assert!(report2.has_violations());
1360    }
1361
1362    #[test]
1363    fn test_rule_presets_strict_max_table_count() {
1364        let engine = RulePresets::strict();
1365        let sql =
1366            "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";
1367        let report = engine.check(sql);
1368        assert!(report.has_violations());
1369    }
1370
1371    // ---- MaxColumnCountRule 测试 ----
1372
1373    #[test]
1374    fn test_max_column_count_pass() {
1375        let rule = MaxColumnCountRule { max_columns: 5 };
1376        let ctx = RuleContext::from_sql("SELECT a, b, c FROM users");
1377        assert!(rule.check(&ctx).is_none());
1378    }
1379
1380    #[test]
1381    fn test_max_column_count_fail() {
1382        let rule = MaxColumnCountRule { max_columns: 2 };
1383        let ctx = RuleContext::from_sql("SELECT a, b, c, d FROM users");
1384        assert!(rule.check(&ctx).is_some());
1385    }
1386
1387    #[test]
1388    fn test_max_column_count_default() {
1389        let rule = MaxColumnCountRule::default();
1390        assert_eq!(rule.max_columns, 20);
1391    }
1392
1393    // ---- NoSubqueryRule 测试 ----
1394
1395    #[test]
1396    fn test_no_subquery_clean() {
1397        let rule = NoSubqueryRule;
1398        let ctx = RuleContext::from_sql("SELECT id FROM users WHERE id = 1");
1399        assert!(rule.check(&ctx).is_none());
1400    }
1401
1402    #[test]
1403    fn test_no_subquery_detected() {
1404        let rule = NoSubqueryRule;
1405        let ctx = RuleContext::from_sql("SELECT * FROM (SELECT id FROM users) AS sub");
1406        assert!(rule.check(&ctx).is_some());
1407    }
1408}