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 = [
247 ("'; DROP TABLE", "DROP TABLE injection"),
248 ("' OR '1'='1", "classic OR injection"),
249 ("' OR 1=1", "OR 1=1 injection"),
250 (" OR 1=1", "OR 1=1 injection (bare)"),
251 (" OR '1'='1", "OR constant injection (bare)"),
252 ("; DROP", "multi-statement DROP injection"),
253 ("; DELETE", "multi-statement DELETE injection"),
254 ("; INSERT", "multi-statement INSERT injection"),
255 ("; UPDATE", "multi-statement UPDATE injection"),
256 (
257 "UNION SELECT",
258 "UNION SELECT injection (not allowed in simple queries)",
259 ),
260 ("--", "comment injection (not allowed)"),
261 ("/*", "block comment (not allowed)"),
262 ];
263
264 for (pattern, desc) in &suspicious_patterns {
265 if sql_upper.contains(pattern) {
266 return Err(SqlValidationError::InjectionDetected(format!(
267 "{}: {}",
268 desc, pattern
269 )));
270 }
271 }
272
273 Ok(())
274}
275
276pub fn detect_statement_type(sql: &str) -> SqlStatementType {
278 let trimmed = sql.trim().to_uppercase();
279
280 if trimmed.starts_with("SELECT") {
281 SqlStatementType::Select
282 } else if trimmed.starts_with("INSERT") {
283 SqlStatementType::Insert
284 } else if trimmed.starts_with("UPDATE") {
285 SqlStatementType::Update
286 } else if trimmed.starts_with("DELETE") {
287 SqlStatementType::Delete
288 } else if trimmed.starts_with("CREATE") {
289 SqlStatementType::Create
290 } else if trimmed.starts_with("DROP") {
291 SqlStatementType::Drop
292 } else if trimmed.starts_with("ALTER") {
293 SqlStatementType::Alter
294 } else if trimmed.starts_with("TRUNCATE") {
295 SqlStatementType::Truncate
296 } else {
297 SqlStatementType::Other
298 }
299}
300
301pub fn validate_parameter_count(sql: &str, expected_params: usize) -> ValidationResult {
303 let param_count = sql.chars().filter(|&c| c == '?').count() + sql.matches('$').count(); if param_count != expected_params {
305 return Err(SqlValidationError::ParameterCountMismatch {
306 expected: expected_params,
307 got: param_count,
308 });
309 }
310 Ok(())
311}
312
313pub fn validate_table_name(name: &str) -> ValidationResult {
315 if name.is_empty() {
316 return Err(SqlValidationError::InvalidTableName(
317 "empty table name".to_string(),
318 ));
319 }
320
321 let cleaned = name.trim_matches('`').trim_matches('"').trim_matches('\'');
324 if cleaned.is_empty() {
325 return Err(SqlValidationError::InvalidTableName(name.to_string()));
326 }
327
328 for ch in cleaned.chars() {
329 if !ch.is_alphanumeric() && ch != '_' {
330 return Err(SqlValidationError::InvalidTableName(format!(
331 "table name '{}' contains invalid character '{}'",
332 name, ch
333 )));
334 }
335 }
336
337 Ok(())
338}
339
340pub fn validate_column_name(name: &str) -> ValidationResult {
342 if name.is_empty() || name == "*" {
343 return Ok(()); }
345
346 let cleaned = name.trim_matches('`').trim_matches('"').trim_matches('\'');
347 if cleaned.is_empty() {
348 return Err(SqlValidationError::InvalidIdentifier(name.to_string()));
349 }
350
351 for ch in cleaned.chars() {
353 if !ch.is_alphanumeric() && ch != '_' && ch != '.' {
354 return Err(SqlValidationError::InvalidIdentifier(format!(
355 "column '{}' contains invalid character '{}'",
356 name, ch
357 )));
358 }
359 }
360
361 Ok(())
362}
363
364pub fn validate(sql: &str) -> ValidationResult {
366 if sql.trim().is_empty() {
367 return Err(SqlValidationError::SyntaxError(
368 "Empty SQL statement".to_string(),
369 ));
370 }
371
372 validate_sql(sql)?;
373 validate_balanced_parentheses(sql)?;
374 validate_string_literals(sql)?;
375 validate_no_injection_patterns(sql)?;
376
377 Ok(())
378}
379
380#[derive(Debug, Clone, PartialEq)]
386pub enum SqlToken {
387 Keyword(String),
389 Identifier(String),
391 StringLiteral(String),
393 NumberLiteral(String),
395 Operator(String),
397 Punctuation(char),
399 Comment(String),
401}
402
403pub fn tokenize(sql: &str) -> Vec<SqlToken> {
408 let mut tokens = Vec::new();
409 let chars: Vec<char> = sql.chars().collect();
410 let mut i = 0;
411
412 while i < chars.len() {
413 let ch = chars[i];
414
415 if ch.is_whitespace() {
417 i += 1;
418 continue;
419 }
420
421 if i + 1 < chars.len() && ch == '-' && chars[i + 1] == '-' {
423 let start = i;
424 while i < chars.len() && chars[i] != '\n' {
425 i += 1;
426 }
427 tokens.push(SqlToken::Comment(chars[start..i].iter().collect()));
428 continue;
429 }
430
431 if i + 1 < chars.len() && ch == '/' && chars[i + 1] == '*' {
433 let start = i;
434 i += 2;
435 while i + 1 < chars.len() {
436 if chars[i] == '*' && chars[i + 1] == '/' {
437 i += 2;
438 break;
439 }
440 i += 1;
441 }
442 tokens.push(SqlToken::Comment(chars[start..i].iter().collect()));
443 continue;
444 }
445
446 if ch == '\'' {
448 let start = i;
449 i += 1;
450 while i < chars.len() {
451 if chars[i] == '\'' {
452 i += 1;
453 break;
454 }
455 i += 1;
456 }
457 tokens.push(SqlToken::StringLiteral(chars[start..i].iter().collect()));
458 continue;
459 }
460
461 if ch == '"' {
463 let start = i;
464 i += 1;
465 while i < chars.len() && chars[i] != '"' {
466 i += 1;
467 }
468 if i < chars.len() {
469 i += 1;
470 }
471 tokens.push(SqlToken::Identifier(chars[start..i].iter().collect()));
472 continue;
473 }
474
475 if ch.is_ascii_digit() {
477 let start = i;
478 while i < chars.len() && (chars[i].is_ascii_digit() || chars[i] == '.') {
479 i += 1;
480 }
481 tokens.push(SqlToken::NumberLiteral(chars[start..i].iter().collect()));
482 continue;
483 }
484
485 if ch.is_alphabetic() || ch == '_' {
487 let start = i;
488 while i < chars.len() && (chars[i].is_alphanumeric() || chars[i] == '_') {
489 i += 1;
490 }
491 let word: String = chars[start..i].iter().collect();
492 let upper = word.to_uppercase();
493 const KEYWORDS: &[&str] = &[
494 "SELECT",
495 "FROM",
496 "WHERE",
497 "INSERT",
498 "INTO",
499 "VALUES",
500 "UPDATE",
501 "SET",
502 "DELETE",
503 "CREATE",
504 "TABLE",
505 "DROP",
506 "ALTER",
507 "TRUNCATE",
508 "JOIN",
509 "INNER",
510 "LEFT",
511 "RIGHT",
512 "OUTER",
513 "ON",
514 "AND",
515 "OR",
516 "NOT",
517 "NULL",
518 "IS",
519 "IN",
520 "LIKE",
521 "BETWEEN",
522 "ORDER",
523 "BY",
524 "GROUP",
525 "HAVING",
526 "LIMIT",
527 "OFFSET",
528 "DISTINCT",
529 "AS",
530 "UNION",
531 "ALL",
532 "INTERSECT",
533 "EXCEPT",
534 "CASE",
535 "WHEN",
536 "THEN",
537 "ELSE",
538 "END",
539 "IF",
540 "EXISTS",
541 "PRIMARY",
542 "KEY",
543 "FOREIGN",
544 "REFERENCES",
545 "INDEX",
546 "VIEW",
547 "DATABASE",
548 "SCHEMA",
549 "GRANT",
550 "REVOKE",
551 "EXEC",
552 "EXECUTE",
553 "PROCEDURE",
554 "FUNCTION",
555 "BEGIN",
556 "COMMIT",
557 "ROLLBACK",
558 "SAVEPOINT",
559 "RELEASE",
560 "TRANSACTION",
561 "START",
562 "WITH",
563 "RECURSIVE",
564 ];
565 if KEYWORDS.contains(&upper.as_str()) {
566 tokens.push(SqlToken::Keyword(upper));
567 } else {
568 tokens.push(SqlToken::Identifier(word));
569 }
570 continue;
571 }
572
573 if "+-*/=<>!".contains(ch) {
575 let start = i;
576 i += 1;
577 while i < chars.len() && "+-*/=<>!".contains(chars[i]) {
578 i += 1;
579 }
580 tokens.push(SqlToken::Operator(chars[start..i].iter().collect()));
581 continue;
582 }
583
584 if "().,;".contains(ch) {
586 tokens.push(SqlToken::Punctuation(ch));
587 i += 1;
588 continue;
589 }
590
591 i += 1;
593 }
594
595 tokens
596}
597
598pub fn detect_injection_ast(sql: &str) -> ValidationResult {
606 let tokens = tokenize(sql);
607 let keywords: Vec<&str> = tokens
608 .iter()
609 .filter_map(|t| match t {
610 SqlToken::Keyword(k) => Some(k.as_str()),
611 _ => None,
612 })
613 .collect();
614
615 if keywords.is_empty() {
616 return Ok(());
617 }
618
619 let mut after_semicolon = false;
621 for token in &tokens {
622 match token {
623 SqlToken::Punctuation(';') => {
624 after_semicolon = true;
625 }
626 SqlToken::Keyword(k) if after_semicolon => {
627 match k.as_str() {
629 "DROP" | "ALTER" | "TRUNCATE" | "DELETE" | "INSERT" | "UPDATE" | "CREATE"
630 | "GRANT" | "REVOKE" | "EXEC" | "EXECUTE" => {
631 return Err(SqlValidationError::InjectionDetected(format!(
632 "multi-statement injection: semicolon followed by {} keyword",
633 k
634 )));
635 }
636 _ => {
637 after_semicolon = false;
638 }
639 }
640 }
641 SqlToken::Keyword(_) => {
642 after_semicolon = false;
643 }
644 _ => {}
645 }
646 }
647
648 if keywords.iter().any(|k| *k == "EXEC" || *k == "EXECUTE") {
650 return Err(SqlValidationError::InjectionDetected(
651 "EXEC/EXECUTE call detected (potential injection)".to_string(),
652 ));
653 }
654
655 if keywords.iter().any(|k| *k == "GRANT" || *k == "REVOKE") {
657 return Err(SqlValidationError::InjectionDetected(
658 "GRANT/REVOKE statement detected (potential privilege escalation)".to_string(),
659 ));
660 }
661
662 let upper_sql = sql.to_uppercase();
664 if upper_sql.contains(" OR 1=1")
665 || upper_sql.contains(" OR 1 = 1")
666 || upper_sql.contains(" OR '1'='1'")
667 || upper_sql.contains(" OR TRUE")
668 || upper_sql.contains(" OR 1<>0")
669 {
670 return Err(SqlValidationError::InjectionDetected(
671 "boolean blind injection pattern detected (OR with tautology)".to_string(),
672 ));
673 }
674
675 Ok(())
676}
677
678#[derive(Debug, Clone, Default)]
680pub struct WhitelistValidator {
681 allowed_tables: std::collections::HashSet<String>,
683 allowed_columns: Option<std::collections::HashSet<String>>,
685}
686
687impl WhitelistValidator {
688 pub fn new() -> Self {
690 Self {
691 allowed_tables: std::collections::HashSet::new(),
692 allowed_columns: None,
693 }
694 }
695
696 pub fn allow_table(mut self, table: &str) -> Self {
698 self.allowed_tables.insert(table.to_lowercase());
699 self
700 }
701
702 pub fn allow_tables(mut self, tables: &[&str]) -> Self {
704 for t in tables {
705 self.allowed_tables.insert(t.to_lowercase());
706 }
707 self
708 }
709
710 pub fn allow_columns(mut self, columns: &[&str]) -> Self {
712 let set: std::collections::HashSet<String> =
713 columns.iter().map(|c| c.to_lowercase()).collect();
714 self.allowed_columns = Some(set);
715 self
716 }
717
718 pub fn validate_tables(&self, sql: &str) -> ValidationResult {
723 if self.allowed_tables.is_empty() {
724 return Ok(()); }
726
727 let tokens = tokenize(sql);
728 let mut check_next_identifier = false;
729
730 for token in &tokens {
731 match token {
732 SqlToken::Keyword(k)
733 if matches!(k.as_str(), "FROM" | "JOIN" | "INTO" | "UPDATE" | "TABLE") =>
734 {
735 check_next_identifier = true;
736 }
737 SqlToken::Identifier(name) if check_next_identifier => {
738 let lower = name.to_lowercase();
739 if !self.allowed_tables.contains(&lower) {
740 return Err(SqlValidationError::InvalidTableName(format!(
741 "table '{}' is not in the whitelist",
742 name
743 )));
744 }
745 check_next_identifier = false;
746 }
747 _ if check_next_identifier => {
748 #[allow(clippy::collapsible_match)]
750 if !matches!(token, SqlToken::Punctuation('.')) {
751 check_next_identifier = false;
752 }
753 }
754 _ => {}
755 }
756 }
757
758 Ok(())
759 }
760
761 pub fn validate_columns(&self, sql: &str) -> ValidationResult {
766 let allowed = match &self.allowed_columns {
767 Some(c) => c,
768 None => return Ok(()), };
770
771 let tokens = tokenize(sql);
772 for token in &tokens {
773 if let SqlToken::Identifier(name) = token {
774 let lower = name.to_lowercase();
775 if self.allowed_tables.contains(&lower) {
777 continue;
778 }
779 if lower == "*" {
781 continue;
782 }
783 if !allowed.contains(&lower) && !name.contains('.') {
785 }
788 }
789 }
790
791 Ok(())
792 }
793
794 pub fn validate(&self, sql: &str) -> ValidationResult {
796 self.validate_tables(sql)?;
797 self.validate_columns(sql)
798 }
799}
800
801#[derive(Debug, Clone)]
803pub struct SqlComplexityScore {
804 pub score: u32,
806 pub join_count: u32,
808 pub subquery_count: u32,
810 pub where_condition_count: u32,
812 pub set_operation_count: u32,
814 pub group_by_count: u32,
816 pub has_having: bool,
818 pub has_window_function: bool,
820 pub has_cte: bool,
822 pub token_count: u32,
824}
825
826impl SqlComplexityScore {
827 fn calculate(&mut self) {
829 let mut score: u32 = 0;
830 score += self.join_count * 5;
831 score += self.subquery_count * 10;
832 score += self.where_condition_count * 3;
833 score += self.set_operation_count * 8;
834 score += self.group_by_count * 3;
835 if self.has_having {
836 score += 5;
837 }
838 if self.has_window_function {
839 score += 8;
840 }
841 if self.has_cte {
842 score += 6;
843 }
844 score += (self.token_count / 50).min(20);
846 self.score = score.min(100);
847 }
848
849 pub fn level(&self) -> ComplexityLevel {
851 match self.score {
852 0..=20 => ComplexityLevel::Simple,
853 21..=40 => ComplexityLevel::Moderate,
854 41..=60 => ComplexityLevel::Complex,
855 _ => ComplexityLevel::VeryComplex,
856 }
857 }
858}
859
860#[derive(Debug, Clone, Copy, PartialEq, Eq)]
862pub enum ComplexityLevel {
863 Simple,
865 Moderate,
867 Complex,
869 VeryComplex,
871}
872
873impl ComplexityLevel {
874 pub fn description(&self) -> &'static str {
876 match self {
877 ComplexityLevel::Simple => "简单",
878 ComplexityLevel::Moderate => "中等",
879 ComplexityLevel::Complex => "复杂",
880 ComplexityLevel::VeryComplex => "非常复杂",
881 }
882 }
883}
884
885pub fn score_complexity(sql: &str) -> SqlComplexityScore {
890 let tokens = tokenize(sql);
891 let token_count = tokens.len() as u32;
892
893 let mut join_count = 0u32;
894 let mut subquery_count = 0u32;
895 let mut where_condition_count = 0u32;
896 let mut set_operation_count = 0u32;
897 let mut group_by_count = 0u32;
898 let mut has_having = false;
899 let mut has_window_function = false;
900 let mut has_cte = false;
901
902 let mut in_where = false;
903 let mut in_group_by = false;
904 let mut paren_depth: i32 = 0;
905
906 for token in &tokens {
907 match token {
908 SqlToken::Keyword(k) => {
909 match k.as_str() {
910 "JOIN" | "INNER" | "LEFT" | "RIGHT" | "OUTER" => {
911 if k == "JOIN" {
912 join_count += 1;
913 }
914 }
915 "WHERE" => {
916 in_where = true;
917 where_condition_count += 1;
918 }
919 "AND" | "OR" if in_where => {
920 where_condition_count += 1;
921 }
922 "GROUP" => {
923 in_group_by = true;
924 }
925 "HAVING" => {
926 has_having = true;
927 in_where = false;
928 in_group_by = false;
929 }
930 "UNION" | "INTERSECT" | "EXCEPT" => {
931 set_operation_count += 1;
932 }
933 "WITH" => {
934 has_cte = true;
935 }
936 "SELECT" if paren_depth > 0 => {
937 subquery_count += 1;
938 }
939 _ => {}
940 }
941 if k != "WHERE" && k != "AND" && k != "OR" {
942 in_where = false;
943 }
944 if k != "GROUP" && k != "BY" && in_group_by {
945 in_group_by = false;
946 }
947 }
948 SqlToken::Identifier(name) => {
949 let upper = name.to_uppercase();
950 if upper.contains("OVER") || upper.contains("ROW_NUMBER") || upper.contains("RANK")
951 {
952 has_window_function = true;
953 }
954 if in_group_by {
955 group_by_count += 1;
956 }
957 }
958 SqlToken::Punctuation('(') => {
959 paren_depth += 1;
960 }
961 SqlToken::Punctuation(')') => {
962 paren_depth -= 1;
963 }
964 _ => {}
965 }
966 }
967
968 let mut score = SqlComplexityScore {
969 score: 0,
970 join_count,
971 subquery_count,
972 where_condition_count,
973 set_operation_count,
974 group_by_count,
975 has_having,
976 has_window_function,
977 has_cte,
978 token_count,
979 };
980 score.calculate();
981 score
982}
983
984#[derive(Debug, Clone, Default)]
986pub struct DdlPolicy {
987 pub allow_create: bool,
989 pub allow_drop: bool,
991 pub allow_alter: bool,
993 pub allow_truncate: bool,
995 pub allow_create_index: bool,
997 pub allow_drop_index: bool,
999}
1000
1001impl DdlPolicy {
1002 pub fn permissive() -> Self {
1004 Self {
1005 allow_create: true,
1006 allow_drop: true,
1007 allow_alter: true,
1008 allow_truncate: true,
1009 allow_create_index: true,
1010 allow_drop_index: true,
1011 }
1012 }
1013
1014 pub fn read_only() -> Self {
1016 Self::default()
1017 }
1018
1019 pub fn safe_evolution() -> Self {
1021 Self {
1022 allow_create: true,
1023 allow_drop: false,
1024 allow_alter: true,
1025 allow_truncate: false,
1026 allow_create_index: true,
1027 allow_drop_index: false,
1028 }
1029 }
1030
1031 pub fn validate(&self, sql: &str) -> ValidationResult {
1033 let stmt_type = detect_statement_type(sql);
1034 let upper = sql.to_uppercase();
1035
1036 match stmt_type {
1037 SqlStatementType::Create => {
1038 if !self.allow_create {
1039 return Err(SqlValidationError::SyntaxError(
1040 "CREATE operations are not allowed by DDL policy".to_string(),
1041 ));
1042 }
1043 if upper.contains("INDEX") && !self.allow_create_index {
1044 return Err(SqlValidationError::SyntaxError(
1045 "CREATE INDEX operations are not allowed by DDL policy".to_string(),
1046 ));
1047 }
1048 Ok(())
1049 }
1050 SqlStatementType::Drop => {
1051 if !self.allow_drop {
1052 return Err(SqlValidationError::SyntaxError(
1053 "DROP operations are not allowed by DDL policy".to_string(),
1054 ));
1055 }
1056 if upper.contains("INDEX") && !self.allow_drop_index {
1057 return Err(SqlValidationError::SyntaxError(
1058 "DROP INDEX operations are not allowed by DDL policy".to_string(),
1059 ));
1060 }
1061 Ok(())
1062 }
1063 SqlStatementType::Alter => {
1064 if !self.allow_alter {
1065 return Err(SqlValidationError::SyntaxError(
1066 "ALTER operations are not allowed by DDL policy".to_string(),
1067 ));
1068 }
1069 Ok(())
1070 }
1071 SqlStatementType::Truncate => {
1072 if !self.allow_truncate {
1073 return Err(SqlValidationError::SyntaxError(
1074 "TRUNCATE operations are not allowed by DDL policy".to_string(),
1075 ));
1076 }
1077 Ok(())
1078 }
1079 _ => Ok(()), }
1081 }
1082}
1083
1084#[cfg(test)]
1085mod tests {
1086 use super::*;
1087
1088 #[test]
1089 fn test_validate_select_basic() {
1090 assert!(validate_select("SELECT * FROM users").is_ok());
1091 assert!(validate_select("SELECT id, name FROM users WHERE id = 1").is_ok());
1092 assert!(validate_select(
1093 "SELECT u.id, u.name FROM users u INNER JOIN orders o ON u.id = o.user_id"
1094 )
1095 .is_ok());
1096 }
1097
1098 #[test]
1099 fn test_validate_select_missing_from() {
1100 let result = validate_select("SELECT *");
1101 assert!(result.is_err());
1102 }
1103
1104 #[test]
1105 fn test_validate_insert_basic() {
1106 assert!(validate_insert("INSERT INTO users (name) VALUES ('alice')").is_ok());
1107 assert!(validate_insert("INSERT INTO users (name, age) VALUES ('bob', 25)").is_ok());
1108 }
1109
1110 #[test]
1111 fn test_validate_insert_missing_values() {
1112 let result = validate_insert("INSERT INTO users (name)");
1113 assert!(result.is_err());
1114 }
1115
1116 #[test]
1117 fn test_validate_update_basic() {
1118 assert!(validate_update("UPDATE users SET name = 'alice' WHERE id = 1").is_ok());
1119 }
1120
1121 #[test]
1122 fn test_validate_update_missing_set() {
1123 let result = validate_update("UPDATE users WHERE id = 1");
1124 assert!(result.is_err());
1125 }
1126
1127 #[test]
1128 fn test_validate_delete_basic() {
1129 assert!(validate_delete("DELETE FROM users WHERE id = 1").is_ok());
1130 }
1131
1132 #[test]
1133 fn test_validate_delete_missing_from() {
1134 let result = validate_delete("DELETE users");
1135 assert!(result.is_err());
1136 }
1137
1138 #[test]
1139 fn test_balanced_parentheses() {
1140 assert!(validate_balanced_parentheses("SELECT * FROM (SELECT * FROM users) t").is_ok());
1141 assert!(validate_balanced_parentheses("FUNC(a, b, c)").is_ok());
1142 assert!(
1143 validate_balanced_parentheses("SELECT * FROM users WHERE (a=1 AND (b=2 OR c=3))")
1144 .is_ok()
1145 );
1146 }
1147
1148 #[test]
1149 fn test_unbalanced_parentheses() {
1150 assert!(validate_balanced_parentheses("SELECT * FROM (SELECT * FROM users").is_err());
1151 assert!(validate_balanced_parentheses("SELECT * FROM users)").is_err());
1152 }
1153
1154 #[test]
1155 fn test_string_literals_closed() {
1156 assert!(validate_string_literals("SELECT * FROM users WHERE name = 'alice'").is_ok());
1157 assert!(validate_string_literals("INSERT INTO users (name) VALUES ('bob')").is_ok());
1158 }
1159
1160 #[test]
1161 fn test_unclosed_string_literal() {
1162 assert!(validate_string_literals("SELECT * FROM users WHERE name = 'alice").is_err());
1163 }
1164
1165 #[test]
1166 fn test_injection_detection() {
1167 assert!(validate_no_injection_patterns("SELECT * FROM users WHERE name = 'alice'").is_ok());
1168 assert!(validate_no_injection_patterns(
1169 "SELECT * FROM users WHERE name = 'alice' OR '1'='1'"
1170 )
1171 .is_err());
1172 assert!(validate_no_injection_patterns("'; DROP TABLE users; --").is_err());
1173 assert!(validate_no_injection_patterns("1 UNION SELECT * FROM users").is_err());
1174 }
1175
1176 #[test]
1177 fn test_detect_statement_type() {
1178 assert_eq!(
1179 detect_statement_type("SELECT * FROM users"),
1180 SqlStatementType::Select
1181 );
1182 assert_eq!(
1183 detect_statement_type("INSERT INTO users VALUES (1)"),
1184 SqlStatementType::Insert
1185 );
1186 assert_eq!(
1187 detect_statement_type("UPDATE users SET a=1"),
1188 SqlStatementType::Update
1189 );
1190 assert_eq!(
1191 detect_statement_type("DELETE FROM users"),
1192 SqlStatementType::Delete
1193 );
1194 assert_eq!(
1195 detect_statement_type("CREATE TABLE users"),
1196 SqlStatementType::Create
1197 );
1198 assert_eq!(
1199 detect_statement_type("DROP TABLE users"),
1200 SqlStatementType::Drop
1201 );
1202 assert_eq!(
1203 detect_statement_type("ALTER TABLE users ADD COLUMN a"),
1204 SqlStatementType::Alter
1205 );
1206 assert_eq!(
1207 detect_statement_type("TRUNCATE TABLE users"),
1208 SqlStatementType::Truncate
1209 );
1210 assert_eq!(
1211 detect_statement_type("EXPLAIN SELECT * FROM users"),
1212 SqlStatementType::Other
1213 );
1214 }
1215
1216 #[test]
1217 fn test_parameter_count() {
1218 assert!(
1219 validate_parameter_count("SELECT * FROM users WHERE id = ? AND name = ?", 2).is_ok()
1220 );
1221 assert!(validate_parameter_count("SELECT * FROM users WHERE id = ?", 1).is_ok());
1222 assert!(validate_parameter_count("SELECT * FROM users WHERE id = ?", 2).is_err());
1223 }
1224
1225 #[test]
1226 fn test_validate_table_name() {
1227 assert!(validate_table_name("users").is_ok());
1228 assert!(validate_table_name("user_orders").is_ok());
1229 assert!(validate_table_name("").is_err());
1230 assert!(validate_table_name("users; DROP TABLE").is_err());
1231 }
1232
1233 #[test]
1234 fn test_validate_column_name() {
1235 assert!(validate_column_name("id").is_ok());
1236 assert!(validate_column_name("*").is_ok());
1237 assert!(validate_column_name("users.name").is_ok());
1238 assert!(validate_column_name("").is_ok()); }
1240
1241 #[test]
1242 fn test_validate_empty_sql() {
1243 assert!(validate("").is_err());
1244 assert!(validate(" ").is_err());
1245 }
1246
1247 #[test]
1248 fn test_validate_complex_queries() {
1249 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());
1250 }
1251
1252 #[test]
1253 fn test_empty_insert_data() {
1254 let sql = "INSERT INTO users () VALUES ()";
1255 assert!(validate_sql(sql).is_ok());
1256 }
1257
1258 #[test]
1259 fn test_create_table_validation() {
1260 assert!(validate_sql("CREATE TABLE users (id INT PRIMARY KEY, name VARCHAR(100))").is_ok());
1261 }
1262
1263 #[test]
1264 fn test_double_quoted_identifiers() {
1265 assert!(
1266 validate_string_literals("SELECT * FROM \"users\" WHERE \"name\" = 'alice'").is_ok()
1267 );
1268 }
1269
1270 #[test]
1271 fn test_nested_function_calls() {
1272 assert!(validate_balanced_parentheses(
1273 "SELECT MAX(COUNT(*)) FROM (SELECT COUNT(*) FROM users GROUP BY status) t"
1274 )
1275 .is_ok());
1276 }
1277
1278 #[test]
1285 fn test_tokenize_select_basic() {
1286 let tokens = tokenize("SELECT id FROM users");
1287 assert!(tokens.iter().any(|t| matches!(
1289 t,
1290 SqlToken::Keyword(k) if k == "SELECT"
1291 )));
1292 assert!(tokens.iter().any(|t| matches!(
1293 t,
1294 SqlToken::Identifier(name) if name == "id"
1295 )));
1296 assert!(tokens.iter().any(|t| matches!(
1297 t,
1298 SqlToken::Keyword(k) if k == "FROM"
1299 )));
1300 assert!(tokens.iter().any(|t| matches!(
1301 t,
1302 SqlToken::Identifier(name) if name == "users"
1303 )));
1304 }
1305
1306 #[test]
1307 fn test_tokenize_string_literal() {
1308 let tokens = tokenize("SELECT * FROM users WHERE name = 'alice'");
1309 assert!(tokens.iter().any(|t| matches!(
1310 t,
1311 SqlToken::StringLiteral(s) if s.contains("alice")
1312 )));
1313 }
1314
1315 #[test]
1316 fn test_tokenize_number_literal() {
1317 let tokens = tokenize("SELECT * FROM users WHERE age > 25");
1318 assert!(tokens.iter().any(|t| matches!(
1319 t,
1320 SqlToken::NumberLiteral(n) if n == "25"
1321 )));
1322 }
1323
1324 #[test]
1325 fn test_tokenize_line_comment() {
1326 let tokens = tokenize("SELECT * FROM users -- this is a comment");
1327 assert!(tokens.iter().any(|t| matches!(
1328 t,
1329 SqlToken::Comment(c) if c.contains("this is a comment")
1330 )));
1331 }
1332
1333 #[test]
1334 fn test_tokenize_block_comment() {
1335 let tokens = tokenize("SELECT * /* block comment */ FROM users");
1336 assert!(tokens.iter().any(|t| matches!(
1337 t,
1338 SqlToken::Comment(c) if c.contains("block comment")
1339 )));
1340 }
1341
1342 #[test]
1343 fn test_tokenize_punctuation() {
1344 let tokens = tokenize("INSERT INTO users (a, b) VALUES (1, 2)");
1345 assert!(tokens
1346 .iter()
1347 .any(|t| matches!(t, SqlToken::Punctuation('('))));
1348 assert!(tokens
1349 .iter()
1350 .any(|t| matches!(t, SqlToken::Punctuation(')'))));
1351 assert!(tokens
1352 .iter()
1353 .any(|t| matches!(t, SqlToken::Punctuation(','))));
1354 }
1355
1356 #[test]
1357 fn test_tokenize_operator() {
1358 let tokens = tokenize("SELECT * FROM users WHERE age >= 18 AND age <= 65");
1359 assert!(tokens.iter().any(|t| matches!(
1360 t,
1361 SqlToken::Operator(op) if op == ">="
1362 )));
1363 assert!(tokens.iter().any(|t| matches!(
1364 t,
1365 SqlToken::Operator(op) if op == "<="
1366 )));
1367 }
1368
1369 #[test]
1370 fn test_tokenize_double_quoted_identifier() {
1371 let tokens = tokenize("SELECT * FROM \"my table\"");
1372 assert!(tokens.iter().any(|t| matches!(
1373 t,
1374 SqlToken::Identifier(s) if s.contains("my table")
1375 )));
1376 }
1377
1378 #[test]
1379 fn test_tokenize_empty_string() {
1380 let tokens = tokenize("");
1381 assert!(tokens.is_empty());
1382 }
1383
1384 #[test]
1385 fn test_tokenize_whitespace_only() {
1386 let tokens = tokenize(" \t\n ");
1387 assert!(tokens.is_empty());
1388 }
1389
1390 #[test]
1393 fn test_ast_injection_clean_sql() {
1394 assert!(detect_injection_ast("SELECT id, name FROM users WHERE age > 18").is_ok());
1395 assert!(detect_injection_ast("INSERT INTO users (name) VALUES ('alice')").is_ok());
1396 assert!(detect_injection_ast("UPDATE users SET name = 'bob' WHERE id = 1").is_ok());
1397 }
1398
1399 #[test]
1400 fn test_ast_injection_multi_statement_drop() {
1401 let result = detect_injection_ast("SELECT * FROM users; DROP TABLE users");
1402 assert!(result.is_err());
1403 assert!(matches!(
1404 result.unwrap_err(),
1405 SqlValidationError::InjectionDetected(_)
1406 ));
1407 }
1408
1409 #[test]
1410 fn test_ast_injection_multi_statement_delete() {
1411 let result = detect_injection_ast("SELECT * FROM users; DELETE FROM users");
1412 assert!(result.is_err());
1413 }
1414
1415 #[test]
1416 fn test_ast_injection_multi_statement_insert() {
1417 let result = detect_injection_ast("SELECT 1; INSERT INTO admin VALUES (1, 'hacker')");
1418 assert!(result.is_err());
1419 }
1420
1421 #[test]
1422 fn test_ast_injection_exec_call() {
1423 let result = detect_injection_ast("EXEC sp_executesql 'DROP TABLE users'");
1424 assert!(result.is_err());
1425 let result2 = detect_injection_ast("EXECUTE sp_executesql 'DELETE FROM users'");
1426 assert!(result2.is_err());
1427 }
1428
1429 #[test]
1430 fn test_ast_injection_grant_revoke() {
1431 assert!(detect_injection_ast("GRANT ALL ON users TO hacker").is_err());
1432 assert!(detect_injection_ast("REVOKE SELECT ON users FROM app_user").is_err());
1433 }
1434
1435 #[test]
1436 fn test_ast_injection_boolean_tautology_or_1_eq_1() {
1437 let result = detect_injection_ast("SELECT * FROM users WHERE name = 'a' OR 1=1");
1438 assert!(result.is_err());
1439 }
1440
1441 #[test]
1442 fn test_ast_injection_boolean_tautology_spaces() {
1443 let result = detect_injection_ast("SELECT * FROM users WHERE name = 'a' OR 1 = 1");
1444 assert!(result.is_err());
1445 }
1446
1447 #[test]
1448 fn test_ast_injection_boolean_tautology_true() {
1449 let result = detect_injection_ast("SELECT * FROM users WHERE name = 'a' OR TRUE");
1450 assert!(result.is_err());
1451 }
1452
1453 #[test]
1454 fn test_ast_injection_no_keywords() {
1455 assert!(detect_injection_ast("12345").is_ok());
1457 assert!(detect_injection_ast("").is_ok());
1458 }
1459
1460 #[test]
1461 fn test_ast_injection_safe_semicolon() {
1462 assert!(detect_injection_ast("SELECT * FROM users; BEGIN").is_ok());
1464 }
1465
1466 #[test]
1469 fn test_whitelist_empty_allows_all() {
1470 let validator = WhitelistValidator::new();
1471 assert!(validator.validate("SELECT * FROM any_table").is_ok());
1472 assert!(validator.validate("SELECT * FROM secret_table").is_ok());
1473 }
1474
1475 #[test]
1476 fn test_whitelist_table_allowed() {
1477 let validator = WhitelistValidator::new().allow_table("users");
1478 assert!(validator.validate_tables("SELECT * FROM users").is_ok());
1479 assert!(validator
1480 .validate_tables("SELECT * FROM users WHERE id = 1")
1481 .is_ok());
1482 }
1483
1484 #[test]
1485 fn test_whitelist_table_blocked() {
1486 let validator = WhitelistValidator::new().allow_table("users");
1487 let result = validator.validate_tables("SELECT * FROM secret_table");
1488 assert!(result.is_err());
1489 assert!(matches!(
1490 result.unwrap_err(),
1491 SqlValidationError::InvalidTableName(_)
1492 ));
1493 }
1494
1495 #[test]
1496 fn test_whitelist_multiple_tables() {
1497 let validator = WhitelistValidator::new().allow_tables(&["users", "orders", "products"]);
1498 assert!(validator.validate_tables("SELECT * FROM users").is_ok());
1499 assert!(validator.validate_tables("SELECT * FROM orders").is_ok());
1500 assert!(validator.validate_tables("SELECT * FROM products").is_ok());
1501 assert!(validator
1502 .validate_tables("SELECT * FROM forbidden")
1503 .is_err());
1504 }
1505
1506 #[test]
1507 fn test_whitelist_case_insensitive() {
1508 let validator = WhitelistValidator::new().allow_table("Users");
1509 assert!(validator.validate_tables("SELECT * FROM users").is_ok());
1510 assert!(validator.validate_tables("SELECT * FROM USERS").is_ok());
1511 assert!(validator.validate_tables("SELECT * FROM Users").is_ok());
1512 }
1513
1514 #[test]
1515 fn test_whitelist_join_table() {
1516 let validator = WhitelistValidator::new().allow_tables(&["users", "orders"]);
1517 assert!(validator
1518 .validate_tables("SELECT * FROM users u JOIN orders o ON u.id = o.user_id")
1519 .is_ok());
1520 }
1521
1522 #[test]
1523 fn test_whitelist_join_blocked_table() {
1524 let validator = WhitelistValidator::new().allow_tables(&["users"]);
1525 assert!(validator
1527 .validate_tables("SELECT * FROM users u JOIN orders o ON u.id = o.user_id")
1528 .is_err());
1529 }
1530
1531 #[test]
1532 fn test_whitelist_insert_table() {
1533 let validator = WhitelistValidator::new().allow_table("users");
1534 assert!(validator
1535 .validate_tables("INSERT INTO users (name) VALUES ('a')")
1536 .is_ok());
1537 assert!(validator
1538 .validate_tables("INSERT INTO secret (name) VALUES ('a')")
1539 .is_err());
1540 }
1541
1542 #[test]
1543 fn test_whitelist_update_table() {
1544 let validator = WhitelistValidator::new().allow_table("users");
1545 assert!(validator
1546 .validate_tables("UPDATE users SET name = 'a' WHERE id = 1")
1547 .is_ok());
1548 assert!(validator
1549 .validate_tables("UPDATE admin SET role = 'super' WHERE id = 1")
1550 .is_err());
1551 }
1552
1553 #[test]
1554 fn test_whitelist_columns_not_set_passes() {
1555 let validator = WhitelistValidator::new().allow_table("users");
1556 assert!(validator
1558 .validate_columns("SELECT id, name, password FROM users")
1559 .is_ok());
1560 }
1561
1562 #[test]
1563 fn test_whitelist_combined_validate() {
1564 let validator = WhitelistValidator::new().allow_table("users");
1565 assert!(validator
1566 .validate("SELECT * FROM users WHERE id = 1")
1567 .is_ok());
1568 assert!(validator.validate("SELECT * FROM forbidden").is_err());
1569 }
1570
1571 #[test]
1574 fn test_complexity_simple_query() {
1575 let score = score_complexity("SELECT * FROM users");
1576 assert_eq!(score.level(), ComplexityLevel::Simple);
1577 assert_eq!(score.join_count, 0);
1578 assert_eq!(score.subquery_count, 0);
1579 assert_eq!(score.where_condition_count, 0);
1580 assert!(!score.has_having);
1581 assert!(!score.has_window_function);
1582 assert!(!score.has_cte);
1583 }
1584
1585 #[test]
1586 fn test_complexity_with_where() {
1587 let score = score_complexity("SELECT * FROM users WHERE age > 18 AND status = 'active'");
1588 assert!(score.where_condition_count >= 2);
1589 }
1590
1591 #[test]
1592 fn test_complexity_with_join() {
1593 let score =
1594 score_complexity("SELECT u.id, o.total FROM users u JOIN orders o ON u.id = o.user_id");
1595 assert_eq!(score.join_count, 1);
1596 }
1597
1598 #[test]
1599 fn test_complexity_with_subquery() {
1600 let score = score_complexity(
1601 "SELECT * FROM (SELECT id, name FROM users WHERE age > 18) t WHERE t.id > 0",
1602 );
1603 assert!(score.subquery_count >= 1);
1604 }
1605
1606 #[test]
1607 fn test_complexity_with_group_by_having() {
1608 let score =
1609 score_complexity("SELECT dept, COUNT(*) FROM users GROUP BY dept HAVING COUNT(*) > 5");
1610 assert!(score.group_by_count >= 1);
1611 assert!(score.has_having);
1612 }
1613
1614 #[test]
1615 fn test_complexity_with_cte() {
1616 let score = score_complexity(
1617 "WITH active_users AS (SELECT id FROM users WHERE status = 'active') SELECT * FROM active_users",
1618 );
1619 assert!(score.has_cte);
1620 }
1621
1622 #[test]
1623 fn test_complexity_with_window_function() {
1624 let score = score_complexity(
1625 "SELECT id, ROW_NUMBER() OVER (PARTITION BY dept ORDER BY salary DESC) FROM users",
1626 );
1627 assert!(score.has_window_function);
1628 }
1629
1630 #[test]
1631 fn test_complexity_with_set_operation() {
1632 let score = score_complexity("SELECT id FROM users UNION SELECT id FROM archived_users");
1633 assert!(score.set_operation_count >= 1);
1634 }
1635
1636 #[test]
1637 fn test_complexity_level_thresholds() {
1638 let simple = score_complexity("SELECT * FROM users");
1640 assert_eq!(simple.level(), ComplexityLevel::Simple);
1641
1642 let complex = score_complexity(
1644 "WITH t1 AS (SELECT id FROM users WHERE a = 1 AND b = 2 OR c = 3) \
1645 SELECT t1.id, t2.name, ROW_NUMBER() OVER (PARTITION BY t1.id ORDER BY t2.name) \
1646 FROM t1 JOIN orders t2 ON t1.id = t2.user_id \
1647 GROUP BY t1.id, t2.name HAVING COUNT(*) > 1 \
1648 UNION SELECT id, name, 1 FROM archived",
1649 );
1650 assert!(complex.score > simple.score);
1651 }
1652
1653 #[test]
1654 fn test_complexity_level_descriptions() {
1655 assert_eq!(ComplexityLevel::Simple.description(), "简单");
1656 assert_eq!(ComplexityLevel::Moderate.description(), "中等");
1657 assert_eq!(ComplexityLevel::Complex.description(), "复杂");
1658 assert_eq!(ComplexityLevel::VeryComplex.description(), "非常复杂");
1659 }
1660
1661 #[test]
1662 fn test_complexity_score_max_100() {
1663 let mut sql = String::from("SELECT * FROM users");
1665 for i in 0..20 {
1666 sql.push_str(&format!(" JOIN orders o{} ON users.id = o{}.user_id", i, i));
1667 }
1668 let score = score_complexity(&sql);
1669 assert!(score.score <= 100);
1670 }
1671
1672 #[test]
1673 fn test_complexity_empty_sql() {
1674 let score = score_complexity("");
1675 assert_eq!(score.score, 0);
1676 assert_eq!(score.level(), ComplexityLevel::Simple);
1677 }
1678
1679 #[test]
1682 fn test_ddl_policy_default_all_denied() {
1683 let policy = DdlPolicy::default();
1684 assert!(policy.validate("CREATE TABLE users (id INT)").is_err());
1685 assert!(policy.validate("DROP TABLE users").is_err());
1686 assert!(policy
1687 .validate("ALTER TABLE users ADD COLUMN name TEXT")
1688 .is_err());
1689 assert!(policy.validate("TRUNCATE TABLE users").is_err());
1690 }
1691
1692 #[test]
1693 fn test_ddl_policy_read_only() {
1694 let policy = DdlPolicy::read_only();
1695 assert!(policy.validate("CREATE TABLE users (id INT)").is_err());
1696 assert!(policy.validate("DROP TABLE users").is_err());
1697 assert!(policy.validate("TRUNCATE TABLE users").is_err());
1698 }
1699
1700 #[test]
1701 fn test_ddl_policy_permissive_allows_all() {
1702 let policy = DdlPolicy::permissive();
1703 assert!(policy.validate("CREATE TABLE users (id INT)").is_ok());
1704 assert!(policy.validate("DROP TABLE users").is_ok());
1705 assert!(policy
1706 .validate("ALTER TABLE users ADD COLUMN name TEXT")
1707 .is_ok());
1708 assert!(policy.validate("TRUNCATE TABLE users").is_ok());
1709 }
1710
1711 #[test]
1712 fn test_ddl_policy_safe_evolution() {
1713 let policy = DdlPolicy::safe_evolution();
1714 assert!(policy.validate("CREATE TABLE users (id INT)").is_ok());
1716 assert!(policy
1717 .validate("ALTER TABLE users ADD COLUMN name TEXT")
1718 .is_ok());
1719 assert!(policy.validate("DROP TABLE users").is_err());
1721 assert!(policy.validate("TRUNCATE TABLE users").is_err());
1722 }
1723
1724 #[test]
1725 fn test_ddl_policy_create_index() {
1726 let mut policy = DdlPolicy::permissive();
1727 policy.allow_create_index = false;
1728 assert!(policy
1730 .validate("CREATE INDEX idx_name ON users (name)")
1731 .is_err());
1732 assert!(policy.validate("CREATE TABLE users (id INT)").is_ok());
1734 }
1735
1736 #[test]
1737 fn test_ddl_policy_drop_index() {
1738 let mut policy = DdlPolicy::permissive();
1739 policy.allow_drop_index = false;
1740 assert!(policy.validate("DROP INDEX idx_name").is_err());
1742 assert!(policy.validate("DROP TABLE users").is_ok());
1744 }
1745
1746 #[test]
1747 fn test_ddl_policy_non_ddl_passes() {
1748 let policy = DdlPolicy::read_only();
1749 assert!(policy.validate("SELECT * FROM users").is_ok());
1751 assert!(policy.validate("INSERT INTO users VALUES (1)").is_ok());
1752 assert!(policy.validate("UPDATE users SET name = 'a'").is_ok());
1753 assert!(policy.validate("DELETE FROM users").is_ok());
1754 }
1755
1756 #[test]
1757 fn test_ddl_policy_custom() {
1758 let policy = DdlPolicy {
1759 allow_create: true,
1760 allow_drop: false,
1761 allow_alter: true,
1762 allow_truncate: false,
1763 allow_create_index: true,
1764 allow_drop_index: false,
1765 };
1766 assert!(policy.validate("CREATE TABLE t (id INT)").is_ok());
1767 assert!(policy.validate("DROP TABLE t").is_err());
1768 assert!(policy.validate("ALTER TABLE t ADD COLUMN c INT").is_ok());
1769 assert!(policy.validate("TRUNCATE TABLE t").is_err());
1770 assert!(policy.validate("CREATE INDEX i ON t (c)").is_ok());
1771 assert!(policy.validate("DROP INDEX i").is_err());
1772 }
1773}