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