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