Skip to main content

sz_orm_core/
query.rs

1//! 查询构造器
2//!
3//! 提供类似 ThinkORM 的链式查询构造 API
4
5use crate::dialect::Dialect;
6use crate::model::Model;
7use crate::value::Value;
8use std::fmt;
9
10/// 用于构造 SQL 查询的查询构造器
11pub 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    /// 设置 SELECT 列。
85    ///
86    /// **M-3 安全警告**:本方法直接拼接 `columns` 到 SQL,**不**进行标识符校验或 quote。
87    /// 调用方必须确保 `columns` 来自可信来源(硬编码或经 `sql_safety::validate_identifier`
88    /// 校验)。若列名可能来自不可信输入,请使用 [`QueryBuilder::select_quoted`]。
89    ///
90    /// 本方法保留原行为以兼容复杂表达式(如 `COUNT(*)`、`users.id AS uid`)。
91    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    /// M-3 修复:安全的 SELECT 列设置,自动校验每个列名并 quote。
97    ///
98    /// 每个 `column` 必须通过 `sql_safety::validate_identifier` 校验
99    /// (仅允许 ASCII 字母数字 + 下划线,不以数字开头,长度 1-63)。
100    /// 校验失败时返回 `DbError::InvalidInput`。
101    ///
102    /// 对于复杂表达式(如 `COUNT(*)`、`users.id AS uid`),请使用 [`QueryBuilder::select`]
103    /// 并自行确保安全。
104    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    /// 构建 SELECT SQL 语句
248    ///
249    /// L-5 修复:补充示例文档
250    ///
251    /// 根据 `table`、`select_columns`、`where_conditions`、`joins`、`order_by`、
252    /// `group_by`、`having`、`limit`、`offset` 等条件拼装最终 SQL。
253    /// 若未通过 `table()` 指定表名,则使用 `M::table_name()`。
254    ///
255    /// # 示例
256    ///
257    /// ```ignore
258    /// use sz_orm_core::query::QueryBuilder;
259    /// use sz_orm_core::dialect::MySqlDialect;
260    /// use sz_orm_core::model::Model;
261    ///
262    /// #[derive(Default)]
263    /// struct User;
264    /// impl Model for User {
265    ///     type PrimaryKey = i64;
266    ///     fn table_name() -> &'static str { "users" }
267    ///     fn pk(&self) -> Self::PrimaryKey { 0 }
268    ///     fn set_pk(&mut self, _: Self::PrimaryKey) {}
269    /// }
270    ///
271    /// let sql = QueryBuilder::<User>::new(Box::new(MySqlDialect))
272    ///     .select(vec!["id", "name"])
273    ///     .where_cond("age > 18")
274    ///     .order_by("id DESC")
275    ///     .limit(10)
276    ///     .build_select();
277    /// // sql => "SELECT id, name FROM `users` WHERE age > 18 ORDER BY id DESC LIMIT 10"
278    /// ```
279    #[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    /// 构建 WHERE 子句(处理所有条件类型:And/Or/In/NotIn/Between/Null 等)
384    /// 返回空字符串表示无 WHERE 子句
385    fn build_where_clause(&self) -> String {
386        if self.where_conditions.is_empty() {
387            return String::new();
388        }
389
390        // 将每个条件转换为字符串,OR 条件标记前缀
391        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                    // v0.2.2 修复 H-1:使用方言感知的转义
399                    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        // OR 分组逻辑:将相邻的 OR 条件组合成 (cond1 OR cond2) 形式
440        // 边界处理:如果第一个条件就是 OR(不合理但需防御),当作 AND 处理
441        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                // OR 条件:无论是否首个,都把 OR 前缀去掉当作普通条件加入当前组
446                current_group.push(stripped.to_string());
447            } else {
448                // AND 条件:如果当前组非空,先保存
449                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        // v0.2.2 修复 H-1:使用方言感知的转义
486        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    // ===================== 参数绑定版本(v1.1.0 新增) =====================
544
545    /// 构建 WHERE 子句(参数绑定版本)。
546    ///
547    /// 将 `In`/`NotIn`/`Between`/`NotBetween` 条件中的值替换为 `?` 占位符,
548    /// 值收集到 `params` 向量中。`And`/`Or`/`Exists` 等原始字符串条件不提取参数
549    /// (调用方负责安全)。
550    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        // OR 分组逻辑:与 build_where_clause 相同
599        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    /// 构建 SELECT SQL(参数绑定版本)。
630    ///
631    /// WHERE 子句中的值使用 `?` 占位符,值通过 `params` 返回。
632    /// 适用于 `Connection::query_with_params()`。
633    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    /// 构建 INSERT SQL(参数绑定版本)。
738    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    /// 构建 UPDATE SQL(参数绑定版本)。
768    /// 参数顺序:SET 参数在前,WHERE 参数在后。
769    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    /// 构建 DELETE SQL(参数绑定版本)。
804    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    /// 校验生成的 SELECT SQL 语句
908    /// 检查 SQL 语法、JOIN 列名、表名合法性
909    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        // 校验 JOIN 子句产生的 SQL 是否合法
918        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        // 校验表名合法性
937        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    /// 校验生成的 INSERT SQL 语句
953    /// 含空数据检测(EmptyInsertData 错误)
954    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    /// 校验生成的 UPDATE SQL 语句
978    /// 含空数据检测(EmptyUpdateData 错误)
979    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    /// 校验生成的 DELETE SQL 语句
1003    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        // DELETE without WHERE still produces valid SQL (just no filter)
1365        let result = builder.table("users").validate_delete();
1366        assert!(result.is_ok());
1367    }
1368
1369    // ==================== M-3 select_quoted 测试 ====================
1370
1371    #[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        // 应自动 quote 列名
1381        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        // SQL 注入尝试:分号 + DROP TABLE
1391        let result = builder
1392            .table("users")
1393            .select_quoted(vec!["id; DROP TABLE users"]);
1394        assert!(result.is_err());
1395
1396        // 含引号
1397        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        // 数字开头
1403        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        // 含空格
1409        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        // PostgreSQL 使用双引号
1425        assert!(sql.contains("SELECT \"id\", \"name\" FROM"));
1426        assert!(sql.contains("\"users\""));
1427    }
1428}