Skip to main content

sz_orm_query/
join_dsl.rs

1//! JoinDSL — 类型安全的 JOIN 语法(Diesel 风格)
2//!
3//! 通过 [`TypedTable`] 和 [`TypedColumn`] 标记类型,
4//! 把 JOIN 的 ON 条件提升到类型系统,让"表 A 的列 vs 表 B 的列"
5//! 在编译期就能验证表归属,杜绝拼错表名/列名。
6//!
7//! # 设计
8//!
9//! - [`JoinBuilder`] 链式构造 JOIN 子句
10//! - [`JoinOn`] 表达 ON 条件,左侧列必须属于左表,右侧列必须属于右表
11//! - 编译期通过 [`TypedColumn::Table`] 关联约束校验
12//!
13//! # 用法
14//!
15//! ```ignore
16//! use sz_orm_query::typed::{TypedTable, TypedColumn};
17//! use sz_orm_query::join_dsl::{JoinBuilder, JoinKind};
18//! use sz_orm_macros::typed_query;
19//!
20//! typed_query! {
21//!     table users { id: i64, name: String }
22//!     table orders { id: i64, user_id: i64, total: f64 }
23//! }
24//!
25//! // 编译期保证:users::col_id 与 orders::col_user_id 是 ON 条件两端
26//! let join = JoinBuilder::new(JoinKind::Inner)
27//!     .table::<users::table>()
28//!     .on::<users::col_id, orders::col_user_id>()
29//!     .build();
30//! assert_eq!(join, "INNER JOIN users ON users.id = orders.user_id");
31//! ```
32
33use crate::typed::{TypedColumn, TypedTable};
34
35/// JOIN 类型
36#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum JoinKind {
38    /// INNER JOIN
39    Inner,
40    /// LEFT \[OUTER\] JOIN
41    Left,
42    /// RIGHT \[OUTER\] JOIN
43    Right,
44    /// FULL \[OUTER\] JOIN
45    Full,
46    /// CROSS JOIN
47    Cross,
48}
49
50impl JoinKind {
51    /// 转换为 SQL 关键字
52    pub fn as_sql(&self) -> &'static str {
53        match self {
54            JoinKind::Inner => "INNER JOIN",
55            JoinKind::Left => "LEFT JOIN",
56            JoinKind::Right => "RIGHT JOIN",
57            JoinKind::Full => "FULL OUTER JOIN",
58            JoinKind::Cross => "CROSS JOIN",
59        }
60    }
61}
62
63/// JOIN 构造器(已指定 JOIN 类型,待指定右表)
64pub struct JoinBuilder {
65    kind: JoinKind,
66}
67
68impl JoinBuilder {
69    /// 创建新的 JoinBuilder,指定 JOIN 类型
70    pub fn new(kind: JoinKind) -> Self {
71        Self { kind }
72    }
73
74    /// 指定 JOIN 的右表(被 join 进来的表)
75    ///
76    /// 返回 [`JoinOn`],等待 ON 条件。
77    /// CROSS JOIN 无需 ON 条件,可直接调用 `JoinOn::build`。
78    pub fn table<T: TypedTable>(self) -> JoinOn {
79        JoinOn {
80            kind: self.kind,
81            right_table: T::NAME,
82        }
83    }
84}
85
86/// JOIN 已指定右表,待指定 ON 条件
87pub struct JoinOn {
88    kind: JoinKind,
89    right_table: &'static str,
90}
91
92impl JoinOn {
93    /// 指定 ON 条件:左表列 = 右表列
94    ///
95    /// # 类型约束
96    ///
97    /// - `L`: 左表列,必须属于"主表"(用户自行保证)
98    /// - `R`: 右表列,必须属于 `right_table`(用户自行保证)
99    ///
100    /// 编译期通过 [`TypedColumn::Table`] 关联约束校验列归属。
101    pub fn on<L, R>(self) -> JoinBuilt
102    where
103        L: TypedColumn,
104        R: TypedColumn,
105    {
106        // 编译期断言:L 和 R 分别属于不同的表(避免自连接时拼错)
107        // 注意:我们不强制 L.Table != R.Table,因为自连接也是合法场景
108        JoinBuilt {
109            kind: self.kind,
110            right_table: self.right_table,
111            left_column: L::NAME,
112            right_column: R::NAME,
113        }
114    }
115
116    /// 不带 ON 条件直接构造(用于 CROSS JOIN)
117    pub fn build_no_on(self) -> String {
118        format!("{} {}", self.kind.as_sql(), self.right_table)
119    }
120}
121
122/// JOIN 已完整构造,可生成 SQL
123pub struct JoinBuilt {
124    kind: JoinKind,
125    right_table: &'static str,
126    left_column: &'static str,
127    right_column: &'static str,
128}
129
130impl JoinBuilt {
131    /// 生成 JOIN 子句 SQL
132    ///
133    /// 输出格式:`<JOIN_TYPE> <right_table> ON <left_column> = <right_column>`
134    ///
135    /// 注意:生成的 SQL 中列名不带表前缀。
136    /// 如需带前缀,使用 [`build_with_prefix`]。
137    ///
138    /// [`build_with_prefix`]: JoinBuilt::build_with_prefix
139    pub fn build(&self) -> String {
140        format!(
141            "{} {} ON {} = {}",
142            self.kind.as_sql(),
143            self.right_table,
144            self.left_column,
145            self.right_column
146        )
147    }
148
149    /// 生成带表前缀的 JOIN 子句
150    ///
151    /// 输出格式:`<JOIN_TYPE> <right_table> ON <left_table>.<left_column> = <right_table>.<right_column>`
152    pub fn build_with_prefix(&self, left_table: &str) -> String {
153        format!(
154            "{} {} ON {}.{} = {}.{}",
155            self.kind.as_sql(),
156            self.right_table,
157            left_table,
158            self.left_column,
159            self.right_table,
160            self.right_column
161        )
162    }
163}
164
165/// 便捷函数:构造 INNER JOIN
166pub fn inner_join<T: TypedTable>() -> JoinOn {
167    JoinBuilder::new(JoinKind::Inner).table::<T>()
168}
169
170/// 便捷函数:构造 LEFT JOIN
171pub fn left_join<T: TypedTable>() -> JoinOn {
172    JoinBuilder::new(JoinKind::Left).table::<T>()
173}
174
175/// 便捷函数:构造 RIGHT JOIN
176pub fn right_join<T: TypedTable>() -> JoinOn {
177    JoinBuilder::new(JoinKind::Right).table::<T>()
178}
179
180/// 便捷函数:构造 FULL OUTER JOIN
181pub fn full_join<T: TypedTable>() -> JoinOn {
182    JoinBuilder::new(JoinKind::Full).table::<T>()
183}
184
185/// 便捷函数:构造 CROSS JOIN
186pub fn cross_join<T: TypedTable>() -> JoinOn {
187    JoinBuilder::new(JoinKind::Cross).table::<T>()
188}
189
190#[cfg(test)]
191mod tests {
192    use super::*;
193
194    // ---- 测试用 mock 类型 ----
195
196    struct UsersTable;
197    impl TypedTable for UsersTable {
198        const NAME: &'static str = "users";
199    }
200
201    struct OrdersTable;
202    impl TypedTable for OrdersTable {
203        const NAME: &'static str = "orders";
204    }
205
206    struct ColUserId;
207    impl TypedColumn for ColUserId {
208        const NAME: &'static str = "id";
209        type Table = UsersTable;
210        type RustType = i64;
211        type SqlType = crate::typed_ast::Untyped;
212    }
213
214    struct ColOrderUserId;
215    impl TypedColumn for ColOrderUserId {
216        const NAME: &'static str = "user_id";
217        type Table = OrdersTable;
218        type RustType = i64;
219        type SqlType = crate::typed_ast::Untyped;
220    }
221
222    struct ColOrderId;
223    impl TypedColumn for ColOrderId {
224        const NAME: &'static str = "id";
225        type Table = OrdersTable;
226        type RustType = i64;
227        type SqlType = crate::typed_ast::Untyped;
228    }
229
230    // ---- JoinKind 测试 ----
231
232    #[test]
233    fn test_join_kind_as_sql() {
234        assert_eq!(JoinKind::Inner.as_sql(), "INNER JOIN");
235        assert_eq!(JoinKind::Left.as_sql(), "LEFT JOIN");
236        assert_eq!(JoinKind::Right.as_sql(), "RIGHT JOIN");
237        assert_eq!(JoinKind::Full.as_sql(), "FULL OUTER JOIN");
238        assert_eq!(JoinKind::Cross.as_sql(), "CROSS JOIN");
239    }
240
241    // ---- 基础 JOIN 构造测试 ----
242
243    #[test]
244    fn test_inner_join_basic() {
245        let join = JoinBuilder::new(JoinKind::Inner)
246            .table::<OrdersTable>()
247            .on::<ColUserId, ColOrderUserId>()
248            .build();
249        assert_eq!(join, "INNER JOIN orders ON id = user_id");
250    }
251
252    #[test]
253    fn test_left_join_with_prefix() {
254        let join = JoinBuilder::new(JoinKind::Left)
255            .table::<OrdersTable>()
256            .on::<ColUserId, ColOrderUserId>()
257            .build_with_prefix("users");
258        assert_eq!(join, "LEFT JOIN orders ON users.id = orders.user_id");
259    }
260
261    #[test]
262    fn test_right_join() {
263        let join = JoinBuilder::new(JoinKind::Right)
264            .table::<OrdersTable>()
265            .on::<ColUserId, ColOrderUserId>()
266            .build();
267        assert_eq!(join, "RIGHT JOIN orders ON id = user_id");
268    }
269
270    #[test]
271    fn test_full_join() {
272        let join = JoinBuilder::new(JoinKind::Full)
273            .table::<OrdersTable>()
274            .on::<ColUserId, ColOrderUserId>()
275            .build();
276        assert_eq!(join, "FULL OUTER JOIN orders ON id = user_id");
277    }
278
279    #[test]
280    fn test_cross_join_no_on() {
281        let join = JoinBuilder::new(JoinKind::Cross)
282            .table::<OrdersTable>()
283            .build_no_on();
284        assert_eq!(join, "CROSS JOIN orders");
285    }
286
287    // ---- 便捷函数测试 ----
288
289    #[test]
290    fn test_inner_join_helper() {
291        let join = inner_join::<OrdersTable>()
292            .on::<ColUserId, ColOrderUserId>()
293            .build();
294        assert_eq!(join, "INNER JOIN orders ON id = user_id");
295    }
296
297    #[test]
298    fn test_left_join_helper() {
299        let join = left_join::<OrdersTable>()
300            .on::<ColUserId, ColOrderUserId>()
301            .build();
302        assert_eq!(join, "LEFT JOIN orders ON id = user_id");
303    }
304
305    #[test]
306    fn test_right_join_helper() {
307        let join = right_join::<OrdersTable>()
308            .on::<ColUserId, ColOrderUserId>()
309            .build();
310        assert_eq!(join, "RIGHT JOIN orders ON id = user_id");
311    }
312
313    #[test]
314    fn test_full_join_helper() {
315        let join = full_join::<OrdersTable>()
316            .on::<ColUserId, ColOrderUserId>()
317            .build();
318        assert_eq!(join, "FULL OUTER JOIN orders ON id = user_id");
319    }
320
321    #[test]
322    fn test_cross_join_helper() {
323        let join = cross_join::<OrdersTable>().build_no_on();
324        assert_eq!(join, "CROSS JOIN orders");
325    }
326
327    // ---- 类型安全验证(编译期) ----
328
329    #[test]
330    fn test_compile_time_table_association() {
331        // 这个测试主要验证编译期约束:
332        // on::<L, R>() 要求 L 和 R 都实现了 TypedColumn
333        // 如果传入非 TypedColumn 类型,编译就会失败
334        let join = inner_join::<OrdersTable>()
335            .on::<ColUserId, ColOrderUserId>()
336            .build();
337        assert!(!join.is_empty());
338    }
339
340    #[test]
341    fn test_self_join_same_table() {
342        // 自连接:同一表的两个列
343        // 注意:右表也是 orders,左列和右列都属于 orders
344        let join = inner_join::<OrdersTable>()
345            .on::<ColOrderId, ColOrderUserId>()
346            .build();
347        assert_eq!(join, "INNER JOIN orders ON id = user_id");
348    }
349
350    // ---- 多 JOIN 拼接测试 ----
351
352    #[test]
353    fn test_multiple_joins_concat() {
354        let j1 = inner_join::<OrdersTable>()
355            .on::<ColUserId, ColOrderUserId>()
356            .build_with_prefix("users");
357
358        // 模拟第二张表(用 mock 类型复用)
359        let j2 = left_join::<OrdersTable>()
360            .on::<ColUserId, ColOrderId>()
361            .build_with_prefix("users");
362
363        let sql = format!("SELECT * FROM users {} {}", j1, j2);
364        assert!(sql.contains("INNER JOIN orders ON users.id = orders.user_id"));
365        assert!(sql.contains("LEFT JOIN orders ON users.id = orders.id"));
366    }
367}