1pub mod firewall;
7
8use thiserror::Error;
9
10#[derive(Error, Debug, Clone, PartialEq)]
12pub enum SqlValidationError {
13 #[error("SQL syntax error: {0}")]
14 SyntaxError(String),
15
16 #[error("Unbalanced parentheses: {0}")]
17 UnbalancedParentheses(String),
18
19 #[error("Unclosed string literal at position {0}")]
20 UnclosedString(usize),
21
22 #[error("Missing required keyword: {0}")]
23 MissingKeyword(String),
24
25 #[error("Invalid parameter count: expected {expected}, got {got}")]
26 ParameterCountMismatch { expected: usize, got: usize },
27
28 #[error("Invalid table name: {0}")]
29 InvalidTableName(String),
30
31 #[error("Empty SELECT columns")]
32 EmptySelectColumns,
33
34 #[error("Empty INSERT data")]
35 EmptyInsertData,
36
37 #[error("Empty UPDATE data")]
38 EmptyUpdateData,
39
40 #[error("DELETE without WHERE clause")]
41 DeleteWithoutWhere,
42
43 #[error("Invalid identifier: {0}")]
44 InvalidIdentifier(String),
45
46 #[error("SQL injection detected: {0}")]
47 InjectionDetected(String),
48}
49
50pub type ValidationResult = Result<(), SqlValidationError>;
52
53#[derive(Debug, Clone, Copy, PartialEq)]
55pub enum SqlStatementType {
56 Select,
57 Insert,
58 Update,
59 Delete,
60 Create,
61 Drop,
62 Alter,
63 Truncate,
64 Other,
65}
66
67pub fn validate_select(sql: &str) -> ValidationResult {
69 let sql_upper = sql.to_uppercase();
70
71 if !sql_upper.trim_start().starts_with("SELECT") {
72 return Err(SqlValidationError::SyntaxError(
73 "SELECT statement must start with SELECT".to_string(),
74 ));
75 }
76
77 if !sql_upper.contains("FROM") {
78 return Err(SqlValidationError::MissingKeyword("FROM".to_string()));
79 }
80
81 validate_balanced_parentheses(sql)?;
82 validate_string_literals(sql)?;
83 validate_no_injection_patterns(sql)?;
84
85 Ok(())
86}
87
88pub fn validate_insert(sql: &str) -> ValidationResult {
90 let sql_upper = sql.to_uppercase();
91
92 if !sql_upper.trim_start().starts_with("INSERT") {
93 return Err(SqlValidationError::SyntaxError(
94 "INSERT statement must start with INSERT".to_string(),
95 ));
96 }
97
98 if !sql_upper.contains("INTO") {
99 return Err(SqlValidationError::MissingKeyword("INTO".to_string()));
100 }
101
102 if !sql_upper.contains("VALUES") {
103 return Err(SqlValidationError::MissingKeyword("VALUES".to_string()));
104 }
105
106 validate_balanced_parentheses(sql)?;
107 validate_string_literals(sql)?;
108 validate_no_injection_patterns(sql)?;
109
110 Ok(())
111}
112
113pub fn validate_update(sql: &str) -> ValidationResult {
115 let sql_upper = sql.to_uppercase();
116
117 if !sql_upper.trim_start().starts_with("UPDATE") {
118 return Err(SqlValidationError::SyntaxError(
119 "UPDATE statement must start with UPDATE".to_string(),
120 ));
121 }
122
123 if !sql_upper.contains("SET") {
124 return Err(SqlValidationError::MissingKeyword("SET".to_string()));
125 }
126
127 validate_balanced_parentheses(sql)?;
128 validate_string_literals(sql)?;
129 validate_no_injection_patterns(sql)?;
130
131 Ok(())
132}
133
134pub fn validate_delete(sql: &str) -> ValidationResult {
136 let sql_upper = sql.to_uppercase();
137
138 if !sql_upper.trim_start().starts_with("DELETE") {
139 return Err(SqlValidationError::SyntaxError(
140 "DELETE statement must start with DELETE".to_string(),
141 ));
142 }
143
144 if !sql_upper.contains("FROM") {
145 return Err(SqlValidationError::MissingKeyword("FROM".to_string()));
146 }
147
148 validate_balanced_parentheses(sql)?;
149 validate_string_literals(sql)?;
150 validate_no_injection_patterns(sql)?;
151
152 Ok(())
153}
154
155pub fn validate_sql(sql: &str) -> ValidationResult {
157 let trimmed = sql.trim();
158 if trimmed.is_empty() {
159 return Err(SqlValidationError::SyntaxError(
160 "Empty SQL statement".to_string(),
161 ));
162 }
163
164 let sql_type = detect_statement_type(trimmed);
165 match sql_type {
166 SqlStatementType::Select => validate_select(trimmed),
167 SqlStatementType::Insert => validate_insert(trimmed),
168 SqlStatementType::Update => validate_update(trimmed),
169 SqlStatementType::Delete => validate_delete(trimmed),
170 _ => {
171 validate_balanced_parentheses(trimmed)?;
172 validate_string_literals(trimmed)?;
173 validate_no_injection_patterns(trimmed)?;
174 Ok(())
175 }
176 }
177}
178
179fn validate_balanced_parentheses(sql: &str) -> ValidationResult {
181 let mut depth: i32 = 0;
182 for (i, ch) in sql.char_indices() {
183 match ch {
184 '(' => depth += 1,
185 ')' => {
186 depth -= 1;
187 if depth < 0 {
188 return Err(SqlValidationError::UnbalancedParentheses(format!(
189 "Unexpected ')' at position {}",
190 i
191 )));
192 }
193 }
194 _ => {}
195 }
196 }
197 if depth != 0 {
198 return Err(SqlValidationError::UnbalancedParentheses(format!(
199 "{} unclosed '(' parentheses",
200 depth
201 )));
202 }
203 Ok(())
204}
205
206fn validate_string_literals(sql: &str) -> ValidationResult {
208 let mut in_single_quote = false;
209 let mut in_double_quote = false;
210 let mut prev_ch = '\0';
211
212 for (_i, ch) in sql.char_indices() {
213 if prev_ch == '\\' {
214 prev_ch = ch;
215 continue;
216 }
217
218 match ch {
219 '\'' if !in_double_quote => {
220 in_single_quote = !in_single_quote;
221 }
222 '"' if !in_single_quote => {
223 in_double_quote = !in_double_quote;
224 }
225 _ => {}
226 }
227 prev_ch = ch;
228 }
229
230 if in_single_quote {
231 return Err(SqlValidationError::UnclosedString(sql.len()));
232 }
233 if in_double_quote {
234 return Err(SqlValidationError::UnclosedString(sql.len()));
235 }
236
237 Ok(())
238}
239
240fn validate_no_injection_patterns(sql: &str) -> ValidationResult {
242 let sql_upper = sql.to_uppercase();
243
244 let suspicious_patterns = [
246 ("'; DROP TABLE", "DROP TABLE injection"),
247 ("' OR '1'='1", "classic OR injection"),
248 ("' OR 1=1", "OR 1=1 injection"),
249 (
250 "UNION SELECT",
251 "UNION SELECT injection (not allowed in simple queries)",
252 ),
253 ("--", "comment injection (not allowed)"),
254 ("/*", "block comment (not allowed)"),
255 ];
256
257 for (pattern, desc) in &suspicious_patterns {
258 if sql_upper.contains(pattern) {
259 return Err(SqlValidationError::InjectionDetected(format!(
260 "{}: {}",
261 desc, pattern
262 )));
263 }
264 }
265
266 Ok(())
267}
268
269pub fn detect_statement_type(sql: &str) -> SqlStatementType {
271 let trimmed = sql.trim().to_uppercase();
272
273 if trimmed.starts_with("SELECT") {
274 SqlStatementType::Select
275 } else if trimmed.starts_with("INSERT") {
276 SqlStatementType::Insert
277 } else if trimmed.starts_with("UPDATE") {
278 SqlStatementType::Update
279 } else if trimmed.starts_with("DELETE") {
280 SqlStatementType::Delete
281 } else if trimmed.starts_with("CREATE") {
282 SqlStatementType::Create
283 } else if trimmed.starts_with("DROP") {
284 SqlStatementType::Drop
285 } else if trimmed.starts_with("ALTER") {
286 SqlStatementType::Alter
287 } else if trimmed.starts_with("TRUNCATE") {
288 SqlStatementType::Truncate
289 } else {
290 SqlStatementType::Other
291 }
292}
293
294pub fn validate_parameter_count(sql: &str, expected_params: usize) -> ValidationResult {
296 let param_count = sql.chars().filter(|&c| c == '?').count() + sql.matches('$').count(); if param_count != expected_params {
298 return Err(SqlValidationError::ParameterCountMismatch {
299 expected: expected_params,
300 got: param_count,
301 });
302 }
303 Ok(())
304}
305
306pub fn validate_table_name(name: &str) -> ValidationResult {
308 if name.is_empty() {
309 return Err(SqlValidationError::InvalidTableName(
310 "empty table name".to_string(),
311 ));
312 }
313
314 let cleaned = name.trim_matches('`').trim_matches('"').trim_matches('\'');
317 if cleaned.is_empty() {
318 return Err(SqlValidationError::InvalidTableName(name.to_string()));
319 }
320
321 for ch in cleaned.chars() {
322 if !ch.is_alphanumeric() && ch != '_' {
323 return Err(SqlValidationError::InvalidTableName(format!(
324 "table name '{}' contains invalid character '{}'",
325 name, ch
326 )));
327 }
328 }
329
330 Ok(())
331}
332
333pub fn validate_column_name(name: &str) -> ValidationResult {
335 if name.is_empty() || name == "*" {
336 return Ok(()); }
338
339 let cleaned = name.trim_matches('`').trim_matches('"').trim_matches('\'');
340 if cleaned.is_empty() {
341 return Err(SqlValidationError::InvalidIdentifier(name.to_string()));
342 }
343
344 for ch in cleaned.chars() {
346 if !ch.is_alphanumeric() && ch != '_' && ch != '.' {
347 return Err(SqlValidationError::InvalidIdentifier(format!(
348 "column '{}' contains invalid character '{}'",
349 name, ch
350 )));
351 }
352 }
353
354 Ok(())
355}
356
357pub fn validate(sql: &str) -> ValidationResult {
359 if sql.trim().is_empty() {
360 return Err(SqlValidationError::SyntaxError(
361 "Empty SQL statement".to_string(),
362 ));
363 }
364
365 validate_sql(sql)?;
366 validate_balanced_parentheses(sql)?;
367 validate_string_literals(sql)?;
368 validate_no_injection_patterns(sql)?;
369
370 Ok(())
371}
372
373#[derive(Debug, Clone, PartialEq)]
379pub enum SqlToken {
380 Keyword(String),
382 Identifier(String),
384 StringLiteral(String),
386 NumberLiteral(String),
388 Operator(String),
390 Punctuation(char),
392 Comment(String),
394}
395
396pub fn tokenize(sql: &str) -> Vec<SqlToken> {
401 let mut tokens = Vec::new();
402 let chars: Vec<char> = sql.chars().collect();
403 let mut i = 0;
404
405 while i < chars.len() {
406 let ch = chars[i];
407
408 if ch.is_whitespace() {
410 i += 1;
411 continue;
412 }
413
414 if i + 1 < chars.len() && ch == '-' && chars[i + 1] == '-' {
416 let start = i;
417 while i < chars.len() && chars[i] != '\n' {
418 i += 1;
419 }
420 tokens.push(SqlToken::Comment(chars[start..i].iter().collect()));
421 continue;
422 }
423
424 if i + 1 < chars.len() && ch == '/' && chars[i + 1] == '*' {
426 let start = i;
427 i += 2;
428 while i + 1 < chars.len() {
429 if chars[i] == '*' && chars[i + 1] == '/' {
430 i += 2;
431 break;
432 }
433 i += 1;
434 }
435 tokens.push(SqlToken::Comment(chars[start..i].iter().collect()));
436 continue;
437 }
438
439 if ch == '\'' {
441 let start = i;
442 i += 1;
443 while i < chars.len() {
444 if chars[i] == '\'' {
445 i += 1;
446 break;
447 }
448 i += 1;
449 }
450 tokens.push(SqlToken::StringLiteral(chars[start..i].iter().collect()));
451 continue;
452 }
453
454 if ch == '"' {
456 let start = i;
457 i += 1;
458 while i < chars.len() && chars[i] != '"' {
459 i += 1;
460 }
461 if i < chars.len() {
462 i += 1;
463 }
464 tokens.push(SqlToken::Identifier(chars[start..i].iter().collect()));
465 continue;
466 }
467
468 if ch.is_ascii_digit() {
470 let start = i;
471 while i < chars.len() && (chars[i].is_ascii_digit() || chars[i] == '.') {
472 i += 1;
473 }
474 tokens.push(SqlToken::NumberLiteral(chars[start..i].iter().collect()));
475 continue;
476 }
477
478 if ch.is_alphabetic() || ch == '_' {
480 let start = i;
481 while i < chars.len() && (chars[i].is_alphanumeric() || chars[i] == '_') {
482 i += 1;
483 }
484 let word: String = chars[start..i].iter().collect();
485 let upper = word.to_uppercase();
486 const KEYWORDS: &[&str] = &[
487 "SELECT",
488 "FROM",
489 "WHERE",
490 "INSERT",
491 "INTO",
492 "VALUES",
493 "UPDATE",
494 "SET",
495 "DELETE",
496 "CREATE",
497 "TABLE",
498 "DROP",
499 "ALTER",
500 "TRUNCATE",
501 "JOIN",
502 "INNER",
503 "LEFT",
504 "RIGHT",
505 "OUTER",
506 "ON",
507 "AND",
508 "OR",
509 "NOT",
510 "NULL",
511 "IS",
512 "IN",
513 "LIKE",
514 "BETWEEN",
515 "ORDER",
516 "BY",
517 "GROUP",
518 "HAVING",
519 "LIMIT",
520 "OFFSET",
521 "DISTINCT",
522 "AS",
523 "UNION",
524 "ALL",
525 "INTERSECT",
526 "EXCEPT",
527 "CASE",
528 "WHEN",
529 "THEN",
530 "ELSE",
531 "END",
532 "IF",
533 "EXISTS",
534 "PRIMARY",
535 "KEY",
536 "FOREIGN",
537 "REFERENCES",
538 "INDEX",
539 "VIEW",
540 "DATABASE",
541 "SCHEMA",
542 "GRANT",
543 "REVOKE",
544 "EXEC",
545 "EXECUTE",
546 "PROCEDURE",
547 "FUNCTION",
548 "BEGIN",
549 "COMMIT",
550 "ROLLBACK",
551 "SAVEPOINT",
552 "RELEASE",
553 "TRANSACTION",
554 "START",
555 "WITH",
556 "RECURSIVE",
557 ];
558 if KEYWORDS.contains(&upper.as_str()) {
559 tokens.push(SqlToken::Keyword(upper));
560 } else {
561 tokens.push(SqlToken::Identifier(word));
562 }
563 continue;
564 }
565
566 if "+-*/=<>!".contains(ch) {
568 let start = i;
569 i += 1;
570 while i < chars.len() && "+-*/=<>!".contains(chars[i]) {
571 i += 1;
572 }
573 tokens.push(SqlToken::Operator(chars[start..i].iter().collect()));
574 continue;
575 }
576
577 if "().,;".contains(ch) {
579 tokens.push(SqlToken::Punctuation(ch));
580 i += 1;
581 continue;
582 }
583
584 i += 1;
586 }
587
588 tokens
589}
590
591pub fn detect_injection_ast(sql: &str) -> ValidationResult {
599 let tokens = tokenize(sql);
600 let keywords: Vec<&str> = tokens
601 .iter()
602 .filter_map(|t| match t {
603 SqlToken::Keyword(k) => Some(k.as_str()),
604 _ => None,
605 })
606 .collect();
607
608 if keywords.is_empty() {
609 return Ok(());
610 }
611
612 let mut after_semicolon = false;
614 for token in &tokens {
615 match token {
616 SqlToken::Punctuation(';') => {
617 after_semicolon = true;
618 }
619 SqlToken::Keyword(k) if after_semicolon => {
620 match k.as_str() {
622 "DROP" | "ALTER" | "TRUNCATE" | "DELETE" | "INSERT" | "UPDATE" | "CREATE"
623 | "GRANT" | "REVOKE" | "EXEC" | "EXECUTE" => {
624 return Err(SqlValidationError::InjectionDetected(format!(
625 "multi-statement injection: semicolon followed by {} keyword",
626 k
627 )));
628 }
629 _ => {
630 after_semicolon = false;
631 }
632 }
633 }
634 SqlToken::Keyword(_) => {
635 after_semicolon = false;
636 }
637 _ => {}
638 }
639 }
640
641 if keywords.iter().any(|k| *k == "EXEC" || *k == "EXECUTE") {
643 return Err(SqlValidationError::InjectionDetected(
644 "EXEC/EXECUTE call detected (potential injection)".to_string(),
645 ));
646 }
647
648 if keywords.iter().any(|k| *k == "GRANT" || *k == "REVOKE") {
650 return Err(SqlValidationError::InjectionDetected(
651 "GRANT/REVOKE statement detected (potential privilege escalation)".to_string(),
652 ));
653 }
654
655 let upper_sql = sql.to_uppercase();
657 if upper_sql.contains(" OR 1=1")
658 || upper_sql.contains(" OR 1 = 1")
659 || upper_sql.contains(" OR '1'='1'")
660 || upper_sql.contains(" OR TRUE")
661 || upper_sql.contains(" OR 1<>0")
662 {
663 return Err(SqlValidationError::InjectionDetected(
664 "boolean blind injection pattern detected (OR with tautology)".to_string(),
665 ));
666 }
667
668 Ok(())
669}
670
671#[derive(Debug, Clone, Default)]
673pub struct WhitelistValidator {
674 allowed_tables: std::collections::HashSet<String>,
676 allowed_columns: Option<std::collections::HashSet<String>>,
678}
679
680impl WhitelistValidator {
681 pub fn new() -> Self {
683 Self {
684 allowed_tables: std::collections::HashSet::new(),
685 allowed_columns: None,
686 }
687 }
688
689 pub fn allow_table(mut self, table: &str) -> Self {
691 self.allowed_tables.insert(table.to_lowercase());
692 self
693 }
694
695 pub fn allow_tables(mut self, tables: &[&str]) -> Self {
697 for t in tables {
698 self.allowed_tables.insert(t.to_lowercase());
699 }
700 self
701 }
702
703 pub fn allow_columns(mut self, columns: &[&str]) -> Self {
705 let set: std::collections::HashSet<String> =
706 columns.iter().map(|c| c.to_lowercase()).collect();
707 self.allowed_columns = Some(set);
708 self
709 }
710
711 pub fn validate_tables(&self, sql: &str) -> ValidationResult {
716 if self.allowed_tables.is_empty() {
717 return Ok(()); }
719
720 let tokens = tokenize(sql);
721 let mut check_next_identifier = false;
722
723 for token in &tokens {
724 match token {
725 SqlToken::Keyword(k)
726 if matches!(k.as_str(), "FROM" | "JOIN" | "INTO" | "UPDATE" | "TABLE") =>
727 {
728 check_next_identifier = true;
729 }
730 SqlToken::Identifier(name) if check_next_identifier => {
731 let lower = name.to_lowercase();
732 if !self.allowed_tables.contains(&lower) {
733 return Err(SqlValidationError::InvalidTableName(format!(
734 "table '{}' is not in the whitelist",
735 name
736 )));
737 }
738 check_next_identifier = false;
739 }
740 _ if check_next_identifier => {
741 #[allow(clippy::collapsible_match)]
743 if !matches!(token, SqlToken::Punctuation('.')) {
744 check_next_identifier = false;
745 }
746 }
747 _ => {}
748 }
749 }
750
751 Ok(())
752 }
753
754 pub fn validate_columns(&self, sql: &str) -> ValidationResult {
759 let allowed = match &self.allowed_columns {
760 Some(c) => c,
761 None => return Ok(()), };
763
764 let tokens = tokenize(sql);
765 for token in &tokens {
766 if let SqlToken::Identifier(name) = token {
767 let lower = name.to_lowercase();
768 if self.allowed_tables.contains(&lower) {
770 continue;
771 }
772 if lower == "*" {
774 continue;
775 }
776 if !allowed.contains(&lower) && !name.contains('.') {
778 }
781 }
782 }
783
784 Ok(())
785 }
786
787 pub fn validate(&self, sql: &str) -> ValidationResult {
789 self.validate_tables(sql)?;
790 self.validate_columns(sql)
791 }
792}
793
794#[derive(Debug, Clone)]
796pub struct SqlComplexityScore {
797 pub score: u32,
799 pub join_count: u32,
801 pub subquery_count: u32,
803 pub where_condition_count: u32,
805 pub set_operation_count: u32,
807 pub group_by_count: u32,
809 pub has_having: bool,
811 pub has_window_function: bool,
813 pub has_cte: bool,
815 pub token_count: u32,
817}
818
819impl SqlComplexityScore {
820 fn calculate(&mut self) {
822 let mut score: u32 = 0;
823 score += self.join_count * 5;
824 score += self.subquery_count * 10;
825 score += self.where_condition_count * 3;
826 score += self.set_operation_count * 8;
827 score += self.group_by_count * 3;
828 if self.has_having {
829 score += 5;
830 }
831 if self.has_window_function {
832 score += 8;
833 }
834 if self.has_cte {
835 score += 6;
836 }
837 score += (self.token_count / 50).min(20);
839 self.score = score.min(100);
840 }
841
842 pub fn level(&self) -> ComplexityLevel {
844 match self.score {
845 0..=20 => ComplexityLevel::Simple,
846 21..=40 => ComplexityLevel::Moderate,
847 41..=60 => ComplexityLevel::Complex,
848 _ => ComplexityLevel::VeryComplex,
849 }
850 }
851}
852
853#[derive(Debug, Clone, Copy, PartialEq, Eq)]
855pub enum ComplexityLevel {
856 Simple,
858 Moderate,
860 Complex,
862 VeryComplex,
864}
865
866impl ComplexityLevel {
867 pub fn description(&self) -> &'static str {
869 match self {
870 ComplexityLevel::Simple => "简单",
871 ComplexityLevel::Moderate => "中等",
872 ComplexityLevel::Complex => "复杂",
873 ComplexityLevel::VeryComplex => "非常复杂",
874 }
875 }
876}
877
878pub fn score_complexity(sql: &str) -> SqlComplexityScore {
883 let tokens = tokenize(sql);
884 let token_count = tokens.len() as u32;
885
886 let mut join_count = 0u32;
887 let mut subquery_count = 0u32;
888 let mut where_condition_count = 0u32;
889 let mut set_operation_count = 0u32;
890 let mut group_by_count = 0u32;
891 let mut has_having = false;
892 let mut has_window_function = false;
893 let mut has_cte = false;
894
895 let mut in_where = false;
896 let mut in_group_by = false;
897 let mut paren_depth: i32 = 0;
898
899 for token in &tokens {
900 match token {
901 SqlToken::Keyword(k) => {
902 match k.as_str() {
903 "JOIN" | "INNER" | "LEFT" | "RIGHT" | "OUTER" => {
904 if k == "JOIN" {
905 join_count += 1;
906 }
907 }
908 "WHERE" => {
909 in_where = true;
910 where_condition_count += 1;
911 }
912 "AND" | "OR" if in_where => {
913 where_condition_count += 1;
914 }
915 "GROUP" => {
916 in_group_by = true;
917 }
918 "HAVING" => {
919 has_having = true;
920 in_where = false;
921 in_group_by = false;
922 }
923 "UNION" | "INTERSECT" | "EXCEPT" => {
924 set_operation_count += 1;
925 }
926 "WITH" => {
927 has_cte = true;
928 }
929 "SELECT" if paren_depth > 0 => {
930 subquery_count += 1;
931 }
932 _ => {}
933 }
934 if k != "WHERE" && k != "AND" && k != "OR" {
935 in_where = false;
936 }
937 if k != "GROUP" && k != "BY" && in_group_by {
938 in_group_by = false;
939 }
940 }
941 SqlToken::Identifier(name) => {
942 let upper = name.to_uppercase();
943 if upper.contains("OVER") || upper.contains("ROW_NUMBER") || upper.contains("RANK")
944 {
945 has_window_function = true;
946 }
947 if in_group_by {
948 group_by_count += 1;
949 }
950 }
951 SqlToken::Punctuation('(') => {
952 paren_depth += 1;
953 }
954 SqlToken::Punctuation(')') => {
955 paren_depth -= 1;
956 }
957 _ => {}
958 }
959 }
960
961 let mut score = SqlComplexityScore {
962 score: 0,
963 join_count,
964 subquery_count,
965 where_condition_count,
966 set_operation_count,
967 group_by_count,
968 has_having,
969 has_window_function,
970 has_cte,
971 token_count,
972 };
973 score.calculate();
974 score
975}
976
977#[derive(Debug, Clone, Default)]
979pub struct DdlPolicy {
980 pub allow_create: bool,
982 pub allow_drop: bool,
984 pub allow_alter: bool,
986 pub allow_truncate: bool,
988 pub allow_create_index: bool,
990 pub allow_drop_index: bool,
992}
993
994impl DdlPolicy {
995 pub fn permissive() -> Self {
997 Self {
998 allow_create: true,
999 allow_drop: true,
1000 allow_alter: true,
1001 allow_truncate: true,
1002 allow_create_index: true,
1003 allow_drop_index: true,
1004 }
1005 }
1006
1007 pub fn read_only() -> Self {
1009 Self::default()
1010 }
1011
1012 pub fn safe_evolution() -> Self {
1014 Self {
1015 allow_create: true,
1016 allow_drop: false,
1017 allow_alter: true,
1018 allow_truncate: false,
1019 allow_create_index: true,
1020 allow_drop_index: false,
1021 }
1022 }
1023
1024 pub fn validate(&self, sql: &str) -> ValidationResult {
1026 let stmt_type = detect_statement_type(sql);
1027 let upper = sql.to_uppercase();
1028
1029 match stmt_type {
1030 SqlStatementType::Create => {
1031 if !self.allow_create {
1032 return Err(SqlValidationError::SyntaxError(
1033 "CREATE operations are not allowed by DDL policy".to_string(),
1034 ));
1035 }
1036 if upper.contains("INDEX") && !self.allow_create_index {
1037 return Err(SqlValidationError::SyntaxError(
1038 "CREATE INDEX operations are not allowed by DDL policy".to_string(),
1039 ));
1040 }
1041 Ok(())
1042 }
1043 SqlStatementType::Drop => {
1044 if !self.allow_drop {
1045 return Err(SqlValidationError::SyntaxError(
1046 "DROP operations are not allowed by DDL policy".to_string(),
1047 ));
1048 }
1049 if upper.contains("INDEX") && !self.allow_drop_index {
1050 return Err(SqlValidationError::SyntaxError(
1051 "DROP INDEX operations are not allowed by DDL policy".to_string(),
1052 ));
1053 }
1054 Ok(())
1055 }
1056 SqlStatementType::Alter => {
1057 if !self.allow_alter {
1058 return Err(SqlValidationError::SyntaxError(
1059 "ALTER operations are not allowed by DDL policy".to_string(),
1060 ));
1061 }
1062 Ok(())
1063 }
1064 SqlStatementType::Truncate => {
1065 if !self.allow_truncate {
1066 return Err(SqlValidationError::SyntaxError(
1067 "TRUNCATE operations are not allowed by DDL policy".to_string(),
1068 ));
1069 }
1070 Ok(())
1071 }
1072 _ => Ok(()), }
1074 }
1075}
1076
1077#[cfg(test)]
1078mod tests {
1079 use super::*;
1080
1081 #[test]
1082 fn test_validate_select_basic() {
1083 assert!(validate_select("SELECT * FROM users").is_ok());
1084 assert!(validate_select("SELECT id, name FROM users WHERE id = 1").is_ok());
1085 assert!(validate_select(
1086 "SELECT u.id, u.name FROM users u INNER JOIN orders o ON u.id = o.user_id"
1087 )
1088 .is_ok());
1089 }
1090
1091 #[test]
1092 fn test_validate_select_missing_from() {
1093 let result = validate_select("SELECT *");
1094 assert!(result.is_err());
1095 }
1096
1097 #[test]
1098 fn test_validate_insert_basic() {
1099 assert!(validate_insert("INSERT INTO users (name) VALUES ('alice')").is_ok());
1100 assert!(validate_insert("INSERT INTO users (name, age) VALUES ('bob', 25)").is_ok());
1101 }
1102
1103 #[test]
1104 fn test_validate_insert_missing_values() {
1105 let result = validate_insert("INSERT INTO users (name)");
1106 assert!(result.is_err());
1107 }
1108
1109 #[test]
1110 fn test_validate_update_basic() {
1111 assert!(validate_update("UPDATE users SET name = 'alice' WHERE id = 1").is_ok());
1112 }
1113
1114 #[test]
1115 fn test_validate_update_missing_set() {
1116 let result = validate_update("UPDATE users WHERE id = 1");
1117 assert!(result.is_err());
1118 }
1119
1120 #[test]
1121 fn test_validate_delete_basic() {
1122 assert!(validate_delete("DELETE FROM users WHERE id = 1").is_ok());
1123 }
1124
1125 #[test]
1126 fn test_validate_delete_missing_from() {
1127 let result = validate_delete("DELETE users");
1128 assert!(result.is_err());
1129 }
1130
1131 #[test]
1132 fn test_balanced_parentheses() {
1133 assert!(validate_balanced_parentheses("SELECT * FROM (SELECT * FROM users) t").is_ok());
1134 assert!(validate_balanced_parentheses("FUNC(a, b, c)").is_ok());
1135 assert!(
1136 validate_balanced_parentheses("SELECT * FROM users WHERE (a=1 AND (b=2 OR c=3))")
1137 .is_ok()
1138 );
1139 }
1140
1141 #[test]
1142 fn test_unbalanced_parentheses() {
1143 assert!(validate_balanced_parentheses("SELECT * FROM (SELECT * FROM users").is_err());
1144 assert!(validate_balanced_parentheses("SELECT * FROM users)").is_err());
1145 }
1146
1147 #[test]
1148 fn test_string_literals_closed() {
1149 assert!(validate_string_literals("SELECT * FROM users WHERE name = 'alice'").is_ok());
1150 assert!(validate_string_literals("INSERT INTO users (name) VALUES ('bob')").is_ok());
1151 }
1152
1153 #[test]
1154 fn test_unclosed_string_literal() {
1155 assert!(validate_string_literals("SELECT * FROM users WHERE name = 'alice").is_err());
1156 }
1157
1158 #[test]
1159 fn test_injection_detection() {
1160 assert!(validate_no_injection_patterns("SELECT * FROM users WHERE name = 'alice'").is_ok());
1161 assert!(validate_no_injection_patterns(
1162 "SELECT * FROM users WHERE name = 'alice' OR '1'='1'"
1163 )
1164 .is_err());
1165 assert!(validate_no_injection_patterns("'; DROP TABLE users; --").is_err());
1166 assert!(validate_no_injection_patterns("1 UNION SELECT * FROM users").is_err());
1167 }
1168
1169 #[test]
1170 fn test_detect_statement_type() {
1171 assert_eq!(
1172 detect_statement_type("SELECT * FROM users"),
1173 SqlStatementType::Select
1174 );
1175 assert_eq!(
1176 detect_statement_type("INSERT INTO users VALUES (1)"),
1177 SqlStatementType::Insert
1178 );
1179 assert_eq!(
1180 detect_statement_type("UPDATE users SET a=1"),
1181 SqlStatementType::Update
1182 );
1183 assert_eq!(
1184 detect_statement_type("DELETE FROM users"),
1185 SqlStatementType::Delete
1186 );
1187 assert_eq!(
1188 detect_statement_type("CREATE TABLE users"),
1189 SqlStatementType::Create
1190 );
1191 assert_eq!(
1192 detect_statement_type("DROP TABLE users"),
1193 SqlStatementType::Drop
1194 );
1195 assert_eq!(
1196 detect_statement_type("ALTER TABLE users ADD COLUMN a"),
1197 SqlStatementType::Alter
1198 );
1199 assert_eq!(
1200 detect_statement_type("TRUNCATE TABLE users"),
1201 SqlStatementType::Truncate
1202 );
1203 assert_eq!(
1204 detect_statement_type("EXPLAIN SELECT * FROM users"),
1205 SqlStatementType::Other
1206 );
1207 }
1208
1209 #[test]
1210 fn test_parameter_count() {
1211 assert!(
1212 validate_parameter_count("SELECT * FROM users WHERE id = ? AND name = ?", 2).is_ok()
1213 );
1214 assert!(validate_parameter_count("SELECT * FROM users WHERE id = ?", 1).is_ok());
1215 assert!(validate_parameter_count("SELECT * FROM users WHERE id = ?", 2).is_err());
1216 }
1217
1218 #[test]
1219 fn test_validate_table_name() {
1220 assert!(validate_table_name("users").is_ok());
1221 assert!(validate_table_name("user_orders").is_ok());
1222 assert!(validate_table_name("").is_err());
1223 assert!(validate_table_name("users; DROP TABLE").is_err());
1224 }
1225
1226 #[test]
1227 fn test_validate_column_name() {
1228 assert!(validate_column_name("id").is_ok());
1229 assert!(validate_column_name("*").is_ok());
1230 assert!(validate_column_name("users.name").is_ok());
1231 assert!(validate_column_name("").is_ok()); }
1233
1234 #[test]
1235 fn test_validate_empty_sql() {
1236 assert!(validate("").is_err());
1237 assert!(validate(" ").is_err());
1238 }
1239
1240 #[test]
1241 fn test_validate_complex_queries() {
1242 assert!(validate("SELECT u.*, o.total FROM users u LEFT JOIN orders o ON u.id = o.user_id WHERE u.status = 'active' AND u.created_at > '2024-01-01' GROUP BY u.id HAVING COUNT(o.id) > 5 ORDER BY u.name ASC LIMIT 10 OFFSET 20").is_ok());
1243 }
1244
1245 #[test]
1246 fn test_empty_insert_data() {
1247 let sql = "INSERT INTO users () VALUES ()";
1248 assert!(validate_sql(sql).is_ok());
1249 }
1250
1251 #[test]
1252 fn test_create_table_validation() {
1253 assert!(validate_sql("CREATE TABLE users (id INT PRIMARY KEY, name VARCHAR(100))").is_ok());
1254 }
1255
1256 #[test]
1257 fn test_double_quoted_identifiers() {
1258 assert!(
1259 validate_string_literals("SELECT * FROM \"users\" WHERE \"name\" = 'alice'").is_ok()
1260 );
1261 }
1262
1263 #[test]
1264 fn test_nested_function_calls() {
1265 assert!(validate_balanced_parentheses(
1266 "SELECT MAX(COUNT(*)) FROM (SELECT COUNT(*) FROM users GROUP BY status) t"
1267 )
1268 .is_ok());
1269 }
1270
1271 #[test]
1278 fn test_tokenize_select_basic() {
1279 let tokens = tokenize("SELECT id FROM users");
1280 assert!(tokens.iter().any(|t| matches!(
1282 t,
1283 SqlToken::Keyword(k) if k == "SELECT"
1284 )));
1285 assert!(tokens.iter().any(|t| matches!(
1286 t,
1287 SqlToken::Identifier(name) if name == "id"
1288 )));
1289 assert!(tokens.iter().any(|t| matches!(
1290 t,
1291 SqlToken::Keyword(k) if k == "FROM"
1292 )));
1293 assert!(tokens.iter().any(|t| matches!(
1294 t,
1295 SqlToken::Identifier(name) if name == "users"
1296 )));
1297 }
1298
1299 #[test]
1300 fn test_tokenize_string_literal() {
1301 let tokens = tokenize("SELECT * FROM users WHERE name = 'alice'");
1302 assert!(tokens.iter().any(|t| matches!(
1303 t,
1304 SqlToken::StringLiteral(s) if s.contains("alice")
1305 )));
1306 }
1307
1308 #[test]
1309 fn test_tokenize_number_literal() {
1310 let tokens = tokenize("SELECT * FROM users WHERE age > 25");
1311 assert!(tokens.iter().any(|t| matches!(
1312 t,
1313 SqlToken::NumberLiteral(n) if n == "25"
1314 )));
1315 }
1316
1317 #[test]
1318 fn test_tokenize_line_comment() {
1319 let tokens = tokenize("SELECT * FROM users -- this is a comment");
1320 assert!(tokens.iter().any(|t| matches!(
1321 t,
1322 SqlToken::Comment(c) if c.contains("this is a comment")
1323 )));
1324 }
1325
1326 #[test]
1327 fn test_tokenize_block_comment() {
1328 let tokens = tokenize("SELECT * /* block comment */ FROM users");
1329 assert!(tokens.iter().any(|t| matches!(
1330 t,
1331 SqlToken::Comment(c) if c.contains("block comment")
1332 )));
1333 }
1334
1335 #[test]
1336 fn test_tokenize_punctuation() {
1337 let tokens = tokenize("INSERT INTO users (a, b) VALUES (1, 2)");
1338 assert!(tokens
1339 .iter()
1340 .any(|t| matches!(t, SqlToken::Punctuation('('))));
1341 assert!(tokens
1342 .iter()
1343 .any(|t| matches!(t, SqlToken::Punctuation(')'))));
1344 assert!(tokens
1345 .iter()
1346 .any(|t| matches!(t, SqlToken::Punctuation(','))));
1347 }
1348
1349 #[test]
1350 fn test_tokenize_operator() {
1351 let tokens = tokenize("SELECT * FROM users WHERE age >= 18 AND age <= 65");
1352 assert!(tokens.iter().any(|t| matches!(
1353 t,
1354 SqlToken::Operator(op) if op == ">="
1355 )));
1356 assert!(tokens.iter().any(|t| matches!(
1357 t,
1358 SqlToken::Operator(op) if op == "<="
1359 )));
1360 }
1361
1362 #[test]
1363 fn test_tokenize_double_quoted_identifier() {
1364 let tokens = tokenize("SELECT * FROM \"my table\"");
1365 assert!(tokens.iter().any(|t| matches!(
1366 t,
1367 SqlToken::Identifier(s) if s.contains("my table")
1368 )));
1369 }
1370
1371 #[test]
1372 fn test_tokenize_empty_string() {
1373 let tokens = tokenize("");
1374 assert!(tokens.is_empty());
1375 }
1376
1377 #[test]
1378 fn test_tokenize_whitespace_only() {
1379 let tokens = tokenize(" \t\n ");
1380 assert!(tokens.is_empty());
1381 }
1382
1383 #[test]
1386 fn test_ast_injection_clean_sql() {
1387 assert!(detect_injection_ast("SELECT id, name FROM users WHERE age > 18").is_ok());
1388 assert!(detect_injection_ast("INSERT INTO users (name) VALUES ('alice')").is_ok());
1389 assert!(detect_injection_ast("UPDATE users SET name = 'bob' WHERE id = 1").is_ok());
1390 }
1391
1392 #[test]
1393 fn test_ast_injection_multi_statement_drop() {
1394 let result = detect_injection_ast("SELECT * FROM users; DROP TABLE users");
1395 assert!(result.is_err());
1396 assert!(matches!(
1397 result.unwrap_err(),
1398 SqlValidationError::InjectionDetected(_)
1399 ));
1400 }
1401
1402 #[test]
1403 fn test_ast_injection_multi_statement_delete() {
1404 let result = detect_injection_ast("SELECT * FROM users; DELETE FROM users");
1405 assert!(result.is_err());
1406 }
1407
1408 #[test]
1409 fn test_ast_injection_multi_statement_insert() {
1410 let result = detect_injection_ast("SELECT 1; INSERT INTO admin VALUES (1, 'hacker')");
1411 assert!(result.is_err());
1412 }
1413
1414 #[test]
1415 fn test_ast_injection_exec_call() {
1416 let result = detect_injection_ast("EXEC sp_executesql 'DROP TABLE users'");
1417 assert!(result.is_err());
1418 let result2 = detect_injection_ast("EXECUTE sp_executesql 'DELETE FROM users'");
1419 assert!(result2.is_err());
1420 }
1421
1422 #[test]
1423 fn test_ast_injection_grant_revoke() {
1424 assert!(detect_injection_ast("GRANT ALL ON users TO hacker").is_err());
1425 assert!(detect_injection_ast("REVOKE SELECT ON users FROM app_user").is_err());
1426 }
1427
1428 #[test]
1429 fn test_ast_injection_boolean_tautology_or_1_eq_1() {
1430 let result = detect_injection_ast("SELECT * FROM users WHERE name = 'a' OR 1=1");
1431 assert!(result.is_err());
1432 }
1433
1434 #[test]
1435 fn test_ast_injection_boolean_tautology_spaces() {
1436 let result = detect_injection_ast("SELECT * FROM users WHERE name = 'a' OR 1 = 1");
1437 assert!(result.is_err());
1438 }
1439
1440 #[test]
1441 fn test_ast_injection_boolean_tautology_true() {
1442 let result = detect_injection_ast("SELECT * FROM users WHERE name = 'a' OR TRUE");
1443 assert!(result.is_err());
1444 }
1445
1446 #[test]
1447 fn test_ast_injection_no_keywords() {
1448 assert!(detect_injection_ast("12345").is_ok());
1450 assert!(detect_injection_ast("").is_ok());
1451 }
1452
1453 #[test]
1454 fn test_ast_injection_safe_semicolon() {
1455 assert!(detect_injection_ast("SELECT * FROM users; BEGIN").is_ok());
1457 }
1458
1459 #[test]
1462 fn test_whitelist_empty_allows_all() {
1463 let validator = WhitelistValidator::new();
1464 assert!(validator.validate("SELECT * FROM any_table").is_ok());
1465 assert!(validator.validate("SELECT * FROM secret_table").is_ok());
1466 }
1467
1468 #[test]
1469 fn test_whitelist_table_allowed() {
1470 let validator = WhitelistValidator::new().allow_table("users");
1471 assert!(validator.validate_tables("SELECT * FROM users").is_ok());
1472 assert!(validator
1473 .validate_tables("SELECT * FROM users WHERE id = 1")
1474 .is_ok());
1475 }
1476
1477 #[test]
1478 fn test_whitelist_table_blocked() {
1479 let validator = WhitelistValidator::new().allow_table("users");
1480 let result = validator.validate_tables("SELECT * FROM secret_table");
1481 assert!(result.is_err());
1482 assert!(matches!(
1483 result.unwrap_err(),
1484 SqlValidationError::InvalidTableName(_)
1485 ));
1486 }
1487
1488 #[test]
1489 fn test_whitelist_multiple_tables() {
1490 let validator = WhitelistValidator::new().allow_tables(&["users", "orders", "products"]);
1491 assert!(validator.validate_tables("SELECT * FROM users").is_ok());
1492 assert!(validator.validate_tables("SELECT * FROM orders").is_ok());
1493 assert!(validator.validate_tables("SELECT * FROM products").is_ok());
1494 assert!(validator
1495 .validate_tables("SELECT * FROM forbidden")
1496 .is_err());
1497 }
1498
1499 #[test]
1500 fn test_whitelist_case_insensitive() {
1501 let validator = WhitelistValidator::new().allow_table("Users");
1502 assert!(validator.validate_tables("SELECT * FROM users").is_ok());
1503 assert!(validator.validate_tables("SELECT * FROM USERS").is_ok());
1504 assert!(validator.validate_tables("SELECT * FROM Users").is_ok());
1505 }
1506
1507 #[test]
1508 fn test_whitelist_join_table() {
1509 let validator = WhitelistValidator::new().allow_tables(&["users", "orders"]);
1510 assert!(validator
1511 .validate_tables("SELECT * FROM users u JOIN orders o ON u.id = o.user_id")
1512 .is_ok());
1513 }
1514
1515 #[test]
1516 fn test_whitelist_join_blocked_table() {
1517 let validator = WhitelistValidator::new().allow_tables(&["users"]);
1518 assert!(validator
1520 .validate_tables("SELECT * FROM users u JOIN orders o ON u.id = o.user_id")
1521 .is_err());
1522 }
1523
1524 #[test]
1525 fn test_whitelist_insert_table() {
1526 let validator = WhitelistValidator::new().allow_table("users");
1527 assert!(validator
1528 .validate_tables("INSERT INTO users (name) VALUES ('a')")
1529 .is_ok());
1530 assert!(validator
1531 .validate_tables("INSERT INTO secret (name) VALUES ('a')")
1532 .is_err());
1533 }
1534
1535 #[test]
1536 fn test_whitelist_update_table() {
1537 let validator = WhitelistValidator::new().allow_table("users");
1538 assert!(validator
1539 .validate_tables("UPDATE users SET name = 'a' WHERE id = 1")
1540 .is_ok());
1541 assert!(validator
1542 .validate_tables("UPDATE admin SET role = 'super' WHERE id = 1")
1543 .is_err());
1544 }
1545
1546 #[test]
1547 fn test_whitelist_columns_not_set_passes() {
1548 let validator = WhitelistValidator::new().allow_table("users");
1549 assert!(validator
1551 .validate_columns("SELECT id, name, password FROM users")
1552 .is_ok());
1553 }
1554
1555 #[test]
1556 fn test_whitelist_combined_validate() {
1557 let validator = WhitelistValidator::new().allow_table("users");
1558 assert!(validator
1559 .validate("SELECT * FROM users WHERE id = 1")
1560 .is_ok());
1561 assert!(validator.validate("SELECT * FROM forbidden").is_err());
1562 }
1563
1564 #[test]
1567 fn test_complexity_simple_query() {
1568 let score = score_complexity("SELECT * FROM users");
1569 assert_eq!(score.level(), ComplexityLevel::Simple);
1570 assert_eq!(score.join_count, 0);
1571 assert_eq!(score.subquery_count, 0);
1572 assert_eq!(score.where_condition_count, 0);
1573 assert!(!score.has_having);
1574 assert!(!score.has_window_function);
1575 assert!(!score.has_cte);
1576 }
1577
1578 #[test]
1579 fn test_complexity_with_where() {
1580 let score = score_complexity("SELECT * FROM users WHERE age > 18 AND status = 'active'");
1581 assert!(score.where_condition_count >= 2);
1582 }
1583
1584 #[test]
1585 fn test_complexity_with_join() {
1586 let score =
1587 score_complexity("SELECT u.id, o.total FROM users u JOIN orders o ON u.id = o.user_id");
1588 assert_eq!(score.join_count, 1);
1589 }
1590
1591 #[test]
1592 fn test_complexity_with_subquery() {
1593 let score = score_complexity(
1594 "SELECT * FROM (SELECT id, name FROM users WHERE age > 18) t WHERE t.id > 0",
1595 );
1596 assert!(score.subquery_count >= 1);
1597 }
1598
1599 #[test]
1600 fn test_complexity_with_group_by_having() {
1601 let score =
1602 score_complexity("SELECT dept, COUNT(*) FROM users GROUP BY dept HAVING COUNT(*) > 5");
1603 assert!(score.group_by_count >= 1);
1604 assert!(score.has_having);
1605 }
1606
1607 #[test]
1608 fn test_complexity_with_cte() {
1609 let score = score_complexity(
1610 "WITH active_users AS (SELECT id FROM users WHERE status = 'active') SELECT * FROM active_users",
1611 );
1612 assert!(score.has_cte);
1613 }
1614
1615 #[test]
1616 fn test_complexity_with_window_function() {
1617 let score = score_complexity(
1618 "SELECT id, ROW_NUMBER() OVER (PARTITION BY dept ORDER BY salary DESC) FROM users",
1619 );
1620 assert!(score.has_window_function);
1621 }
1622
1623 #[test]
1624 fn test_complexity_with_set_operation() {
1625 let score = score_complexity("SELECT id FROM users UNION SELECT id FROM archived_users");
1626 assert!(score.set_operation_count >= 1);
1627 }
1628
1629 #[test]
1630 fn test_complexity_level_thresholds() {
1631 let simple = score_complexity("SELECT * FROM users");
1633 assert_eq!(simple.level(), ComplexityLevel::Simple);
1634
1635 let complex = score_complexity(
1637 "WITH t1 AS (SELECT id FROM users WHERE a = 1 AND b = 2 OR c = 3) \
1638 SELECT t1.id, t2.name, ROW_NUMBER() OVER (PARTITION BY t1.id ORDER BY t2.name) \
1639 FROM t1 JOIN orders t2 ON t1.id = t2.user_id \
1640 GROUP BY t1.id, t2.name HAVING COUNT(*) > 1 \
1641 UNION SELECT id, name, 1 FROM archived",
1642 );
1643 assert!(complex.score > simple.score);
1644 }
1645
1646 #[test]
1647 fn test_complexity_level_descriptions() {
1648 assert_eq!(ComplexityLevel::Simple.description(), "简单");
1649 assert_eq!(ComplexityLevel::Moderate.description(), "中等");
1650 assert_eq!(ComplexityLevel::Complex.description(), "复杂");
1651 assert_eq!(ComplexityLevel::VeryComplex.description(), "非常复杂");
1652 }
1653
1654 #[test]
1655 fn test_complexity_score_max_100() {
1656 let mut sql = String::from("SELECT * FROM users");
1658 for i in 0..20 {
1659 sql.push_str(&format!(" JOIN orders o{} ON users.id = o{}.user_id", i, i));
1660 }
1661 let score = score_complexity(&sql);
1662 assert!(score.score <= 100);
1663 }
1664
1665 #[test]
1666 fn test_complexity_empty_sql() {
1667 let score = score_complexity("");
1668 assert_eq!(score.score, 0);
1669 assert_eq!(score.level(), ComplexityLevel::Simple);
1670 }
1671
1672 #[test]
1675 fn test_ddl_policy_default_all_denied() {
1676 let policy = DdlPolicy::default();
1677 assert!(policy.validate("CREATE TABLE users (id INT)").is_err());
1678 assert!(policy.validate("DROP TABLE users").is_err());
1679 assert!(policy
1680 .validate("ALTER TABLE users ADD COLUMN name TEXT")
1681 .is_err());
1682 assert!(policy.validate("TRUNCATE TABLE users").is_err());
1683 }
1684
1685 #[test]
1686 fn test_ddl_policy_read_only() {
1687 let policy = DdlPolicy::read_only();
1688 assert!(policy.validate("CREATE TABLE users (id INT)").is_err());
1689 assert!(policy.validate("DROP TABLE users").is_err());
1690 assert!(policy.validate("TRUNCATE TABLE users").is_err());
1691 }
1692
1693 #[test]
1694 fn test_ddl_policy_permissive_allows_all() {
1695 let policy = DdlPolicy::permissive();
1696 assert!(policy.validate("CREATE TABLE users (id INT)").is_ok());
1697 assert!(policy.validate("DROP TABLE users").is_ok());
1698 assert!(policy
1699 .validate("ALTER TABLE users ADD COLUMN name TEXT")
1700 .is_ok());
1701 assert!(policy.validate("TRUNCATE TABLE users").is_ok());
1702 }
1703
1704 #[test]
1705 fn test_ddl_policy_safe_evolution() {
1706 let policy = DdlPolicy::safe_evolution();
1707 assert!(policy.validate("CREATE TABLE users (id INT)").is_ok());
1709 assert!(policy
1710 .validate("ALTER TABLE users ADD COLUMN name TEXT")
1711 .is_ok());
1712 assert!(policy.validate("DROP TABLE users").is_err());
1714 assert!(policy.validate("TRUNCATE TABLE users").is_err());
1715 }
1716
1717 #[test]
1718 fn test_ddl_policy_create_index() {
1719 let mut policy = DdlPolicy::permissive();
1720 policy.allow_create_index = false;
1721 assert!(policy
1723 .validate("CREATE INDEX idx_name ON users (name)")
1724 .is_err());
1725 assert!(policy.validate("CREATE TABLE users (id INT)").is_ok());
1727 }
1728
1729 #[test]
1730 fn test_ddl_policy_drop_index() {
1731 let mut policy = DdlPolicy::permissive();
1732 policy.allow_drop_index = false;
1733 assert!(policy.validate("DROP INDEX idx_name").is_err());
1735 assert!(policy.validate("DROP TABLE users").is_ok());
1737 }
1738
1739 #[test]
1740 fn test_ddl_policy_non_ddl_passes() {
1741 let policy = DdlPolicy::read_only();
1742 assert!(policy.validate("SELECT * FROM users").is_ok());
1744 assert!(policy.validate("INSERT INTO users VALUES (1)").is_ok());
1745 assert!(policy.validate("UPDATE users SET name = 'a'").is_ok());
1746 assert!(policy.validate("DELETE FROM users").is_ok());
1747 }
1748
1749 #[test]
1750 fn test_ddl_policy_custom() {
1751 let policy = DdlPolicy {
1752 allow_create: true,
1753 allow_drop: false,
1754 allow_alter: true,
1755 allow_truncate: false,
1756 allow_create_index: true,
1757 allow_drop_index: false,
1758 };
1759 assert!(policy.validate("CREATE TABLE t (id INT)").is_ok());
1760 assert!(policy.validate("DROP TABLE t").is_err());
1761 assert!(policy.validate("ALTER TABLE t ADD COLUMN c INT").is_ok());
1762 assert!(policy.validate("TRUNCATE TABLE t").is_err());
1763 assert!(policy.validate("CREATE INDEX i ON t (c)").is_ok());
1764 assert!(policy.validate("DROP INDEX i").is_err());
1765 }
1766}