use std::collections::HashMap;
use std::marker::PhantomData;
use super::{Db, Executor, Model, ModelKey, Query, ToDbValue, quote, sql};
use crate::Result;
pub trait ForeignKey<K>: foreign_key::Sealed<K> {
fn key(self) -> Option<K>;
}
mod foreign_key {
pub trait Sealed<K> {}
impl<K: super::ModelKey> Sealed<K> for K {}
impl<K: super::ModelKey> Sealed<K> for Option<K> {}
}
impl<K: ModelKey> ForeignKey<K> for K {
fn key(self) -> Option<K> {
Some(self)
}
}
impl<K: ModelKey> ForeignKey<K> for Option<K> {
fn key(self) -> Option<K> {
self
}
}
pub fn belongs_to<'a, P: Model, C, K: ForeignKey<P::Key>>(
db: &'a Db,
children: &[C],
foreign_key: impl Fn(&C) -> K,
) -> impl Future<Output = Result<HashMap<P::Key, P>>> + Send + 'a {
let mut ids: Vec<P::Key> = children
.iter()
.filter_map(|c| foreign_key(c).key())
.collect();
ids.sort_unstable();
ids.dedup();
async move {
if ids.is_empty() {
return Ok(HashMap::new());
}
Ok(P::find_many(db, ids)
.await?
.into_iter()
.map(|parent| (parent.id(), parent))
.collect())
}
}
pub fn has_many<'a, C: Model, P: Model, K: ForeignKey<P::Key>>(
db: &'a Db,
parents: &[P],
children: Query<C>,
column: &str,
foreign_key: impl Fn(&C) -> K + Send + 'a,
) -> impl Future<Output = Result<HashMap<P::Key, Vec<C>>>> + Send + 'a {
let ids: Vec<P::Key> = parents.iter().map(Model::id).collect();
let query = (!ids.is_empty()).then(|| children.where_in(column, ids));
async move {
let mut grouped: HashMap<P::Key, Vec<C>> = HashMap::new();
let Some(query) = query else {
return Ok(grouped);
};
let rows = query.get(db).await?;
for child in rows {
if let Some(parent) = foreign_key(&child).key() {
grouped.entry(parent).or_default().push(child);
}
}
Ok(grouped)
}
}
#[allow(clippy::too_many_arguments)]
pub fn has_many_through<'a, C, T, P, KT, KC>(
db: &'a Db,
parents: &[P],
through: Query<T>,
through_column: &str,
through_key: impl Fn(&T) -> KT + Send + 'a,
children: Query<C>,
column: &str,
foreign_key: impl Fn(&C) -> KC + Send + 'a,
) -> impl Future<Output = Result<HashMap<P::Key, Vec<C>>>> + Send + 'a
where
C: Model,
T: Model,
P: Model,
KT: ForeignKey<P::Key>,
KC: ForeignKey<T::Key>,
{
let ids: Vec<P::Key> = parents.iter().map(Model::id).collect();
let middle = (!ids.is_empty()).then(|| through.where_in(through_column, ids));
let column = column.to_owned();
async move {
let mut grouped: HashMap<P::Key, Vec<C>> = HashMap::new();
let Some(middle) = middle else {
return Ok(grouped);
};
let parent_of: HashMap<T::Key, P::Key> = middle
.get(db)
.await?
.iter()
.filter_map(|row| Some((row.id(), through_key(row).key()?)))
.collect();
if parent_of.is_empty() {
return Ok(grouped);
}
let middle_ids: Vec<T::Key> = parent_of.keys().cloned().collect();
for child in children.where_in(&column, middle_ids).get(db).await? {
let parent = foreign_key(&child)
.key()
.and_then(|middle| parent_of.get(&middle).cloned());
if let Some(parent) = parent {
grouped.entry(parent).or_default().push(child);
}
}
Ok(grouped)
}
}
pub fn count_many<'a, C: Model, P: Model>(
db: &'a Db,
parents: &[P],
children: Query<C>,
column: &str,
) -> impl Future<Output = Result<HashMap<P::Key, i64>>> + Send + 'a {
grouped(db, parents, children, column, "COUNT(*)".to_owned())
}
pub fn sum_many<'a, T: super::Number + Default + Send + 'a, C: Model, P: Model>(
db: &'a Db,
parents: &[P],
children: Query<C>,
column: &str,
sum_column: &str,
) -> impl Future<Output = Result<HashMap<P::Key, T>>> + Send + 'a {
let expression = if C::COLUMNS.contains(&sum_column) {
format!(
"CAST(COALESCE(SUM({}), 0) AS {})",
quote(sum_column),
T::SQL_TYPE
)
} else {
format!("SUM({})", quote(sum_column))
};
let checked = children.check_column(sum_column);
grouped(db, parents, checked, column, expression)
}
fn grouped<'a, T: crate::db::FromDb + Default + Send + 'a, C: Model, P: Model>(
db: &'a Db,
parents: &[P],
children: Query<C>,
column: &str,
expression: String,
) -> impl Future<Output = Result<HashMap<P::Key, T>>> + Send + 'a {
let ids: Vec<P::Key> = parents.iter().map(Model::id).collect();
let column = column.to_owned();
async move {
let mut totals: HashMap<P::Key, T> =
ids.iter().map(|id| (id.clone(), T::default())).collect();
if ids.is_empty() {
return Ok(totals);
}
let rows: Vec<(P::Key, T)> = children
.where_in(&column, ids)
.group_by(&column)
.select_as(db, &format!("{}, {expression}", quote(&column)))
.await?;
totals.extend(rows);
Ok(totals)
}
}
pub struct Pivot<L = i64, R = i64> {
table: &'static str,
left: &'static str,
right: &'static str,
timestamps: bool,
keys: PhantomData<fn() -> (L, R)>,
}
impl<L, R> Clone for Pivot<L, R> {
fn clone(&self) -> Self {
*self
}
}
impl<L, R> Copy for Pivot<L, R> {}
impl<L, R> std::fmt::Debug for Pivot<L, R> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Pivot")
.field("table", &self.table)
.field("left", &self.left)
.field("right", &self.right)
.field("timestamps", &self.timestamps)
.finish()
}
}
pub type PivotData<'a> = &'a [(&'a str, &'a (dyn super::ToDbValue + Sync))];
impl<L, R> Pivot<L, R> {
pub const fn new(table: &'static str, left: &'static str, right: &'static str) -> Self {
Self {
table,
left,
right,
timestamps: false,
keys: PhantomData,
}
}
pub const fn with_timestamps(mut self) -> Self {
self.timestamps = true;
self
}
pub const fn inverse(self) -> Pivot<R, L> {
Pivot {
table: self.table,
left: self.right,
right: self.left,
timestamps: self.timestamps,
keys: PhantomData,
}
}
}
impl<L: ModelKey, R: ModelKey> Pivot<L, R> {
async fn insert_link(
&self,
conn: &mut super::Conn<'_>,
left: &L,
right: &R,
data: PivotData<'_>,
) -> Result<bool> {
let mut columns = vec![quote(self.left), quote(self.right)];
let mut values = vec![left.to_db_value(), right.to_db_value()];
for (column, value) in data {
columns.push(quote(column));
values.push(value.to_db_value());
}
if self.timestamps {
let at = super::ToDbValue::to_db_value(&super::now());
columns.extend([quote("created_at"), quote("updated_at")]);
values.extend([at.clone(), at]);
}
let marks = vec!["?"; columns.len()].join(", ");
let added = sql(format!(
"INSERT INTO {table} ({columns}) SELECT {marks} \
WHERE NOT EXISTS (SELECT 1 FROM {table} WHERE {l} = ? AND {r} = ?)",
table = quote(self.table),
columns = columns.join(", "),
l = quote(self.left),
r = quote(self.right),
))
.bind_all(values)
.bind(left.to_db_value())
.bind(right.to_db_value())
.execute(conn.reborrow())
.await?;
Ok(added > 0)
}
pub async fn ids<'c>(&self, db: impl Executor<'c>, left: L) -> Result<Vec<R>> {
Ok(sql(format!(
"SELECT {} FROM {} WHERE {} = ? ORDER BY {}",
quote(self.right),
quote(self.table),
quote(self.left),
quote(self.right)
))
.bind(left)
.scalars(db)
.await?)
}
pub async fn attach<'c>(
&self,
db: impl Executor<'c>,
left: L,
rights: impl IntoIterator<Item = R>,
) -> Result<u64> {
let mut conn = db.into_conn();
let mut added = 0;
for right in rights {
added += u64::from(self.insert_link(&mut conn, &left, &right, &[]).await?);
}
Ok(added)
}
pub async fn attach_with<'c>(
&self,
db: impl Executor<'c>,
left: L,
right: R,
data: PivotData<'_>,
) -> Result<bool> {
let mut conn = db.into_conn();
self.insert_link(&mut conn, &left, &right, data).await
}
pub async fn update_pivot<'c>(
&self,
db: impl Executor<'c>,
left: L,
right: R,
data: PivotData<'_>,
) -> Result<bool> {
let mut sets = Vec::new();
let mut values = Vec::new();
for (column, value) in data {
sets.push(format!("{} = ?", quote(column)));
values.push(value.to_db_value());
}
if self.timestamps {
sets.push(format!("{} = ?", quote("updated_at")));
values.push(super::ToDbValue::to_db_value(&super::now()));
}
if sets.is_empty() {
return Ok(false);
}
let changed = sql(format!(
"UPDATE {} SET {} WHERE {} = ? AND {} = ?",
quote(self.table),
sets.join(", "),
quote(self.left),
quote(self.right)
))
.bind_all(values)
.bind(left)
.bind(right)
.execute(db)
.await?;
Ok(changed > 0)
}
pub async fn toggle(
&self,
db: &Db,
left: L,
rights: impl IntoIterator<Item = R>,
) -> Result<(Vec<R>, Vec<R>)> {
let mut tx = db.begin().await?;
let current = self.ids(&mut tx, left.clone()).await?;
let (detach, attach): (Vec<R>, Vec<R>) =
rights.into_iter().partition(|id| current.contains(id));
self.detach(&mut tx, left.clone(), detach.iter().cloned())
.await?;
self.attach(&mut tx, left, attach.iter().cloned()).await?;
tx.commit().await?;
Ok((attach, detach))
}
pub async fn detach<'c>(
&self,
db: impl Executor<'c>,
left: L,
rights: impl IntoIterator<Item = R>,
) -> Result<u64> {
let rights: Vec<R> = rights.into_iter().collect();
if rights.is_empty() {
return Ok(0);
}
let marks = vec!["?"; rights.len()].join(", ");
Ok(sql(format!(
"DELETE FROM {} WHERE {} = ? AND {} IN ({marks})",
quote(self.table),
quote(self.left),
quote(self.right)
))
.bind(left)
.bind_all(rights.iter().map(ToDbValue::to_db_value))
.execute(db)
.await?)
}
pub async fn sync(&self, db: &Db, left: L, rights: impl IntoIterator<Item = R>) -> Result {
let wanted: Vec<R> = rights.into_iter().collect();
let mut tx = db.begin().await?;
let current = self.ids(&mut tx, left.clone()).await?;
let gone: Vec<R> = current
.into_iter()
.filter(|id| !wanted.contains(id))
.collect();
self.detach(&mut tx, left.clone(), gone).await?;
self.attach(&mut tx, left, wanted).await?;
tx.commit().await?;
Ok(())
}
pub fn load<'a, T: Model<Key = R> + Clone>(
&'a self,
db: &'a Db,
lefts: impl IntoIterator<Item = L>,
) -> impl Future<Output = Result<HashMap<L, Vec<T>>>> + Send + 'a {
let lefts: Vec<L> = lefts.into_iter().collect();
self.load_ids(db, lefts)
}
async fn load_ids<T: Model<Key = R> + Clone>(
&self,
db: &Db,
lefts: Vec<L>,
) -> Result<HashMap<L, Vec<T>>> {
let mut grouped: HashMap<L, Vec<T>> = HashMap::new();
if lefts.is_empty() {
return Ok(grouped);
}
let marks = vec!["?"; lefts.len()].join(", ");
let links: Vec<(L, R)> = sql(format!(
"SELECT {}, {} FROM {} WHERE {} IN ({marks})",
quote(self.left),
quote(self.right),
quote(self.table),
quote(self.left)
))
.bind_all(lefts.iter().map(ToDbValue::to_db_value))
.fetch_as(db)
.await?;
let mut rights: Vec<R> = links.iter().map(|(_, right)| right.clone()).collect();
rights.sort_unstable();
rights.dedup();
let models: HashMap<R, T> = T::find_many(db, rights)
.await?
.into_iter()
.map(|model| (model.id(), model))
.collect();
for (left, right) in links {
if let Some(model) = models.get(&right) {
grouped.entry(left).or_default().push(model.clone());
}
}
Ok(grouped)
}
pub fn load_with_pivot<'a, T: Model<Key = R> + Clone, D: super::FromRow + Send + 'a>(
&'a self,
db: &'a Db,
lefts: impl IntoIterator<Item = L>,
) -> impl Future<Output = Result<HashMap<L, Vec<(T, D)>>>> + Send + 'a {
let lefts: Vec<L> = lefts.into_iter().collect();
async move {
let mut grouped: HashMap<L, Vec<(T, D)>> = HashMap::new();
if lefts.is_empty() {
return Ok(grouped);
}
let marks = vec!["?"; lefts.len()].join(", ");
let rows = sql(format!(
"SELECT * FROM {} WHERE {} IN ({marks})",
quote(self.table),
quote(self.left)
))
.bind_all(lefts.iter().map(ToDbValue::to_db_value))
.fetch_all(db)
.await?;
let mut links = Vec::with_capacity(rows.len());
for row in &rows {
let left: L = row.try_get(self.left)?;
let right: R = row.try_get(self.right)?;
links.push((left, right, D::from_row(row)?));
}
let mut rights: Vec<R> = links.iter().map(|(_, right, _)| right.clone()).collect();
rights.sort_unstable();
rights.dedup();
let models: HashMap<R, T> = T::find_many(db, rights)
.await?
.into_iter()
.map(|model| (model.id(), model))
.collect();
for (left, right, data) in links {
if let Some(model) = models.get(&right) {
grouped.entry(left).or_default().push((model.clone(), data));
}
}
Ok(grouped)
}
}
pub fn load_for<'a, T: Model<Key = R> + Clone, P: Model<Key = L>>(
&'a self,
db: &'a Db,
parents: &[P],
) -> impl Future<Output = Result<HashMap<L, Vec<T>>>> + Send + 'a {
let ids: Vec<L> = parents.iter().map(Model::id).collect();
self.load_ids(db, ids)
}
}
#[derive(Debug, Clone, Copy)]
pub struct Morph {
type_column: &'static str,
id_column: &'static str,
}
impl Morph {
pub const fn new(type_column: &'static str, id_column: &'static str) -> Self {
Self {
type_column,
id_column,
}
}
pub fn of<C: Model, P: Model>(&self, parent: &P, children: Query<C>) -> Query<C> {
children
.where_eq(self.type_column, P::TABLE)
.where_eq(self.id_column, parent.id())
}
pub fn load_many<'a, C: Model, P: Model>(
&self,
db: &'a Db,
parents: &[P],
children: Query<C>,
foreign_key: impl Fn(&C) -> P::Key + Send + 'a,
) -> impl Future<Output = Result<HashMap<P::Key, Vec<C>>>> + Send + 'a {
let children = children.where_eq(self.type_column, P::TABLE);
has_many(db, parents, children, self.id_column, foreign_key)
}
pub fn count_many<'a, C: Model, P: Model>(
&self,
db: &'a Db,
parents: &[P],
children: Query<C>,
) -> impl Future<Output = Result<HashMap<P::Key, i64>>> + Send + 'a {
let children = children.where_eq(self.type_column, P::TABLE);
count_many(db, parents, children, self.id_column)
}
pub fn parents<'a, P: Model, C>(
&self,
db: &'a Db,
children: &[C],
parent: impl Fn(&C) -> (String, P::Key),
) -> impl Future<Output = Result<HashMap<P::Key, P>>> + Send + 'a {
let mut ids: Vec<P::Key> = children
.iter()
.filter_map(|child| {
let (kind, id) = parent(child);
(kind == P::TABLE).then_some(id)
})
.collect();
ids.sort_unstable();
ids.dedup();
async move {
if ids.is_empty() {
return Ok(HashMap::new());
}
Ok(P::find_many(db, ids)
.await?
.into_iter()
.map(|parent| (parent.id(), parent))
.collect())
}
}
}