1use 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#[async_trait]
34pub trait LoaderTrait {
35 type Entity: EntityTrait;
37
38 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 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 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
88fn key(values: &[Value]) -> String {
93 values
94 .iter()
95 .map(Value::to_literal)
96 .collect::<Vec<_>>()
97 .join("\u{1f}")
98}
99
100fn 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
112fn 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
120fn 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
145async 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 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
182async 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 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 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}