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