use aro_core::error::RepoError;
use aro_core::pagination::{Page, PageRequest};
use aro_core::repository::{BoxFuture, Paginatable, Queryable};
use fletch_orm::Entity as FletchEntityTrait;
use fletch_orm::column::Column;
use fletch_orm::dialect::Dialect;
use fletch_orm::filter::Filter;
use fletch_orm::query_builder::{Order, QueryBuilder};
use crate::error::from_fletch_err;
use crate::repo::{FletchEntity, FletchRepository};
#[derive(Debug, Clone, Default)]
pub struct EntityQuery {
filters: Vec<Filter>,
order_by: Vec<(String, Order)>,
}
impl EntityQuery {
pub fn new() -> Self {
Self::default()
}
pub fn filter(mut self, filter: Filter) -> Self {
self.filters.push(filter);
self
}
pub fn order_by(mut self, column: impl Into<String>, order: Order) -> Self {
self.order_by.push((column.into(), order));
self
}
fn with_default_order(mut self, column: impl Into<String>, order: Order) -> Self {
if self.order_by.is_empty() {
self.order_by.push((column.into(), order));
}
self
}
}
#[derive(Debug, sqlx::FromRow)]
struct CountRow {
count: i64,
}
fn column_names<E: FletchEntityTrait>() -> Vec<&'static str> {
E::columns().iter().map(Column::name).collect()
}
fn build_select<E, D>(
dialect: &D,
query: &EntityQuery,
limit: Option<u64>,
offset: Option<u64>,
) -> fletch_orm::BuiltQuery
where
E: FletchEntityTrait,
D: Dialect,
{
let mut builder = QueryBuilder::select(dialect, E::table_name()).columns(&column_names::<E>());
for filter in &query.filters {
builder = builder.filter(filter.clone());
}
for (column, order) in &query.order_by {
builder = builder.order_by(column, *order);
}
if let Some(limit) = limit {
builder = builder.limit(limit);
}
if let Some(offset) = offset {
builder = builder.offset(offset);
}
builder.build()
}
fn build_count<E, D>(dialect: &D, query: &EntityQuery) -> fletch_orm::BuiltQuery
where
E: FletchEntityTrait,
D: Dialect,
{
let mut builder = QueryBuilder::select(dialect, E::table_name());
for filter in &query.filters {
builder = builder.filter(filter.clone());
}
let inner = builder.build();
let sql = format!(
"SELECT COUNT(*) AS count FROM ({}) AS _count_sub",
inner.sql
);
fletch_orm::BuiltQuery {
sql,
values: inner.values,
}
}
macro_rules! impl_query {
($db:ty, $dialect:expr) => {
impl<T> FletchRepository<T, $db>
where
T: FletchEntity<$db>,
{
async fn find_by_query(&self, filter: EntityQuery) -> Result<Vec<T>, RepoError> {
let built = build_select::<T, _>($dialect, &filter, None, None);
self.pool()
.fetch_all::<T>(&built.sql, &built.values)
.await
.map_err(from_fletch_err)
}
async fn find_page_query(
&self,
query: EntityQuery,
request: PageRequest,
) -> Result<Page<T>, RepoError> {
let entity_query = query.with_default_order(T::id_column(), Order::Asc);
let count_built = build_count::<T, _>($dialect, &entity_query);
let count_rows: Vec<CountRow> = self
.pool()
.fetch_all(&count_built.sql, &count_built.values)
.await
.map_err(from_fletch_err)?;
let total = u64::try_from(count_rows.first().map_or(0, |r| r.count)).unwrap_or(0);
let offset = (request.page.saturating_sub(1)) * request.per_page;
let select_built = build_select::<T, _>(
$dialect,
&entity_query,
Some(request.per_page),
Some(offset),
);
let items = self
.pool()
.fetch_all::<T>(&select_built.sql, &select_built.values)
.await
.map_err(from_fletch_err)?;
Ok(Page {
items,
total,
page: request.page,
per_page: request.per_page,
})
}
}
impl<T> Queryable<T> for FletchRepository<T, $db>
where
T: FletchEntity<$db>,
{
type Filter = EntityQuery;
fn find_by(&self, filter: EntityQuery) -> BoxFuture<'_, Result<Vec<T>, RepoError>> {
Box::pin(self.find_by_query(filter))
}
fn find_page_by(
&self,
filter: EntityQuery,
request: PageRequest,
) -> BoxFuture<'_, Result<Page<T>, RepoError>> {
Box::pin(self.find_page_query(filter, request))
}
}
impl<T> Paginatable<T> for FletchRepository<T, $db>
where
T: FletchEntity<$db>,
{
fn find_page(&self, request: PageRequest) -> BoxFuture<'_, Result<Page<T>, RepoError>> {
Box::pin(self.find_page_query(EntityQuery::new(), request))
}
}
};
}
#[cfg(feature = "sqlite")]
impl_query!(sqlx::Sqlite, &fletch_orm::SqliteDialect);
#[cfg(feature = "postgres")]
impl_query!(sqlx::Postgres, &fletch_orm::PostgresDialect);
#[cfg(feature = "mysql")]
impl_query!(sqlx::MySql, &fletch_orm::MySqlDialect);