1use crate::dialect::Dialect;
6use crate::model::Model;
7use crate::value::Value;
8use std::fmt;
9
10pub struct QueryBuilder<M: Model> {
12 table: Option<String>,
13 select_columns: Vec<String>,
14 where_conditions: Vec<WhereCondition>,
15 order_by: Vec<OrderClause>,
16 group_by: Vec<String>,
17 having_conditions: Vec<WhereCondition>,
18 limit_value: Option<usize>,
19 offset_value: Option<usize>,
20 joins: Vec<JoinClause>,
21 dialect: Box<dyn Dialect>,
22 #[allow(dead_code)]
23 model: std::marker::PhantomData<M>,
24}
25
26#[derive(Debug, Clone)]
27#[allow(dead_code)]
28enum WhereCondition {
29 And(String),
30 Or(String),
31 In(String, Vec<Value>),
32 NotIn(String, Vec<Value>),
33 Between(String, Value, Value),
34 NotBetween(String, Value, Value),
35 Null(String),
36 NotNull(String),
37 Exists(String),
38 NotExists(String),
39}
40
41#[derive(Debug, Clone)]
42struct OrderClause {
43 field: String,
44 direction: OrderDirection,
45}
46
47#[derive(Debug, Clone)]
48enum OrderDirection {
49 Asc,
50 Desc,
51}
52
53#[derive(Debug, Clone)]
54#[allow(dead_code)]
55enum JoinClause {
56 Inner(String, String, String),
57 Left(String, String, String),
58 Right(String, String, String),
59 Cross(String, String),
60}
61
62impl<M: Model> QueryBuilder<M> {
63 pub fn new(dialect: Box<dyn Dialect>) -> Self {
64 Self {
65 table: None,
66 select_columns: vec!["*".to_string()],
67 where_conditions: Vec::new(),
68 order_by: Vec::new(),
69 group_by: Vec::new(),
70 having_conditions: Vec::new(),
71 limit_value: None,
72 offset_value: None,
73 joins: Vec::new(),
74 dialect,
75 model: std::marker::PhantomData,
76 }
77 }
78
79 pub fn table(mut self, table: impl Into<String>) -> Self {
80 self.table = Some(table.into());
81 self
82 }
83
84 pub fn select(mut self, columns: Vec<&str>) -> Self {
92 self.select_columns = columns.into_iter().map(|s| s.to_string()).collect();
93 self
94 }
95
96 pub fn select_quoted(mut self, columns: Vec<&str>) -> Result<Self, crate::DbError> {
105 let mut quoted = Vec::with_capacity(columns.len());
106 for col in columns {
107 crate::sql_safety::validate_identifier(col, "select column")?;
108 quoted.push(self.dialect.quote(col));
109 }
110 self.select_columns = quoted;
111 Ok(self)
112 }
113
114 pub fn where_cond(mut self, condition: impl Into<String>) -> Self {
115 self.where_conditions
116 .push(WhereCondition::And(condition.into()));
117 self
118 }
119
120 pub fn or_where(mut self, condition: impl Into<String>) -> Self {
121 self.where_conditions
122 .push(WhereCondition::Or(condition.into()));
123 self
124 }
125
126 pub fn where_in(mut self, field: impl Into<String>, values: Vec<Value>) -> Self {
127 self.where_conditions
128 .push(WhereCondition::In(field.into(), values));
129 self
130 }
131
132 pub fn where_not_in(mut self, field: impl Into<String>, values: Vec<Value>) -> Self {
133 self.where_conditions
134 .push(WhereCondition::NotIn(field.into(), values));
135 self
136 }
137
138 pub fn where_between(mut self, field: impl Into<String>, start: Value, end: Value) -> Self {
139 self.where_conditions
140 .push(WhereCondition::Between(field.into(), start, end));
141 self
142 }
143
144 pub fn where_not_between(mut self, field: impl Into<String>, start: Value, end: Value) -> Self {
145 self.where_conditions
146 .push(WhereCondition::NotBetween(field.into(), start, end));
147 self
148 }
149
150 pub fn where_null(mut self, field: impl Into<String>) -> Self {
151 self.where_conditions
152 .push(WhereCondition::Null(field.into()));
153 self
154 }
155
156 pub fn where_not_null(mut self, field: impl Into<String>) -> Self {
157 self.where_conditions
158 .push(WhereCondition::NotNull(field.into()));
159 self
160 }
161
162 pub fn order_by(mut self, field: impl Into<String>) -> Self {
163 self.order_by.push(OrderClause {
164 field: field.into(),
165 direction: OrderDirection::Asc,
166 });
167 self
168 }
169
170 pub fn order_desc(mut self, field: impl Into<String>) -> Self {
171 self.order_by.push(OrderClause {
172 field: field.into(),
173 direction: OrderDirection::Desc,
174 });
175 self
176 }
177
178 pub fn group_by(mut self, field: impl Into<String>) -> Self {
179 self.group_by.push(field.into());
180 self
181 }
182
183 pub fn having(mut self, condition: impl Into<String>) -> Self {
184 self.having_conditions
185 .push(WhereCondition::And(condition.into()));
186 self
187 }
188
189 pub fn limit(mut self, limit: usize) -> Self {
190 self.limit_value = Some(limit);
191 self
192 }
193
194 pub fn offset(mut self, offset: usize) -> Self {
195 self.offset_value = Some(offset);
196 self
197 }
198
199 pub fn page(mut self, page: usize, page_size: usize) -> Self {
200 self.limit_value = Some(page_size);
201 self.offset_value = Some((page.saturating_sub(1)) * page_size);
202 self
203 }
204
205 pub fn join_inner(
206 mut self,
207 table: impl Into<String>,
208 on_left: impl Into<String>,
209 on_right: impl Into<String>,
210 ) -> Self {
211 self.joins.push(JoinClause::Inner(
212 table.into(),
213 on_left.into(),
214 on_right.into(),
215 ));
216 self
217 }
218
219 pub fn join_left(
220 mut self,
221 table: impl Into<String>,
222 on_left: impl Into<String>,
223 on_right: impl Into<String>,
224 ) -> Self {
225 self.joins.push(JoinClause::Left(
226 table.into(),
227 on_left.into(),
228 on_right.into(),
229 ));
230 self
231 }
232
233 pub fn join_right(
234 mut self,
235 table: impl Into<String>,
236 on_left: impl Into<String>,
237 on_right: impl Into<String>,
238 ) -> Self {
239 self.joins.push(JoinClause::Right(
240 table.into(),
241 on_left.into(),
242 on_right.into(),
243 ));
244 self
245 }
246
247 #[tracing::instrument(skip(self), fields(op = "select"))]
280 pub fn build_select(&self) -> String {
281 let table = self
282 .table
283 .clone()
284 .unwrap_or_else(|| M::table_name().to_string());
285
286 let columns = if self.select_columns.is_empty() {
287 "*".to_string()
288 } else {
289 self.select_columns.join(", ")
290 };
291
292 let mut sql = format!("SELECT {} FROM {}", columns, self.dialect.quote(&table));
293
294 for join in &self.joins {
295 match join {
296 JoinClause::Inner(t, l, r) => {
297 sql.push_str(&format!(
298 " INNER JOIN {} ON {} = {}",
299 self.dialect.quote(t),
300 self.dialect.quote(l),
301 self.dialect.quote(r)
302 ));
303 }
304 JoinClause::Left(t, l, r) => {
305 sql.push_str(&format!(
306 " LEFT JOIN {} ON {} = {}",
307 self.dialect.quote(t),
308 self.dialect.quote(l),
309 self.dialect.quote(r)
310 ));
311 }
312 JoinClause::Right(t, l, r) => {
313 sql.push_str(&format!(
314 " RIGHT JOIN {} ON {} = {}",
315 self.dialect.quote(t),
316 self.dialect.quote(l),
317 self.dialect.quote(r)
318 ));
319 }
320 JoinClause::Cross(t, on) => {
321 sql.push_str(&format!(
322 " CROSS JOIN {} ON {}",
323 self.dialect.quote(t),
324 self.dialect.quote(on)
325 ));
326 }
327 }
328 }
329
330 if !self.where_conditions.is_empty() {
331 sql.push_str(&self.build_where_clause());
332 }
333
334 if !self.group_by.is_empty() {
335 let cols: Vec<String> = self
336 .group_by
337 .iter()
338 .map(|c| self.dialect.quote(c))
339 .collect();
340 sql.push_str(" GROUP BY ");
341 sql.push_str(&cols.join(", "));
342 }
343
344 if !self.having_conditions.is_empty() {
345 sql.push_str(" HAVING ");
346 for (i, cond) in self.having_conditions.iter().enumerate() {
347 if i > 0 {
348 sql.push_str(" AND ");
349 }
350 if let WhereCondition::And(c) = cond {
351 sql.push_str(c);
352 }
353 }
354 }
355
356 if !self.order_by.is_empty() {
357 let order_cols: Vec<String> = self
358 .order_by
359 .iter()
360 .map(|o| {
361 let dir = match o.direction {
362 OrderDirection::Asc => " ASC",
363 OrderDirection::Desc => " DESC",
364 };
365 format!("{}{}", self.dialect.quote(&o.field), dir)
366 })
367 .collect();
368 sql.push_str(" ORDER BY ");
369 sql.push_str(&order_cols.join(", "));
370 }
371
372 if let Some(limit) = self.limit_value {
373 sql.push_str(&format!(" LIMIT {}", limit));
374 }
375
376 if let Some(offset) = self.offset_value {
377 sql.push_str(&format!(" OFFSET {}", offset));
378 }
379
380 sql
381 }
382
383 fn build_where_clause(&self) -> String {
386 if self.where_conditions.is_empty() {
387 return String::new();
388 }
389
390 let conditions: Vec<String> = self
392 .where_conditions
393 .iter()
394 .map(|cond| match cond {
395 WhereCondition::And(c) => c.clone(),
396 WhereCondition::Or(c) => format!("OR {}", c),
397 WhereCondition::In(f, vals) => {
398 let vals_str: Vec<String> = vals
400 .iter()
401 .map(|v| v.to_param_with_dialect(&*self.dialect).to_string())
402 .collect();
403 format!("{} IN ({})", self.dialect.quote(f), vals_str.join(", "))
404 }
405 WhereCondition::NotIn(f, vals) => {
406 let vals_str: Vec<String> = vals
407 .iter()
408 .map(|v| v.to_param_with_dialect(&*self.dialect).to_string())
409 .collect();
410 format!("{} NOT IN ({})", self.dialect.quote(f), vals_str.join(", "))
411 }
412 WhereCondition::Between(f, start, end) => {
413 format!(
414 "{} BETWEEN {} AND {}",
415 self.dialect.quote(f),
416 start.to_param_with_dialect(&*self.dialect),
417 end.to_param_with_dialect(&*self.dialect)
418 )
419 }
420 WhereCondition::NotBetween(f, start, end) => {
421 format!(
422 "{} NOT BETWEEN {} AND {}",
423 self.dialect.quote(f),
424 start.to_param_with_dialect(&*self.dialect),
425 end.to_param_with_dialect(&*self.dialect)
426 )
427 }
428 WhereCondition::Null(f) => format!("{} IS NULL", self.dialect.quote(f)),
429 WhereCondition::NotNull(f) => format!("{} IS NOT NULL", self.dialect.quote(f)),
430 WhereCondition::Exists(s) => format!("EXISTS ({})", s),
431 WhereCondition::NotExists(s) => format!("NOT EXISTS ({})", s),
432 })
433 .collect();
434
435 if conditions.is_empty() {
436 return String::new();
437 }
438
439 let mut groups: Vec<Vec<String>> = Vec::new();
442 let mut current_group: Vec<String> = Vec::new();
443 for cond in conditions.iter() {
444 if let Some(stripped) = cond.strip_prefix("OR ") {
445 current_group.push(stripped.to_string());
447 } else {
448 if !current_group.is_empty() {
450 groups.push(std::mem::take(&mut current_group));
451 }
452 current_group.push(cond.clone());
453 }
454 }
455 if !current_group.is_empty() {
456 groups.push(current_group);
457 }
458
459 let group_strs: Vec<String> = groups
460 .iter()
461 .map(|g| {
462 if g.len() == 1 {
463 g[0].clone()
464 } else {
465 format!("({})", g.join(" OR "))
466 }
467 })
468 .collect();
469
470 format!(" WHERE {}", group_strs.join(" AND "))
471 }
472
473 #[tracing::instrument(skip(self, data), fields(op = "insert"))]
474 pub fn build_insert(&self, data: &std::collections::HashMap<String, Value>) -> String {
475 let table = self
476 .table
477 .clone()
478 .unwrap_or_else(|| M::table_name().to_string());
479
480 if data.is_empty() {
481 return String::new();
482 }
483
484 let columns: Vec<String> = data.keys().map(|k| self.dialect.quote(k)).collect();
485 let values: Vec<String> = data
487 .values()
488 .map(|v| v.to_param_with_dialect(&*self.dialect).to_string())
489 .collect();
490
491 format!(
492 "INSERT INTO {} ({}) VALUES ({})",
493 self.dialect.quote(&table),
494 columns.join(", "),
495 values.join(", ")
496 )
497 }
498
499 #[tracing::instrument(skip(self, data), fields(op = "update"))]
500 pub fn build_update(&self, data: &std::collections::HashMap<String, Value>) -> String {
501 let table = self
502 .table
503 .clone()
504 .unwrap_or_else(|| M::table_name().to_string());
505
506 if data.is_empty() {
507 return String::new();
508 }
509
510 let set_clauses: Vec<String> = data
511 .iter()
512 .map(|(k, v)| {
513 format!(
514 "{} = {}",
515 self.dialect.quote(k),
516 v.to_param_with_dialect(&*self.dialect)
517 )
518 })
519 .collect();
520
521 let mut sql = format!(
522 "UPDATE {} SET {}",
523 self.dialect.quote(&table),
524 set_clauses.join(", ")
525 );
526
527 sql.push_str(&self.build_where_clause());
528 sql
529 }
530
531 #[tracing::instrument(skip(self), fields(op = "delete"))]
532 pub fn build_delete(&self) -> String {
533 let table = self
534 .table
535 .clone()
536 .unwrap_or_else(|| M::table_name().to_string());
537
538 let mut sql = format!("DELETE FROM {}", self.dialect.quote(&table));
539 sql.push_str(&self.build_where_clause());
540 sql
541 }
542
543 fn build_where_clause_with_params(&self) -> (String, Vec<Value>) {
551 if self.where_conditions.is_empty() {
552 return (String::new(), Vec::new());
553 }
554
555 let mut params = Vec::new();
556
557 let conditions: Vec<String> = self
558 .where_conditions
559 .iter()
560 .map(|cond| match cond {
561 WhereCondition::And(c) => c.clone(),
562 WhereCondition::Or(c) => format!("OR {}", c),
563 WhereCondition::In(f, vals) => {
564 let placeholders: Vec<&str> = vals.iter().map(|_| "?").collect();
565 params.extend(vals.iter().cloned());
566 format!("{} IN ({})", self.dialect.quote(f), placeholders.join(", "))
567 }
568 WhereCondition::NotIn(f, vals) => {
569 let placeholders: Vec<&str> = vals.iter().map(|_| "?").collect();
570 params.extend(vals.iter().cloned());
571 format!(
572 "{} NOT IN ({})",
573 self.dialect.quote(f),
574 placeholders.join(", ")
575 )
576 }
577 WhereCondition::Between(f, start, end) => {
578 params.push(start.clone());
579 params.push(end.clone());
580 format!("{} BETWEEN ? AND ?", self.dialect.quote(f))
581 }
582 WhereCondition::NotBetween(f, start, end) => {
583 params.push(start.clone());
584 params.push(end.clone());
585 format!("{} NOT BETWEEN ? AND ?", self.dialect.quote(f))
586 }
587 WhereCondition::Null(f) => format!("{} IS NULL", self.dialect.quote(f)),
588 WhereCondition::NotNull(f) => format!("{} IS NOT NULL", self.dialect.quote(f)),
589 WhereCondition::Exists(s) => format!("EXISTS ({})", s),
590 WhereCondition::NotExists(s) => format!("NOT EXISTS ({})", s),
591 })
592 .collect();
593
594 if conditions.is_empty() {
595 return (String::new(), params);
596 }
597
598 let mut groups: Vec<Vec<String>> = Vec::new();
600 let mut current_group: Vec<String> = Vec::new();
601 for cond in conditions.iter() {
602 if let Some(stripped) = cond.strip_prefix("OR ") {
603 current_group.push(stripped.to_string());
604 } else {
605 if !current_group.is_empty() {
606 groups.push(std::mem::take(&mut current_group));
607 }
608 current_group.push(cond.clone());
609 }
610 }
611 if !current_group.is_empty() {
612 groups.push(current_group);
613 }
614
615 let group_strs: Vec<String> = groups
616 .iter()
617 .map(|g| {
618 if g.len() == 1 {
619 g[0].clone()
620 } else {
621 format!("({})", g.join(" OR "))
622 }
623 })
624 .collect();
625
626 (format!(" WHERE {}", group_strs.join(" AND ")), params)
627 }
628
629 pub fn build_select_with_params(&self) -> (String, Vec<Value>) {
634 let table = self
635 .table
636 .clone()
637 .unwrap_or_else(|| M::table_name().to_string());
638 let columns = if self.select_columns.is_empty() {
639 "*".to_string()
640 } else {
641 self.select_columns.join(", ")
642 };
643
644 let mut sql = format!("SELECT {} FROM {}", columns, self.dialect.quote(&table));
645
646 for join in &self.joins {
647 match join {
648 JoinClause::Inner(t, l, r) => {
649 sql.push_str(&format!(
650 " INNER JOIN {} ON {} = {}",
651 self.dialect.quote(t),
652 self.dialect.quote(l),
653 self.dialect.quote(r)
654 ));
655 }
656 JoinClause::Left(t, l, r) => {
657 sql.push_str(&format!(
658 " LEFT JOIN {} ON {} = {}",
659 self.dialect.quote(t),
660 self.dialect.quote(l),
661 self.dialect.quote(r)
662 ));
663 }
664 JoinClause::Right(t, l, r) => {
665 sql.push_str(&format!(
666 " RIGHT JOIN {} ON {} = {}",
667 self.dialect.quote(t),
668 self.dialect.quote(l),
669 self.dialect.quote(r)
670 ));
671 }
672 JoinClause::Cross(t, on) => {
673 sql.push_str(&format!(
674 " CROSS JOIN {} ON {}",
675 self.dialect.quote(t),
676 self.dialect.quote(on)
677 ));
678 }
679 }
680 }
681
682 let mut params = Vec::new();
683 if !self.where_conditions.is_empty() {
684 let (where_clause, where_params) = self.build_where_clause_with_params();
685 sql.push_str(&where_clause);
686 params = where_params;
687 }
688
689 if !self.group_by.is_empty() {
690 let cols: Vec<String> = self
691 .group_by
692 .iter()
693 .map(|c| self.dialect.quote(c))
694 .collect();
695 sql.push_str(" GROUP BY ");
696 sql.push_str(&cols.join(", "));
697 }
698
699 if !self.having_conditions.is_empty() {
700 sql.push_str(" HAVING ");
701 for (i, cond) in self.having_conditions.iter().enumerate() {
702 if i > 0 {
703 sql.push_str(" AND ");
704 }
705 if let WhereCondition::And(c) = cond {
706 sql.push_str(c);
707 }
708 }
709 }
710
711 if !self.order_by.is_empty() {
712 let order_cols: Vec<String> = self
713 .order_by
714 .iter()
715 .map(|o| {
716 let dir = match o.direction {
717 OrderDirection::Asc => " ASC",
718 OrderDirection::Desc => " DESC",
719 };
720 format!("{}{}", self.dialect.quote(&o.field), dir)
721 })
722 .collect();
723 sql.push_str(" ORDER BY ");
724 sql.push_str(&order_cols.join(", "));
725 }
726
727 if let Some(limit) = self.limit_value {
728 sql.push_str(&format!(" LIMIT {}", limit));
729 }
730 if let Some(offset) = self.offset_value {
731 sql.push_str(&format!(" OFFSET {}", offset));
732 }
733
734 (sql, params)
735 }
736
737 pub fn build_insert_with_params(
739 &self,
740 data: &std::collections::HashMap<String, Value>,
741 ) -> (String, Vec<Value>) {
742 let table = self
743 .table
744 .clone()
745 .unwrap_or_else(|| M::table_name().to_string());
746 if data.is_empty() {
747 return (String::new(), Vec::new());
748 }
749
750 let mut columns = Vec::with_capacity(data.len());
751 let mut params = Vec::with_capacity(data.len());
752 let placeholders: Vec<&str> = data.iter().map(|_| "?").collect();
753 for (k, v) in data.iter() {
754 columns.push(self.dialect.quote(k));
755 params.push(v.clone());
756 }
757
758 let sql = format!(
759 "INSERT INTO {} ({}) VALUES ({})",
760 self.dialect.quote(&table),
761 columns.join(", "),
762 placeholders.join(", ")
763 );
764 (sql, params)
765 }
766
767 pub fn build_update_with_params(
770 &self,
771 data: &std::collections::HashMap<String, Value>,
772 ) -> (String, Vec<Value>) {
773 let table = self
774 .table
775 .clone()
776 .unwrap_or_else(|| M::table_name().to_string());
777 if data.is_empty() {
778 return (String::new(), Vec::new());
779 }
780
781 let mut set_clauses = Vec::with_capacity(data.len());
782 let mut params = Vec::with_capacity(data.len());
783 for (k, v) in data.iter() {
784 set_clauses.push(format!("{} = ?", self.dialect.quote(k)));
785 params.push(v.clone());
786 }
787
788 let mut sql = format!(
789 "UPDATE {} SET {}",
790 self.dialect.quote(&table),
791 set_clauses.join(", ")
792 );
793
794 if !self.where_conditions.is_empty() {
795 let (where_clause, where_params) = self.build_where_clause_with_params();
796 sql.push_str(&where_clause);
797 params.extend(where_params);
798 }
799
800 (sql, params)
801 }
802
803 pub fn build_delete_with_params(&self) -> (String, Vec<Value>) {
805 let table = self
806 .table
807 .clone()
808 .unwrap_or_else(|| M::table_name().to_string());
809 let mut sql = format!("DELETE FROM {}", self.dialect.quote(&table));
810 let mut params = Vec::new();
811
812 if !self.where_conditions.is_empty() {
813 let (where_clause, where_params) = self.build_where_clause_with_params();
814 sql.push_str(&where_clause);
815 params = where_params;
816 }
817
818 (sql, params)
819 }
820
821 pub fn build_count(&self) -> String {
822 let table = self
823 .table
824 .clone()
825 .unwrap_or_else(|| M::table_name().to_string());
826
827 let mut sql = format!(
828 "SELECT COUNT(*) as total FROM {}",
829 self.dialect.quote(&table)
830 );
831 sql.push_str(&self.build_where_clause());
832 sql
833 }
834
835 pub fn build_exists(&self) -> String {
836 let table = self
837 .table
838 .clone()
839 .unwrap_or_else(|| M::table_name().to_string());
840
841 let mut sql = format!("SELECT 1 FROM {}", self.dialect.quote(&table));
842 sql.push_str(&self.build_where_clause());
843 sql.push_str(" LIMIT 1");
844 format!("SELECT EXISTS({})", sql)
845 }
846
847 pub fn build_max(&self, field: &str) -> String {
848 let table = self
849 .table
850 .clone()
851 .unwrap_or_else(|| M::table_name().to_string());
852
853 let mut sql = format!(
854 "SELECT MAX({}) as max_val FROM {}",
855 self.dialect.quote(field),
856 self.dialect.quote(&table)
857 );
858 sql.push_str(&self.build_where_clause());
859 sql
860 }
861
862 pub fn build_min(&self, field: &str) -> String {
863 let table = self
864 .table
865 .clone()
866 .unwrap_or_else(|| M::table_name().to_string());
867
868 let mut sql = format!(
869 "SELECT MIN({}) as min_val FROM {}",
870 self.dialect.quote(field),
871 self.dialect.quote(&table)
872 );
873 sql.push_str(&self.build_where_clause());
874 sql
875 }
876
877 pub fn build_sum(&self, field: &str) -> String {
878 let table = self
879 .table
880 .clone()
881 .unwrap_or_else(|| M::table_name().to_string());
882
883 let mut sql = format!(
884 "SELECT SUM({}) as sum_val FROM {}",
885 self.dialect.quote(field),
886 self.dialect.quote(&table)
887 );
888 sql.push_str(&self.build_where_clause());
889 sql
890 }
891
892 pub fn build_avg(&self, field: &str) -> String {
893 let table = self
894 .table
895 .clone()
896 .unwrap_or_else(|| M::table_name().to_string());
897
898 let mut sql = format!(
899 "SELECT AVG({}) as avg_val FROM {}",
900 self.dialect.quote(field),
901 self.dialect.quote(&table)
902 );
903 sql.push_str(&self.build_where_clause());
904 sql
905 }
906
907 pub fn validate(&self) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
910 let sql = self.build_select();
911 let mut errors = Vec::new();
912
913 if let Err(e) = sz_orm_sql_validator::validate_select(&sql) {
914 errors.push(e);
915 }
916
917 if !self.joins.is_empty() {
919 for join in &self.joins {
920 match join {
921 JoinClause::Inner(_, left, right)
922 | JoinClause::Left(_, left, right)
923 | JoinClause::Right(_, left, right) => {
924 if let Err(e) = sz_orm_sql_validator::validate_column_name(left) {
925 errors.push(e);
926 }
927 if let Err(e) = sz_orm_sql_validator::validate_column_name(right) {
928 errors.push(e);
929 }
930 }
931 _ => {}
932 }
933 }
934 }
935
936 let table = self
938 .table
939 .clone()
940 .unwrap_or_else(|| M::table_name().to_string());
941 if let Err(e) = sz_orm_sql_validator::validate_table_name(&table) {
942 errors.push(e);
943 }
944
945 if errors.is_empty() {
946 Ok(())
947 } else {
948 Err(errors)
949 }
950 }
951
952 pub fn validate_insert(
955 &self,
956 data: &std::collections::HashMap<String, Value>,
957 ) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
958 let sql = self.build_insert(data);
959 let mut errors = Vec::new();
960
961 if sql.is_empty() {
962 errors.push(sz_orm_sql_validator::SqlValidationError::EmptyInsertData);
963 return Err(errors);
964 }
965
966 if let Err(e) = sz_orm_sql_validator::validate_insert(&sql) {
967 errors.push(e);
968 }
969
970 if errors.is_empty() {
971 Ok(())
972 } else {
973 Err(errors)
974 }
975 }
976
977 pub fn validate_update(
980 &self,
981 data: &std::collections::HashMap<String, Value>,
982 ) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
983 let sql = self.build_update(data);
984 let mut errors = Vec::new();
985
986 if sql.is_empty() {
987 errors.push(sz_orm_sql_validator::SqlValidationError::EmptyUpdateData);
988 return Err(errors);
989 }
990
991 if let Err(e) = sz_orm_sql_validator::validate_update(&sql) {
992 errors.push(e);
993 }
994
995 if errors.is_empty() {
996 Ok(())
997 } else {
998 Err(errors)
999 }
1000 }
1001
1002 pub fn validate_delete(&self) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
1004 let sql = self.build_delete();
1005 let mut errors = Vec::new();
1006
1007 if let Err(e) = sz_orm_sql_validator::validate_delete(&sql) {
1008 errors.push(e);
1009 }
1010
1011 if errors.is_empty() {
1012 Ok(())
1013 } else {
1014 Err(errors)
1015 }
1016 }
1017}
1018
1019impl<M: Model> fmt::Debug for QueryBuilder<M> {
1020 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1021 f.debug_struct("QueryBuilder")
1022 .field("table", &self.table)
1023 .field("select_columns", &self.select_columns)
1024 .field("where_conditions", &self.where_conditions.len())
1025 .field("limit", &self.limit_value)
1026 .finish()
1027 }
1028}
1029
1030#[cfg(test)]
1031mod tests {
1032 use super::*;
1033 use crate::db_type::DbType;
1034 use crate::dialect::get_dialect;
1035
1036 struct TestModel;
1037 impl Model for TestModel {
1038 type PrimaryKey = i64;
1039
1040 fn table_name() -> &'static str {
1041 "test_models"
1042 }
1043
1044 fn pk(&self) -> Self::PrimaryKey {
1045 1
1046 }
1047
1048 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1049 }
1050
1051 #[test]
1052 fn test_query_builder_select() {
1053 let dialect = get_dialect(DbType::MySQL).unwrap();
1054 let builder = QueryBuilder::<TestModel>::new(dialect);
1055
1056 let sql = builder
1057 .table("users")
1058 .select(vec!["id", "name"])
1059 .build_select();
1060 assert!(sql.contains("SELECT id, name FROM"));
1061 assert!(sql.contains("`users`"));
1062 }
1063
1064 #[test]
1065 fn test_query_builder_where() {
1066 let dialect = get_dialect(DbType::MySQL).unwrap();
1067 let builder = QueryBuilder::<TestModel>::new(dialect);
1068
1069 let sql = builder
1070 .table("users")
1071 .where_cond("status = 'active'")
1072 .where_cond("age > 18")
1073 .build_select();
1074
1075 assert!(sql.contains("WHERE"));
1076 assert!(sql.contains("status = 'active'"));
1077 assert!(sql.contains("age > 18"));
1078 }
1079
1080 #[test]
1081 fn test_query_builder_order_by() {
1082 let dialect = get_dialect(DbType::MySQL).unwrap();
1083 let builder = QueryBuilder::<TestModel>::new(dialect);
1084
1085 let sql = builder
1086 .table("users")
1087 .order_by("created_at")
1088 .order_desc("id")
1089 .build_select();
1090
1091 assert!(sql.contains("ORDER BY"));
1092 assert!(sql.contains("`created_at` ASC"));
1093 assert!(sql.contains("`id` DESC"));
1094 }
1095
1096 #[test]
1097 fn test_query_builder_limit_offset() {
1098 let dialect = get_dialect(DbType::MySQL).unwrap();
1099 let builder = QueryBuilder::<TestModel>::new(dialect);
1100
1101 let sql = builder.table("users").limit(10).offset(20).build_select();
1102
1103 assert!(sql.contains("LIMIT 10"));
1104 assert!(sql.contains("OFFSET 20"));
1105 }
1106
1107 #[test]
1108 fn test_query_builder_page() {
1109 let dialect = get_dialect(DbType::MySQL).unwrap();
1110 let builder = QueryBuilder::<TestModel>::new(dialect);
1111
1112 let sql = builder.table("users").page(3, 20).build_select();
1113
1114 assert!(sql.contains("LIMIT 20"));
1115 assert!(sql.contains("OFFSET 40"));
1116 }
1117
1118 #[test]
1119 fn test_query_builder_insert() {
1120 let dialect = get_dialect(DbType::MySQL).unwrap();
1121 let builder = QueryBuilder::<TestModel>::new(dialect);
1122
1123 let mut data = std::collections::HashMap::new();
1124 data.insert("name".to_string(), Value::String("test".to_string()));
1125 data.insert("age".to_string(), Value::I64(25));
1126
1127 let sql = builder.table("users").build_insert(&data);
1128
1129 assert!(sql.contains("INSERT INTO"));
1130 assert!(sql.contains("`name`"));
1131 assert!(sql.contains("'test'"));
1132 }
1133
1134 #[test]
1135 fn test_query_builder_update() {
1136 let dialect = get_dialect(DbType::MySQL).unwrap();
1137 let builder = QueryBuilder::<TestModel>::new(dialect);
1138
1139 let mut data = std::collections::HashMap::new();
1140 data.insert("name".to_string(), Value::String("updated".to_string()));
1141
1142 let sql = builder
1143 .table("users")
1144 .where_cond("id = 1")
1145 .build_update(&data);
1146
1147 assert!(sql.contains("UPDATE"));
1148 assert!(sql.contains("`name` = 'updated'"));
1149 assert!(sql.contains("WHERE"));
1150 }
1151
1152 #[test]
1153 fn test_query_builder_delete() {
1154 let dialect = get_dialect(DbType::MySQL).unwrap();
1155 let builder = QueryBuilder::<TestModel>::new(dialect);
1156
1157 let sql = builder.table("users").where_cond("id = 1").build_delete();
1158
1159 assert!(sql.contains("DELETE FROM"));
1160 assert!(sql.contains("WHERE"));
1161 }
1162
1163 #[test]
1164 fn test_query_builder_count() {
1165 let dialect = get_dialect(DbType::MySQL).unwrap();
1166 let builder = QueryBuilder::<TestModel>::new(dialect);
1167
1168 let sql = builder.table("users").build_count();
1169
1170 assert!(sql.contains("SELECT COUNT(*)"));
1171 assert!(sql.contains("FROM"));
1172 }
1173
1174 #[test]
1175 fn test_query_builder_where_in() {
1176 let dialect = get_dialect(DbType::MySQL).unwrap();
1177 let builder = QueryBuilder::<TestModel>::new(dialect);
1178
1179 let sql = builder
1180 .table("users")
1181 .where_in("id", vec![Value::I64(1), Value::I64(2), Value::I64(3)])
1182 .build_select();
1183
1184 assert!(sql.contains("IN ("));
1185 }
1186
1187 #[test]
1188 fn test_query_builder_where_between() {
1189 let dialect = get_dialect(DbType::MySQL).unwrap();
1190 let builder = QueryBuilder::<TestModel>::new(dialect);
1191
1192 let sql = builder
1193 .table("users")
1194 .where_between("age", Value::I64(18), Value::I64(30))
1195 .build_select();
1196
1197 assert!(sql.contains("BETWEEN"));
1198 }
1199
1200 #[test]
1201 fn test_query_builder_where_null() {
1202 let dialect = get_dialect(DbType::MySQL).unwrap();
1203 let builder = QueryBuilder::<TestModel>::new(dialect);
1204
1205 let sql = builder
1206 .table("users")
1207 .where_null("deleted_at")
1208 .build_select();
1209
1210 assert!(sql.contains("IS NULL"));
1211 }
1212
1213 #[test]
1214 fn test_query_builder_join() {
1215 let dialect = get_dialect(DbType::MySQL).unwrap();
1216 let builder = QueryBuilder::<TestModel>::new(dialect);
1217
1218 let sql = builder
1219 .table("users")
1220 .join_inner("posts", "users.id", "posts.user_id")
1221 .build_select();
1222
1223 assert!(sql.contains("INNER JOIN"));
1224 assert!(sql.contains("`posts`"));
1225 }
1226
1227 #[test]
1228 fn test_query_builder_group_by() {
1229 let dialect = get_dialect(DbType::MySQL).unwrap();
1230 let builder = QueryBuilder::<TestModel>::new(dialect);
1231
1232 let sql = builder.table("users").group_by("status").build_select();
1233
1234 assert!(sql.contains("GROUP BY"));
1235 assert!(sql.contains("`status`"));
1236 }
1237
1238 #[test]
1239 fn test_query_builder_max() {
1240 let dialect = get_dialect(DbType::MySQL).unwrap();
1241 let builder = QueryBuilder::<TestModel>::new(dialect);
1242
1243 let sql = builder.table("users").build_max("score");
1244
1245 assert!(sql.contains("MAX("));
1246 assert!(sql.contains("`score`"));
1247 }
1248
1249 #[test]
1250 fn test_query_builder_min() {
1251 let dialect = get_dialect(DbType::MySQL).unwrap();
1252 let builder = QueryBuilder::<TestModel>::new(dialect);
1253
1254 let sql = builder.table("users").build_min("price");
1255
1256 assert!(sql.contains("MIN("));
1257 assert!(sql.contains("`price`"));
1258 }
1259
1260 #[test]
1261 fn test_query_builder_sum() {
1262 let dialect = get_dialect(DbType::MySQL).unwrap();
1263 let builder = QueryBuilder::<TestModel>::new(dialect);
1264
1265 let sql = builder.table("orders").build_sum("amount");
1266
1267 assert!(sql.contains("SUM("));
1268 assert!(sql.contains("`amount`"));
1269 }
1270
1271 #[test]
1272 fn test_query_builder_avg() {
1273 let dialect = get_dialect(DbType::MySQL).unwrap();
1274 let builder = QueryBuilder::<TestModel>::new(dialect);
1275
1276 let sql = builder.table("scores").build_avg("value");
1277
1278 assert!(sql.contains("AVG("));
1279 assert!(sql.contains("`value`"));
1280 }
1281
1282 #[test]
1283 fn test_validator_select() {
1284 let dialect = get_dialect(DbType::MySQL).unwrap();
1285 let builder = QueryBuilder::<TestModel>::new(dialect);
1286
1287 let result = builder.table("users").select(vec!["id", "name"]).validate();
1288 assert!(result.is_ok());
1289 }
1290
1291 #[test]
1292 fn test_validator_select_with_join() {
1293 let dialect = get_dialect(DbType::MySQL).unwrap();
1294 let builder = QueryBuilder::<TestModel>::new(dialect);
1295
1296 let result = builder
1297 .table("users")
1298 .join_inner("posts", "users.id", "posts.user_id")
1299 .validate();
1300 assert!(result.is_ok());
1301 }
1302
1303 #[test]
1304 fn test_validator_insert() {
1305 let dialect = get_dialect(DbType::MySQL).unwrap();
1306 let builder = QueryBuilder::<TestModel>::new(dialect);
1307
1308 let mut data = std::collections::HashMap::new();
1309 data.insert("name".to_string(), Value::String("test".to_string()));
1310
1311 let result = builder.table("users").validate_insert(&data);
1312 assert!(result.is_ok());
1313 }
1314
1315 #[test]
1316 fn test_validator_insert_empty_data() {
1317 let dialect = get_dialect(DbType::MySQL).unwrap();
1318 let builder = QueryBuilder::<TestModel>::new(dialect);
1319
1320 let data = std::collections::HashMap::new();
1321 let result = builder.table("users").validate_insert(&data);
1322 assert!(result.is_err());
1323 }
1324
1325 #[test]
1326 fn test_validator_update() {
1327 let dialect = get_dialect(DbType::MySQL).unwrap();
1328 let builder = QueryBuilder::<TestModel>::new(dialect);
1329
1330 let mut data = std::collections::HashMap::new();
1331 data.insert("name".to_string(), Value::String("updated".to_string()));
1332
1333 let result = builder.table("users").validate_update(&data);
1334 assert!(result.is_ok());
1335 }
1336
1337 #[test]
1338 fn test_validator_update_empty_data() {
1339 let dialect = get_dialect(DbType::MySQL).unwrap();
1340 let builder = QueryBuilder::<TestModel>::new(dialect);
1341
1342 let data = std::collections::HashMap::new();
1343 let result = builder.table("users").validate_update(&data);
1344 assert!(result.is_err());
1345 }
1346
1347 #[test]
1348 fn test_validator_delete() {
1349 let dialect = get_dialect(DbType::MySQL).unwrap();
1350 let builder = QueryBuilder::<TestModel>::new(dialect);
1351
1352 let result = builder
1353 .table("users")
1354 .where_cond("id = 1")
1355 .validate_delete();
1356 assert!(result.is_ok());
1357 }
1358
1359 #[test]
1360 fn test_validator_delete_no_where() {
1361 let dialect = get_dialect(DbType::MySQL).unwrap();
1362 let builder = QueryBuilder::<TestModel>::new(dialect);
1363
1364 let result = builder.table("users").validate_delete();
1366 assert!(result.is_ok());
1367 }
1368
1369 #[test]
1372 fn test_m3_select_quoted_valid_columns() {
1373 let dialect = get_dialect(DbType::MySQL).unwrap();
1374 let builder = QueryBuilder::<TestModel>::new(dialect);
1375 let builder = builder
1376 .table("users")
1377 .select_quoted(vec!["id", "name"])
1378 .expect("valid columns should succeed");
1379 let sql = builder.build_select();
1380 assert!(sql.contains("SELECT `id`, `name` FROM"));
1382 assert!(sql.contains("`users`"));
1383 }
1384
1385 #[test]
1386 fn test_m3_select_quoted_rejects_sql_injection() {
1387 let dialect = get_dialect(DbType::MySQL).unwrap();
1388 let builder = QueryBuilder::<TestModel>::new(dialect);
1389
1390 let result = builder
1392 .table("users")
1393 .select_quoted(vec!["id; DROP TABLE users"]);
1394 assert!(result.is_err());
1395
1396 let dialect = get_dialect(DbType::MySQL).unwrap();
1398 let builder = QueryBuilder::<TestModel>::new(dialect);
1399 let result = builder.table("users").select_quoted(vec!["name'"]);
1400 assert!(result.is_err());
1401
1402 let dialect = get_dialect(DbType::MySQL).unwrap();
1404 let builder = QueryBuilder::<TestModel>::new(dialect);
1405 let result = builder.table("users").select_quoted(vec!["1col"]);
1406 assert!(result.is_err());
1407
1408 let dialect = get_dialect(DbType::MySQL).unwrap();
1410 let builder = QueryBuilder::<TestModel>::new(dialect);
1411 let result = builder.table("users").select_quoted(vec!["col name"]);
1412 assert!(result.is_err());
1413 }
1414
1415 #[test]
1416 fn test_m3_select_quoted_postgresql_dialect() {
1417 let dialect = get_dialect(DbType::PostgreSQL).unwrap();
1418 let builder = QueryBuilder::<TestModel>::new(dialect);
1419 let builder = builder
1420 .table("users")
1421 .select_quoted(vec!["id", "name"])
1422 .expect("valid columns should succeed");
1423 let sql = builder.build_select();
1424 assert!(sql.contains("SELECT \"id\", \"name\" FROM"));
1426 assert!(sql.contains("\"users\""));
1427 }
1428}