1use crate::typed::{TypedColumn, TypedTable};
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum JoinKind {
38 Inner,
40 Left,
42 Right,
44 Full,
46 Cross,
48}
49
50impl JoinKind {
51 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
63pub struct JoinBuilder {
65 kind: JoinKind,
66}
67
68impl JoinBuilder {
69 pub fn new(kind: JoinKind) -> Self {
71 Self { kind }
72 }
73
74 pub fn table<T: TypedTable>(self) -> JoinOn {
79 JoinOn {
80 kind: self.kind,
81 right_table: T::NAME,
82 }
83 }
84}
85
86pub struct JoinOn {
88 kind: JoinKind,
89 right_table: &'static str,
90}
91
92impl JoinOn {
93 pub fn on<L, R>(self) -> JoinBuilt
102 where
103 L: TypedColumn,
104 R: TypedColumn,
105 {
106 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 pub fn build_no_on(self) -> String {
118 format!("{} {}", self.kind.as_sql(), self.right_table)
119 }
120}
121
122pub 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 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 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
165pub fn inner_join<T: TypedTable>() -> JoinOn {
167 JoinBuilder::new(JoinKind::Inner).table::<T>()
168}
169
170pub fn left_join<T: TypedTable>() -> JoinOn {
172 JoinBuilder::new(JoinKind::Left).table::<T>()
173}
174
175pub fn right_join<T: TypedTable>() -> JoinOn {
177 JoinBuilder::new(JoinKind::Right).table::<T>()
178}
179
180pub fn full_join<T: TypedTable>() -> JoinOn {
182 JoinBuilder::new(JoinKind::Full).table::<T>()
183}
184
185pub 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 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 #[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 #[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 #[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 #[test]
330 fn test_compile_time_table_association() {
331 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 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 #[test]
353 fn test_multiple_joins_concat() {
354 let j1 = inner_join::<OrdersTable>()
355 .on::<ColUserId, ColOrderUserId>()
356 .build_with_prefix("users");
357
358 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}