use std::collections::HashMap;
use std::marker::PhantomData;
use std::pin::Pin;
use futures_util::{Stream, TryStreamExt};
use turso_orm_driver::{ConnectionTrait, Row, StreamTrait};
use turso_sql::{
Build, Condition, Expr, Func, IntoCondition, IntoIden, JoinType, Order, Statement, TableRef,
Value,
};
use crate::entity::relation::column_of;
use crate::entity::{
EntityTrait, FromQueryResult, IdenStatic, Iterable, Linked, ModelTrait, PartialModelTrait,
PrimaryKeyToColumn, Related, RelationDef,
};
use crate::{DbErr, Result};
pub type ModelStream<'a, M> = Pin<Box<dyn Stream<Item = Result<M>> + Send + 'a>>;
#[derive(Clone, Debug)]
pub struct Select<E: EntityTrait> {
query: turso_sql::Select,
_e: PhantomData<E>,
}
fn qualified<E: EntityTrait>(column: E::Column) -> Expr {
Expr::col((E::TABLE_NAME, column.as_str()))
}
fn related_condition<S: EntityTrait>(
rel: &RelationDef,
model: &S::Model,
target_ref: &str,
) -> Condition {
let mut cond = Condition::all();
for (f, t) in rel.from_col.iter().zip(&rel.to_col) {
let value = column_of::<S>(f).map_or(Value::Null, |c| model.get(c));
cond = cond.add(Expr::col((target_ref.to_owned(), *t)).eq(Expr::val(value)));
}
cond
}
impl<E: EntityTrait> Select<E> {
pub(crate) fn new() -> Self {
let mut query = turso_sql::Select::new().from(E::TABLE_NAME);
for c in E::Column::iter() {
query = query.expr(qualified::<E>(c));
}
Self {
query,
_e: PhantomData,
}
}
pub(crate) fn filter_by_pk(mut self, values: Vec<Value>) -> Self {
for (pk, value) in E::PrimaryKey::iter().zip(values) {
let column = pk.into_column();
self.query = self
.query
.and_where(qualified::<E>(column).eq(Expr::val(value)));
}
self
}
pub(crate) fn find_related_to<S>(model: &S::Model) -> Self
where
S: EntityTrait + Related<E>,
{
let to = <S as Related<E>>::to();
let mut select = Self::new();
match <S as Related<E>>::via() {
Some(via) => {
let junction = select.join_table(JoinType::Inner, to.from_tbl, |r| {
to.join_condition_refs(r, E::TABLE_NAME)
});
select.query = select
.query
.and_where(related_condition::<S>(&via, model, &junction));
}
None => {
select.query =
select
.query
.and_where(related_condition::<S>(&to, model, E::TABLE_NAME));
}
}
select
}
pub(crate) fn find_linked_to<L>(link: &L, model: &<L::FromEntity as EntityTrait>::Model) -> Self
where
L: Linked<ToEntity = E>,
{
let mut select = Self::new();
let mut previous = E::TABLE_NAME.to_owned();
for hop in link.link().into_iter().rev() {
previous = select.join_table(JoinType::Inner, hop.from_tbl, |r| {
hop.join_condition_refs(r, &previous)
});
}
for pk in <L::FromEntity as EntityTrait>::PrimaryKey::iter() {
let column = pk.into_column();
select.query = select.query.and_where(
Expr::col((previous.clone(), column.as_str())).eq(Expr::val(model.get(column))),
);
}
select
}
pub(crate) fn join_table(
&mut self,
kind: JoinType,
table: &'static str,
on: impl FnOnce(&str) -> Expr,
) -> String {
let occurrences = self
.query
.from_tables()
.iter()
.chain(self.query.joins().iter().map(|j| &j.table))
.filter(|t| t.name.name() == table)
.count();
let (table_ref, reference) = if occurrences == 0 {
(TableRef::new(table), table.to_owned())
} else {
let alias = format!("{table}_{occurrences}");
(TableRef::new(table).alias(alias.clone()), alias)
};
let cond = on(&reference);
self.query = std::mem::take(&mut self.query).join(kind, table_ref, cond);
reference
}
fn join_related<R: EntityTrait>(&mut self, kind: JoinType) -> String
where
E: Related<R>,
{
let to = <E as Related<R>>::to();
match <E as Related<R>>::via() {
Some(via) => {
let junction = self.join_table(kind, via.to_tbl, |r| {
via.join_condition_refs(via.from_tbl, r)
});
self.join_table(kind, to.to_tbl, |r| to.join_condition_refs(&junction, r))
}
None => self.join_table(kind, to.to_tbl, |r| to.join_condition_refs(to.from_tbl, r)),
}
}
fn join_linked<L: Linked<FromEntity = E>>(&mut self, kind: JoinType, link: &L) -> String {
let mut previous = E::TABLE_NAME.to_owned();
for hop in link.link() {
previous = self.join_table(kind, hop.to_tbl, |r| hop.join_condition_refs(&previous, r));
}
previous
}
fn into_select_two<R: EntityTrait>(self, right_ref: &str) -> SelectTwo<E, R> {
let mut query = self.query.clear_items();
for c in E::Column::iter() {
query = query.expr_as(qualified::<E>(c), format!("A_{}", c.as_str()));
}
for c in R::Column::iter() {
query = query.expr_as(
Expr::col((right_ref.to_owned(), c.as_str())),
format!("B_{}", c.as_str()),
);
}
SelectTwo {
query,
_e: PhantomData,
}
}
#[must_use]
pub fn filter(mut self, cond: impl IntoCondition) -> Self {
self.query = self.query.and_where(cond);
self
}
#[must_use]
pub fn filter_option(self, cond: Option<impl IntoCondition>) -> Self {
match cond {
Some(c) => self.filter(c),
None => self,
}
}
#[must_use]
pub fn related_to<M: ModelTrait>(mut self, rel: &RelationDef, model: &M) -> Self {
self.query =
self.query
.and_where(related_condition::<M::Entity>(rel, model, E::TABLE_NAME));
self
}
#[must_use]
pub fn order_by(mut self, column: E::Column, order: Order) -> Self {
self.query = self.query.order_by_expr(qualified::<E>(column), order);
self
}
#[must_use]
pub fn order_by_asc(self, column: E::Column) -> Self {
self.order_by(column, Order::Asc)
}
#[must_use]
pub fn order_by_desc(self, column: E::Column) -> Self {
self.order_by(column, Order::Desc)
}
#[must_use]
pub fn order_by_expr(mut self, expr: Expr, order: Order) -> Self {
self.query = self.query.order_by_expr(expr, order);
self
}
#[must_use]
pub fn limit(mut self, limit: u64) -> Self {
self.query = self.query.limit(limit);
self
}
#[must_use]
pub fn offset(mut self, offset: u64) -> Self {
self.query = self.query.offset(offset);
self
}
#[must_use]
pub fn distinct(mut self) -> Self {
self.query = self.query.distinct();
self
}
#[must_use]
pub fn group_by(mut self, column: E::Column) -> Self {
self.query = self.query.group_by(qualified::<E>(column));
self
}
#[must_use]
pub fn group_by_expr(mut self, expr: Expr) -> Self {
self.query = self.query.group_by(expr);
self
}
#[must_use]
pub fn having(mut self, cond: impl IntoCondition) -> Self {
self.query = self.query.and_having(cond);
self
}
#[must_use]
pub fn select_only(mut self) -> Self {
self.query = self.query.clear_items();
self
}
#[must_use]
pub fn column(mut self, column: E::Column) -> Self {
self.query = self.query.expr(qualified::<E>(column));
self
}
#[must_use]
pub fn column_as(mut self, column: E::Column, alias: impl IntoIden) -> Self {
self.query = self.query.expr_as(qualified::<E>(column), alias);
self
}
#[must_use]
pub fn expr(mut self, expr: Expr) -> Self {
self.query = self.query.expr(expr);
self
}
#[must_use]
pub fn expr_as(mut self, expr: Expr, alias: impl IntoIden) -> Self {
self.query = self.query.expr_as(expr, alias);
self
}
#[must_use]
pub fn join(mut self, kind: JoinType, rel: &RelationDef) -> Self {
self.join_table(kind, rel.to_tbl, |r| {
rel.join_condition_refs(rel.from_tbl, r)
});
self
}
#[must_use]
pub fn join_as(mut self, kind: JoinType, rel: &RelationDef, alias: &'static str) -> Self {
let on = rel.join_condition_refs(rel.from_tbl, alias);
self.query = self
.query
.join(kind, TableRef::new(rel.to_tbl).alias(alias), on);
self
}
#[must_use]
pub fn join_rev(mut self, kind: JoinType, rel: &RelationDef) -> Self {
self.join_table(kind, rel.from_tbl, |r| {
rel.join_condition_refs(r, rel.to_tbl)
});
self
}
#[must_use]
pub fn inner_join<R: EntityTrait>(mut self, _: R) -> Self
where
E: Related<R>,
{
self.join_related::<R>(JoinType::Inner);
self
}
#[must_use]
pub fn left_join<R: EntityTrait>(mut self, _: R) -> Self
where
E: Related<R>,
{
self.join_related::<R>(JoinType::Left);
self
}
#[must_use]
pub fn find_also_related<R: EntityTrait>(mut self, _: R) -> SelectTwo<E, R>
where
E: Related<R>,
{
let right = self.join_related::<R>(JoinType::Left);
self.into_select_two::<R>(&right)
}
#[must_use]
pub fn find_with_related<R: EntityTrait>(self, r: R) -> SelectTwoMany<E, R>
where
E: Related<R>,
{
SelectTwoMany {
inner: self.find_also_related(r),
}
}
#[must_use]
pub fn find_also_linked<L>(mut self, link: &L) -> SelectTwo<E, L::ToEntity>
where
L: Linked<FromEntity = E>,
{
let right = self.join_linked(JoinType::Left, link);
self.into_select_two::<L::ToEntity>(&right)
}
#[must_use]
pub fn find_with_linked<L>(self, link: &L) -> SelectTwoMany<E, L::ToEntity>
where
L: Linked<FromEntity = E>,
{
SelectTwoMany {
inner: self.find_also_linked(link),
}
}
pub fn into_model<M: FromQueryResult>(self) -> Selector<M> {
Selector {
query: self.query,
_m: PhantomData,
}
}
pub fn into_partial_model<P: PartialModelTrait>(self) -> Selector<P> {
Selector {
query: P::select_cols(self.query.clear_items()),
_m: PhantomData,
}
}
pub fn into_tuple<T: FromQueryResult>(self) -> Selector<T> {
self.into_model::<T>()
}
#[cfg(feature = "with-json")]
#[cfg_attr(docsrs, doc(cfg(feature = "with-json")))]
pub fn into_json(self) -> Selector<serde_json::Value> {
self.into_model::<serde_json::Value>()
}
pub fn from_raw_sql(self, statement: Statement) -> RawSelector<E::Model> {
Selector::<E::Model>::from_statement(statement)
}
pub fn as_query(&self) -> &turso_sql::Select {
&self.query
}
pub fn query_mut(&mut self) -> &mut turso_sql::Select {
&mut self.query
}
pub fn into_query(self) -> turso_sql::Select {
self.query
}
pub fn build(&self) -> Statement {
self.query.to_statement()
}
pub async fn one<C: ConnectionTrait>(self, db: &C) -> Result<Option<E::Model>> {
self.into_model::<E::Model>().one(db).await
}
pub async fn all<C: ConnectionTrait>(self, db: &C) -> Result<Vec<E::Model>> {
self.into_model::<E::Model>().all(db).await
}
pub async fn stream<C: StreamTrait>(self, db: &C) -> Result<ModelStream<'_, E::Model>> {
self.into_model::<E::Model>().stream(db).await
}
pub async fn count<C: ConnectionTrait>(self, db: &C) -> Result<u64> {
self.into_model::<E::Model>().count(db).await
}
pub async fn exists<C: ConnectionTrait>(self, db: &C) -> Result<bool> {
self.into_model::<E::Model>().exists(db).await
}
pub fn paginate<C: ConnectionTrait>(
self,
db: &C,
page_size: u64,
) -> Paginator<'_, C, E::Model> {
self.into_model::<E::Model>().paginate(db, page_size)
}
pub fn cursor_by<const N: usize>(
self,
columns: [E::Column; N],
) -> crate::query::Cursor<E, E::Model> {
crate::query::Cursor::new(self.query, columns.to_vec())
}
}
#[derive(Clone, Debug)]
pub struct Selector<M> {
query: turso_sql::Select,
_m: PhantomData<M>,
}
impl<M: FromQueryResult> Selector<M> {
pub fn from_query(query: turso_sql::Select) -> Self {
Self {
query,
_m: PhantomData,
}
}
pub fn from_statement(statement: Statement) -> RawSelector<M> {
RawSelector {
statement,
_m: PhantomData,
}
}
pub fn build(&self) -> Statement {
self.query.to_statement()
}
pub async fn one<C: ConnectionTrait>(mut self, db: &C) -> Result<Option<M>> {
self.query = self.query.limit(1);
let row = db.query_one(self.build()).await?;
row.map(|r| M::from_query_result(&r, "")).transpose()
}
pub async fn all<C: ConnectionTrait>(self, db: &C) -> Result<Vec<M>> {
let rows = db.query_all(self.build()).await?;
rows.iter().map(|r| M::from_query_result(r, "")).collect()
}
#[allow(
clippy::needless_lifetimes,
reason = "the stream borrows `db`, not `self`, and the explicit lifetime says so"
)]
pub async fn stream<'a, C: StreamTrait>(self, db: &'a C) -> Result<ModelStream<'a, M>>
where
M: 'a,
{
let stream = db.stream(self.build()).await?;
Ok(Box::pin(stream.map_err(DbErr::from).and_then(
|row| async move { M::from_query_result(&row, "") },
)))
}
fn unordered(&self) -> turso_sql::Select {
self.query
.clone()
.clear_order_by()
.reset_limit()
.reset_offset()
}
pub async fn count<C: ConnectionTrait>(self, db: &C) -> Result<u64> {
let outer = turso_sql::Select::new()
.expr_as(Func::count_star(), "num_items")
.from_subquery(self.unordered(), "sub");
let row = db
.query_one(outer.to_statement())
.await?
.ok_or(DbErr::RecordNotFound("count".into()))?;
Ok(row.get::<u64>("num_items")?)
}
pub async fn exists<C: ConnectionTrait>(self, db: &C) -> Result<bool> {
let probe = turso_sql::Select::new().expr_as(Expr::exists(self.unordered()), "found");
let row = db
.query_one(probe.to_statement())
.await?
.ok_or(DbErr::RecordNotFound("exists".into()))?;
Ok(row.get::<bool>("found")?)
}
pub fn paginate<C: ConnectionTrait>(self, db: &C, page_size: u64) -> Paginator<'_, C, M> {
Paginator {
query: self.query,
page_size: page_size.max(1),
db,
_m: PhantomData,
}
}
}
#[derive(Clone, Debug)]
pub struct RawSelector<M> {
statement: Statement,
_m: PhantomData<M>,
}
impl<M: FromQueryResult> RawSelector<M> {
pub async fn one<C: ConnectionTrait>(self, db: &C) -> Result<Option<M>> {
let row = db.query_one(self.statement).await?;
row.map(|r| M::from_query_result(&r, "")).transpose()
}
pub async fn all<C: ConnectionTrait>(self, db: &C) -> Result<Vec<M>> {
let rows = db.query_all(self.statement).await?;
rows.iter().map(|r| M::from_query_result(r, "")).collect()
}
}
#[derive(Debug)]
pub struct Paginator<'db, C, M> {
query: turso_sql::Select,
page_size: u64,
db: &'db C,
_m: PhantomData<M>,
}
impl<C: ConnectionTrait, M: FromQueryResult> Paginator<'_, C, M> {
pub async fn fetch_page(&self, page: u64) -> Result<Vec<M>> {
let query = self
.query
.clone()
.limit(self.page_size)
.offset(page.saturating_mul(self.page_size));
Selector::<M>::from_query(query).all(self.db).await
}
pub async fn num_items(&self) -> Result<u64> {
Selector::<M>::from_query(self.query.clone())
.count(self.db)
.await
}
pub async fn num_pages(&self) -> Result<u64> {
let items = self.num_items().await?;
Ok(items.div_ceil(self.page_size))
}
pub async fn num_items_and_pages(&self) -> Result<(u64, u64)> {
let items = self.num_items().await?;
Ok((items, items.div_ceil(self.page_size)))
}
pub async fn for_each_page(&self, mut f: impl FnMut(Vec<M>) -> bool) -> Result<()> {
let mut page = 0;
loop {
let rows = self.fetch_page(page).await?;
let full = u64::try_from(rows.len()).unwrap_or(u64::MAX) == self.page_size;
if !f(rows) || !full {
return Ok(());
}
page += 1;
}
}
}
#[derive(Clone, Debug)]
pub struct SelectTwo<E: EntityTrait, R: EntityTrait> {
query: turso_sql::Select,
_e: PhantomData<(E, R)>,
}
impl<E: EntityTrait, R: EntityTrait> SelectTwo<E, R> {
#[must_use]
pub fn filter(mut self, cond: impl IntoCondition) -> Self {
self.query = self.query.and_where(cond);
self
}
#[must_use]
pub fn order_by(mut self, column: E::Column, order: Order) -> Self {
self.query = self.query.order_by_expr(qualified::<E>(column), order);
self
}
#[must_use]
pub fn order_by_related(mut self, column: R::Column, order: Order) -> Self {
self.query = self.query.order_by_expr(qualified::<R>(column), order);
self
}
#[must_use]
pub fn order_by_expr(mut self, expr: Expr, order: Order) -> Self {
self.query = self.query.order_by_expr(expr, order);
self
}
#[must_use]
pub fn limit(mut self, limit: u64) -> Self {
self.query = self.query.limit(limit);
self
}
#[must_use]
pub fn offset(mut self, offset: u64) -> Self {
self.query = self.query.offset(offset);
self
}
pub fn query_mut(&mut self) -> &mut turso_sql::Select {
&mut self.query
}
pub fn build(&self) -> Statement {
self.query.to_statement()
}
fn decode(row: &Row) -> Result<(E::Model, Option<R::Model>)> {
let a = E::Model::from_query_result(row, "A_")?;
let b = R::Model::from_query_result_optional(row, "B_")?;
Ok((a, b))
}
pub async fn one<C: ConnectionTrait>(
mut self,
db: &C,
) -> Result<Option<(E::Model, Option<R::Model>)>> {
self.query = self.query.limit(1);
db.query_one(self.build())
.await?
.as_ref()
.map(Self::decode)
.transpose()
}
pub async fn all<C: ConnectionTrait>(
self,
db: &C,
) -> Result<Vec<(E::Model, Option<R::Model>)>> {
db.query_all(self.build())
.await?
.iter()
.map(Self::decode)
.collect()
}
}
#[derive(Clone, Debug)]
pub struct SelectTwoMany<E: EntityTrait, R: EntityTrait> {
inner: SelectTwo<E, R>,
}
impl<E: EntityTrait, R: EntityTrait> SelectTwoMany<E, R> {
#[must_use]
pub fn filter(mut self, cond: impl IntoCondition) -> Self {
self.inner = self.inner.filter(cond);
self
}
#[must_use]
pub fn order_by(mut self, column: E::Column, order: Order) -> Self {
self.inner = self.inner.order_by(column, order);
self
}
#[must_use]
pub fn order_by_related(mut self, column: R::Column, order: Order) -> Self {
self.inner = self.inner.order_by_related(column, order);
self
}
#[must_use]
pub fn order_by_expr(mut self, expr: Expr, order: Order) -> Self {
self.inner = self.inner.order_by_expr(expr, order);
self
}
pub fn build(&self) -> Statement {
self.inner.build()
}
pub async fn all<C: ConnectionTrait>(self, db: &C) -> Result<Vec<(E::Model, Vec<R::Model>)>> {
let pairs = self.inner.all(db).await?;
let mut groups: Vec<(E::Model, Vec<R::Model>)> = Vec::new();
let mut index: HashMap<String, usize> = HashMap::new();
for (left, right) in pairs {
let key = E::PrimaryKey::iter()
.map(|pk| left.get(pk.into_column()).to_literal())
.collect::<Vec<_>>()
.join("\u{1f}");
let at = if let Some(&at) = index.get(&key) {
at
} else {
groups.push((left, Vec::new()));
index.insert(key, groups.len() - 1);
groups.len() - 1
};
if let Some(r) = right {
groups[at].1.push(r);
}
}
Ok(groups)
}
}
pub(crate) fn pk_condition<E: EntityTrait>(values: Vec<Value>) -> Condition {
let mut cond = Condition::all();
for (pk, value) in E::PrimaryKey::iter().zip(values) {
cond = cond.add(qualified::<E>(pk.into_column()).eq(Expr::val(value)));
}
cond
}