Skip to main content

sz_orm_core/
linq.rs

1//! LINQ 风格查询 API
2//!
3//! 对标 C# LINQ / EF Core `IQueryable<T>`。
4//!
5//! 提供流式查询构建器,方法名与 LINQ 标准操作符一致。
6//!
7//! # 使用示例
8//!
9//! ```
10//! use sz_orm_core::linq::LinqQuery;
11//! use sz_orm_core::Value;
12//!
13//! let query = LinqQuery::from("users")
14//!     .select(vec!["id", "name", "age"])
15//!     .where_eq("age", Value::I64(25))
16//!     .order_by("name")
17//!     .take(10)
18//!     .skip(0);
19//!
20//! let sql = query.build();
21//! assert!(sql.contains("SELECT id, name, age FROM `users`"));
22//! assert!(sql.contains("WHERE"));
23//! assert!(sql.contains("ORDER BY `name` ASC"));
24//! assert!(sql.contains("LIMIT 10"));
25//! ```
26
27use crate::dialect::{Dialect, MySqlDialect};
28use crate::value::Value;
29
30/// LINQ 风格查询构建器
31pub struct LinqQuery {
32    table: String,
33    columns: Vec<String>,
34    conditions: Vec<String>,
35    params: Vec<Value>,
36    order_by: Vec<(String, bool)>,
37    limit: Option<usize>,
38    offset: Option<usize>,
39    distinct: bool,
40    group_by: Vec<String>,
41    dialect: Box<dyn Dialect>,
42}
43
44impl LinqQuery {
45    /// FROM 子句 — 指定表名
46    pub fn from(table: &str) -> Self {
47        Self {
48            table: table.to_string(),
49            columns: Vec::new(),
50            conditions: Vec::new(),
51            params: Vec::new(),
52            order_by: Vec::new(),
53            limit: None,
54            offset: None,
55            distinct: false,
56            group_by: Vec::new(),
57            dialect: Box::new(MySqlDialect),
58        }
59    }
60
61    /// SELECT 子句 — 指定列
62    pub fn select(mut self, cols: Vec<&str>) -> Self {
63        self.columns = cols.into_iter().map(String::from).collect();
64        self
65    }
66
67    /// WHERE 等于条件
68    pub fn where_eq(mut self, field: &str, value: Value) -> Self {
69        self.conditions
70            .push(format!("{} = ?", self.dialect.quote(field)));
71        self.params.push(value);
72        self
73    }
74
75    /// WHERE 不等于条件
76    pub fn where_ne(mut self, field: &str, value: Value) -> Self {
77        self.conditions
78            .push(format!("{} != ?", self.dialect.quote(field)));
79        self.params.push(value);
80        self
81    }
82
83    /// WHERE 大于条件
84    pub fn where_gt(mut self, field: &str, value: Value) -> Self {
85        self.conditions
86            .push(format!("{} > ?", self.dialect.quote(field)));
87        self.params.push(value);
88        self
89    }
90
91    /// WHERE 小于条件
92    pub fn where_lt(mut self, field: &str, value: Value) -> Self {
93        self.conditions
94            .push(format!("{} < ?", self.dialect.quote(field)));
95        self.params.push(value);
96        self
97    }
98
99    /// WHERE 大于等于条件
100    pub fn where_ge(mut self, field: &str, value: Value) -> Self {
101        self.conditions
102            .push(format!("{} >= ?", self.dialect.quote(field)));
103        self.params.push(value);
104        self
105    }
106
107    /// WHERE 小于等于条件
108    pub fn where_le(mut self, field: &str, value: Value) -> Self {
109        self.conditions
110            .push(format!("{} <= ?", self.dialect.quote(field)));
111        self.params.push(value);
112        self
113    }
114
115    /// WHERE LIKE 条件
116    pub fn where_like(mut self, field: &str, pattern: Value) -> Self {
117        self.conditions
118            .push(format!("{} LIKE ?", self.dialect.quote(field)));
119        self.params.push(pattern);
120        self
121    }
122
123    /// WHERE IN 条件
124    pub fn where_in(mut self, field: &str, values: Vec<Value>) -> Self {
125        if values.is_empty() {
126            self.conditions.push("1 = 0".to_string());
127            return self;
128        }
129        let placeholders: Vec<String> = (0..values.len()).map(|_| "?".to_string()).collect();
130        self.conditions.push(format!(
131            "{} IN ({})",
132            self.dialect.quote(field),
133            placeholders.join(", ")
134        ));
135        self.params.extend(values);
136        self
137    }
138
139    /// WHERE IS NULL
140    pub fn where_null(mut self, field: &str) -> Self {
141        self.conditions
142            .push(format!("{} IS NULL", self.dialect.quote(field)));
143        self
144    }
145
146    /// WHERE IS NOT NULL
147    pub fn where_not_null(mut self, field: &str) -> Self {
148        self.conditions
149            .push(format!("{} IS NOT NULL", self.dialect.quote(field)));
150        self
151    }
152
153    /// ORDER BY 升序
154    pub fn order_by(mut self, field: &str) -> Self {
155        self.order_by.push((field.to_string(), false));
156        self
157    }
158
159    /// ORDER BY 降序
160    pub fn order_by_desc(mut self, field: &str) -> Self {
161        self.order_by.push((field.to_string(), true));
162        self
163    }
164
165    /// TAKE — 取前 N 条(等价于 LIMIT)
166    pub fn take(mut self, n: usize) -> Self {
167        self.limit = Some(n);
168        self
169    }
170
171    /// SKIP — 跳过 N 条(等价于 OFFSET)
172    pub fn skip(mut self, n: usize) -> Self {
173        if n > 0 {
174            self.offset = Some(n);
175        }
176        self
177    }
178
179    /// DISTINCT — 去重
180    pub fn distinct(mut self) -> Self {
181        self.distinct = true;
182        self
183    }
184
185    /// GROUP BY — 分组
186    pub fn group_by(mut self, cols: Vec<&str>) -> Self {
187        self.group_by = cols.into_iter().map(String::from).collect();
188        self
189    }
190
191    /// 构建 SQL
192    pub fn build(&self) -> String {
193        let mut sql = String::new();
194
195        let cols = if self.columns.is_empty() {
196            "*".to_string()
197        } else {
198            self.columns.join(", ")
199        };
200
201        if self.distinct {
202            sql.push_str(&format!(
203                "SELECT DISTINCT {} FROM {}",
204                cols,
205                self.dialect.quote(&self.table)
206            ));
207        } else {
208            sql.push_str(&format!(
209                "SELECT {} FROM {}",
210                cols,
211                self.dialect.quote(&self.table)
212            ));
213        }
214
215        if !self.conditions.is_empty() {
216            sql.push_str(" WHERE ");
217            sql.push_str(&self.conditions.join(" AND "));
218        }
219
220        if !self.group_by.is_empty() {
221            sql.push_str(" GROUP BY ");
222            sql.push_str(&self.group_by.join(", "));
223        }
224
225        if !self.order_by.is_empty() {
226            sql.push_str(" ORDER BY ");
227            let orders: Vec<String> = self
228                .order_by
229                .iter()
230                .map(|(f, desc)| {
231                    if *desc {
232                        format!("{} DESC", self.dialect.quote(f))
233                    } else {
234                        format!("{} ASC", self.dialect.quote(f))
235                    }
236                })
237                .collect();
238            sql.push_str(&orders.join(", "));
239        }
240
241        if let Some(limit) = self.limit {
242            sql.push_str(&format!(" LIMIT {}", limit));
243        }
244
245        if let Some(offset) = self.offset {
246            sql.push_str(&format!(" OFFSET {}", offset));
247        }
248
249        sql
250    }
251
252    /// 获取参数
253    pub fn params(&self) -> &[Value] {
254        &self.params
255    }
256
257    /// 构建 COUNT 查询
258    pub fn build_count(&self) -> String {
259        let mut sql = format!("SELECT COUNT(*) FROM {}", self.dialect.quote(&self.table));
260
261        if !self.conditions.is_empty() {
262            sql.push_str(" WHERE ");
263            sql.push_str(&self.conditions.join(" AND "));
264        }
265
266        sql
267    }
268
269    /// 构建 EXISTS 查询
270    pub fn build_exists(&self) -> String {
271        let mut sql = format!(
272            "SELECT EXISTS(SELECT 1 FROM {}",
273            self.dialect.quote(&self.table)
274        );
275
276        if !self.conditions.is_empty() {
277            sql.push_str(" WHERE ");
278            sql.push_str(&self.conditions.join(" AND "));
279        }
280
281        sql.push(')');
282        sql
283    }
284}
285
286#[cfg(test)]
287mod tests {
288    use super::*;
289
290    #[test]
291    fn test_linq_basic_select() {
292        let q = LinqQuery::from("users").select(vec!["id", "name"]);
293        let sql = q.build();
294        assert!(sql.contains("SELECT id, name FROM"));
295        assert!(sql.contains("users"));
296    }
297
298    #[test]
299    fn test_linq_select_all() {
300        let q = LinqQuery::from("users");
301        let sql = q.build();
302        assert!(sql.contains("SELECT * FROM"));
303        assert!(sql.contains("users"));
304    }
305
306    #[test]
307    fn test_linq_where_eq() {
308        let q = LinqQuery::from("users").where_eq("age", Value::I64(25));
309        let sql = q.build();
310        assert!(sql.contains("WHERE `age` = ?"));
311        assert_eq!(q.params(), &[Value::I64(25)]);
312    }
313
314    #[test]
315    fn test_linq_where_in() {
316        let q = LinqQuery::from("users")
317            .where_in("id", vec![Value::I64(1), Value::I64(2), Value::I64(3)]);
318        let sql = q.build();
319        assert!(sql.contains("IN (?, ?, ?)"));
320        assert_eq!(q.params().len(), 3);
321    }
322
323    #[test]
324    fn test_linq_where_in_empty() {
325        let q = LinqQuery::from("users").where_in("id", vec![]);
326        let sql = q.build();
327        assert!(sql.contains("1 = 0"));
328    }
329
330    #[test]
331    fn test_linq_order_by() {
332        let q = LinqQuery::from("users").order_by("name");
333        let sql = q.build();
334        assert!(sql.contains("ORDER BY `name` ASC"));
335    }
336
337    #[test]
338    fn test_linq_order_by_desc() {
339        let q = LinqQuery::from("users").order_by_desc("age");
340        let sql = q.build();
341        assert!(sql.contains("ORDER BY `age` DESC"));
342    }
343
344    #[test]
345    fn test_linq_take_skip() {
346        let q = LinqQuery::from("users").take(10).skip(5);
347        let sql = q.build();
348        assert!(sql.contains("LIMIT 10"));
349        assert!(sql.contains("OFFSET 5"));
350    }
351
352    #[test]
353    fn test_linq_distinct() {
354        let q = LinqQuery::from("users").select(vec!["city"]).distinct();
355        let sql = q.build();
356        assert!(sql.starts_with("SELECT DISTINCT"));
357    }
358
359    #[test]
360    fn test_linq_group_by() {
361        let q = LinqQuery::from("users")
362            .select(vec!["city"])
363            .group_by(vec!["city"]);
364        let sql = q.build();
365        assert!(sql.contains("GROUP BY city"));
366    }
367
368    #[test]
369    fn test_linq_chained() {
370        let q = LinqQuery::from("users")
371            .select(vec!["id", "name", "age"])
372            .where_eq("age", Value::I64(25))
373            .where_ne("name", Value::String("admin".into()))
374            .order_by("name")
375            .take(10);
376
377        let sql = q.build();
378        assert!(sql.contains("SELECT id, name, age FROM"));
379        assert!(sql.contains("users"));
380        assert!(sql.contains("WHERE"));
381        assert!(sql.contains("AND"));
382        assert!(sql.contains("ORDER BY `name` ASC"));
383        assert!(sql.contains("LIMIT 10"));
384        assert_eq!(q.params().len(), 2);
385    }
386
387    #[test]
388    fn test_linq_build_count() {
389        let q = LinqQuery::from("users").where_eq("age", Value::I64(25));
390        let sql = q.build_count();
391        assert!(sql.starts_with("SELECT COUNT(*)"));
392        assert!(sql.contains("WHERE"));
393    }
394
395    #[test]
396    fn test_linq_build_exists() {
397        let q = LinqQuery::from("users").where_eq("email", Value::String("test@test.com".into()));
398        let sql = q.build_exists();
399        assert!(sql.starts_with("SELECT EXISTS("));
400    }
401
402    #[test]
403    fn test_linq_where_null() {
404        let q = LinqQuery::from("users").where_null("deleted_at");
405        let sql = q.build();
406        assert!(sql.contains("IS NULL"));
407    }
408
409    #[test]
410    fn test_linq_where_not_null() {
411        let q = LinqQuery::from("users").where_not_null("email");
412        let sql = q.build();
413        assert!(sql.contains("IS NOT NULL"));
414    }
415
416    #[test]
417    fn test_linq_where_gt_lt() {
418        let q = LinqQuery::from("users")
419            .where_gt("age", Value::I64(18))
420            .where_lt("age", Value::I64(65));
421        let sql = q.build();
422        assert!(sql.contains("> ?"));
423        assert!(sql.contains("< ?"));
424        assert_eq!(q.params().len(), 2);
425    }
426
427    #[test]
428    fn test_e2e_linq_realistic_user_query() {
429        let q = LinqQuery::from("users")
430            .select(vec!["id", "name", "email", "age"])
431            .where_eq("status", Value::String("active".into()))
432            .where_gt("age", Value::I64(18))
433            .where_like("email", Value::String("%@gmail.com".into()))
434            .order_by("name")
435            .take(20)
436            .skip(0);
437
438        let sql = q.build();
439        assert!(sql.contains("SELECT id, name, email, age FROM `users`"));
440        assert!(sql.contains("WHERE"));
441        assert!(sql.contains("`status` = ?"));
442        assert!(sql.contains("`age` > ?"));
443        assert!(sql.contains("`email` LIKE ?"));
444        assert!(sql.contains("AND"));
445        assert!(sql.contains("ORDER BY `name` ASC"));
446        assert!(sql.contains("LIMIT 20"));
447
448        assert_eq!(
449            q.params(),
450            &[
451                Value::String("active".into()),
452                Value::I64(18),
453                Value::String("%@gmail.com".into()),
454            ]
455        );
456    }
457
458    #[test]
459    fn test_e2e_linq_pagination_with_count() {
460        let make_base = || {
461            LinqQuery::from("orders")
462                .where_eq("user_id", Value::I64(42))
463                .where_eq("status", Value::String("paid".into()))
464        };
465
466        let page1 = make_base()
467            .select(vec!["id", "amount"])
468            .order_by_desc("created_at")
469            .take(10)
470            .skip(0);
471        let page2 = make_base()
472            .select(vec!["id", "amount"])
473            .order_by_desc("created_at")
474            .take(10)
475            .skip(10);
476        let count = make_base().build_count();
477
478        let sql1 = page1.build();
479        let sql2 = page2.build();
480
481        assert!(sql1.contains("LIMIT 10"));
482        assert!(!sql1.contains("OFFSET"));
483        assert!(sql2.contains("LIMIT 10"));
484        assert!(sql2.contains("OFFSET 10"));
485        assert!(count.starts_with("SELECT COUNT(*)"));
486        assert!(count.contains("`user_id` = ?"));
487        assert!(count.contains("`status` = ?"));
488    }
489
490    #[test]
491    fn test_e2e_linq_exists_check() {
492        let q = LinqQuery::from("users")
493            .where_eq("email", Value::String("alice@example.com".into()))
494            .where_null("deleted_at");
495
496        let exists_sql = q.build_exists();
497        assert!(exists_sql.starts_with("SELECT EXISTS(SELECT 1 FROM `users`"));
498        assert!(exists_sql.contains("`email` = ?"));
499        assert!(exists_sql.contains("`deleted_at` IS NULL"));
500        assert!(exists_sql.ends_with(')'));
501        assert_eq!(q.params().len(), 1);
502    }
503
504    #[test]
505    fn test_e2e_linq_in_clause_batch_lookup() {
506        let ids: Vec<Value> = (1..=5).map(Value::I64).collect();
507        let q = LinqQuery::from("products")
508            .select(vec!["id", "name", "price"])
509            .where_in("id", ids)
510            .where_eq("active", Value::Bool(true))
511            .order_by("price");
512
513        let sql = q.build();
514        assert!(sql.contains("IN (?, ?, ?, ?, ?)"));
515        assert!(sql.contains("`active` = ?"));
516        assert_eq!(q.params().len(), 6);
517    }
518
519    #[test]
520    fn test_e2e_linq_distinct_cities() {
521        let q = LinqQuery::from("users")
522            .select(vec!["city"])
523            .distinct()
524            .where_not_null("city")
525            .order_by("city");
526
527        let sql = q.build();
528        assert!(sql.starts_with("SELECT DISTINCT city FROM `users`"));
529        assert!(sql.contains("`city` IS NOT NULL"));
530        assert!(sql.contains("ORDER BY `city` ASC"));
531    }
532}