Skip to main content

turso_orm/entity/
loader.rs

1//! Batch loading of related rows, modeled by [`LoaderTrait`].
2//!
3//! Given a list of models, the loaders fetch every related row in a single
4//! query and hand them back aligned with the input, which avoids the N+1
5//! pattern of calling `find_related` in a loop. The trait is implemented on
6//! `Vec<M>` and `[M]` so that the result of `.all(db)` can be passed
7//! straight in.
8//!
9//! Rows are grouped by the literal rendering of their key values rather
10//! than by the values themselves. That sidesteps the need for `Hash` and
11//! `Eq` on every key type — `f64` has neither — at the cost of a small
12//! string per row, and the literal form is unambiguous because
13//! `Value::to_literal` quotes text and distinguishes `1` from `1.0`.
14//!
15//! A many-to-many relation is loaded through its junction table in one
16//! query too: the related rows are selected together with the junction
17//! columns that point back at the input models, aliased `__via_<col>`, and
18//! grouped on those.
19
20use std::collections::HashMap;
21
22use async_trait::async_trait;
23use turso_orm_driver::{ConnectionTrait, Row};
24use turso_sql::{Condition, Expr, JoinType, Value};
25
26use super::base_entity::EntityTrait;
27use super::column::ColumnTrait;
28use super::model::{FromQueryResult, ModelTrait};
29use super::relation::{Related, RelationDef, column_of};
30use crate::{DbErr, Result};
31
32/// Loads related rows for a batch of models in one query.
33#[async_trait]
34pub trait LoaderTrait {
35    /// The entity of the models in the batch.
36    type Entity: EntityTrait;
37
38    /// Loads, for each model, the related `R` row if any; for `has_one` and `belongs_to`.
39    ///
40    /// The result has one entry per input model, in input order.
41    ///
42    /// # Errors
43    ///
44    /// Returns [`DbErr::Custom`] when the relation names a column the
45    /// entity does not have; [`DbErr::Driver`] when the query fails or a
46    /// row cannot be decoded.
47    async fn load_one<R, C>(&self, _: R, db: &C) -> Result<Vec<Option<R::Model>>>
48    where
49        R: EntityTrait,
50        Self::Entity: Related<R>,
51        C: ConnectionTrait;
52
53    /// Loads, for each model, the related `R` rows; for `has_many`, and for
54    /// a many-to-many relation, which is followed through its junction.
55    ///
56    /// The result has one entry per input model, in input order.
57    ///
58    /// # Errors
59    ///
60    /// Returns [`DbErr::Custom`] when the relation names a column the
61    /// entity does not have; [`DbErr::Driver`] when the query fails or a
62    /// row cannot be decoded.
63    async fn load_many<R, C>(&self, _: R, db: &C) -> Result<Vec<Vec<R::Model>>>
64    where
65        R: EntityTrait,
66        Self::Entity: Related<R>,
67        C: ConnectionTrait;
68
69    /// Loads, for each model, the `R` rows reached through the junction
70    /// table of a many-to-many relation.
71    ///
72    /// [`load_many`](Self::load_many) does the same when the relation
73    /// declares a junction; this method exists to make the intent explicit
74    /// and fails on a direct relation.
75    ///
76    /// # Errors
77    ///
78    /// Returns [`DbErr::Custom`] when the relation has no junction or names
79    /// a column the entity does not have; [`DbErr::Driver`] when the query
80    /// fails or a row cannot be decoded.
81    async fn load_many_to_many<R, C>(&self, _: R, db: &C) -> Result<Vec<Vec<R::Model>>>
82    where
83        R: EntityTrait,
84        Self::Entity: Related<R>,
85        C: ConnectionTrait;
86}
87
88/// Renders key values as a grouping key.
89///
90/// The unit separator keeps composite keys unambiguous even when a text
91/// component could itself contain the literal of another component.
92fn key(values: &[Value]) -> String {
93    values
94        .iter()
95        .map(Value::to_literal)
96        .collect::<Vec<_>>()
97        .join("\u{1f}")
98}
99
100/// Maps the column names of one side of a relation to column variants of `E`.
101///
102/// # Errors
103///
104/// Returns [`DbErr::Custom`] when a name is not a column of `E`.
105fn columns<E: EntityTrait>(names: &[&'static str]) -> Result<Vec<E::Column>> {
106    names
107        .iter()
108        .map(|n| column_of::<E>(n).ok_or_else(|| DbErr::Custom(format!("unknown column {n}"))))
109        .collect()
110}
111
112/// The grouping key of every model, from its `from_cols`, in input order.
113fn keys_of<E: EntityTrait>(models: &[E::Model], from_cols: &[E::Column]) -> Vec<String> {
114    models
115        .iter()
116        .map(|m| key(&from_cols.iter().map(|c| m.get(*c)).collect::<Vec<_>>()))
117        .collect()
118}
119
120/// Builds the condition matching the keys of `models` on `targets`, which
121/// are column expressions on the related or junction side.
122///
123/// A single key column renders `IN (...)`; SQLite has no row-value `IN`,
124/// so composite keys become an `OR` of one `AND` group per model.
125fn keys_condition<E: EntityTrait>(
126    models: &[E::Model],
127    from_cols: &[E::Column],
128    targets: &[Expr],
129) -> Condition {
130    if targets.len() == 1 {
131        let values: Vec<Value> = models.iter().map(|m| m.get(from_cols[0])).collect();
132        return Condition::all().add(targets[0].clone().is_in(values));
133    }
134    let mut cond = Condition::any();
135    for m in models {
136        let mut c = Condition::all();
137        for (f, t) in from_cols.iter().zip(targets) {
138            c = c.add(t.clone().eq(Expr::val(m.get(*f))));
139        }
140        cond = cond.add(c);
141    }
142    cond
143}
144
145/// Fetches every `R` row directly related to `models` and groups them by key.
146///
147/// Returns the grouping key of each input model, in order, alongside the
148/// grouped rows, so that the callers can align the two without re-deriving
149/// keys.
150///
151/// # Errors
152///
153/// Returns [`DbErr::Custom`] when the relation names an unknown column;
154/// [`DbErr::Driver`] when the query fails or a row cannot be decoded.
155async fn load_direct<E, R, C>(
156    def: &RelationDef,
157    models: &[E::Model],
158    db: &C,
159) -> Result<(Vec<String>, HashMap<String, Vec<R::Model>>)>
160where
161    E: EntityTrait,
162    R: EntityTrait,
163    C: ConnectionTrait,
164{
165    let from_cols = columns::<E>(&def.from_col)?;
166    let to_cols = columns::<R>(&def.to_col)?;
167    let keys = keys_of::<E>(models, &from_cols);
168    let mut grouped: HashMap<String, Vec<R::Model>> = HashMap::new();
169    // An empty batch would render `IN ()`, which is not valid SQL.
170    if models.is_empty() {
171        return Ok((keys, grouped));
172    }
173    let targets: Vec<Expr> = to_cols.iter().map(|c| c.into_expr()).collect();
174    let query = R::find().filter(keys_condition::<E>(models, &from_cols, &targets));
175    for related in query.all(db).await? {
176        let k = key(&to_cols.iter().map(|c| related.get(*c)).collect::<Vec<_>>());
177        grouped.entry(k).or_default().push(related);
178    }
179    Ok((keys, grouped))
180}
181
182/// Fetches every `R` row related to `models` through a junction table and
183/// groups them by the junction columns pointing back at the models.
184///
185/// # Errors
186///
187/// Returns [`DbErr::Custom`] when a relation names an unknown column;
188/// [`DbErr::Driver`] when the query fails or a row cannot be decoded.
189async fn load_via<E, R, C>(
190    via: &RelationDef,
191    to: &RelationDef,
192    models: &[E::Model],
193    db: &C,
194) -> Result<(Vec<String>, HashMap<String, Vec<R::Model>>)>
195where
196    E: EntityTrait,
197    R: EntityTrait,
198    C: ConnectionTrait,
199{
200    let from_cols = columns::<E>(&via.from_col)?;
201    let keys = keys_of::<E>(models, &from_cols);
202    let mut grouped: HashMap<String, Vec<R::Model>> = HashMap::new();
203    if models.is_empty() {
204        return Ok((keys, grouped));
205    }
206    let mut query = R::find();
207    let junction = query.join_table(JoinType::Inner, to.from_tbl, |r| {
208        to.join_condition_refs(r, R::TABLE_NAME)
209    });
210    let targets: Vec<Expr> = via
211        .to_col
212        .iter()
213        .map(|c| Expr::col((junction.clone(), *c)))
214        .collect();
215    // The junction columns ride along under a prefix no entity column can
216    // carry, so that the related model still decodes from the same row.
217    let aliases: Vec<String> = via.to_col.iter().map(|c| format!("__via_{c}")).collect();
218    for (target, alias) in targets.iter().zip(&aliases) {
219        query = query.expr_as(target.clone(), alias.clone());
220    }
221    let query = query.filter(keys_condition::<E>(models, &from_cols, &targets));
222    for row in query.into_model::<Row>().all(db).await? {
223        let related = R::Model::from_query_result(&row, "")?;
224        let values: Vec<Value> = aliases
225            .iter()
226            .map(|a| row.raw(a.as_str()).cloned().unwrap_or(Value::Null))
227            .collect();
228        grouped.entry(key(&values)).or_default().push(related);
229    }
230    Ok((keys, grouped))
231}
232
233#[async_trait]
234impl<M: ModelTrait> LoaderTrait for Vec<M> {
235    type Entity = M::Entity;
236
237    async fn load_one<R, C>(&self, r: R, db: &C) -> Result<Vec<Option<R::Model>>>
238    where
239        R: EntityTrait,
240        Self::Entity: Related<R>,
241        C: ConnectionTrait,
242    {
243        self.as_slice().load_one(r, db).await
244    }
245
246    async fn load_many<R, C>(&self, r: R, db: &C) -> Result<Vec<Vec<R::Model>>>
247    where
248        R: EntityTrait,
249        Self::Entity: Related<R>,
250        C: ConnectionTrait,
251    {
252        self.as_slice().load_many(r, db).await
253    }
254
255    async fn load_many_to_many<R, C>(&self, r: R, db: &C) -> Result<Vec<Vec<R::Model>>>
256    where
257        R: EntityTrait,
258        Self::Entity: Related<R>,
259        C: ConnectionTrait,
260    {
261        self.as_slice().load_many_to_many(r, db).await
262    }
263}
264
265#[async_trait]
266impl<M: ModelTrait> LoaderTrait for [M] {
267    type Entity = M::Entity;
268
269    async fn load_one<R, C>(&self, _: R, db: &C) -> Result<Vec<Option<R::Model>>>
270    where
271        R: EntityTrait,
272        Self::Entity: Related<R>,
273        C: ConnectionTrait,
274    {
275        let def = <M::Entity as Related<R>>::to();
276        let (keys, mut grouped) = load_direct::<M::Entity, R, C>(&def, self, db).await?;
277        // Several input models may share a key, so the first match is cloned
278        // rather than moved out of the group.
279        Ok(keys
280            .iter()
281            .map(|k| {
282                grouped.get_mut(k).and_then(|v| {
283                    if v.is_empty() {
284                        None
285                    } else {
286                        Some(v[0].clone())
287                    }
288                })
289            })
290            .collect())
291    }
292
293    async fn load_many<R, C>(&self, r: R, db: &C) -> Result<Vec<Vec<R::Model>>>
294    where
295        R: EntityTrait,
296        Self::Entity: Related<R>,
297        C: ConnectionTrait,
298    {
299        if <M::Entity as Related<R>>::via().is_some() {
300            return self.load_many_to_many(r, db).await;
301        }
302        let def = <M::Entity as Related<R>>::to();
303        let (keys, grouped) = load_direct::<M::Entity, R, C>(&def, self, db).await?;
304        Ok(keys
305            .iter()
306            .map(|k| grouped.get(k).cloned().unwrap_or_default())
307            .collect())
308    }
309
310    async fn load_many_to_many<R, C>(&self, _: R, db: &C) -> Result<Vec<Vec<R::Model>>>
311    where
312        R: EntityTrait,
313        Self::Entity: Related<R>,
314        C: ConnectionTrait,
315    {
316        let via = <M::Entity as Related<R>>::via().ok_or_else(|| {
317            DbErr::Custom(format!(
318                "{} is not related to {} through a junction table",
319                M::Entity::TABLE_NAME,
320                R::TABLE_NAME
321            ))
322        })?;
323        let to = <M::Entity as Related<R>>::to();
324        let (keys, grouped) = load_via::<M::Entity, R, C>(&via, &to, self, db).await?;
325        Ok(keys
326            .iter()
327            .map(|k| grouped.get(k).cloned().unwrap_or_default())
328            .collect())
329    }
330}