1use crate::{detect_statement_type, tokenize, SqlStatementType, SqlToken};
26use regex::Regex;
27use std::fmt;
28
29#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
38pub enum RuleSeverity {
39 Info,
41 Warning,
43 Error,
45 Critical,
47}
48
49impl RuleSeverity {
50 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 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 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#[derive(Debug, Clone, PartialEq, Eq)]
90pub struct RuleViolation {
91 pub rule_name: String,
93 pub severity: RuleSeverity,
95 pub message: String,
97 pub position: Option<usize>,
99}
100
101impl RuleViolation {
102 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 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 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#[derive(Debug, Clone)]
162pub struct RuleContext {
163 pub sql: String,
165 pub tokens: Vec<SqlToken>,
167 pub statement_type: SqlStatementType,
169 pub sql_upper: String,
171}
172
173impl RuleContext {
174 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 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 pub fn has_keyword(&self, keyword: &str) -> bool {
195 self.keyword_count(keyword) > 0
196 }
197
198 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 let _ = name;
212 count += 1;
213 expect_table = false;
214 }
215 SqlToken::Punctuation('.') if expect_table => {
216 }
218 _ if expect_table => {
219 expect_table = false;
221 }
222 _ => {}
223 }
224 }
225 count
226 }
227}
228
229pub trait SqlRule: Send + Sync {
238 fn name(&self) -> &str;
240
241 fn severity(&self) -> RuleSeverity;
243
244 fn check(&self, ctx: &RuleContext) -> Option<RuleViolation>;
246}
247
248#[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 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 after_select = false;
290 }
291 _ => {}
292 }
293 }
294 None
295 }
296}
297
298#[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#[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#[derive(Debug)]
365pub struct MaxTableCountRule {
366 pub max: usize,
368}
369
370impl MaxTableCountRule {
371 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#[derive(Debug)]
402pub struct MaxJoinCountRule {
403 pub max: usize,
405}
406
407impl MaxJoinCountRule {
408 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#[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#[derive(Debug)]
471pub struct ForbiddenKeywordRule {
472 pub rule_name: String,
474 pub keywords: Vec<String>,
476 pub severity: RuleSeverity,
478}
479
480impl ForbiddenKeywordRule {
481 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 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#[derive(Debug)]
527pub struct RegexRule {
528 pub rule_name: String,
530 pub pattern: Regex,
532 pub severity: RuleSeverity,
534 pub message: String,
536}
537
538impl RegexRule {
539 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#[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 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
609pub struct MaxColumnCountRule {
611 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
662pub 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#[derive(Debug, Clone, Default)]
700pub struct RuleReport {
701 pub violations: Vec<RuleViolation>,
703}
704
705impl RuleReport {
706 pub fn is_clean(&self) -> bool {
708 self.violations.is_empty()
709 }
710
711 pub fn has_violations(&self) -> bool {
713 !self.violations.is_empty()
714 }
715
716 pub fn has_blocking(&self) -> bool {
718 self.violations.iter().any(|v| v.is_blocking())
719 }
720
721 pub fn has_errors(&self) -> bool {
723 self.violations
724 .iter()
725 .any(|v| v.severity == RuleSeverity::Error)
726 }
727
728 pub fn has_critical(&self) -> bool {
730 self.violations
731 .iter()
732 .any(|v| v.severity == RuleSeverity::Critical)
733 }
734
735 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 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 pub fn blocking_violations(&self) -> Vec<&RuleViolation> {
753 self.violations.iter().filter(|v| v.is_blocking()).collect()
754 }
755
756 pub fn violation_count(&self) -> usize {
758 self.violations.len()
759 }
760
761 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
788pub struct RuleEngine {
797 rules: Vec<Box<dyn SqlRule>>,
798}
799
800impl RuleEngine {
801 pub fn new() -> Self {
803 Self { rules: Vec::new() }
804 }
805
806 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 pub fn add_rule(&mut self, rule: Box<dyn SqlRule>) -> &mut Self {
829 self.rules.push(rule);
830 self
831 }
832
833 pub fn rule_count(&self) -> usize {
835 self.rules.len()
836 }
837
838 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 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 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
866pub struct RulePresets;
872
873impl RulePresets {
874 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 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#[cfg(test)]
906mod tests {
907 use super::*;
908
909 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 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 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 #[test]
1343 fn test_rule_presets_strict() {
1344 let engine = RulePresets::strict();
1345 assert!(engine.rule_count() >= 8);
1346 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 let report = engine.check("DELETE FROM users");
1356 assert!(report.is_clean());
1357 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 #[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 #[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}