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    pub fn build_count(&self) -> String {
544        let table = self
545            .table
546            .clone()
547            .unwrap_or_else(|| M::table_name().to_string());
548
549        let mut sql = format!(
550            "SELECT COUNT(*) as total FROM {}",
551            self.dialect.quote(&table)
552        );
553        sql.push_str(&self.build_where_clause());
554        sql
555    }
556
557    pub fn build_exists(&self) -> String {
558        let table = self
559            .table
560            .clone()
561            .unwrap_or_else(|| M::table_name().to_string());
562
563        let mut sql = format!("SELECT 1 FROM {}", self.dialect.quote(&table));
564        sql.push_str(&self.build_where_clause());
565        sql.push_str(" LIMIT 1");
566        format!("SELECT EXISTS({})", sql)
567    }
568
569    pub fn build_max(&self, field: &str) -> String {
570        let table = self
571            .table
572            .clone()
573            .unwrap_or_else(|| M::table_name().to_string());
574
575        let mut sql = format!(
576            "SELECT MAX({}) as max_val FROM {}",
577            self.dialect.quote(field),
578            self.dialect.quote(&table)
579        );
580        sql.push_str(&self.build_where_clause());
581        sql
582    }
583
584    pub fn build_min(&self, field: &str) -> String {
585        let table = self
586            .table
587            .clone()
588            .unwrap_or_else(|| M::table_name().to_string());
589
590        let mut sql = format!(
591            "SELECT MIN({}) as min_val FROM {}",
592            self.dialect.quote(field),
593            self.dialect.quote(&table)
594        );
595        sql.push_str(&self.build_where_clause());
596        sql
597    }
598
599    pub fn build_sum(&self, field: &str) -> String {
600        let table = self
601            .table
602            .clone()
603            .unwrap_or_else(|| M::table_name().to_string());
604
605        let mut sql = format!(
606            "SELECT SUM({}) as sum_val FROM {}",
607            self.dialect.quote(field),
608            self.dialect.quote(&table)
609        );
610        sql.push_str(&self.build_where_clause());
611        sql
612    }
613
614    pub fn build_avg(&self, field: &str) -> String {
615        let table = self
616            .table
617            .clone()
618            .unwrap_or_else(|| M::table_name().to_string());
619
620        let mut sql = format!(
621            "SELECT AVG({}) as avg_val FROM {}",
622            self.dialect.quote(field),
623            self.dialect.quote(&table)
624        );
625        sql.push_str(&self.build_where_clause());
626        sql
627    }
628
629    /// 校验生成的 SELECT SQL 语句
630    /// 检查 SQL 语法、JOIN 列名、表名合法性
631    pub fn validate(&self) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
632        let sql = self.build_select();
633        let mut errors = Vec::new();
634
635        if let Err(e) = sz_orm_sql_validator::validate_select(&sql) {
636            errors.push(e);
637        }
638
639        // 校验 JOIN 子句产生的 SQL 是否合法
640        if !self.joins.is_empty() {
641            for join in &self.joins {
642                match join {
643                    JoinClause::Inner(_, left, right)
644                    | JoinClause::Left(_, left, right)
645                    | JoinClause::Right(_, left, right) => {
646                        if let Err(e) = sz_orm_sql_validator::validate_column_name(left) {
647                            errors.push(e);
648                        }
649                        if let Err(e) = sz_orm_sql_validator::validate_column_name(right) {
650                            errors.push(e);
651                        }
652                    }
653                    _ => {}
654                }
655            }
656        }
657
658        // 校验表名合法性
659        let table = self
660            .table
661            .clone()
662            .unwrap_or_else(|| M::table_name().to_string());
663        if let Err(e) = sz_orm_sql_validator::validate_table_name(&table) {
664            errors.push(e);
665        }
666
667        if errors.is_empty() {
668            Ok(())
669        } else {
670            Err(errors)
671        }
672    }
673
674    /// 校验生成的 INSERT SQL 语句
675    /// 含空数据检测(EmptyInsertData 错误)
676    pub fn validate_insert(
677        &self,
678        data: &std::collections::HashMap<String, Value>,
679    ) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
680        let sql = self.build_insert(data);
681        let mut errors = Vec::new();
682
683        if sql.is_empty() {
684            errors.push(sz_orm_sql_validator::SqlValidationError::EmptyInsertData);
685            return Err(errors);
686        }
687
688        if let Err(e) = sz_orm_sql_validator::validate_insert(&sql) {
689            errors.push(e);
690        }
691
692        if errors.is_empty() {
693            Ok(())
694        } else {
695            Err(errors)
696        }
697    }
698
699    /// 校验生成的 UPDATE SQL 语句
700    /// 含空数据检测(EmptyUpdateData 错误)
701    pub fn validate_update(
702        &self,
703        data: &std::collections::HashMap<String, Value>,
704    ) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
705        let sql = self.build_update(data);
706        let mut errors = Vec::new();
707
708        if sql.is_empty() {
709            errors.push(sz_orm_sql_validator::SqlValidationError::EmptyUpdateData);
710            return Err(errors);
711        }
712
713        if let Err(e) = sz_orm_sql_validator::validate_update(&sql) {
714            errors.push(e);
715        }
716
717        if errors.is_empty() {
718            Ok(())
719        } else {
720            Err(errors)
721        }
722    }
723
724    /// 校验生成的 DELETE SQL 语句
725    pub fn validate_delete(&self) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
726        let sql = self.build_delete();
727        let mut errors = Vec::new();
728
729        if let Err(e) = sz_orm_sql_validator::validate_delete(&sql) {
730            errors.push(e);
731        }
732
733        if errors.is_empty() {
734            Ok(())
735        } else {
736            Err(errors)
737        }
738    }
739}
740
741impl<M: Model> fmt::Debug for QueryBuilder<M> {
742    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
743        f.debug_struct("QueryBuilder")
744            .field("table", &self.table)
745            .field("select_columns", &self.select_columns)
746            .field("where_conditions", &self.where_conditions.len())
747            .field("limit", &self.limit_value)
748            .finish()
749    }
750}
751
752#[cfg(test)]
753mod tests {
754    use super::*;
755    use crate::db_type::DbType;
756    use crate::dialect::get_dialect;
757
758    struct TestModel;
759    impl Model for TestModel {
760        type PrimaryKey = i64;
761
762        fn table_name() -> &'static str {
763            "test_models"
764        }
765
766        fn pk(&self) -> Self::PrimaryKey {
767            1
768        }
769
770        fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
771    }
772
773    #[test]
774    fn test_query_builder_select() {
775        let dialect = get_dialect(DbType::MySQL).unwrap();
776        let builder = QueryBuilder::<TestModel>::new(dialect);
777
778        let sql = builder
779            .table("users")
780            .select(vec!["id", "name"])
781            .build_select();
782        assert!(sql.contains("SELECT id, name FROM"));
783        assert!(sql.contains("`users`"));
784    }
785
786    #[test]
787    fn test_query_builder_where() {
788        let dialect = get_dialect(DbType::MySQL).unwrap();
789        let builder = QueryBuilder::<TestModel>::new(dialect);
790
791        let sql = builder
792            .table("users")
793            .where_cond("status = 'active'")
794            .where_cond("age > 18")
795            .build_select();
796
797        assert!(sql.contains("WHERE"));
798        assert!(sql.contains("status = 'active'"));
799        assert!(sql.contains("age > 18"));
800    }
801
802    #[test]
803    fn test_query_builder_order_by() {
804        let dialect = get_dialect(DbType::MySQL).unwrap();
805        let builder = QueryBuilder::<TestModel>::new(dialect);
806
807        let sql = builder
808            .table("users")
809            .order_by("created_at")
810            .order_desc("id")
811            .build_select();
812
813        assert!(sql.contains("ORDER BY"));
814        assert!(sql.contains("`created_at` ASC"));
815        assert!(sql.contains("`id` DESC"));
816    }
817
818    #[test]
819    fn test_query_builder_limit_offset() {
820        let dialect = get_dialect(DbType::MySQL).unwrap();
821        let builder = QueryBuilder::<TestModel>::new(dialect);
822
823        let sql = builder.table("users").limit(10).offset(20).build_select();
824
825        assert!(sql.contains("LIMIT 10"));
826        assert!(sql.contains("OFFSET 20"));
827    }
828
829    #[test]
830    fn test_query_builder_page() {
831        let dialect = get_dialect(DbType::MySQL).unwrap();
832        let builder = QueryBuilder::<TestModel>::new(dialect);
833
834        let sql = builder.table("users").page(3, 20).build_select();
835
836        assert!(sql.contains("LIMIT 20"));
837        assert!(sql.contains("OFFSET 40"));
838    }
839
840    #[test]
841    fn test_query_builder_insert() {
842        let dialect = get_dialect(DbType::MySQL).unwrap();
843        let builder = QueryBuilder::<TestModel>::new(dialect);
844
845        let mut data = std::collections::HashMap::new();
846        data.insert("name".to_string(), Value::String("test".to_string()));
847        data.insert("age".to_string(), Value::I64(25));
848
849        let sql = builder.table("users").build_insert(&data);
850
851        assert!(sql.contains("INSERT INTO"));
852        assert!(sql.contains("`name`"));
853        assert!(sql.contains("'test'"));
854    }
855
856    #[test]
857    fn test_query_builder_update() {
858        let dialect = get_dialect(DbType::MySQL).unwrap();
859        let builder = QueryBuilder::<TestModel>::new(dialect);
860
861        let mut data = std::collections::HashMap::new();
862        data.insert("name".to_string(), Value::String("updated".to_string()));
863
864        let sql = builder
865            .table("users")
866            .where_cond("id = 1")
867            .build_update(&data);
868
869        assert!(sql.contains("UPDATE"));
870        assert!(sql.contains("`name` = 'updated'"));
871        assert!(sql.contains("WHERE"));
872    }
873
874    #[test]
875    fn test_query_builder_delete() {
876        let dialect = get_dialect(DbType::MySQL).unwrap();
877        let builder = QueryBuilder::<TestModel>::new(dialect);
878
879        let sql = builder.table("users").where_cond("id = 1").build_delete();
880
881        assert!(sql.contains("DELETE FROM"));
882        assert!(sql.contains("WHERE"));
883    }
884
885    #[test]
886    fn test_query_builder_count() {
887        let dialect = get_dialect(DbType::MySQL).unwrap();
888        let builder = QueryBuilder::<TestModel>::new(dialect);
889
890        let sql = builder.table("users").build_count();
891
892        assert!(sql.contains("SELECT COUNT(*)"));
893        assert!(sql.contains("FROM"));
894    }
895
896    #[test]
897    fn test_query_builder_where_in() {
898        let dialect = get_dialect(DbType::MySQL).unwrap();
899        let builder = QueryBuilder::<TestModel>::new(dialect);
900
901        let sql = builder
902            .table("users")
903            .where_in("id", vec![Value::I64(1), Value::I64(2), Value::I64(3)])
904            .build_select();
905
906        assert!(sql.contains("IN ("));
907    }
908
909    #[test]
910    fn test_query_builder_where_between() {
911        let dialect = get_dialect(DbType::MySQL).unwrap();
912        let builder = QueryBuilder::<TestModel>::new(dialect);
913
914        let sql = builder
915            .table("users")
916            .where_between("age", Value::I64(18), Value::I64(30))
917            .build_select();
918
919        assert!(sql.contains("BETWEEN"));
920    }
921
922    #[test]
923    fn test_query_builder_where_null() {
924        let dialect = get_dialect(DbType::MySQL).unwrap();
925        let builder = QueryBuilder::<TestModel>::new(dialect);
926
927        let sql = builder
928            .table("users")
929            .where_null("deleted_at")
930            .build_select();
931
932        assert!(sql.contains("IS NULL"));
933    }
934
935    #[test]
936    fn test_query_builder_join() {
937        let dialect = get_dialect(DbType::MySQL).unwrap();
938        let builder = QueryBuilder::<TestModel>::new(dialect);
939
940        let sql = builder
941            .table("users")
942            .join_inner("posts", "users.id", "posts.user_id")
943            .build_select();
944
945        assert!(sql.contains("INNER JOIN"));
946        assert!(sql.contains("`posts`"));
947    }
948
949    #[test]
950    fn test_query_builder_group_by() {
951        let dialect = get_dialect(DbType::MySQL).unwrap();
952        let builder = QueryBuilder::<TestModel>::new(dialect);
953
954        let sql = builder.table("users").group_by("status").build_select();
955
956        assert!(sql.contains("GROUP BY"));
957        assert!(sql.contains("`status`"));
958    }
959
960    #[test]
961    fn test_query_builder_max() {
962        let dialect = get_dialect(DbType::MySQL).unwrap();
963        let builder = QueryBuilder::<TestModel>::new(dialect);
964
965        let sql = builder.table("users").build_max("score");
966
967        assert!(sql.contains("MAX("));
968        assert!(sql.contains("`score`"));
969    }
970
971    #[test]
972    fn test_query_builder_min() {
973        let dialect = get_dialect(DbType::MySQL).unwrap();
974        let builder = QueryBuilder::<TestModel>::new(dialect);
975
976        let sql = builder.table("users").build_min("price");
977
978        assert!(sql.contains("MIN("));
979        assert!(sql.contains("`price`"));
980    }
981
982    #[test]
983    fn test_query_builder_sum() {
984        let dialect = get_dialect(DbType::MySQL).unwrap();
985        let builder = QueryBuilder::<TestModel>::new(dialect);
986
987        let sql = builder.table("orders").build_sum("amount");
988
989        assert!(sql.contains("SUM("));
990        assert!(sql.contains("`amount`"));
991    }
992
993    #[test]
994    fn test_query_builder_avg() {
995        let dialect = get_dialect(DbType::MySQL).unwrap();
996        let builder = QueryBuilder::<TestModel>::new(dialect);
997
998        let sql = builder.table("scores").build_avg("value");
999
1000        assert!(sql.contains("AVG("));
1001        assert!(sql.contains("`value`"));
1002    }
1003
1004    #[test]
1005    fn test_validator_select() {
1006        let dialect = get_dialect(DbType::MySQL).unwrap();
1007        let builder = QueryBuilder::<TestModel>::new(dialect);
1008
1009        let result = builder.table("users").select(vec!["id", "name"]).validate();
1010        assert!(result.is_ok());
1011    }
1012
1013    #[test]
1014    fn test_validator_select_with_join() {
1015        let dialect = get_dialect(DbType::MySQL).unwrap();
1016        let builder = QueryBuilder::<TestModel>::new(dialect);
1017
1018        let result = builder
1019            .table("users")
1020            .join_inner("posts", "users.id", "posts.user_id")
1021            .validate();
1022        assert!(result.is_ok());
1023    }
1024
1025    #[test]
1026    fn test_validator_insert() {
1027        let dialect = get_dialect(DbType::MySQL).unwrap();
1028        let builder = QueryBuilder::<TestModel>::new(dialect);
1029
1030        let mut data = std::collections::HashMap::new();
1031        data.insert("name".to_string(), Value::String("test".to_string()));
1032
1033        let result = builder.table("users").validate_insert(&data);
1034        assert!(result.is_ok());
1035    }
1036
1037    #[test]
1038    fn test_validator_insert_empty_data() {
1039        let dialect = get_dialect(DbType::MySQL).unwrap();
1040        let builder = QueryBuilder::<TestModel>::new(dialect);
1041
1042        let data = std::collections::HashMap::new();
1043        let result = builder.table("users").validate_insert(&data);
1044        assert!(result.is_err());
1045    }
1046
1047    #[test]
1048    fn test_validator_update() {
1049        let dialect = get_dialect(DbType::MySQL).unwrap();
1050        let builder = QueryBuilder::<TestModel>::new(dialect);
1051
1052        let mut data = std::collections::HashMap::new();
1053        data.insert("name".to_string(), Value::String("updated".to_string()));
1054
1055        let result = builder.table("users").validate_update(&data);
1056        assert!(result.is_ok());
1057    }
1058
1059    #[test]
1060    fn test_validator_update_empty_data() {
1061        let dialect = get_dialect(DbType::MySQL).unwrap();
1062        let builder = QueryBuilder::<TestModel>::new(dialect);
1063
1064        let data = std::collections::HashMap::new();
1065        let result = builder.table("users").validate_update(&data);
1066        assert!(result.is_err());
1067    }
1068
1069    #[test]
1070    fn test_validator_delete() {
1071        let dialect = get_dialect(DbType::MySQL).unwrap();
1072        let builder = QueryBuilder::<TestModel>::new(dialect);
1073
1074        let result = builder
1075            .table("users")
1076            .where_cond("id = 1")
1077            .validate_delete();
1078        assert!(result.is_ok());
1079    }
1080
1081    #[test]
1082    fn test_validator_delete_no_where() {
1083        let dialect = get_dialect(DbType::MySQL).unwrap();
1084        let builder = QueryBuilder::<TestModel>::new(dialect);
1085
1086        // DELETE without WHERE still produces valid SQL (just no filter)
1087        let result = builder.table("users").validate_delete();
1088        assert!(result.is_ok());
1089    }
1090
1091    // ==================== M-3 select_quoted 测试 ====================
1092
1093    #[test]
1094    fn test_m3_select_quoted_valid_columns() {
1095        let dialect = get_dialect(DbType::MySQL).unwrap();
1096        let builder = QueryBuilder::<TestModel>::new(dialect);
1097        let builder = builder
1098            .table("users")
1099            .select_quoted(vec!["id", "name"])
1100            .expect("valid columns should succeed");
1101        let sql = builder.build_select();
1102        // 应自动 quote 列名
1103        assert!(sql.contains("SELECT `id`, `name` FROM"));
1104        assert!(sql.contains("`users`"));
1105    }
1106
1107    #[test]
1108    fn test_m3_select_quoted_rejects_sql_injection() {
1109        let dialect = get_dialect(DbType::MySQL).unwrap();
1110        let builder = QueryBuilder::<TestModel>::new(dialect);
1111
1112        // SQL 注入尝试:分号 + DROP TABLE
1113        let result = builder
1114            .table("users")
1115            .select_quoted(vec!["id; DROP TABLE users"]);
1116        assert!(result.is_err());
1117
1118        // 含引号
1119        let dialect = get_dialect(DbType::MySQL).unwrap();
1120        let builder = QueryBuilder::<TestModel>::new(dialect);
1121        let result = builder.table("users").select_quoted(vec!["name'"]);
1122        assert!(result.is_err());
1123
1124        // 数字开头
1125        let dialect = get_dialect(DbType::MySQL).unwrap();
1126        let builder = QueryBuilder::<TestModel>::new(dialect);
1127        let result = builder.table("users").select_quoted(vec!["1col"]);
1128        assert!(result.is_err());
1129
1130        // 含空格
1131        let dialect = get_dialect(DbType::MySQL).unwrap();
1132        let builder = QueryBuilder::<TestModel>::new(dialect);
1133        let result = builder.table("users").select_quoted(vec!["col name"]);
1134        assert!(result.is_err());
1135    }
1136
1137    #[test]
1138    fn test_m3_select_quoted_postgresql_dialect() {
1139        let dialect = get_dialect(DbType::PostgreSQL).unwrap();
1140        let builder = QueryBuilder::<TestModel>::new(dialect);
1141        let builder = builder
1142            .table("users")
1143            .select_quoted(vec!["id", "name"])
1144            .expect("valid columns should succeed");
1145        let sql = builder.build_select();
1146        // PostgreSQL 使用双引号
1147        assert!(sql.contains("SELECT \"id\", \"name\" FROM"));
1148        assert!(sql.contains("\"users\""));
1149    }
1150}