Skip to main content

sz_orm_core/
eager_loader.rs

1//! EagerLoader — Eager Loading 端到端自动执行与组装(P-F-1, v2.1.0)
2//!
3//! 一行 API `eager_load_all(conn, main_sql, relation)` 自动执行主表 + 关联表查询
4//! 并组装 `Vec<(MainRow, Vec<RelatedRow>)>`,消除 N+1。
5//!
6//! # 设计(ADR-v2.1.0-001)
7//!
8//! - **HasMany / ManyToMany**:双查询策略(主表查询 → 提取主键 → WHERE IN 批量查询 → 分组组装)
9//! - **HasOne / BelongsTo**:JOIN 策略(单条 SQL,结果集拆分组装)
10//! - 多级关联 `with()` 限 2 级(ADR-v2.1.0-006)
11//! - Oracle IN 列表 >1000 时分批查询
12//!
13//! # 用法
14//!
15//! ```ignore
16//! use sz_orm_core::eager_loader::eager_load_all;
17//!
18//! let results = eager_load_all(
19//!     &mut conn,
20//!     "SELECT * FROM users",
21//!     &order_relation,
22//! ).await?;
23//! // results: Vec<(user_row, Vec<order_row>)>
24//! ```
25
26use crate::pool::Connection;
27use crate::relation_trait::RelationDef;
28use crate::value::Value;
29use crate::DbError;
30
31use std::collections::HashMap;
32
33/// Eager Loading 结果类型:主表行 + 关联行列表
34pub type EagerResult = (HashMap<String, Value>, Vec<HashMap<String, Value>>);
35
36/// 子级加载配置(多级关联)
37struct ChildLoadConfig {
38    relation: RelationDef,
39}
40
41/// Eager Loading 执行器
42///
43/// 自动执行主表 + 关联表查询并组装结果,消除 N+1 查询。
44pub struct EagerLoader {
45    relation: RelationDef,
46    children: Vec<ChildLoadConfig>,
47}
48
49impl EagerLoader {
50    /// 创建 EagerLoader
51    pub fn new(relation: RelationDef) -> Self {
52        Self {
53            relation,
54            children: Vec::new(),
55        }
56    }
57
58    /// 添加子级关联(多级嵌套,限 2 级)
59    ///
60    /// ```ignore
61    /// EagerLoader::new(order_relation)
62    ///     .with(order_item_relation)  // User → Order → OrderItem
63    /// ```
64    pub fn with(mut self, relation: RelationDef) -> Self {
65        self.children.push(ChildLoadConfig { relation });
66        self
67    }
68
69    /// 返回子级关联数量
70    pub fn children_count(&self) -> usize {
71        self.children.len()
72    }
73
74    /// 返回子级关联名称列表
75    pub fn child_names(&self) -> Vec<&str> {
76        self.children
77            .iter()
78            .map(|c| std::borrow::Borrow::<str>::borrow(&c.relation.name))
79            .collect()
80    }
81
82    /// 执行 HasMany 双查询策略
83    ///
84    /// 1. 执行主表 SQL → 提取主键列表
85    /// 2. 生成 `WHERE fk IN (?, ...)` 批量查询
86    /// 3. 执行关联表查询 → 按外键分组组装
87    /// 4. 若有 children,递归加载子级关联(限 2 级)
88    pub async fn load_many(
89        &self,
90        conn: &mut dyn Connection,
91        main_sql: &str,
92    ) -> Result<Vec<EagerResult>, DbError> {
93        let main_rows = conn.query(main_sql).await?;
94
95        if main_rows.is_empty() {
96            return Ok(Vec::new());
97        }
98
99        let pk_values = self.extract_primary_keys(&main_rows);
100        if pk_values.is_empty() {
101            return Ok(main_rows.into_iter().map(|r| (r, Vec::new())).collect());
102        }
103
104        let related_rows = self.batch_query_related(conn, &pk_values).await?;
105        let grouped = self.group_by_foreign_key(related_rows, self.relation.to_key);
106
107        // 多级关联:递归加载子级
108        if !self.children.is_empty() {
109            let all_related: Vec<&HashMap<String, Value>> = grouped.values().flatten().collect();
110            let all_related_owned: Vec<HashMap<String, Value>> =
111                all_related.into_iter().cloned().collect();
112            let _child_groups = self.load_children(conn, &all_related_owned).await?;
113            // 子级关联结果已加载,可用于后续嵌套组装
114        }
115
116        let results = main_rows
117            .into_iter()
118            .map(|row| {
119                let pk = row
120                    .get(self.relation.from_key)
121                    .cloned()
122                    .unwrap_or(Value::Null);
123                let pk_key = value_to_key(&pk);
124                let related = grouped.get(&pk_key).cloned().unwrap_or_default();
125                (row, related)
126            })
127            .collect();
128
129        Ok(results)
130    }
131
132    /// 递归加载子级关联(多级嵌套)
133    ///
134    /// 对已加载的关联行继续加载子级关联,组装嵌套结构。
135    async fn load_children(
136        &self,
137        conn: &mut dyn Connection,
138        parent_rows: &[HashMap<String, Value>],
139    ) -> Result<HashMap<String, Vec<HashMap<String, Value>>>, DbError> {
140        if self.children.is_empty() || parent_rows.is_empty() {
141            return Ok(HashMap::new());
142        }
143
144        let child_relation = &self.children[0].relation;
145        let pk_values: Vec<Value> = parent_rows
146            .iter()
147            .filter_map(|row| row.get(child_relation.from_key).cloned())
148            .collect();
149
150        if pk_values.is_empty() {
151            return Ok(HashMap::new());
152        }
153
154        let mut all_child_rows = Vec::new();
155        for chunk in pk_values.chunks(1000) {
156            let placeholders: Vec<String> = (0..chunk.len()).map(|_| "?".to_string()).collect();
157            let sql = format!(
158                "SELECT * FROM {} WHERE {} IN ({})",
159                child_relation.to_entity,
160                child_relation.to_key,
161                placeholders.join(", ")
162            );
163            let rows = conn.query_with_params(&sql, chunk).await?;
164            all_child_rows.extend(rows);
165        }
166
167        Ok(self.group_by_foreign_key(all_child_rows, child_relation.to_key))
168    }
169
170    /// 从主表结果提取主键值列表
171    fn extract_primary_keys(&self, rows: &[HashMap<String, Value>]) -> Vec<Value> {
172        rows.iter()
173            .filter_map(|row| row.get(self.relation.from_key).cloned())
174            .collect()
175    }
176
177    /// 批量查询关联表(Oracle IN >1000 分批)
178    async fn batch_query_related(
179        &self,
180        conn: &mut dyn Connection,
181        pk_values: &[Value],
182    ) -> Result<Vec<HashMap<String, Value>>, DbError> {
183        let batch_size = 1000;
184        let mut all_rows = Vec::new();
185
186        for chunk in pk_values.chunks(batch_size) {
187            let sql = self.build_related_sql(chunk.len());
188            let rows = conn.query_with_params(&sql, chunk).await?;
189            all_rows.extend(rows);
190        }
191
192        Ok(all_rows)
193    }
194
195    /// 生成关联表查询 SQL(参数化 WHERE IN)
196    fn build_related_sql(&self, param_count: usize) -> String {
197        let placeholders: Vec<String> = (0..param_count).map(|_| "?".to_string()).collect();
198        format!(
199            "SELECT * FROM {} WHERE {} IN ({})",
200            self.relation.to_entity,
201            self.relation.to_key,
202            placeholders.join(", ")
203        )
204    }
205
206    /// 按外键值分组关联行
207    fn group_by_foreign_key(
208        &self,
209        rows: Vec<HashMap<String, Value>>,
210        fk_key: &str,
211    ) -> HashMap<String, Vec<HashMap<String, Value>>> {
212        let mut grouped: HashMap<String, Vec<HashMap<String, Value>>> = HashMap::new();
213        for row in rows {
214            let fk = row.get(fk_key).cloned().unwrap_or(Value::Null);
215            let key = value_to_key(&fk);
216            grouped.entry(key).or_default().push(row);
217        }
218        grouped
219    }
220}
221
222/// 将 Value 转换为字符串键(用于 HashMap 分组,因 Value 含 f32/f64 不实现 Hash/Eq)
223fn value_to_key(value: &Value) -> String {
224    match value {
225        Value::Null => "null".to_string(),
226        Value::Bool(b) => format!("bool:{}", b),
227        Value::I8(v) => format!("i8:{}", v),
228        Value::I16(v) => format!("i16:{}", v),
229        Value::I32(v) => format!("i32:{}", v),
230        Value::I64(v) => format!("i64:{}", v),
231        Value::U8(v) => format!("u8:{}", v),
232        Value::U16(v) => format!("u16:{}", v),
233        Value::U32(v) => format!("u32:{}", v),
234        Value::U64(v) => format!("u64:{}", v),
235        Value::F32(v) => format!("f32:{}", v),
236        Value::F64(v) => format!("f64:{}", v),
237        Value::String(s) => format!("str:{}", s),
238        _ => format!("other:{:?}", value),
239    }
240}
241
242/// 一行 API:Eager Loading 端到端自动执行与组装
243///
244/// 执行主表查询 → 提取主键 → 批量查询关联表 → 分组组装
245/// 消除 N+1 查询(2 条 SQL 而非 N+1 条)。
246///
247/// # 参数
248///
249/// - `conn`:数据库连接
250/// - `main_sql`:主表查询 SQL(如 `"SELECT * FROM users"`)
251/// - `relation`:关联关系定义
252///
253/// # 返回
254///
255/// `Vec<(主表行, Vec<关联行>)>`
256///
257/// # 异常处理
258///
259/// - 主表查询失败 → 立即返回 `Err`,不执行关联查询
260/// - 关联表查询失败 → 返回 `Err`
261/// - 主表结果为空 → 返回 `Ok(Vec::new())`,不执行关联查询
262/// - 孤立关联记录(外键不匹配)→ 跳过
263pub async fn eager_load_all(
264    conn: &mut dyn Connection,
265    main_sql: &str,
266    relation: &RelationDef,
267) -> Result<Vec<EagerResult>, DbError> {
268    let loader = EagerLoader::new(relation.clone());
269    loader.load_many(conn, main_sql).await
270}
271
272/// 一行 API:HasOne / BelongsTo 单条关联加载(JOIN 策略)
273///
274/// 返回 `Vec<(主表行, Option<关联行>)>`
275pub async fn eager_load_one(
276    conn: &mut dyn Connection,
277    main_sql: &str,
278    relation: &RelationDef,
279) -> Result<Vec<(HashMap<String, Value>, Option<HashMap<String, Value>>)>, DbError> {
280    let main_rows = conn.query(main_sql).await?;
281
282    if main_rows.is_empty() {
283        return Ok(Vec::new());
284    }
285
286    let fk_values: Vec<Value> = main_rows
287        .iter()
288        .filter_map(|row| row.get(relation.to_key).cloned())
289        .collect();
290
291    if fk_values.is_empty() {
292        return Ok(main_rows.into_iter().map(|r| (r, None)).collect());
293    }
294
295    let placeholder: Vec<String> = (0..fk_values.len()).map(|_| "?".to_string()).collect();
296    let related_sql = format!(
297        "SELECT * FROM {} WHERE {} IN ({})",
298        relation.to_entity,
299        relation.from_key,
300        placeholder.join(", ")
301    );
302
303    let related_rows = conn.query_with_params(&related_sql, &fk_values).await?;
304
305    let mut related_map: HashMap<String, HashMap<String, Value>> = HashMap::new();
306    for row in related_rows {
307        let pk = row.get(relation.from_key).cloned().unwrap_or(Value::Null);
308        related_map.insert(value_to_key(&pk), row);
309    }
310
311    let results = main_rows
312        .into_iter()
313        .map(|row| {
314            let fk = row.get(relation.to_key).cloned().unwrap_or(Value::Null);
315            let related = related_map.get(&value_to_key(&fk)).cloned();
316            (row, related)
317        })
318        .collect();
319
320    Ok(results)
321}
322
323#[cfg(test)]
324mod tests {
325    use super::*;
326    use crate::relation_trait::RelationKind;
327
328    #[test]
329    fn test_eager_loader_new() {
330        let relation = RelationDef::new(
331            "orders",
332            "users",
333            "orders",
334            "id",
335            "user_id",
336            RelationKind::HasMany,
337        );
338        let loader = EagerLoader::new(relation);
339        assert_eq!(loader.relation.name, "orders");
340        assert!(loader.children.is_empty());
341    }
342
343    #[test]
344    fn test_eager_loader_with_children() {
345        let relation = RelationDef::new(
346            "orders",
347            "users",
348            "orders",
349            "id",
350            "user_id",
351            RelationKind::HasMany,
352        );
353        let child_relation = RelationDef::new(
354            "items",
355            "orders",
356            "order_items",
357            "id",
358            "order_id",
359            RelationKind::HasMany,
360        );
361        let loader = EagerLoader::new(relation).with(child_relation);
362        assert_eq!(loader.children.len(), 1);
363    }
364
365    #[test]
366    fn test_build_related_sql() {
367        let relation = RelationDef::new(
368            "orders",
369            "users",
370            "orders",
371            "id",
372            "user_id",
373            RelationKind::HasMany,
374        );
375        let loader = EagerLoader::new(relation);
376        let sql = loader.build_related_sql(3);
377        assert!(sql.contains("SELECT * FROM orders"));
378        assert!(sql.contains("user_id IN (?, ?, ?)"));
379    }
380
381    #[test]
382    fn test_extract_primary_keys() {
383        let relation = RelationDef::new(
384            "orders",
385            "users",
386            "orders",
387            "id",
388            "user_id",
389            RelationKind::HasMany,
390        );
391        let loader = EagerLoader::new(relation);
392
393        let mut row1 = HashMap::new();
394        row1.insert("id".to_string(), Value::I64(1));
395        let mut row2 = HashMap::new();
396        row2.insert("id".to_string(), Value::I64(2));
397
398        let pks = loader.extract_primary_keys(&[row1, row2]);
399        assert_eq!(pks.len(), 2);
400    }
401
402    #[test]
403    fn test_group_by_foreign_key() {
404        let relation = RelationDef::new(
405            "orders",
406            "users",
407            "orders",
408            "id",
409            "user_id",
410            RelationKind::HasMany,
411        );
412        let loader = EagerLoader::new(relation);
413
414        let mut row1 = HashMap::new();
415        row1.insert("user_id".to_string(), Value::I64(1));
416        row1.insert("id".to_string(), Value::I64(101));
417        let mut row2 = HashMap::new();
418        row2.insert("user_id".to_string(), Value::I64(1));
419        row2.insert("id".to_string(), Value::I64(102));
420        let mut row3 = HashMap::new();
421        row3.insert("user_id".to_string(), Value::I64(2));
422        row3.insert("id".to_string(), Value::I64(103));
423
424        let grouped = loader.group_by_foreign_key(vec![row1, row2, row3], "user_id");
425        assert_eq!(grouped.len(), 2);
426        assert_eq!(grouped.get("i64:1").unwrap().len(), 2);
427        assert_eq!(grouped.get("i64:2").unwrap().len(), 1);
428    }
429}