1use thiserror::Error;
7
8#[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
48pub type ValidationResult = Result<(), SqlValidationError>;
50
51#[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
65pub 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
86pub 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
111pub 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
132pub 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
153pub 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
177fn 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
204fn 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
238fn validate_no_injection_patterns(sql: &str) -> ValidationResult {
240 let sql_upper = sql.to_uppercase();
241
242 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
267pub 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
292pub fn validate_parameter_count(sql: &str, expected_params: usize) -> ValidationResult {
294 let param_count = sql.chars().filter(|&c| c == '?').count() + sql.matches('$').count(); if param_count != expected_params {
296 return Err(SqlValidationError::ParameterCountMismatch {
297 expected: expected_params,
298 got: param_count,
299 });
300 }
301 Ok(())
302}
303
304pub 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 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
331pub fn validate_column_name(name: &str) -> ValidationResult {
333 if name.is_empty() || name == "*" {
334 return Ok(()); }
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 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
355pub 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()); }
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}