Skip to main content

aro_fletch/
query.rs

1//! Query and pagination adapters for [`FletchRepository`].
2//!
3//! Implements the [`Queryable`] and [`Paginatable`] ports from `aro-core`
4//! using fletch's dialect-aware [`QueryBuilder`].
5
6use aro_core::error::RepoError;
7use aro_core::pagination::{Page, PageRequest};
8use aro_core::repository::{BoxFuture, Paginatable, Queryable};
9use fletch_orm::Entity as FletchEntityTrait;
10use fletch_orm::column::Column;
11use fletch_orm::dialect::Dialect;
12use fletch_orm::filter::Filter;
13use fletch_orm::query_builder::{Order, QueryBuilder};
14
15use crate::error::from_fletch_err;
16use crate::repo::{FletchEntity, FletchRepository};
17
18/// Filter and sort criteria for entity queries.
19///
20/// Build with [`EntityQuery::new`], add [`Filter`] conditions via
21/// [`EntityQuery::filter`], and optional [`Order`] clauses via
22/// [`EntityQuery::order_by`].
23#[derive(Debug, Clone, Default)]
24pub struct EntityQuery {
25    filters: Vec<Filter>,
26    order_by: Vec<(String, Order)>,
27}
28
29impl EntityQuery {
30    /// Create an empty query (matches all rows).
31    pub fn new() -> Self {
32        Self::default()
33    }
34
35    /// Add a filter condition to the WHERE clause.
36    pub fn filter(mut self, filter: Filter) -> Self {
37        self.filters.push(filter);
38        self
39    }
40
41    /// Add an ORDER BY clause.
42    pub fn order_by(mut self, column: impl Into<String>, order: Order) -> Self {
43        self.order_by.push((column.into(), order));
44        self
45    }
46
47    fn with_default_order(mut self, column: impl Into<String>, order: Order) -> Self {
48        if self.order_by.is_empty() {
49            self.order_by.push((column.into(), order));
50        }
51        self
52    }
53}
54
55#[derive(Debug, sqlx::FromRow)]
56struct CountRow {
57    count: i64,
58}
59
60fn column_names<E: FletchEntityTrait>() -> Vec<&'static str> {
61    E::columns().iter().map(Column::name).collect()
62}
63
64fn build_select<E, D>(
65    dialect: &D,
66    query: &EntityQuery,
67    limit: Option<u64>,
68    offset: Option<u64>,
69) -> fletch_orm::BuiltQuery
70where
71    E: FletchEntityTrait,
72    D: Dialect,
73{
74    let mut builder = QueryBuilder::select(dialect, E::table_name()).columns(&column_names::<E>());
75    for filter in &query.filters {
76        builder = builder.filter(filter.clone());
77    }
78    for (column, order) in &query.order_by {
79        builder = builder.order_by(column, *order);
80    }
81    if let Some(limit) = limit {
82        builder = builder.limit(limit);
83    }
84    if let Some(offset) = offset {
85        builder = builder.offset(offset);
86    }
87    builder.build()
88}
89
90fn build_count<E, D>(dialect: &D, query: &EntityQuery) -> fletch_orm::BuiltQuery
91where
92    E: FletchEntityTrait,
93    D: Dialect,
94{
95    let mut builder = QueryBuilder::select(dialect, E::table_name());
96    for filter in &query.filters {
97        builder = builder.filter(filter.clone());
98    }
99    let inner = builder.build();
100    let sql = format!(
101        "SELECT COUNT(*) AS count FROM ({}) AS _count_sub",
102        inner.sql
103    );
104    fletch_orm::BuiltQuery {
105        sql,
106        values: inner.values,
107    }
108}
109
110macro_rules! impl_query {
111    ($db:ty, $dialect:expr) => {
112        impl<T> FletchRepository<T, $db>
113        where
114            T: FletchEntity<$db>,
115        {
116            async fn find_by_query(&self, filter: EntityQuery) -> Result<Vec<T>, RepoError> {
117                let built = build_select::<T, _>($dialect, &filter, None, None);
118                self.pool()
119                    .fetch_all::<T>(&built.sql, &built.values)
120                    .await
121                    .map_err(from_fletch_err)
122            }
123
124            async fn find_page_query(
125                &self,
126                query: EntityQuery,
127                request: PageRequest,
128            ) -> Result<Page<T>, RepoError> {
129                let entity_query = query.with_default_order(T::id_column(), Order::Asc);
130                let count_built = build_count::<T, _>($dialect, &entity_query);
131                let count_rows: Vec<CountRow> = self
132                    .pool()
133                    .fetch_all(&count_built.sql, &count_built.values)
134                    .await
135                    .map_err(from_fletch_err)?;
136                let total = u64::try_from(count_rows.first().map_or(0, |r| r.count)).unwrap_or(0);
137
138                let offset = (request.page.saturating_sub(1)) * request.per_page;
139                let select_built = build_select::<T, _>(
140                    $dialect,
141                    &entity_query,
142                    Some(request.per_page),
143                    Some(offset),
144                );
145                let items = self
146                    .pool()
147                    .fetch_all::<T>(&select_built.sql, &select_built.values)
148                    .await
149                    .map_err(from_fletch_err)?;
150
151                Ok(Page {
152                    items,
153                    total,
154                    page: request.page,
155                    per_page: request.per_page,
156                })
157            }
158        }
159
160        impl<T> Queryable<T> for FletchRepository<T, $db>
161        where
162            T: FletchEntity<$db>,
163        {
164            type Filter = EntityQuery;
165
166            fn find_by(&self, filter: EntityQuery) -> BoxFuture<'_, Result<Vec<T>, RepoError>> {
167                Box::pin(self.find_by_query(filter))
168            }
169
170            fn find_page_by(
171                &self,
172                filter: EntityQuery,
173                request: PageRequest,
174            ) -> BoxFuture<'_, Result<Page<T>, RepoError>> {
175                Box::pin(self.find_page_query(filter, request))
176            }
177        }
178
179        impl<T> Paginatable<T> for FletchRepository<T, $db>
180        where
181            T: FletchEntity<$db>,
182        {
183            fn find_page(&self, request: PageRequest) -> BoxFuture<'_, Result<Page<T>, RepoError>> {
184                Box::pin(self.find_page_query(EntityQuery::new(), request))
185            }
186        }
187    };
188}
189
190#[cfg(feature = "sqlite")]
191impl_query!(sqlx::Sqlite, &fletch_orm::SqliteDialect);
192
193#[cfg(feature = "postgres")]
194impl_query!(sqlx::Postgres, &fletch_orm::PostgresDialect);
195
196#[cfg(feature = "mysql")]
197impl_query!(sqlx::MySql, &fletch_orm::MySqlDialect);