1use crate::error::QueryError;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10pub enum StatementType {
11 Select,
13 Insert,
15 Update,
17 Delete,
19 Ddl,
21 Transaction,
23 Other,
25}
26
27impl StatementType {
28 pub fn from_sql(sql: &str) -> Self {
30 let trimmed = sql.trim_start().to_uppercase();
31
32 if trimmed.starts_with("SELECT") || trimmed.starts_with("WITH") {
33 Self::Select
34 } else if trimmed.starts_with("INSERT") {
35 Self::Insert
36 } else if trimmed.starts_with("UPDATE") {
37 Self::Update
38 } else if trimmed.starts_with("DELETE") {
39 Self::Delete
40 } else if trimmed.starts_with("CREATE")
41 || trimmed.starts_with("ALTER")
42 || trimmed.starts_with("DROP")
43 || trimmed.starts_with("TRUNCATE")
44 {
45 Self::Ddl
46 } else if trimmed.starts_with("BEGIN")
47 || trimmed.starts_with("COMMIT")
48 || trimmed.starts_with("ROLLBACK")
49 {
50 Self::Transaction
51 } else {
52 Self::Other
53 }
54 }
55
56 pub fn returns_result_set(&self) -> bool {
58 matches!(self, Self::Select)
59 }
60
61 pub fn returns_row_count(&self) -> bool {
63 matches!(self, Self::Insert | Self::Update | Self::Delete)
64 }
65}
66
67#[derive(Debug, Clone)]
69pub enum Parameter {
70 Null,
72 Boolean(bool),
74 Integer(i64),
76 Float(f64),
78 String(String),
80 Binary(Vec<u8>),
82}
83
84impl Parameter {
85 pub fn to_sql_literal(&self) -> Result<String, QueryError> {
90 match self {
91 Parameter::Null => Ok("NULL".to_string()),
92 Parameter::Boolean(b) => Ok(if *b { "TRUE" } else { "FALSE" }.to_string()),
93 Parameter::Integer(i) => Ok(i.to_string()),
94 Parameter::Float(f) => {
95 if f.is_nan() || f.is_infinite() {
96 Err(QueryError::ParameterBindingError {
97 index: 0,
98 message: "NaN and Infinity are not supported".to_string(),
99 })
100 } else {
101 Ok(f.to_string())
102 }
103 }
104 Parameter::String(s) => {
105 if Self::contains_sql_injection_pattern(s) {
107 return Err(QueryError::SqlInjectionDetected);
108 }
109
110 let escaped = s.replace('\'', "''");
112
113 Ok(format!("'{}'", escaped))
114 }
115 Parameter::Binary(b) => {
116 Ok(format!("'{}'", hex::encode(b)))
118 }
119 }
120 }
121
122 fn contains_sql_injection_pattern(s: &str) -> bool {
124 let upper = s.to_uppercase();
125
126 let patterns = [
128 "'; DROP",
129 "'; DELETE",
130 "'; UPDATE",
131 "'; INSERT",
132 "' OR '1'='1",
133 "' OR 1=1",
134 "' OR TRUE",
135 "UNION SELECT",
136 "EXEC(",
137 "EXECUTE(",
138 ];
139
140 patterns.iter().any(|pattern| upper.contains(pattern))
141 }
142}
143
144impl From<bool> for Parameter {
145 fn from(value: bool) -> Self {
146 Parameter::Boolean(value)
147 }
148}
149
150impl From<i32> for Parameter {
151 fn from(value: i32) -> Self {
152 Parameter::Integer(value as i64)
153 }
154}
155
156impl From<i64> for Parameter {
157 fn from(value: i64) -> Self {
158 Parameter::Integer(value)
159 }
160}
161
162impl From<f64> for Parameter {
163 fn from(value: f64) -> Self {
164 Parameter::Float(value)
165 }
166}
167
168impl From<String> for Parameter {
169 fn from(value: String) -> Self {
170 Parameter::String(value)
171 }
172}
173
174impl From<&str> for Parameter {
175 fn from(value: &str) -> Self {
176 Parameter::String(value.to_string())
177 }
178}
179
180impl From<Vec<u8>> for Parameter {
181 fn from(value: Vec<u8>) -> Self {
182 Parameter::Binary(value)
183 }
184}
185
186#[derive(Debug, Clone, Copy, PartialEq, Eq)]
188enum ScanState {
189 Normal,
190 SingleQuoted,
191 DoubleQuoted,
192 LineComment,
193 BlockComment,
194}
195
196fn scan_placeholders(sql: &str) -> Vec<usize> {
209 let mut positions = Vec::new();
210 let mut state = ScanState::Normal;
211 let bytes = sql.as_bytes();
212 let mut iter = sql.char_indices();
213
214 while let Some((i, c)) = iter.next() {
215 match state {
216 ScanState::Normal => match c {
217 '?' => positions.push(i),
218 '\'' => state = ScanState::SingleQuoted,
219 '"' => state = ScanState::DoubleQuoted,
220 '-' if bytes.get(i + 1) == Some(&b'-') => {
221 iter.next();
222 state = ScanState::LineComment;
223 }
224 '/' if bytes.get(i + 1) == Some(&b'*') => {
225 iter.next();
226 state = ScanState::BlockComment;
227 }
228 _ => {}
229 },
230 ScanState::SingleQuoted => {
231 if c == '\'' {
232 if bytes.get(i + 1) == Some(&b'\'') {
233 iter.next();
234 } else {
235 state = ScanState::Normal;
236 }
237 }
238 }
239 ScanState::DoubleQuoted => {
240 if c == '"' {
241 if bytes.get(i + 1) == Some(&b'"') {
242 iter.next();
243 } else {
244 state = ScanState::Normal;
245 }
246 }
247 }
248 ScanState::LineComment => {
249 if c == '\n' {
250 state = ScanState::Normal;
251 }
252 }
253 ScanState::BlockComment => {
254 if c == '*' && bytes.get(i + 1) == Some(&b'/') {
255 iter.next();
256 state = ScanState::Normal;
257 }
258 }
259 }
260 }
261
262 positions
263}
264
265pub struct Statement {
273 sql: String,
275 parameters: Vec<Option<Parameter>>,
277 timeout_ms: Option<u64>,
280 statement_type: StatementType,
282}
283
284impl Statement {
285 pub fn new(sql: impl Into<String>) -> Self {
287 let sql = sql.into();
288 let statement_type = StatementType::from_sql(&sql);
289
290 Self {
291 sql,
292 parameters: Vec::new(),
293 timeout_ms: None,
294 statement_type,
295 }
296 }
297
298 pub fn sql(&self) -> &str {
300 &self.sql
301 }
302
303 pub fn statement_type(&self) -> StatementType {
305 self.statement_type
306 }
307
308 pub fn timeout_ms(&self) -> Option<u64> {
310 self.timeout_ms
311 }
312
313 pub fn set_timeout(&mut self, timeout_ms: u64) {
315 self.timeout_ms = Some(timeout_ms);
316 }
317
318 pub fn bind<T: Into<Parameter>>(&mut self, index: usize, value: T) -> Result<(), QueryError> {
327 if index >= self.parameters.len() {
329 self.parameters.resize(index + 1, None);
330 }
331
332 self.parameters[index] = Some(value.into());
333 Ok(())
334 }
335
336 pub fn bind_all<T: Into<Parameter> + Clone>(&mut self, params: &[T]) -> Result<(), QueryError> {
338 for (index, param) in params.iter().enumerate() {
339 self.bind(index, param.clone())?;
340 }
341 Ok(())
342 }
343
344 pub fn clear_parameters(&mut self) {
346 self.parameters.clear();
347 }
348
349 pub fn parameters(&self) -> &[Option<Parameter>] {
351 &self.parameters
352 }
353
354 pub fn build_sql(&self) -> Result<String, QueryError> {
360 let positions = scan_placeholders(&self.sql);
361
362 if positions.len() > self.parameters.len() {
363 return Err(QueryError::ParameterBindingError {
364 index: self.parameters.len(),
365 message: "Not enough parameters bound".to_string(),
366 });
367 }
368
369 let mut sql = self.sql.clone();
370
371 for (param_index, &pos) in positions.iter().enumerate().rev() {
372 let param = self.parameters[param_index].as_ref().ok_or_else(|| {
373 QueryError::ParameterBindingError {
374 index: param_index,
375 message: "Parameter not bound".to_string(),
376 }
377 })?;
378
379 let literal = param.to_sql_literal()?;
380 sql.replace_range(pos..pos + 1, &literal);
381 }
382
383 Ok(sql)
384 }
385}
386
387impl std::fmt::Debug for Statement {
388 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
389 f.debug_struct("Statement")
390 .field("sql", &self.sql)
391 .field("statement_type", &self.statement_type)
392 .field("timeout_ms", &self.timeout_ms)
393 .finish()
394 }
395}
396
397impl std::fmt::Display for Statement {
398 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
399 write!(f, "Statement({})", self.sql)
400 }
401}
402
403#[cfg(test)]
404#[allow(clippy::approx_constant)]
405mod tests {
406 use super::*;
407
408 #[test]
409 fn test_statement_type_detection() {
410 assert_eq!(
411 StatementType::from_sql("SELECT * FROM users"),
412 StatementType::Select
413 );
414 assert_eq!(
415 StatementType::from_sql(" select * from users"),
416 StatementType::Select
417 );
418 assert_eq!(
419 StatementType::from_sql("WITH cte AS (SELECT 1) SELECT * FROM cte"),
420 StatementType::Select
421 );
422 assert_eq!(
423 StatementType::from_sql("INSERT INTO users VALUES (1)"),
424 StatementType::Insert
425 );
426 assert_eq!(
427 StatementType::from_sql("UPDATE users SET name = 'John'"),
428 StatementType::Update
429 );
430 assert_eq!(
431 StatementType::from_sql("DELETE FROM users WHERE id = 1"),
432 StatementType::Delete
433 );
434 assert_eq!(
435 StatementType::from_sql("CREATE TABLE test (id INT)"),
436 StatementType::Ddl
437 );
438 assert_eq!(
439 StatementType::from_sql("DROP TABLE test"),
440 StatementType::Ddl
441 );
442 assert_eq!(StatementType::from_sql("BEGIN"), StatementType::Transaction);
443 assert_eq!(
444 StatementType::from_sql("COMMIT"),
445 StatementType::Transaction
446 );
447 assert_eq!(
448 StatementType::from_sql("ROLLBACK"),
449 StatementType::Transaction
450 );
451 }
452
453 #[test]
454 fn test_statement_type_returns_result_set() {
455 assert!(StatementType::Select.returns_result_set());
456 assert!(!StatementType::Insert.returns_result_set());
457 assert!(!StatementType::Update.returns_result_set());
458 assert!(!StatementType::Delete.returns_result_set());
459 }
460
461 #[test]
462 fn test_parameter_to_sql_literal() {
463 assert_eq!(Parameter::Null.to_sql_literal().unwrap(), "NULL");
464 assert_eq!(Parameter::Boolean(true).to_sql_literal().unwrap(), "TRUE");
465 assert_eq!(Parameter::Boolean(false).to_sql_literal().unwrap(), "FALSE");
466 assert_eq!(Parameter::Integer(42).to_sql_literal().unwrap(), "42");
467 assert_eq!(Parameter::Float(3.14).to_sql_literal().unwrap(), "3.14");
468 assert_eq!(
469 Parameter::String("hello".to_string())
470 .to_sql_literal()
471 .unwrap(),
472 "'hello'"
473 );
474 }
475
476 #[test]
477 fn test_parameter_string_escaping() {
478 let param = Parameter::String("O'Reilly".to_string());
479 assert_eq!(param.to_sql_literal().unwrap(), "'O''Reilly'");
480 }
481
482 #[test]
483 fn test_parameter_sql_injection_detection() {
484 let dangerous = Parameter::String("'; DROP TABLE users; --".to_string());
485 assert!(dangerous.to_sql_literal().is_err());
486
487 let malicious = Parameter::String("' OR '1'='1".to_string());
488 assert!(malicious.to_sql_literal().is_err());
489
490 let safe = Parameter::String("It's a nice day".to_string());
491 assert!(safe.to_sql_literal().is_ok());
492 }
493
494 #[test]
495 fn test_parameter_conversions() {
496 let _p: Parameter = true.into();
497 let _p: Parameter = 42i32.into();
498 let _p: Parameter = 42i64.into();
499 let _p: Parameter = 3.14f64.into();
500 let _p: Parameter = "test".into();
501 let _p: Parameter = String::from("test").into();
502 let _p: Parameter = vec![1u8, 2, 3].into();
503 }
504
505 #[test]
506 fn test_statement_creation() {
507 let stmt = Statement::new("SELECT * FROM users");
508
509 assert_eq!(stmt.sql(), "SELECT * FROM users");
510 assert_eq!(stmt.statement_type(), StatementType::Select);
511 assert_eq!(stmt.timeout_ms(), None);
512 }
513
514 #[test]
515 fn test_statement_parameter_binding() {
516 let mut stmt = Statement::new("SELECT * FROM users WHERE id = ?");
517
518 stmt.bind(0, 42).unwrap();
519
520 let final_sql = stmt.build_sql().unwrap();
521 assert_eq!(final_sql, "SELECT * FROM users WHERE id = 42");
522 }
523
524 #[test]
525 fn test_statement_multiple_parameters() {
526 let mut stmt = Statement::new("SELECT * FROM users WHERE age > ? AND name = ?");
527
528 stmt.bind(0, 18).unwrap();
529 stmt.bind(1, "John").unwrap();
530
531 let final_sql = stmt.build_sql().unwrap();
532 assert_eq!(
533 final_sql,
534 "SELECT * FROM users WHERE age > 18 AND name = 'John'"
535 );
536 }
537
538 #[test]
539 fn test_statement_set_timeout() {
540 let mut stmt = Statement::new("SELECT * FROM users");
541 stmt.set_timeout(30_000);
542 assert_eq!(stmt.timeout_ms(), Some(30_000));
543 }
544
545 #[test]
546 fn test_statement_clear_parameters() {
547 let mut stmt = Statement::new("SELECT * FROM users WHERE id = ?");
548 stmt.bind(0, 42).unwrap();
549 stmt.clear_parameters();
550 assert!(stmt.parameters().is_empty());
551 }
552
553 #[test]
554 fn test_statement_display() {
555 let stmt = Statement::new("SELECT 1");
556 let display = format!("{}", stmt);
557 assert!(display.contains("SELECT 1"));
558 }
559
560 #[test]
563 fn build_sql_substitutes_question_mark_in_normal_text() {
564 let mut stmt = Statement::new("SELECT ?");
565 stmt.bind(0, 42i64).unwrap();
566 assert_eq!(stmt.build_sql().unwrap(), "SELECT 42");
567 }
568
569 #[test]
570 fn build_sql_does_not_substitute_question_mark_in_single_quoted_string() {
571 let stmt = Statement::new("SELECT 'a?b' AS v");
573 assert_eq!(stmt.build_sql().unwrap(), "SELECT 'a?b' AS v");
574 }
575
576 #[test]
577 fn build_sql_does_not_substitute_question_mark_in_double_quoted_identifier() {
578 let stmt = Statement::new("SELECT 1 AS \"col?name\"");
579 assert_eq!(stmt.build_sql().unwrap(), "SELECT 1 AS \"col?name\"");
580 }
581
582 #[test]
583 fn build_sql_does_not_substitute_question_mark_in_line_comment() {
584 let mut stmt = Statement::new("SELECT 1 -- has ?\n WHERE x = ?");
586 stmt.bind(0, 7i64).unwrap();
587 assert_eq!(stmt.build_sql().unwrap(), "SELECT 1 -- has ?\n WHERE x = 7");
588 }
589
590 #[test]
591 fn build_sql_does_not_substitute_question_mark_in_block_comment() {
592 let mut stmt = Statement::new("SELECT /* what? */ ?");
593 stmt.bind(0, 99i64).unwrap();
594 assert_eq!(stmt.build_sql().unwrap(), "SELECT /* what? */ 99");
595 }
596
597 #[test]
598 fn build_sql_mixed_placeholder_and_literal_question_mark() {
599 let mut stmt = Statement::new("SELECT 'a?b', ?");
601 stmt.bind(0, 5i64).unwrap();
602 assert_eq!(stmt.build_sql().unwrap(), "SELECT 'a?b', 5");
603 }
604
605 #[test]
606 fn build_sql_escaped_single_quote_in_string_with_question_mark_stays_literal() {
607 let stmt = Statement::new("SELECT 'O''Reilly?'");
609 assert_eq!(stmt.build_sql().unwrap(), "SELECT 'O''Reilly?'");
610 }
611
612 #[test]
613 fn build_sql_empty_sql_with_no_params_returns_empty_string() {
614 let stmt = Statement::new("");
615 assert_eq!(stmt.build_sql().unwrap(), "");
616 }
617
618 #[test]
619 fn build_sql_sql_ending_mid_string_literal_does_not_panic() {
620 let stmt = Statement::new("SELECT 'unclosed");
622 assert_eq!(stmt.build_sql().unwrap(), "SELECT 'unclosed");
623 }
624
625 #[test]
626 fn build_sql_round_trip_select_by_id() {
627 let mut stmt = Statement::new("SELECT * FROM users WHERE id = ?");
628 stmt.bind(0, 42i64).unwrap();
629 assert_eq!(
630 stmt.build_sql().unwrap(),
631 "SELECT * FROM users WHERE id = 42"
632 );
633 }
634
635 #[test]
636 fn build_sql_not_enough_parameters_returns_error() {
637 let stmt = Statement::new("SELECT ?, ?");
638 let err = stmt.build_sql().unwrap_err();
640 assert!(matches!(err, QueryError::ParameterBindingError { .. }));
641 }
642
643 #[test]
644 fn build_sql_parameter_not_bound_returns_error() {
645 let mut stmt = Statement::new("SELECT ?, ?");
646 stmt.bind(1, 99i64).unwrap();
648 let err = stmt.build_sql().unwrap_err();
649 assert!(matches!(
650 err,
651 QueryError::ParameterBindingError { index: 0, .. }
652 ));
653 }
654
655 #[test]
658 fn scan_placeholders_empty_sql() {
659 assert_eq!(scan_placeholders(""), Vec::<usize>::new());
660 }
661
662 #[test]
663 fn scan_placeholders_returns_byte_offsets_in_normal_text() {
664 let positions = scan_placeholders("SELECT ? , ?");
666 assert_eq!(positions, vec![7, 11]);
667 }
668
669 #[test]
670 fn scan_placeholders_ignores_question_mark_in_single_quoted_string() {
671 assert!(scan_placeholders("SELECT 'a?b' AS v").is_empty());
673 }
674
675 #[test]
676 fn scan_placeholders_ignores_question_mark_in_double_quoted_identifier() {
677 assert!(scan_placeholders("SELECT 1 AS \"col?name\"").is_empty());
678 }
679
680 #[test]
681 fn scan_placeholders_ignores_question_mark_in_line_comment() {
682 let sql = "SELECT 1 -- has ?\n WHERE x = ?";
684 let positions = scan_placeholders(sql);
685 assert_eq!(positions.len(), 1, "got {:?}", positions);
686 assert_eq!(&sql[positions[0]..positions[0] + 1], "?");
687 assert_eq!(positions[0], sql.len() - 1);
689 }
690
691 #[test]
692 fn scan_placeholders_ignores_question_mark_in_block_comment() {
693 let sql = "SELECT /* what? */ ?";
694 let positions = scan_placeholders(sql);
695 assert_eq!(positions.len(), 1);
696 assert_eq!(&sql[positions[0]..positions[0] + 1], "?");
697 }
698
699 #[test]
700 fn scan_placeholders_handles_escaped_single_quote() {
701 let sql = "SELECT 'O''Reilly?' = ?";
704 let positions = scan_placeholders(sql);
705 assert_eq!(positions.len(), 1);
706 assert_eq!(positions[0], sql.len() - 1);
707 }
708
709 #[test]
710 fn scan_placeholders_handles_escaped_double_quote() {
711 let sql = "SELECT \"a\"\"b?\" = ?";
712 let positions = scan_placeholders(sql);
713 assert_eq!(positions.len(), 1);
714 assert_eq!(positions[0], sql.len() - 1);
715 }
716
717 #[test]
718 fn scan_placeholders_unterminated_string_does_not_panic() {
719 let positions = scan_placeholders("SELECT 'a?b");
722 assert!(positions.is_empty());
723 }
724
725 #[test]
726 fn scan_placeholders_unterminated_block_comment_does_not_panic() {
727 let positions = scan_placeholders("SELECT /* what? AND ? then EOF");
728 assert!(positions.is_empty());
729 }
730
731 #[test]
732 fn scan_placeholders_utf8_multibyte_offsets_are_byte_safe() {
733 let sql = "ä?";
736 assert_eq!(sql.len(), 3);
737 let positions = scan_placeholders(sql);
738 assert_eq!(positions, vec![2]);
739 assert_eq!(&sql[positions[0]..positions[0] + 1], "?");
740 }
741
742 #[test]
743 fn scan_placeholders_consecutive_question_marks() {
744 let positions = scan_placeholders("??");
745 assert_eq!(positions, vec![0, 1]);
746 }
747
748 #[test]
749 fn scan_placeholders_block_comment_inside_string_is_ignored() {
750 let sql = "SELECT '/* ? */', ?";
753 let positions = scan_placeholders(sql);
754 assert_eq!(positions, vec![sql.len() - 1]);
755 }
756
757 #[test]
758 fn scan_placeholders_line_comment_inside_string_is_ignored() {
759 let sql = "SELECT '-- still in string ?', ?";
760 let positions = scan_placeholders(sql);
761 assert_eq!(positions, vec![sql.len() - 1]);
762 }
763}