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
6use thiserror::Error;
7
8/// SQL validation errors
9#[derive(Error, Debug, Clone, PartialEq)]
10pub enum SqlValidationError {
11    #[error("SQL syntax error: {0}")]
12    SyntaxError(String),
13
14    #[error("Unbalanced parentheses: {0}")]
15    UnbalancedParentheses(String),
16
17    #[error("Unclosed string literal at position {0}")]
18    UnclosedString(usize),
19
20    #[error("Missing required keyword: {0}")]
21    MissingKeyword(String),
22
23    #[error("Invalid parameter count: expected {expected}, got {got}")]
24    ParameterCountMismatch { expected: usize, got: usize },
25
26    #[error("Invalid table name: {0}")]
27    InvalidTableName(String),
28
29    #[error("Empty SELECT columns")]
30    EmptySelectColumns,
31
32    #[error("Empty INSERT data")]
33    EmptyInsertData,
34
35    #[error("Empty UPDATE data")]
36    EmptyUpdateData,
37
38    #[error("DELETE without WHERE clause")]
39    DeleteWithoutWhere,
40
41    #[error("Invalid identifier: {0}")]
42    InvalidIdentifier(String),
43
44    #[error("SQL injection detected: {0}")]
45    InjectionDetected(String),
46}
47
48/// Result type for SQL validation
49pub type ValidationResult = Result<(), SqlValidationError>;
50
51/// SQL statement type detected by the parser
52#[derive(Debug, Clone, Copy, PartialEq)]
53pub enum SqlStatementType {
54    Select,
55    Insert,
56    Update,
57    Delete,
58    Create,
59    Drop,
60    Alter,
61    Truncate,
62    Other,
63}
64
65/// Validate a SQL SELECT statement
66pub fn validate_select(sql: &str) -> ValidationResult {
67    let sql_upper = sql.to_uppercase();
68
69    if !sql_upper.trim_start().starts_with("SELECT") {
70        return Err(SqlValidationError::SyntaxError(
71            "SELECT statement must start with SELECT".to_string(),
72        ));
73    }
74
75    if !sql_upper.contains("FROM") {
76        return Err(SqlValidationError::MissingKeyword("FROM".to_string()));
77    }
78
79    validate_balanced_parentheses(sql)?;
80    validate_string_literals(sql)?;
81    validate_no_injection_patterns(sql)?;
82
83    Ok(())
84}
85
86/// Validate a SQL INSERT statement
87pub fn validate_insert(sql: &str) -> ValidationResult {
88    let sql_upper = sql.to_uppercase();
89
90    if !sql_upper.trim_start().starts_with("INSERT") {
91        return Err(SqlValidationError::SyntaxError(
92            "INSERT statement must start with INSERT".to_string(),
93        ));
94    }
95
96    if !sql_upper.contains("INTO") {
97        return Err(SqlValidationError::MissingKeyword("INTO".to_string()));
98    }
99
100    if !sql_upper.contains("VALUES") {
101        return Err(SqlValidationError::MissingKeyword("VALUES".to_string()));
102    }
103
104    validate_balanced_parentheses(sql)?;
105    validate_string_literals(sql)?;
106    validate_no_injection_patterns(sql)?;
107
108    Ok(())
109}
110
111/// Validate a SQL UPDATE statement
112pub fn validate_update(sql: &str) -> ValidationResult {
113    let sql_upper = sql.to_uppercase();
114
115    if !sql_upper.trim_start().starts_with("UPDATE") {
116        return Err(SqlValidationError::SyntaxError(
117            "UPDATE statement must start with UPDATE".to_string(),
118        ));
119    }
120
121    if !sql_upper.contains("SET") {
122        return Err(SqlValidationError::MissingKeyword("SET".to_string()));
123    }
124
125    validate_balanced_parentheses(sql)?;
126    validate_string_literals(sql)?;
127    validate_no_injection_patterns(sql)?;
128
129    Ok(())
130}
131
132/// Validate a SQL DELETE statement
133pub fn validate_delete(sql: &str) -> ValidationResult {
134    let sql_upper = sql.to_uppercase();
135
136    if !sql_upper.trim_start().starts_with("DELETE") {
137        return Err(SqlValidationError::SyntaxError(
138            "DELETE statement must start with DELETE".to_string(),
139        ));
140    }
141
142    if !sql_upper.contains("FROM") {
143        return Err(SqlValidationError::MissingKeyword("FROM".to_string()));
144    }
145
146    validate_balanced_parentheses(sql)?;
147    validate_string_literals(sql)?;
148    validate_no_injection_patterns(sql)?;
149
150    Ok(())
151}
152
153/// Validate any SQL statement by detecting its type and applying appropriate rules
154pub fn validate_sql(sql: &str) -> ValidationResult {
155    let trimmed = sql.trim();
156    if trimmed.is_empty() {
157        return Err(SqlValidationError::SyntaxError(
158            "Empty SQL statement".to_string(),
159        ));
160    }
161
162    let sql_type = detect_statement_type(trimmed);
163    match sql_type {
164        SqlStatementType::Select => validate_select(trimmed),
165        SqlStatementType::Insert => validate_insert(trimmed),
166        SqlStatementType::Update => validate_update(trimmed),
167        SqlStatementType::Delete => validate_delete(trimmed),
168        _ => {
169            validate_balanced_parentheses(trimmed)?;
170            validate_string_literals(trimmed)?;
171            validate_no_injection_patterns(trimmed)?;
172            Ok(())
173        }
174    }
175}
176
177/// Validate balanced parentheses (including nested levels)
178fn validate_balanced_parentheses(sql: &str) -> ValidationResult {
179    let mut depth: i32 = 0;
180    for (i, ch) in sql.char_indices() {
181        match ch {
182            '(' => depth += 1,
183            ')' => {
184                depth -= 1;
185                if depth < 0 {
186                    return Err(SqlValidationError::UnbalancedParentheses(format!(
187                        "Unexpected ')' at position {}",
188                        i
189                    )));
190                }
191            }
192            _ => {}
193        }
194    }
195    if depth != 0 {
196        return Err(SqlValidationError::UnbalancedParentheses(format!(
197            "{} unclosed '(' parentheses",
198            depth
199        )));
200    }
201    Ok(())
202}
203
204/// Validate string literals are properly closed
205fn validate_string_literals(sql: &str) -> ValidationResult {
206    let mut in_single_quote = false;
207    let mut in_double_quote = false;
208    let mut prev_ch = '\0';
209
210    for (_i, ch) in sql.char_indices() {
211        if prev_ch == '\\' {
212            prev_ch = ch;
213            continue;
214        }
215
216        match ch {
217            '\'' if !in_double_quote => {
218                in_single_quote = !in_single_quote;
219            }
220            '"' if !in_single_quote => {
221                in_double_quote = !in_double_quote;
222            }
223            _ => {}
224        }
225        prev_ch = ch;
226    }
227
228    if in_single_quote {
229        return Err(SqlValidationError::UnclosedString(sql.len()));
230    }
231    if in_double_quote {
232        return Err(SqlValidationError::UnclosedString(sql.len()));
233    }
234
235    Ok(())
236}
237
238/// Validate no obvious SQL injection patterns
239fn validate_no_injection_patterns(sql: &str) -> ValidationResult {
240    let sql_upper = sql.to_uppercase();
241
242    // Check for suspicious patterns
243    let suspicious_patterns = [
244        ("'; DROP TABLE", "DROP TABLE injection"),
245        ("' OR '1'='1", "classic OR injection"),
246        ("' OR 1=1", "OR 1=1 injection"),
247        (
248            "UNION SELECT",
249            "UNION SELECT injection (not allowed in simple queries)",
250        ),
251        ("--", "comment injection (not allowed)"),
252        ("/*", "block comment (not allowed)"),
253    ];
254
255    for (pattern, desc) in &suspicious_patterns {
256        if sql_upper.contains(pattern) {
257            return Err(SqlValidationError::InjectionDetected(format!(
258                "{}: {}",
259                desc, pattern
260            )));
261        }
262    }
263
264    Ok(())
265}
266
267/// Detect the type of SQL statement
268pub fn detect_statement_type(sql: &str) -> SqlStatementType {
269    let trimmed = sql.trim().to_uppercase();
270
271    if trimmed.starts_with("SELECT") {
272        SqlStatementType::Select
273    } else if trimmed.starts_with("INSERT") {
274        SqlStatementType::Insert
275    } else if trimmed.starts_with("UPDATE") {
276        SqlStatementType::Update
277    } else if trimmed.starts_with("DELETE") {
278        SqlStatementType::Delete
279    } else if trimmed.starts_with("CREATE") {
280        SqlStatementType::Create
281    } else if trimmed.starts_with("DROP") {
282        SqlStatementType::Drop
283    } else if trimmed.starts_with("ALTER") {
284        SqlStatementType::Alter
285    } else if trimmed.starts_with("TRUNCATE") {
286        SqlStatementType::Truncate
287    } else {
288        SqlStatementType::Other
289    }
290}
291
292/// Validate parameter count in prepared statements
293pub fn validate_parameter_count(sql: &str, expected_params: usize) -> ValidationResult {
294    let param_count = sql.chars().filter(|&c| c == '?').count() + sql.matches('$').count(); // PostgreSQL style
295    if param_count != expected_params {
296        return Err(SqlValidationError::ParameterCountMismatch {
297            expected: expected_params,
298            got: param_count,
299        });
300    }
301    Ok(())
302}
303
304/// Validate table name is a valid SQL identifier
305pub fn validate_table_name(name: &str) -> ValidationResult {
306    if name.is_empty() {
307        return Err(SqlValidationError::InvalidTableName(
308            "empty table name".to_string(),
309        ));
310    }
311
312    // Table names should only contain alphanumeric, underscore
313    // Allow backtick-quoted identifiers
314    let cleaned = name.trim_matches('`').trim_matches('"').trim_matches('\'');
315    if cleaned.is_empty() {
316        return Err(SqlValidationError::InvalidTableName(name.to_string()));
317    }
318
319    for ch in cleaned.chars() {
320        if !ch.is_alphanumeric() && ch != '_' {
321            return Err(SqlValidationError::InvalidTableName(format!(
322                "table name '{}' contains invalid character '{}'",
323                name, ch
324            )));
325        }
326    }
327
328    Ok(())
329}
330
331/// Validate column name is a valid SQL identifier
332pub fn validate_column_name(name: &str) -> ValidationResult {
333    if name.is_empty() || name == "*" {
334        return Ok(()); // * is valid for SELECT
335    }
336
337    let cleaned = name.trim_matches('`').trim_matches('"').trim_matches('\'');
338    if cleaned.is_empty() {
339        return Err(SqlValidationError::InvalidIdentifier(name.to_string()));
340    }
341
342    // Allow alphanumeric, underscore, dot (for table.column)
343    for ch in cleaned.chars() {
344        if !ch.is_alphanumeric() && ch != '_' && ch != '.' {
345            return Err(SqlValidationError::InvalidIdentifier(format!(
346                "column '{}' contains invalid character '{}'",
347                name, ch
348            )));
349        }
350    }
351
352    Ok(())
353}
354
355/// Entry-level validation that runs all checks
356pub fn validate(sql: &str) -> ValidationResult {
357    if sql.trim().is_empty() {
358        return Err(SqlValidationError::SyntaxError(
359            "Empty SQL statement".to_string(),
360        ));
361    }
362
363    validate_sql(sql)?;
364    validate_balanced_parentheses(sql)?;
365    validate_string_literals(sql)?;
366    validate_no_injection_patterns(sql)?;
367
368    Ok(())
369}
370
371#[cfg(test)]
372mod tests {
373    use super::*;
374
375    #[test]
376    fn test_validate_select_basic() {
377        assert!(validate_select("SELECT * FROM users").is_ok());
378        assert!(validate_select("SELECT id, name FROM users WHERE id = 1").is_ok());
379        assert!(validate_select(
380            "SELECT u.id, u.name FROM users u INNER JOIN orders o ON u.id = o.user_id"
381        )
382        .is_ok());
383    }
384
385    #[test]
386    fn test_validate_select_missing_from() {
387        let result = validate_select("SELECT *");
388        assert!(result.is_err());
389    }
390
391    #[test]
392    fn test_validate_insert_basic() {
393        assert!(validate_insert("INSERT INTO users (name) VALUES ('alice')").is_ok());
394        assert!(validate_insert("INSERT INTO users (name, age) VALUES ('bob', 25)").is_ok());
395    }
396
397    #[test]
398    fn test_validate_insert_missing_values() {
399        let result = validate_insert("INSERT INTO users (name)");
400        assert!(result.is_err());
401    }
402
403    #[test]
404    fn test_validate_update_basic() {
405        assert!(validate_update("UPDATE users SET name = 'alice' WHERE id = 1").is_ok());
406    }
407
408    #[test]
409    fn test_validate_update_missing_set() {
410        let result = validate_update("UPDATE users WHERE id = 1");
411        assert!(result.is_err());
412    }
413
414    #[test]
415    fn test_validate_delete_basic() {
416        assert!(validate_delete("DELETE FROM users WHERE id = 1").is_ok());
417    }
418
419    #[test]
420    fn test_validate_delete_missing_from() {
421        let result = validate_delete("DELETE users");
422        assert!(result.is_err());
423    }
424
425    #[test]
426    fn test_balanced_parentheses() {
427        assert!(validate_balanced_parentheses("SELECT * FROM (SELECT * FROM users) t").is_ok());
428        assert!(validate_balanced_parentheses("FUNC(a, b, c)").is_ok());
429        assert!(
430            validate_balanced_parentheses("SELECT * FROM users WHERE (a=1 AND (b=2 OR c=3))")
431                .is_ok()
432        );
433    }
434
435    #[test]
436    fn test_unbalanced_parentheses() {
437        assert!(validate_balanced_parentheses("SELECT * FROM (SELECT * FROM users").is_err());
438        assert!(validate_balanced_parentheses("SELECT * FROM users)").is_err());
439    }
440
441    #[test]
442    fn test_string_literals_closed() {
443        assert!(validate_string_literals("SELECT * FROM users WHERE name = 'alice'").is_ok());
444        assert!(validate_string_literals("INSERT INTO users (name) VALUES ('bob')").is_ok());
445    }
446
447    #[test]
448    fn test_unclosed_string_literal() {
449        assert!(validate_string_literals("SELECT * FROM users WHERE name = 'alice").is_err());
450    }
451
452    #[test]
453    fn test_injection_detection() {
454        assert!(validate_no_injection_patterns("SELECT * FROM users WHERE name = 'alice'").is_ok());
455        assert!(validate_no_injection_patterns(
456            "SELECT * FROM users WHERE name = 'alice' OR '1'='1'"
457        )
458        .is_err());
459        assert!(validate_no_injection_patterns("'; DROP TABLE users; --").is_err());
460        assert!(validate_no_injection_patterns("1 UNION SELECT * FROM users").is_err());
461    }
462
463    #[test]
464    fn test_detect_statement_type() {
465        assert_eq!(
466            detect_statement_type("SELECT * FROM users"),
467            SqlStatementType::Select
468        );
469        assert_eq!(
470            detect_statement_type("INSERT INTO users VALUES (1)"),
471            SqlStatementType::Insert
472        );
473        assert_eq!(
474            detect_statement_type("UPDATE users SET a=1"),
475            SqlStatementType::Update
476        );
477        assert_eq!(
478            detect_statement_type("DELETE FROM users"),
479            SqlStatementType::Delete
480        );
481        assert_eq!(
482            detect_statement_type("CREATE TABLE users"),
483            SqlStatementType::Create
484        );
485        assert_eq!(
486            detect_statement_type("DROP TABLE users"),
487            SqlStatementType::Drop
488        );
489        assert_eq!(
490            detect_statement_type("ALTER TABLE users ADD COLUMN a"),
491            SqlStatementType::Alter
492        );
493        assert_eq!(
494            detect_statement_type("TRUNCATE TABLE users"),
495            SqlStatementType::Truncate
496        );
497        assert_eq!(
498            detect_statement_type("EXPLAIN SELECT * FROM users"),
499            SqlStatementType::Other
500        );
501    }
502
503    #[test]
504    fn test_parameter_count() {
505        assert!(
506            validate_parameter_count("SELECT * FROM users WHERE id = ? AND name = ?", 2).is_ok()
507        );
508        assert!(validate_parameter_count("SELECT * FROM users WHERE id = ?", 1).is_ok());
509        assert!(validate_parameter_count("SELECT * FROM users WHERE id = ?", 2).is_err());
510    }
511
512    #[test]
513    fn test_validate_table_name() {
514        assert!(validate_table_name("users").is_ok());
515        assert!(validate_table_name("user_orders").is_ok());
516        assert!(validate_table_name("").is_err());
517        assert!(validate_table_name("users; DROP TABLE").is_err());
518    }
519
520    #[test]
521    fn test_validate_column_name() {
522        assert!(validate_column_name("id").is_ok());
523        assert!(validate_column_name("*").is_ok());
524        assert!(validate_column_name("users.name").is_ok());
525        assert!(validate_column_name("").is_ok()); // * replacement
526    }
527
528    #[test]
529    fn test_validate_empty_sql() {
530        assert!(validate("").is_err());
531        assert!(validate("   ").is_err());
532    }
533
534    #[test]
535    fn test_validate_complex_queries() {
536        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());
537    }
538
539    #[test]
540    fn test_empty_insert_data() {
541        let sql = "INSERT INTO users () VALUES ()";
542        assert!(validate_sql(sql).is_ok());
543    }
544
545    #[test]
546    fn test_create_table_validation() {
547        assert!(validate_sql("CREATE TABLE users (id INT PRIMARY KEY, name VARCHAR(100))").is_ok());
548    }
549
550    #[test]
551    fn test_double_quoted_identifiers() {
552        assert!(
553            validate_string_literals("SELECT * FROM \"users\" WHERE \"name\" = 'alice'").is_ok()
554        );
555    }
556
557    #[test]
558    fn test_nested_function_calls() {
559        assert!(validate_balanced_parentheses(
560            "SELECT MAX(COUNT(*)) FROM (SELECT COUNT(*) FROM users GROUP BY status) t"
561        )
562        .is_ok());
563    }
564}