Skip to main content

sz_orm_core/
lambda.rs

1//! Lambda 类型安全 Wrapper
2//!
3//! 对应文档 6.8 节改进项 38(Lambda 类型安全 Wrapper)。
4//!
5//! # 核心概念
6//!
7//! - **Column**:字段标记 trait,关联 Model 类型 M,提供字段名与表名
8//! - **`LambdaWrapper<M>`**:类型安全的查询构造器,所有字段引用都通过 `Column` 类型而非 `&str`
9//! - **define_columns!**:宏,为 Model 定义所有字段的类型安全标记
10//!
11//! # 设计灵感
12//!
13//! - MyBatis-Plus `LambdaQueryWrapper` / `LambdaUpdateWrapper`
14//! - JOOQ 类型安全 DSL
15//! - Diesel 强类型 schema
16//! - SeaORM `Column` trait
17//!
18//! # 优势
19//!
20//! 1. **编译期字段名检查**:拼写错误直接编译失败
21//! 2. **IDE 自动补全**:`UserColumns::` 后可补全所有字段
22//! 3. **重构友好**:字段重命名后所有引用处编译失败,便于定位修改
23//! 4. **表名隔离**:不同 Model 的字段不会混淆
24//!
25//! # 使用示例
26//!
27//! ```
28//! use sz_orm_core::lambda::{LambdaWrapper, Column};
29//! use sz_orm_core::define_columns;
30//! use sz_orm_core::Value;
31//!
32//! // 1. 定义 Model(示例)
33//! struct User;
34//!
35//! // 2. 为 User 定义字段标记
36//! define_columns! {
37//!     UserColumns for User table = "users" {
38//!         Id => "id",
39//!         Name => "name",
40//!         Age => "age",
41//!     }
42//! }
43//!
44//! // 3. 使用 LambdaWrapper 构造类型安全查询
45//! let mut wrapper = LambdaWrapper::<User>::new("users");
46//! wrapper
47//!     .select(UserColumns::Id)
48//!     .select(UserColumns::Name)
49//!     .eq(UserColumns::Id, Value::I64(1))
50//!     .gt(UserColumns::Age, Value::I64(18));
51//!
52//! let sql = wrapper.build_select();
53//! assert!(sql.contains("SELECT `id`, `name`"));
54//! assert!(sql.contains("FROM `users`"));
55//! assert!(sql.contains("`id` = 1"));
56//! assert!(sql.contains("`age` > 18"));
57//! ```
58
59use crate::dialect::{Dialect, MySqlDialect};
60use crate::Value;
61use std::marker::PhantomData;
62
63// ============================================================================
64// Column trait — 字段标记
65// ============================================================================
66
67/// 字段标记 trait
68///
69/// 实现此 trait 的类型作为 Model 字段的类型安全引用。
70/// 一个 `Column<M>` 实例携带:
71/// - 字段名(`name()`)
72/// - 所属表名(`table()`,从 Model 关联)
73///
74/// # 实现方式
75///
76/// 通常通过 `define_columns!` 宏自动生成实现,无需手动实现。
77pub trait Column<M>: Send + Sync + Clone {
78    /// 返回字段名
79    fn name(&self) -> &'static str;
80
81    /// 返回字段所属的表名
82    fn table(&self) -> &'static str;
83}
84
85// ============================================================================
86// WhereClause — WHERE 条件子句
87// ============================================================================
88
89/// WHERE 条件子句(内部表示)
90#[derive(Debug, Clone)]
91pub enum WhereClause {
92    /// `col = value`
93    Eq(String, Value),
94    /// `col != value`
95    Ne(String, Value),
96    /// `col > value`
97    Gt(String, Value),
98    /// `col >= value`
99    Ge(String, Value),
100    /// `col < value`
101    Lt(String, Value),
102    /// `col <= value`
103    Le(String, Value),
104    /// `col LIKE value`
105    Like(String, Value),
106    /// `col IS NULL`
107    IsNull(String),
108    /// `col IS NOT NULL`
109    IsNotNull(String),
110    /// `col IN (v1, v2, ...)`
111    In(String, Vec<Value>),
112    /// `col NOT IN (v1, v2, ...)`
113    NotIn(String, Vec<Value>),
114    /// `col BETWEEN v1 AND v2`
115    Between(String, Value, Value),
116    /// 原始 SQL(用于 OR 等复杂条件)
117    Raw(String),
118}
119
120impl WhereClause {
121    /// 渲染为 SQL 片段(不带前缀 AND/OR)
122    fn render(&self, dialect: &dyn Dialect) -> String {
123        match self {
124            // v0.2.2 修复 H-1:使用方言感知的转义
125            WhereClause::Eq(col, v) => format!(
126                "{} = {}",
127                dialect.quote(col),
128                v.to_param_with_dialect(dialect)
129            ),
130            WhereClause::Ne(col, v) => format!(
131                "{} != {}",
132                dialect.quote(col),
133                v.to_param_with_dialect(dialect)
134            ),
135            WhereClause::Gt(col, v) => format!(
136                "{} > {}",
137                dialect.quote(col),
138                v.to_param_with_dialect(dialect)
139            ),
140            WhereClause::Ge(col, v) => format!(
141                "{} >= {}",
142                dialect.quote(col),
143                v.to_param_with_dialect(dialect)
144            ),
145            WhereClause::Lt(col, v) => format!(
146                "{} < {}",
147                dialect.quote(col),
148                v.to_param_with_dialect(dialect)
149            ),
150            WhereClause::Le(col, v) => format!(
151                "{} <= {}",
152                dialect.quote(col),
153                v.to_param_with_dialect(dialect)
154            ),
155            WhereClause::Like(col, v) => format!(
156                "{} LIKE {}",
157                dialect.quote(col),
158                v.to_param_with_dialect(dialect)
159            ),
160            WhereClause::IsNull(col) => format!("{} IS NULL", dialect.quote(col)),
161            WhereClause::IsNotNull(col) => format!("{} IS NOT NULL", dialect.quote(col)),
162            WhereClause::In(col, vs) => {
163                let values: Vec<String> = vs
164                    .iter()
165                    .map(|v| v.to_param_with_dialect(dialect).to_string())
166                    .collect();
167                format!("{} IN ({})", dialect.quote(col), values.join(", "))
168            }
169            WhereClause::NotIn(col, vs) => {
170                let values: Vec<String> = vs
171                    .iter()
172                    .map(|v| v.to_param_with_dialect(dialect).to_string())
173                    .collect();
174                format!("{} NOT IN ({})", dialect.quote(col), values.join(", "))
175            }
176            WhereClause::Between(col, a, b) => format!(
177                "{} BETWEEN {} AND {}",
178                dialect.quote(col),
179                a.to_param_with_dialect(dialect),
180                b.to_param_with_dialect(dialect)
181            ),
182            WhereClause::Raw(sql) => sql.clone(),
183        }
184    }
185}
186
187// ============================================================================
188// OrderBy — 排序子句
189// ============================================================================
190
191/// 排序方向
192#[derive(Debug, Clone, Copy, PartialEq, Eq)]
193pub enum OrderDirection {
194    /// 升序
195    Asc,
196    /// 降序
197    Desc,
198}
199
200/// 排序子句
201#[derive(Debug, Clone)]
202pub struct OrderBy {
203    /// 字段名
204    pub column: String,
205    /// 排序方向
206    pub direction: OrderDirection,
207}
208
209// ============================================================================
210// LambdaWrapper — 类型安全查询构造器
211// ============================================================================
212
213/// Lambda 类型安全查询构造器
214///
215/// 泛型参数 `M` 是 Model 类型(仅用于类型隔离,不实际存储实例)。
216///
217/// # 字段引用方式
218///
219/// 与原始 `QueryBuilder` 使用 `&str` 字段名不同,`LambdaWrapper` 接受 `Column<M>` 实例,
220/// 从而在编译期检查字段拼写错误。
221///
222/// # 示例
223///
224/// ```
225/// use sz_orm_core::lambda::{LambdaWrapper, Column};
226/// use sz_orm_core::define_columns;
227/// use sz_orm_core::Value;
228///
229/// struct User;
230/// define_columns! {
231///     UserColumns for User table = "users" {
232///         Id => "id",
233///         Name => "name",
234///     }
235/// }
236///
237/// let mut w = LambdaWrapper::<User>::new("users");
238/// w.eq(UserColumns::Id, Value::I64(1));
239/// ```
240pub struct LambdaWrapper<M> {
241    /// 表名
242    table: String,
243    /// SELECT 字段列表(空表示 SELECT *)
244    selects: Vec<String>,
245    /// WHERE 条件列表(AND 连接)
246    wheres: Vec<WhereClause>,
247    /// ORDER BY 子句
248    orders: Vec<OrderBy>,
249    /// LIMIT
250    limit: Option<u64>,
251    /// OFFSET
252    offset: Option<u64>,
253    /// 数据库方言
254    dialect: Box<dyn Dialect>,
255    _marker: PhantomData<M>,
256}
257
258impl<M> LambdaWrapper<M> {
259    /// 创建 LambdaWrapper,默认使用 MySQL 方言
260    pub fn new(table: impl Into<String>) -> Self {
261        Self {
262            table: table.into(),
263            selects: Vec::new(),
264            wheres: Vec::new(),
265            orders: Vec::new(),
266            limit: None,
267            offset: None,
268            dialect: Box::new(MySqlDialect),
269            _marker: PhantomData,
270        }
271    }
272
273    /// 创建 LambdaWrapper 并指定方言
274    pub fn with_dialect(table: impl Into<String>, dialect: Box<dyn Dialect>) -> Self {
275        Self {
276            table: table.into(),
277            selects: Vec::new(),
278            wheres: Vec::new(),
279            orders: Vec::new(),
280            limit: None,
281            offset: None,
282            dialect,
283            _marker: PhantomData,
284        }
285    }
286
287    // -------------------- SELECT 字段 --------------------
288
289    /// 添加 SELECT 字段(类型安全)
290    ///
291    /// M-4 修复:对列名进行 `validate_identifier` 校验,防止恶意实现 `Column` trait
292    /// 注入非法标识符。
293    pub fn select<C: Column<M>>(&mut self, col: C) -> &mut Self {
294        let name = col.name();
295        // 校验列名为合法 SQL 标识符(非空、仅 ASCII 字母数字+下划线、不以数字开头)
296        // 校验失败时跳过该列(保留向后兼容,不中断调用链)
297        if crate::sql_safety::validate_identifier(name, "lambda select column").is_ok() {
298            self.selects.push(name.to_string());
299        }
300        self
301    }
302
303    /// 批量添加 SELECT 字段
304    ///
305    /// M-4 修复:同 `select`,对每个列名校验。
306    pub fn select_many<C: Column<M>>(&mut self, cols: &[C]) -> &mut Self {
307        for c in cols {
308            let name = c.name();
309            if crate::sql_safety::validate_identifier(name, "lambda select column").is_ok() {
310                self.selects.push(name.to_string());
311            }
312        }
313        self
314    }
315
316    /// SELECT *(清空已有字段选择)
317    pub fn select_all(&mut self) -> &mut Self {
318        self.selects.clear();
319        self
320    }
321
322    // -------------------- WHERE 条件 --------------------
323
324    /// `col = value`
325    pub fn eq<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
326        self.wheres
327            .push(WhereClause::Eq(col.name().to_string(), value));
328        self
329    }
330
331    /// `col != value`
332    pub fn ne<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
333        self.wheres
334            .push(WhereClause::Ne(col.name().to_string(), value));
335        self
336    }
337
338    /// `col > value`
339    pub fn gt<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
340        self.wheres
341            .push(WhereClause::Gt(col.name().to_string(), value));
342        self
343    }
344
345    /// `col >= value`
346    pub fn ge<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
347        self.wheres
348            .push(WhereClause::Ge(col.name().to_string(), value));
349        self
350    }
351
352    /// `col < value`
353    pub fn lt<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
354        self.wheres
355            .push(WhereClause::Lt(col.name().to_string(), value));
356        self
357    }
358
359    /// `col <= value`
360    pub fn le<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
361        self.wheres
362            .push(WhereClause::Le(col.name().to_string(), value));
363        self
364    }
365
366    /// `col LIKE value`
367    pub fn like<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
368        self.wheres
369            .push(WhereClause::Like(col.name().to_string(), value));
370        self
371    }
372
373    /// `col IS NULL`
374    pub fn is_null<C: Column<M>>(&mut self, col: C) -> &mut Self {
375        self.wheres
376            .push(WhereClause::IsNull(col.name().to_string()));
377        self
378    }
379
380    /// `col IS NOT NULL`
381    pub fn is_not_null<C: Column<M>>(&mut self, col: C) -> &mut Self {
382        self.wheres
383            .push(WhereClause::IsNotNull(col.name().to_string()));
384        self
385    }
386
387    /// `col IN (v1, v2, ...)`
388    pub fn r#in<C: Column<M>>(&mut self, col: C, values: Vec<Value>) -> &mut Self {
389        self.wheres
390            .push(WhereClause::In(col.name().to_string(), values));
391        self
392    }
393
394    /// `col NOT IN (v1, v2, ...)`
395    pub fn not_in<C: Column<M>>(&mut self, col: C, values: Vec<Value>) -> &mut Self {
396        self.wheres
397            .push(WhereClause::NotIn(col.name().to_string(), values));
398        self
399    }
400
401    /// `col BETWEEN a AND b`
402    pub fn between<C: Column<M>>(&mut self, col: C, a: Value, b: Value) -> &mut Self {
403        self.wheres
404            .push(WhereClause::Between(col.name().to_string(), a, b));
405        self
406    }
407
408    /// 追加原始 SQL WHERE 条件(用于 OR 等复杂场景)
409    ///
410    /// # 安全警告
411    ///
412    /// 此方法是 escape hatch,传入的 SQL 会**原样拼接**到最终 SQL 中。
413    /// **严禁将用户输入直接拼接**到 `sql` 参数中(会引入 SQL 注入风险)。
414    /// 若需使用用户输入,请改用 `eq` / `ne` / `lt` 等参数化方法。
415    pub fn raw_where(&mut self, sql: impl Into<String>) -> &mut Self {
416        self.wheres.push(WhereClause::Raw(sql.into()));
417        self
418    }
419
420    // -------------------- ORDER BY / LIMIT / OFFSET --------------------
421
422    /// 添加升序排序
423    pub fn order_by_asc<C: Column<M>>(&mut self, col: C) -> &mut Self {
424        self.orders.push(OrderBy {
425            column: col.name().to_string(),
426            direction: OrderDirection::Asc,
427        });
428        self
429    }
430
431    /// 添加降序排序
432    pub fn order_by_desc<C: Column<M>>(&mut self, col: C) -> &mut Self {
433        self.orders.push(OrderBy {
434            column: col.name().to_string(),
435            direction: OrderDirection::Desc,
436        });
437        self
438    }
439
440    /// 设置 LIMIT
441    pub fn limit(&mut self, n: u64) -> &mut Self {
442        self.limit = Some(n);
443        self
444    }
445
446    /// 设置 OFFSET
447    pub fn offset(&mut self, n: u64) -> &mut Self {
448        self.offset = Some(n);
449        self
450    }
451
452    /// 分页(设置 LIMIT + OFFSET)
453    ///
454    /// page 从 1 开始计数,page=1 时无 OFFSET(从第 0 条开始)。
455    pub fn page(&mut self, page: u64, page_size: u64) -> &mut Self {
456        self.limit = Some(page_size);
457        if page > 1 {
458            self.offset = Some((page - 1) * page_size);
459        } else {
460            self.offset = None;
461        }
462        self
463    }
464
465    // -------------------- SQL 生成 --------------------
466
467    /// 生成 SELECT SQL
468    pub fn build_select(&self) -> String {
469        let quoted_table = self.dialect.quote(&self.table);
470
471        // SELECT 字段
472        let select_sql = if self.selects.is_empty() {
473            "*".to_string()
474        } else {
475            self.selects
476                .iter()
477                .map(|c| self.dialect.quote(c))
478                .collect::<Vec<_>>()
479                .join(", ")
480        };
481
482        let mut sql = format!("SELECT {} FROM {}", select_sql, quoted_table);
483
484        // WHERE
485        if !self.wheres.is_empty() {
486            let conditions: Vec<String> = self
487                .wheres
488                .iter()
489                .map(|w| w.render(self.dialect.as_ref()))
490                .collect();
491            sql.push_str(" WHERE ");
492            sql.push_str(&conditions.join(" AND "));
493        }
494
495        // ORDER BY
496        if !self.orders.is_empty() {
497            let orders: Vec<String> = self
498                .orders
499                .iter()
500                .map(|o| {
501                    let dir = match o.direction {
502                        OrderDirection::Asc => "ASC",
503                        OrderDirection::Desc => "DESC",
504                    };
505                    format!("{} {}", self.dialect.quote(&o.column), dir)
506                })
507                .collect();
508            sql.push_str(" ORDER BY ");
509            sql.push_str(&orders.join(", "));
510        }
511
512        // LIMIT / OFFSET
513        if let Some(l) = self.limit {
514            sql.push_str(&format!(" LIMIT {}", l));
515        }
516        if let Some(o) = self.offset {
517            sql.push_str(&format!(" OFFSET {}", o));
518        }
519
520        sql
521    }
522
523    /// 生成 COUNT SQL(SELECT COUNT(*) FROM ... WHERE ...)
524    pub fn build_count(&self) -> String {
525        let quoted_table = self.dialect.quote(&self.table);
526        let mut sql = format!("SELECT COUNT(*) FROM {}", quoted_table);
527
528        if !self.wheres.is_empty() {
529            let conditions: Vec<String> = self
530                .wheres
531                .iter()
532                .map(|w| w.render(self.dialect.as_ref()))
533                .collect();
534            sql.push_str(" WHERE ");
535            sql.push_str(&conditions.join(" AND "));
536        }
537
538        sql
539    }
540
541    /// 生成 EXISTS SQL(SELECT EXISTS(SELECT 1 FROM ... WHERE ...) AS exists_flag)
542    pub fn build_exists(&self) -> String {
543        let inner = self.build_select();
544        // 把 SELECT 字段部分替换为 SELECT 1
545        let inner = if let Some(pos) = inner.find(" FROM ") {
546            format!("SELECT 1{}", &inner[pos..])
547        } else {
548            inner
549        };
550        format!("SELECT EXISTS({}) AS exists_flag", inner)
551    }
552
553    /// 生成 DELETE SQL
554    pub fn build_delete(&self) -> String {
555        let quoted_table = self.dialect.quote(&self.table);
556        let mut sql = format!("DELETE FROM {}", quoted_table);
557
558        if !self.wheres.is_empty() {
559            let conditions: Vec<String> = self
560                .wheres
561                .iter()
562                .map(|w| w.render(self.dialect.as_ref()))
563                .collect();
564            sql.push_str(" WHERE ");
565            sql.push_str(&conditions.join(" AND "));
566        }
567
568        sql
569    }
570
571    // -------------------- 状态查询 --------------------
572
573    /// 返回当前 WHERE 条件数量
574    pub fn where_count(&self) -> usize {
575        self.wheres.len()
576    }
577
578    /// 返回当前 SELECT 字段数量
579    pub fn select_count(&self) -> usize {
580        self.selects.len()
581    }
582
583    /// 返回表名
584    pub fn table(&self) -> &str {
585        &self.table
586    }
587
588    /// 重置所有条件(保留表名和方言)
589    pub fn reset(&mut self) -> &mut Self {
590        self.selects.clear();
591        self.wheres.clear();
592        self.orders.clear();
593        self.limit = None;
594        self.offset = None;
595        self
596    }
597}
598
599// ============================================================================
600// define_columns! 宏 — 为 Model 定义字段标记
601// ============================================================================
602
603/// 为 Model 定义一组类型安全的字段标记
604///
605/// # 示例
606///
607/// ```
608/// use sz_orm_core::define_columns;
609///
610/// struct User;
611///
612/// define_columns! {
613///     UserColumns for User table = "users" {
614///         Id => "id",
615///         Name => "name",
616///         Age => "age",
617///     }
618/// }
619/// ```
620///
621/// 生成如下代码:
622///
623/// ```rust,ignore
624/// #[derive(Clone, Copy)]
625/// pub struct UserColumns;
626///
627/// impl UserColumns {
628///     pub const Id: UserColumn = UserColumn { name: "id", table: "users" };
629///     pub const Name: UserColumn = UserColumn { name: "name", table: "users" };
630///     pub const Age: UserColumn = UserColumn { name: "age", table: "users" };
631/// }
632///
633/// #[derive(Clone, Copy)]
634/// pub struct UserColumn {
635///     pub name: &'static str,
636///     pub table: &'static str,
637/// }
638///
639/// impl Column<User> for UserColumn {
640///     fn name(&self) -> &'static str { self.name }
641///     fn table(&self) -> &'static str { self.table }
642/// }
643/// ```
644#[macro_export]
645macro_rules! define_columns {
646    (
647        $columns_struct:ident for $model:ident table = $table:literal {
648            $( $field:ident => $name:literal ),* $(,)?
649        }
650    ) => {
651        /// 字段标记类型(每个 Model 一组)
652        ///
653        /// 通过 `ModelColumns::FieldName` 引用对应字段,编译期检查拼写错误。
654        #[derive(Debug, Clone, Copy)]
655        pub struct $columns_struct {
656            /// 字段名
657            pub name: &'static str,
658            /// 表名
659            pub table: &'static str,
660        }
661
662        impl $crate::lambda::Column<$model> for $columns_struct {
663            fn name(&self) -> &'static str {
664                self.name
665            }
666
667            fn table(&self) -> &'static str {
668                self.table
669            }
670        }
671
672        impl $columns_struct {
673            // 允许 PascalCase 常量名(如 Id、Name),符合字段命名习惯
674            $(
675                #[allow(non_upper_case_globals, dead_code)]
676                pub const $field: $columns_struct = $columns_struct { name: $name, table: $table };
677            )*
678        }
679    };
680}
681
682// ============================================================================
683// 单元测试
684// ============================================================================
685
686#[cfg(test)]
687mod tests {
688    use super::*;
689    use crate::dialect::PostgreSqlDialect;
690    use crate::get_dialect;
691    use crate::DbType;
692
693    // 测试用 Model
694    struct User;
695
696    // 为 User 定义字段标记
697    define_columns! {
698        UserColumns for User table = "users" {
699            Id => "id",
700            Name => "name",
701            Age => "age",
702            Email => "email",
703        }
704    }
705
706    // 另一个 Model 用于测试表名隔离
707    struct Order;
708
709    define_columns! {
710        OrderColumns for Order table = "orders" {
711            OrderId => "order_id",
712            UserId => "user_id",
713            Total => "total",
714        }
715    }
716
717    // ===== Column trait 测试 =====
718
719    #[test]
720    fn test_column_name_and_table() {
721        assert_eq!(UserColumns::Id.name(), "id");
722        assert_eq!(UserColumns::Id.table(), "users");
723        assert_eq!(UserColumns::Name.name(), "name");
724        assert_eq!(UserColumns::Age.name(), "age");
725        assert_eq!(UserColumns::Email.name(), "email");
726    }
727
728    #[test]
729    fn test_column_for_different_models() {
730        assert_eq!(OrderColumns::OrderId.name(), "order_id");
731        assert_eq!(OrderColumns::OrderId.table(), "orders");
732        assert_eq!(OrderColumns::UserId.name(), "user_id");
733    }
734
735    // ===== LambdaWrapper 基础测试 =====
736
737    #[test]
738    fn test_new_wrapper() {
739        let w = LambdaWrapper::<User>::new("users");
740        assert_eq!(w.table(), "users");
741        assert_eq!(w.where_count(), 0);
742        assert_eq!(w.select_count(), 0);
743    }
744
745    #[test]
746    fn test_select_single() {
747        let mut w = LambdaWrapper::<User>::new("users");
748        w.select(UserColumns::Id);
749        let sql = w.build_select();
750        assert!(sql.contains("SELECT `id` FROM `users`"));
751    }
752
753    #[test]
754    fn test_select_multiple() {
755        let mut w = LambdaWrapper::<User>::new("users");
756        w.select(UserColumns::Id)
757            .select(UserColumns::Name)
758            .select(UserColumns::Age);
759        let sql = w.build_select();
760        assert!(sql.contains("`id`, `name`, `age`"));
761    }
762
763    #[test]
764    fn test_select_many() {
765        let mut w = LambdaWrapper::<User>::new("users");
766        w.select_many(&[UserColumns::Id, UserColumns::Name, UserColumns::Age]);
767        let sql = w.build_select();
768        assert!(sql.contains("`id`, `name`, `age`"));
769    }
770
771    #[test]
772    fn test_select_all_clears_selects() {
773        let mut w = LambdaWrapper::<User>::new("users");
774        w.select(UserColumns::Id);
775        assert_eq!(w.select_count(), 1);
776        w.select_all();
777        assert_eq!(w.select_count(), 0);
778        let sql = w.build_select();
779        assert!(sql.contains("SELECT * FROM"));
780    }
781
782    #[test]
783    fn test_default_select_is_star() {
784        let w = LambdaWrapper::<User>::new("users");
785        let sql = w.build_select();
786        assert!(sql.contains("SELECT * FROM `users`"));
787    }
788
789    // ===== WHERE 条件测试 =====
790
791    #[test]
792    fn test_where_eq() {
793        let mut w = LambdaWrapper::<User>::new("users");
794        w.eq(UserColumns::Id, Value::I64(1));
795        let sql = w.build_select();
796        assert!(sql.contains("WHERE `id` = 1"));
797    }
798
799    #[test]
800    fn test_where_ne() {
801        let mut w = LambdaWrapper::<User>::new("users");
802        w.ne(UserColumns::Id, Value::I64(1));
803        let sql = w.build_select();
804        assert!(sql.contains("`id` != 1"));
805    }
806
807    #[test]
808    fn test_where_gt_ge_lt_le() {
809        let mut w = LambdaWrapper::<User>::new("users");
810        w.gt(UserColumns::Age, Value::I64(18))
811            .ge(UserColumns::Age, Value::I64(20))
812            .lt(UserColumns::Age, Value::I64(65))
813            .le(UserColumns::Age, Value::I64(60));
814        let sql = w.build_select();
815        assert!(sql.contains("`age` > 18"));
816        assert!(sql.contains("`age` >= 20"));
817        assert!(sql.contains("`age` < 65"));
818        assert!(sql.contains("`age` <= 60"));
819    }
820
821    #[test]
822    fn test_where_like() {
823        let mut w = LambdaWrapper::<User>::new("users");
824        w.like(UserColumns::Name, Value::String("%alice%".to_string()));
825        let sql = w.build_select();
826        assert!(sql.contains("`name` LIKE '%alice%'"));
827    }
828
829    #[test]
830    fn test_where_is_null() {
831        let mut w = LambdaWrapper::<User>::new("users");
832        w.is_null(UserColumns::Email);
833        let sql = w.build_select();
834        assert!(sql.contains("`email` IS NULL"));
835    }
836
837    #[test]
838    fn test_where_is_not_null() {
839        let mut w = LambdaWrapper::<User>::new("users");
840        w.is_not_null(UserColumns::Email);
841        let sql = w.build_select();
842        assert!(sql.contains("`email` IS NOT NULL"));
843    }
844
845    #[test]
846    fn test_where_in() {
847        let mut w = LambdaWrapper::<User>::new("users");
848        w.r#in(
849            UserColumns::Id,
850            vec![Value::I64(1), Value::I64(2), Value::I64(3)],
851        );
852        let sql = w.build_select();
853        assert!(sql.contains("`id` IN (1, 2, 3)"));
854    }
855
856    #[test]
857    fn test_where_not_in() {
858        let mut w = LambdaWrapper::<User>::new("users");
859        w.not_in(UserColumns::Id, vec![Value::I64(1), Value::I64(2)]);
860        let sql = w.build_select();
861        assert!(sql.contains("`id` NOT IN (1, 2)"));
862    }
863
864    #[test]
865    fn test_where_between() {
866        let mut w = LambdaWrapper::<User>::new("users");
867        w.between(UserColumns::Age, Value::I64(18), Value::I64(65));
868        let sql = w.build_select();
869        assert!(sql.contains("`age` BETWEEN 18 AND 65"));
870    }
871
872    #[test]
873    fn test_where_multiple_anded() {
874        let mut w = LambdaWrapper::<User>::new("users");
875        w.eq(UserColumns::Id, Value::I64(1))
876            .gt(UserColumns::Age, Value::I64(18))
877            .like(UserColumns::Name, Value::String("alice%".to_string()));
878        let sql = w.build_select();
879        assert!(sql.contains("`id` = 1"));
880        assert!(sql.contains("`age` > 18"));
881        assert!(sql.contains("`name` LIKE 'alice%'"));
882        // 所有条件用 AND 连接
883        assert!(sql.contains(" AND "));
884    }
885
886    #[test]
887    fn test_where_raw() {
888        let mut w = LambdaWrapper::<User>::new("users");
889        w.raw_where("name = 'alice' OR name = 'bob'");
890        let sql = w.build_select();
891        assert!(sql.contains("name = 'alice' OR name = 'bob'"));
892    }
893
894    // ===== ORDER BY / LIMIT / OFFSET 测试 =====
895
896    #[test]
897    fn test_order_by_asc() {
898        let mut w = LambdaWrapper::<User>::new("users");
899        w.order_by_asc(UserColumns::Name);
900        let sql = w.build_select();
901        assert!(sql.contains("ORDER BY `name` ASC"));
902    }
903
904    #[test]
905    fn test_order_by_desc() {
906        let mut w = LambdaWrapper::<User>::new("users");
907        w.order_by_desc(UserColumns::Id);
908        let sql = w.build_select();
909        assert!(sql.contains("ORDER BY `id` DESC"));
910    }
911
912    #[test]
913    fn test_order_by_multiple() {
914        let mut w = LambdaWrapper::<User>::new("users");
915        w.order_by_asc(UserColumns::Name)
916            .order_by_desc(UserColumns::Id);
917        let sql = w.build_select();
918        assert!(sql.contains("ORDER BY `name` ASC, `id` DESC"));
919    }
920
921    #[test]
922    fn test_limit() {
923        let mut w = LambdaWrapper::<User>::new("users");
924        w.limit(10);
925        let sql = w.build_select();
926        assert!(sql.contains("LIMIT 10"));
927    }
928
929    #[test]
930    fn test_offset() {
931        let mut w = LambdaWrapper::<User>::new("users");
932        w.limit(10).offset(20);
933        let sql = w.build_select();
934        assert!(sql.contains("LIMIT 10"));
935        assert!(sql.contains("OFFSET 20"));
936    }
937
938    #[test]
939    fn test_page() {
940        let mut w = LambdaWrapper::<User>::new("users");
941        w.page(3, 20); // 第 3 页,每页 20 条
942        let sql = w.build_select();
943        assert!(sql.contains("LIMIT 20"));
944        assert!(sql.contains("OFFSET 40")); // (3-1) * 20
945    }
946
947    #[test]
948    fn test_page_1_no_offset() {
949        let mut w = LambdaWrapper::<User>::new("users");
950        w.page(1, 10);
951        let sql = w.build_select();
952        assert!(sql.contains("LIMIT 10"));
953        assert!(!sql.contains("OFFSET")); // 第 1 页无 OFFSET
954    }
955
956    // ===== SQL 生成测试 =====
957
958    #[test]
959    fn test_build_count() {
960        let mut w = LambdaWrapper::<User>::new("users");
961        w.gt(UserColumns::Age, Value::I64(18));
962        let sql = w.build_count();
963        assert!(sql.contains("SELECT COUNT(*) FROM `users`"));
964        assert!(sql.contains("`age` > 18"));
965        // 不应包含 ORDER BY / LIMIT
966        assert!(!sql.contains("ORDER BY"));
967        assert!(!sql.contains("LIMIT"));
968    }
969
970    #[test]
971    fn test_build_exists() {
972        let mut w = LambdaWrapper::<User>::new("users");
973        w.eq(UserColumns::Id, Value::I64(1));
974        let sql = w.build_exists();
975        assert!(sql.starts_with("SELECT EXISTS("));
976        assert!(sql.contains("SELECT 1 FROM `users`"));
977        assert!(sql.contains("`id` = 1"));
978        assert!(sql.ends_with(") AS exists_flag"));
979    }
980
981    #[test]
982    fn test_build_delete() {
983        let mut w = LambdaWrapper::<User>::new("users");
984        w.eq(UserColumns::Id, Value::I64(1));
985        let sql = w.build_delete();
986        assert!(sql.starts_with("DELETE FROM `users`"));
987        assert!(sql.contains("WHERE `id` = 1"));
988    }
989
990    // ===== 方言测试 =====
991
992    #[test]
993    fn test_with_postgres_dialect() {
994        let dialect: Box<dyn Dialect> = Box::new(PostgreSqlDialect);
995        let mut w = LambdaWrapper::<User>::with_dialect("users", dialect);
996        w.eq(UserColumns::Id, Value::I64(1));
997        let sql = w.build_select();
998        assert!(sql.contains("\"users\""));
999        assert!(sql.contains("\"id\" = 1"));
1000    }
1001
1002    #[test]
1003    fn test_postgres_select() {
1004        let dialect: Box<dyn Dialect> = Box::new(PostgreSqlDialect);
1005        let mut w = LambdaWrapper::<User>::with_dialect("users", dialect);
1006        w.select(UserColumns::Id).select(UserColumns::Name);
1007        let sql = w.build_select();
1008        assert!(sql.contains("\"id\", \"name\""));
1009    }
1010
1011    // ===== 完整查询测试 =====
1012
1013    #[test]
1014    fn test_complex_query() {
1015        let mut w = LambdaWrapper::<User>::new("users");
1016        w.select(UserColumns::Id)
1017            .select(UserColumns::Name)
1018            .select(UserColumns::Age)
1019            .gt(UserColumns::Age, Value::I64(18))
1020            .like(UserColumns::Name, Value::String("a%".to_string()))
1021            .is_not_null(UserColumns::Email)
1022            .order_by_desc(UserColumns::Id)
1023            .limit(10)
1024            .offset(20);
1025
1026        let sql = w.build_select();
1027        assert!(sql.contains("SELECT `id`, `name`, `age` FROM `users`"));
1028        assert!(sql.contains("`age` > 18"));
1029        assert!(sql.contains("`name` LIKE 'a%'"));
1030        assert!(sql.contains("`email` IS NOT NULL"));
1031        assert!(sql.contains("ORDER BY `id` DESC"));
1032        assert!(sql.contains("LIMIT 10"));
1033        assert!(sql.contains("OFFSET 20"));
1034    }
1035
1036    #[test]
1037    fn test_reset_clears_all() {
1038        let mut w = LambdaWrapper::<User>::new("users");
1039        w.select(UserColumns::Id)
1040            .eq(UserColumns::Id, Value::I64(1))
1041            .order_by_asc(UserColumns::Name)
1042            .limit(10);
1043
1044        w.reset();
1045        assert_eq!(w.select_count(), 0);
1046        assert_eq!(w.where_count(), 0);
1047        let sql = w.build_select();
1048        assert!(sql.contains("SELECT * FROM `users`"));
1049        assert!(!sql.contains("WHERE"));
1050        assert!(!sql.contains("ORDER BY"));
1051        assert!(!sql.contains("LIMIT"));
1052    }
1053
1054    // ===== 跨 Model 类型隔离测试 =====
1055
1056    #[test]
1057    fn test_different_models_dont_share_columns() {
1058        let mut user_w = LambdaWrapper::<User>::new("users");
1059        user_w.eq(UserColumns::Id, Value::I64(1));
1060
1061        let mut order_w = LambdaWrapper::<Order>::new("orders");
1062        order_w.eq(OrderColumns::OrderId, Value::I64(100));
1063
1064        let user_sql = user_w.build_select();
1065        let order_sql = order_w.build_select();
1066
1067        assert!(user_sql.contains("`users`"));
1068        assert!(user_sql.contains("`id` = 1"));
1069        assert!(order_sql.contains("`orders`"));
1070        assert!(order_sql.contains("`order_id` = 100"));
1071    }
1072
1073    // ===== 编译期类型安全演示(运行时验证)=====
1074
1075    #[test]
1076    fn test_type_safety_compile_time_check() {
1077        // 以下代码若取消注释将无法编译:
1078        // let mut w = LambdaWrapper::<User>::new("users");
1079        // w.eq(OrderColumns::OrderId, Value::I64(1)); // OrderColumns 不能用于 User 的 wrapper
1080
1081        // 这里通过正常调用验证类型隔离工作正常
1082        let mut w = LambdaWrapper::<User>::new("users");
1083        w.eq(UserColumns::Id, Value::I64(1));
1084        assert_eq!(w.where_count(), 1);
1085    }
1086
1087    #[test]
1088    fn test_string_value_escape() {
1089        // v0.2.2 修复 H-1:默认使用 MySqlDialect,单引号转义为 \'
1090        let mut w = LambdaWrapper::<User>::new("users");
1091        w.eq(UserColumns::Name, Value::String("O'Brien".to_string()));
1092        let sql = w.build_select();
1093        // MySQL 方言下字符串中的单引号应被转义为 \'
1094        assert!(sql.contains("'O\\'Brien'"));
1095
1096        // 使用 PostgreSQL 方言时,单引号转义为 ''
1097        let pg_dialect: Box<dyn Dialect> = get_dialect(DbType::PostgreSQL).unwrap();
1098        let mut w_pg = LambdaWrapper::<User>::with_dialect("users", pg_dialect);
1099        w_pg.eq(UserColumns::Name, Value::String("O'Brien".to_string()));
1100        let sql_pg = w_pg.build_select();
1101        assert!(sql_pg.contains("'O''Brien'"));
1102    }
1103
1104    // ===== 使用真实方言验证集成 =====
1105
1106    #[test]
1107    fn test_with_real_mysql_dialect() {
1108        let dialect: Box<dyn Dialect> = get_dialect(DbType::MySQL).unwrap();
1109        let mut w = LambdaWrapper::<User>::with_dialect("users", dialect);
1110        w.select(UserColumns::Id)
1111            .select(UserColumns::Name)
1112            .eq(UserColumns::Id, Value::I64(42))
1113            .order_by_desc(UserColumns::Id);
1114
1115        let sql = w.build_select();
1116        assert!(sql.contains("SELECT `id`, `name` FROM `users`"));
1117        assert!(sql.contains("`id` = 42"));
1118        assert!(sql.contains("ORDER BY `id` DESC"));
1119    }
1120}