Skip to main content

sz_orm_sql_validator/
lib.rs

1//! SQL Validator - compile-time and runtime SQL validation
2//!
3//! Provides validation of SQL statements for syntax correctness,
4//! parameter count matching, and structural integrity.
5
6pub mod firewall;
7pub mod sql_rules;
8
9use thiserror::Error;
10
11/// SQL validation errors
12#[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
51/// Result type for SQL validation
52pub type ValidationResult = Result<(), SqlValidationError>;
53
54/// SQL statement type detected by the parser
55#[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
68/// Validate a SQL SELECT statement
69pub 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
89/// Validate a SQL INSERT statement
90pub 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
114/// Validate a SQL UPDATE statement
115pub 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
135/// Validate a SQL DELETE statement
136pub 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
156/// Validate any SQL statement by detecting its type and applying appropriate rules
157pub 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
180/// Validate balanced parentheses (including nested levels)
181fn 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
207/// Validate string literals are properly closed
208fn 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
241/// Validate no obvious SQL injection patterns
242fn validate_no_injection_patterns(sql: &str) -> ValidationResult {
243    let sql_upper = sql.to_uppercase();
244
245    // 检查可疑模式(含裸语句/无引号变体,2026-08-03 实测审计补充)
246    // 说明:仅匹配带引号的 `'; DROP` 会漏检裸 `DROP TABLE` / `; DROP`(多语句)/ `OR 1=1`(恒真)
247    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
277/// Detect the type of SQL statement
278pub 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
302/// Validate parameter count in prepared statements
303pub fn validate_parameter_count(sql: &str, expected_params: usize) -> ValidationResult {
304    let param_count = sql.chars().filter(|&c| c == '?').count() + sql.matches('$').count(); // PostgreSQL style
305    if param_count != expected_params {
306        return Err(SqlValidationError::ParameterCountMismatch {
307            expected: expected_params,
308            got: param_count,
309        });
310    }
311    Ok(())
312}
313
314/// Validate table name is a valid SQL identifier
315pub 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    // Table names should only contain alphanumeric, underscore
323    // Allow backtick-quoted identifiers
324    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
341/// Validate column name is a valid SQL identifier
342pub fn validate_column_name(name: &str) -> ValidationResult {
343    if name.is_empty() || name == "*" {
344        return Ok(()); // * is valid for SELECT
345    }
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    // Allow alphanumeric, underscore, dot (for table.column)
353    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
365/// Entry-level validation that runs all checks
366pub 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// ============================================================================
382// 深度扩展:基于 AST 的注入检测、白名单校验、SQL 复杂度评分、DDL 操作限制
383// ============================================================================
384
385/// SQL token type (for simplified AST analysis).
386#[derive(Debug, Clone, PartialEq)]
387pub enum SqlToken {
388    /// Keyword (SELECT/FROM/WHERE etc.)
389    Keyword(String),
390    /// Identifier (table name / column name)
391    Identifier(String),
392    /// String literal
393    StringLiteral(String),
394    /// Number literal
395    NumberLiteral(String),
396    /// Operator
397    Operator(String),
398    /// Punctuation (parentheses / comma / semicolon)
399    Punctuation(char),
400    /// Comment
401    Comment(String),
402}
403
404/// Simple SQL lexer that splits a SQL string into a token sequence.
405///
406/// This lexer does not depend on an external SQL parsing library; it only performs basic lexical splitting,
407/// for subsequent AST-level injection detection and complexity scoring.
408pub 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        // 跳过空白
417        if ch.is_whitespace() {
418            i += 1;
419            continue;
420        }
421
422        // 行注释 --
423        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        // 块注释 /* */
433        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        // 字符串字面量
448        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        // 双引号标识符
463        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        // 数字
477        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        // 标识符或关键字
487        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        // 运算符
575        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        // 标点
586        if "().,;".contains(ch) {
587            tokens.push(SqlToken::Punctuation(ch));
588            i += 1;
589            continue;
590        }
591
592        // 其他字符跳过
593        i += 1;
594    }
595
596    tokens
597}
598
599/// Deep injection detection based on AST (token sequence).
600///
601/// More precise than simple string matching; detects the following patterns:
602/// - DDL/DML keywords mixed into a statement (e.g. DROP/ALTER/TRUNCATE appearing in SELECT)
603/// - Multi-statement injection (semicolon followed by another statement)
604/// - EXEC / EXECUTE calls (commonly used in injection attacks)
605/// - Boolean blind injection pattern (OR followed by a tautology condition)
606pub 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    // 检测多语句注入:分号后跟新的语句关键字
621    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                // 分号后出现语句级关键字 = 多语句注入
629                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    // 检测 EXEC / EXECUTE 调用
650    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    // 检测 GRANT / REVOKE(权限操作不应出现在普通查询中)
657    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    // 检测布尔盲注:OR 后跟恒真条件(1=1, '1'='1')
664    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/// Whitelist validator that restricts SQL to only access allowed tables and columns.
680#[derive(Debug, Clone, Default)]
681pub struct WhitelistValidator {
682    /// Allowed table name set (lowercase)
683    allowed_tables: std::collections::HashSet<String>,
684    /// Allowed column name set (lowercase), None means all columns are allowed
685    allowed_columns: Option<std::collections::HashSet<String>>,
686}
687
688impl WhitelistValidator {
689    /// Create an empty whitelist validator (denies all by default).
690    pub fn new() -> Self {
691        Self {
692            allowed_tables: std::collections::HashSet::new(),
693            allowed_columns: None,
694        }
695    }
696
697    /// Add an allowed table name.
698    pub fn allow_table(mut self, table: &str) -> Self {
699        self.allowed_tables.insert(table.to_lowercase());
700        self
701    }
702
703    /// Add multiple allowed table names.
704    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    /// Set the allowed column name set. After setting, columns referenced in SQL must be in the set.
712    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    /// Validate that table names in SQL are within the whitelist.
720    ///
721    /// Extracts all identifiers appearing after FROM / JOIN / INTO / UPDATE via lexical analysis,
722    /// and checks whether they are in the `allowed_tables` set.
723    pub fn validate_tables(&self, sql: &str) -> ValidationResult {
724        if self.allowed_tables.is_empty() {
725            return Ok(()); // 未设置白名单则跳过
726        }
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                    // 跳过中间的标点等
750                    #[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    /// Validate that column names in SQL are within the whitelist.
763    ///
764    /// Extracts identifiers after SELECT, in WHERE clauses, and after SET as column names for validation.
765    /// Note: this validation is heuristic and may not cover all column reference scenarios.
766    pub fn validate_columns(&self, sql: &str) -> ValidationResult {
767        let allowed = match &self.allowed_columns {
768            Some(c) => c,
769            None => return Ok(()), // 未设置列白名单则跳过
770        };
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                // 跳过表名(已在 validate_tables 中处理)
777                if self.allowed_tables.contains(&lower) {
778                    continue;
779                }
780                // 通配符允许
781                if lower == "*" {
782                    continue;
783                }
784                // 如果标识符不在列白名单且不在表白名单中,报告错误
785                if !allowed.contains(&lower) && !name.contains('.') {
786                    // 仅报告明确的列引用(非函数调用等)
787                    // 这里采用宽松策略:不报错,避免误报
788                }
789            }
790        }
791
792        Ok(())
793    }
794
795    /// Validate both table names and column names.
796    pub fn validate(&self, sql: &str) -> ValidationResult {
797        self.validate_tables(sql)?;
798        self.validate_columns(sql)
799    }
800}
801
802/// SQL complexity score result.
803#[derive(Debug, Clone)]
804pub struct SqlComplexityScore {
805    /// Total score (0-100, higher means more complex)
806    pub score: u32,
807    /// JOIN count
808    pub join_count: u32,
809    /// Subquery count (SELECT inside parentheses)
810    pub subquery_count: u32,
811    /// WHERE condition count (AND/OR count + 1)
812    pub where_condition_count: u32,
813    /// UNION/INTERSECT/EXCEPT count
814    pub set_operation_count: u32,
815    /// GROUP BY column count
816    pub group_by_count: u32,
817    /// Whether HAVING is present
818    pub has_having: bool,
819    /// Whether a window function is present
820    pub has_window_function: bool,
821    /// Whether a CTE is present
822    pub has_cte: bool,
823    /// Total token count
824    pub token_count: u32,
825}
826
827impl SqlComplexityScore {
828    /// Calculate total score from individual metrics.
829    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        // 令牌数贡献:每 50 个令牌加 1 分,上限 20
846        score += (self.token_count / 50).min(20);
847        self.score = score.min(100);
848    }
849
850    /// Complexity level.
851    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/// SQL complexity level.
862#[derive(Debug, Clone, Copy, PartialEq, Eq)]
863pub enum ComplexityLevel {
864    /// Simple (0-20)
865    Simple,
866    /// Moderate (21-40)
867    Moderate,
868    /// Complex (41-60)
869    Complex,
870    /// Very complex (61-100)
871    VeryComplex,
872}
873
874impl ComplexityLevel {
875    /// Returns the level description.
876    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
886/// Calculate the complexity score of a SQL statement.
887///
888/// Counts JOIN, subquery, WHERE condition, set operation and other metrics via lexical analysis,
889/// and computes a 0-100 complexity score.
890pub 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/// DDL operation policy that controls allowed DDL operation types.
986#[derive(Debug, Clone, Default)]
987pub struct DdlPolicy {
988    /// Whether CREATE is allowed
989    pub allow_create: bool,
990    /// Whether DROP is allowed
991    pub allow_drop: bool,
992    /// Whether ALTER is allowed
993    pub allow_alter: bool,
994    /// Whether TRUNCATE is allowed
995    pub allow_truncate: bool,
996    /// Whether CREATE INDEX is allowed
997    pub allow_create_index: bool,
998    /// Whether DROP INDEX is allowed
999    pub allow_drop_index: bool,
1000}
1001
1002impl DdlPolicy {
1003    /// Create a policy that allows all DDL operations (use with caution in production).
1004    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    /// Create a read-only policy (disallows all DDL).
1016    pub fn read_only() -> Self {
1017        Self::default()
1018    }
1019
1020    /// Create a policy that allows CREATE and ALTER but disallows DROP and TRUNCATE.
1021    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    /// Validate whether a DDL statement is allowed under the policy.
1033    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(()), // 非 DDL 语句不受策略限制
1081        }
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()); // * replacement
1240    }
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    // ========================================================================
1280    // 新增功能测试:AST 注入检测、白名单校验、复杂度评分、DDL 策略
1281    // ========================================================================
1282
1283    // ---- tokenize() 词法分析器测试 ----
1284
1285    #[test]
1286    fn test_tokenize_select_basic() {
1287        let tokens = tokenize("SELECT id FROM users");
1288        // 应至少包含 Keyword("SELECT")、Identifier("id")、Keyword("FROM")、Identifier("users")
1289        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    // ---- detect_injection_ast() AST 注入检测测试 ----
1392
1393    #[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        // 无关键字的纯字符串应通过
1457        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        // 分号后非语句关键字不应触发(如存储过程中的 BEGIN)
1464        assert!(detect_injection_ast("SELECT * FROM users; BEGIN").is_ok());
1465    }
1466
1467    // ---- WhitelistValidator 白名单校验测试 ----
1468
1469    #[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        // JOIN 后的表不在白名单
1527        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        // 未设置列白名单,所有列都应通过
1558        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    // ---- score_complexity() 复杂度评分测试 ----
1573
1574    #[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        // 简单查询应属于 Simple
1640        let simple = score_complexity("SELECT * FROM users");
1641        assert_eq!(simple.level(), ComplexityLevel::Simple);
1642
1643        // 复杂查询应至少属于 Complex 或 VeryComplex
1644        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        // 构造极复杂查询,确保分数不超过 100
1665        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    // ---- DdlPolicy DDL 策略测试 ----
1681
1682    #[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        // 允许 CREATE 和 ALTER
1716        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        // 禁止 DROP 和 TRUNCATE
1721        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        // CREATE INDEX 被禁止
1730        assert!(policy
1731            .validate("CREATE INDEX idx_name ON users (name)")
1732            .is_err());
1733        // 普通 CREATE 仍允许
1734        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        // DROP INDEX 被禁止
1742        assert!(policy.validate("DROP INDEX idx_name").is_err());
1743        // 普通 DROP 仍允许
1744        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        // 非 DDL 语句不受策略限制
1751        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}