use std::collections::HashMap;
use async_trait::async_trait;
use turso_orm_driver::{ConnectionTrait, Row};
use turso_sql::{Condition, Expr, JoinType, Value};
use super::base_entity::EntityTrait;
use super::column::ColumnTrait;
use super::model::{FromQueryResult, ModelTrait};
use super::relation::{Related, RelationDef, column_of};
use crate::{DbErr, Result};
#[async_trait]
pub trait LoaderTrait {
type Entity: EntityTrait;
async fn load_one<R, C>(&self, _: R, db: &C) -> Result<Vec<Option<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait;
async fn load_many<R, C>(&self, _: R, db: &C) -> Result<Vec<Vec<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait;
async fn load_many_to_many<R, C>(&self, _: R, db: &C) -> Result<Vec<Vec<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait;
}
fn key(values: &[Value]) -> String {
values
.iter()
.map(Value::to_literal)
.collect::<Vec<_>>()
.join("\u{1f}")
}
fn columns<E: EntityTrait>(names: &[&'static str]) -> Result<Vec<E::Column>> {
names
.iter()
.map(|n| column_of::<E>(n).ok_or_else(|| DbErr::Custom(format!("unknown column {n}"))))
.collect()
}
fn keys_of<E: EntityTrait>(models: &[E::Model], from_cols: &[E::Column]) -> Vec<String> {
models
.iter()
.map(|m| key(&from_cols.iter().map(|c| m.get(*c)).collect::<Vec<_>>()))
.collect()
}
fn keys_condition<E: EntityTrait>(
models: &[E::Model],
from_cols: &[E::Column],
targets: &[Expr],
) -> Condition {
if targets.len() == 1 {
let values: Vec<Value> = models.iter().map(|m| m.get(from_cols[0])).collect();
return Condition::all().add(targets[0].clone().is_in(values));
}
let mut cond = Condition::any();
for m in models {
let mut c = Condition::all();
for (f, t) in from_cols.iter().zip(targets) {
c = c.add(t.clone().eq(Expr::val(m.get(*f))));
}
cond = cond.add(c);
}
cond
}
async fn load_direct<E, R, C>(
def: &RelationDef,
models: &[E::Model],
db: &C,
) -> Result<(Vec<String>, HashMap<String, Vec<R::Model>>)>
where
E: EntityTrait,
R: EntityTrait,
C: ConnectionTrait,
{
let from_cols = columns::<E>(&def.from_col)?;
let to_cols = columns::<R>(&def.to_col)?;
let keys = keys_of::<E>(models, &from_cols);
let mut grouped: HashMap<String, Vec<R::Model>> = HashMap::new();
if models.is_empty() {
return Ok((keys, grouped));
}
let targets: Vec<Expr> = to_cols.iter().map(|c| c.into_expr()).collect();
let query = R::find().filter(keys_condition::<E>(models, &from_cols, &targets));
for related in query.all(db).await? {
let k = key(&to_cols.iter().map(|c| related.get(*c)).collect::<Vec<_>>());
grouped.entry(k).or_default().push(related);
}
Ok((keys, grouped))
}
async fn load_via<E, R, C>(
via: &RelationDef,
to: &RelationDef,
models: &[E::Model],
db: &C,
) -> Result<(Vec<String>, HashMap<String, Vec<R::Model>>)>
where
E: EntityTrait,
R: EntityTrait,
C: ConnectionTrait,
{
let from_cols = columns::<E>(&via.from_col)?;
let keys = keys_of::<E>(models, &from_cols);
let mut grouped: HashMap<String, Vec<R::Model>> = HashMap::new();
if models.is_empty() {
return Ok((keys, grouped));
}
let mut query = R::find();
let junction = query.join_table(JoinType::Inner, to.from_tbl, |r| {
to.join_condition_refs(r, R::TABLE_NAME)
});
let targets: Vec<Expr> = via
.to_col
.iter()
.map(|c| Expr::col((junction.clone(), *c)))
.collect();
let aliases: Vec<String> = via.to_col.iter().map(|c| format!("__via_{c}")).collect();
for (target, alias) in targets.iter().zip(&aliases) {
query = query.expr_as(target.clone(), alias.clone());
}
let query = query.filter(keys_condition::<E>(models, &from_cols, &targets));
for row in query.into_model::<Row>().all(db).await? {
let related = R::Model::from_query_result(&row, "")?;
let values: Vec<Value> = aliases
.iter()
.map(|a| row.raw(a.as_str()).cloned().unwrap_or(Value::Null))
.collect();
grouped.entry(key(&values)).or_default().push(related);
}
Ok((keys, grouped))
}
#[async_trait]
impl<M: ModelTrait> LoaderTrait for Vec<M> {
type Entity = M::Entity;
async fn load_one<R, C>(&self, r: R, db: &C) -> Result<Vec<Option<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait,
{
self.as_slice().load_one(r, db).await
}
async fn load_many<R, C>(&self, r: R, db: &C) -> Result<Vec<Vec<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait,
{
self.as_slice().load_many(r, db).await
}
async fn load_many_to_many<R, C>(&self, r: R, db: &C) -> Result<Vec<Vec<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait,
{
self.as_slice().load_many_to_many(r, db).await
}
}
#[async_trait]
impl<M: ModelTrait> LoaderTrait for [M] {
type Entity = M::Entity;
async fn load_one<R, C>(&self, _: R, db: &C) -> Result<Vec<Option<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait,
{
let def = <M::Entity as Related<R>>::to();
let (keys, mut grouped) = load_direct::<M::Entity, R, C>(&def, self, db).await?;
Ok(keys
.iter()
.map(|k| {
grouped.get_mut(k).and_then(|v| {
if v.is_empty() {
None
} else {
Some(v[0].clone())
}
})
})
.collect())
}
async fn load_many<R, C>(&self, r: R, db: &C) -> Result<Vec<Vec<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait,
{
if <M::Entity as Related<R>>::via().is_some() {
return self.load_many_to_many(r, db).await;
}
let def = <M::Entity as Related<R>>::to();
let (keys, grouped) = load_direct::<M::Entity, R, C>(&def, self, db).await?;
Ok(keys
.iter()
.map(|k| grouped.get(k).cloned().unwrap_or_default())
.collect())
}
async fn load_many_to_many<R, C>(&self, _: R, db: &C) -> Result<Vec<Vec<R::Model>>>
where
R: EntityTrait,
Self::Entity: Related<R>,
C: ConnectionTrait,
{
let via = <M::Entity as Related<R>>::via().ok_or_else(|| {
DbErr::Custom(format!(
"{} is not related to {} through a junction table",
M::Entity::TABLE_NAME,
R::TABLE_NAME
))
})?;
let to = <M::Entity as Related<R>>::to();
let (keys, grouped) = load_via::<M::Entity, R, C>(&via, &to, self, db).await?;
Ok(keys
.iter()
.map(|k| grouped.get(k).cloned().unwrap_or_default())
.collect())
}
}