Skip to main content

renox_core/db/
query.rs

1use std::marker::PhantomData;
2
3use super::paginate::{CursorPage, SimplePage};
4use super::{Db, DbValue, Dialect, Executor, FromDb, Model, Paginated, ToDbValue, now, quote, sql};
5use crate::Result;
6use anyhow::anyhow;
7
8const OPERATORS: &[&str] = &["=", "!=", "<>", "<", "<=", ">", ">=", "like", "not like"];
9
10#[derive(Clone, Copy, PartialEq)]
11enum Trashed {
12    Without,
13    With,
14    Only,
15}
16
17/// `where_in` lists longer than this are sent as one JSON array.
18const LARGE_IN: usize = 1000;
19
20/// A number type `Query::sum` can return.
21pub trait Number: FromDb + sealed::Sealed {
22    #[doc(hidden)]
23    const SQL_TYPE: &'static str;
24}
25
26impl Number for i64 {
27    const SQL_TYPE: &'static str = "BIGINT";
28}
29
30impl Number for f64 {
31    const SQL_TYPE: &'static str = "DOUBLE PRECISION";
32}
33
34mod sealed {
35    pub trait Sealed {}
36    impl Sealed for i64 {}
37    impl Sealed for f64 {}
38}
39
40/// One condition of a query, rendered per database.
41#[derive(Clone)]
42enum Filter {
43    /// SQL with `?` placeholders.
44    Sql(String),
45    /// `column LIKE ?` (ILIKE on PostgreSQL); `not` for NOT LIKE.
46    Like { column: String, not: bool },
47    /// `column IN (…)` for a long list sent as one JSON array.
48    JsonIn { kind: &'static str, column: String },
49    /// Conditions joined with OR (`any`) or AND, in parentheses.
50    Group { any: bool, filters: Vec<Filter> },
51    /// `NOT (…)`.
52    Not(Box<Filter>),
53    /// `[NOT] EXISTS (SELECT 1 FROM table WHERE correlation AND …)`.
54    Exists {
55        not: bool,
56        table: &'static str,
57        correlation: String,
58        filters: Vec<Filter>,
59    },
60    /// `column IN (SELECT sub_column FROM table WHERE …)`.
61    InQuery {
62        column: String,
63        table: &'static str,
64        sub_column: String,
65        filters: Vec<Filter>,
66    },
67    /// A full-text search on a model's index, already rendered per
68    /// database (`search::filter_sql`), with one `?`.
69    Search { sqlite: String, postgres: String },
70}
71
72/// One `ORDER BY` term.
73#[derive(Clone)]
74enum Order {
75    /// SQL without values.
76    Sql(String),
77    /// Full-text relevance (`search::rank_sql`) per database, with one `?`
78    /// in `order_binds`.
79    Relevance { sqlite: String, postgres: String },
80}
81
82impl Order {
83    fn render(&self, dialect: Dialect) -> &str {
84        match (self, dialect) {
85            (Order::Sql(sql), _) => sql,
86            (Order::Relevance { sqlite, .. }, Dialect::Sqlite) => sqlite,
87            (Order::Relevance { postgres, .. }, Dialect::Postgres) => postgres,
88        }
89    }
90}
91
92impl Filter {
93    fn render(&self, dialect: Dialect) -> String {
94        match self {
95            Filter::Sql(sql) => sql.clone(),
96            Filter::Like { column, not } => {
97                let op = match (dialect, not) {
98                    (Dialect::Postgres, false) => "ILIKE",
99                    (Dialect::Postgres, true) => "NOT ILIKE",
100                    (_, false) => "LIKE",
101                    (_, true) => "NOT LIKE",
102                };
103                format!("{column} {op} ?")
104            }
105            Filter::JsonIn { kind, column } => json_in(kind, column, dialect),
106            Filter::Group { any, filters } => {
107                if filters.is_empty() {
108                    // Nothing to match for `any`, nothing to restrict for AND.
109                    return if *any { "0 = 1".into() } else { "1 = 1".into() };
110                }
111                let joiner = if *any { " OR " } else { " AND " };
112                let parts: Vec<String> = filters.iter().map(|f| f.render(dialect)).collect();
113                format!("({})", parts.join(joiner))
114            }
115            Filter::Not(filter) => format!("NOT ({})", filter.render(dialect)),
116            Filter::Exists {
117                not,
118                table,
119                correlation,
120                filters,
121            } => {
122                let mut parts = vec![correlation.clone()];
123                parts.extend(filters.iter().map(|f| f.render(dialect)));
124                format!(
125                    "{}EXISTS (SELECT 1 FROM {} WHERE {})",
126                    if *not { "NOT " } else { "" },
127                    quote(table),
128                    parts.join(" AND ")
129                )
130            }
131            Filter::InQuery {
132                column,
133                table,
134                sub_column,
135                filters,
136            } => {
137                let condition = if filters.is_empty() {
138                    String::new()
139                } else {
140                    let parts: Vec<String> = filters.iter().map(|f| f.render(dialect)).collect();
141                    format!(" WHERE {}", parts.join(" AND "))
142                };
143                format!(
144                    "{column} IN (SELECT {sub_column} FROM {}{condition})",
145                    quote(table)
146                )
147            }
148            Filter::Search { sqlite, postgres } => match dialect {
149                Dialect::Sqlite => sqlite.clone(),
150                Dialect::Postgres => postgres.clone(),
151            },
152        }
153    }
154}
155
156/// A long list of integers or of strings, as a JSON array.
157fn large_list(values: &[DbValue]) -> Option<(&'static str, serde_json::Value)> {
158    if values.len() <= LARGE_IN {
159        return None;
160    }
161    if let Some(ints) = values
162        .iter()
163        .map(|v| match v {
164            DbValue::Integer(n) => Some(serde_json::Value::from(*n)),
165            _ => None,
166        })
167        .collect::<Option<Vec<_>>>()
168    {
169        return Some(("int", ints.into()));
170    }
171    values
172        .iter()
173        .map(|v| match v {
174            DbValue::Text(s) => Some(serde_json::Value::from(s.clone())),
175            _ => None,
176        })
177        .collect::<Option<Vec<_>>>()
178        .map(|texts| ("text", texts.into()))
179}
180
181fn json_in(kind: &str, column: &str, dialect: Dialect) -> String {
182    match (dialect, kind) {
183        (Dialect::Sqlite, _) => format!("{column} IN (SELECT value FROM json_each(?))"),
184        (Dialect::Postgres, "int") => format!(
185            "{column} IN (SELECT CAST(x AS BIGINT) FROM jsonb_array_elements_text(CAST(? AS JSONB)) AS t(x))"
186        ),
187        (Dialect::Postgres, _) => format!(
188            "{column} IN (SELECT x FROM jsonb_array_elements_text(CAST(? AS JSONB)) AS t(x))"
189        ),
190    }
191}
192
193/// A query on a model's table, built with chained filters.
194///
195/// ```
196/// # #[derive(Model, serde::Serialize, Default)]
197/// # #[model(table = "products")]
198/// # struct Product { id: i64, name: String, price: i64, category: Option<String>, user_id: i64 }
199/// # use renox::prelude::*;
200/// # async fn demo(db: Db, page: u32) -> Result {
201/// let products = Product::query()
202///     .where_eq("category", "coffee")
203///     .where_op("price", "<", 25_000)
204///     .order_by("name")
205///     .paginate(&db, page, 20)
206///     .await?;
207/// # let _ = products; Ok(()) }
208/// ```
209///
210/// Column names are checked against the model; an unknown column or operator
211/// makes the query return an error instead of running.
212pub struct Query<M> {
213    filters: Vec<Filter>,
214    binds: Vec<DbValue>,
215    group: Vec<String>,
216    having: Vec<String>,
217    having_binds: Vec<DbValue>,
218    order: Vec<Order>,
219    /// Values of the `ORDER BY` terms, bound after the WHERE and HAVING ones.
220    order_binds: Vec<DbValue>,
221    limit: Option<u64>,
222    offset: Option<u64>,
223    lock: Option<&'static str>,
224    trashed: Trashed,
225    error: Option<String>,
226    model: PhantomData<fn() -> M>,
227}
228
229impl<M> Clone for Query<M> {
230    fn clone(&self) -> Self {
231        Self {
232            filters: self.filters.clone(),
233            binds: self.binds.clone(),
234            group: self.group.clone(),
235            having: self.having.clone(),
236            having_binds: self.having_binds.clone(),
237            order: self.order.clone(),
238            order_binds: self.order_binds.clone(),
239            limit: self.limit,
240            offset: self.offset,
241            lock: self.lock,
242            trashed: self.trashed,
243            error: self.error.clone(),
244            model: PhantomData,
245        }
246    }
247}
248
249impl<M: Model> Query<M> {
250    pub(crate) fn new() -> Self {
251        Self {
252            filters: Vec::new(),
253            binds: Vec::new(),
254            group: Vec::new(),
255            having: Vec::new(),
256            having_binds: Vec::new(),
257            order: Vec::new(),
258            order_binds: Vec::new(),
259            limit: None,
260            offset: None,
261            lock: None,
262            trashed: Trashed::Without,
263            error: None,
264            model: PhantomData,
265        }
266    }
267
268    fn column(&mut self, column: &str) -> Option<String> {
269        let plain = !column.is_empty()
270            && column
271                .chars()
272                .all(|c| c.is_ascii_alphanumeric() || c == '_');
273        if M::COLUMNS.contains(&column) || (M::SELECT_ALL && plain) {
274            Some(quote(column))
275        } else {
276            self.error
277                .get_or_insert_with(|| format!("`{}` has no column `{column}`", M::TABLE));
278            None
279        }
280    }
281
282    /// Filters on `column = value`.
283    pub fn where_eq(self, column: &str, value: impl ToDbValue) -> Self {
284        self.where_op(column, "=", value)
285    }
286
287    /// Filters with a comparison: `=`, `!=`, `<>`, `<`, `<=`, `>`, `>=`, `like`, `not like`.
288    /// `like` ignores ASCII case on both databases (`ILIKE` on PostgreSQL, as
289    /// SQLite's `LIKE` already does).
290    pub fn where_op(mut self, column: &str, op: &str, value: impl ToDbValue) -> Self {
291        let op = op.to_ascii_lowercase();
292        if !OPERATORS.contains(&op.as_str()) {
293            self.error
294                .get_or_insert_with(|| format!("unsupported operator `{op}`"));
295            return self;
296        }
297        if let Some(column) = self.column(column) {
298            self.filters.push(match op.as_str() {
299                "like" => Filter::Like { column, not: false },
300                "not like" => Filter::Like { column, not: true },
301                _ => Filter::Sql(format!("{column} {} ?", op.to_uppercase())),
302            });
303            self.binds.push(value.to_db_value());
304        }
305        self
306    }
307
308    /// Filters on `column LIKE pattern`, ignoring ASCII case (`%` and `_` are wildcards).
309    pub fn where_like(self, column: &str, pattern: impl ToDbValue) -> Self {
310        self.where_op(column, "like", pattern)
311    }
312
313    /// Filters on `column IS NULL`.
314    pub fn where_null(mut self, column: &str) -> Self {
315        if let Some(column) = self.column(column) {
316            self.filters.push(Filter::Sql(format!("{column} IS NULL")));
317        }
318        self
319    }
320
321    /// Filters on `column IS NOT NULL`.
322    pub fn where_not_null(mut self, column: &str) -> Self {
323        if let Some(column) = self.column(column) {
324            self.filters
325                .push(Filter::Sql(format!("{column} IS NOT NULL")));
326        }
327        self
328    }
329
330    /// Filters on `column IN (…)`; an empty list matches no rows.
331    pub fn where_in<V: ToDbValue>(self, column: &str, values: impl IntoIterator<Item = V>) -> Self {
332        self.in_list(column, values, false)
333    }
334
335    /// Filters on `column NOT IN (…)`; an empty list matches every row.
336    pub fn where_not_in<V: ToDbValue>(
337        self,
338        column: &str,
339        values: impl IntoIterator<Item = V>,
340    ) -> Self {
341        self.in_list(column, values, true)
342    }
343
344    fn in_list<V: ToDbValue>(
345        mut self,
346        column: &str,
347        values: impl IntoIterator<Item = V>,
348        not: bool,
349    ) -> Self {
350        let values: Vec<DbValue> = values.into_iter().map(|v| v.to_db_value()).collect();
351        if let Some(column) = self.column(column) {
352            let filter = if values.is_empty() {
353                // Nothing is in an empty list.
354                Filter::Sql("0 = 1".into())
355            } else if let Some((kind, list)) = large_list(&values) {
356                // Over the databases' bind limits: one JSON array instead.
357                self.binds.push(DbValue::Json(list));
358                Filter::JsonIn { kind, column }
359            } else {
360                let marks = vec!["?"; values.len()].join(", ");
361                self.binds.extend(values);
362                Filter::Sql(format!("{column} IN ({marks})"))
363            };
364            self.filters.push(if not {
365                Filter::Not(Box::new(filter))
366            } else {
367                filter
368            });
369        }
370        self
371    }
372
373    /// `low <= column <= high`.
374    pub fn where_between(
375        mut self,
376        column: &str,
377        low: impl ToDbValue,
378        high: impl ToDbValue,
379    ) -> Self {
380        if let Some(column) = self.column(column) {
381            self.filters
382                .push(Filter::Sql(format!("{column} BETWEEN ? AND ?")));
383            self.binds.push(low.to_db_value());
384            self.binds.push(high.to_db_value());
385        }
386        self
387    }
388
389    /// Rows whose `column` is among `sub_column`'s values in another model's
390    /// query, e.g. products in active categories:
391    ///
392    /// ```
393    /// # use renox::prelude::*;
394    /// # #[derive(Model, serde::Serialize, Default)] struct Product { id: i64, category_id: i64 }
395    /// # #[derive(Model, serde::Serialize, Default)] struct Category { id: i64, active: bool }
396    /// # async fn demo(db: Db) -> Result {
397    /// let products = Product::query()
398    ///     .where_in_query("category_id", Category::where_eq("active", true), "id")
399    ///     .get(&db)
400    ///     .await?;
401    /// # let _ = products; Ok(()) }
402    /// ```
403    pub fn where_in_query<N: Model>(
404        mut self,
405        column: &str,
406        mut sub: Query<N>,
407        sub_column: &str,
408    ) -> Self {
409        let sub_column = sub.column(sub_column);
410        let column = self.column(column);
411        if let Some(error) = sub.error.take() {
412            self.error.get_or_insert(error);
413        }
414        if let (Some(column), Some(sub_column)) = (column, sub_column) {
415            let trashed = sub.trashed_filter();
416            let mut filters = sub.filters;
417            filters.extend(trashed);
418            self.filters.push(Filter::InQuery {
419                column,
420                table: N::TABLE,
421                sub_column,
422                filters,
423            });
424            self.binds.extend(sub.binds);
425        }
426        self
427    }
428
429    /// Like `where_in_query`, keeping the rows whose `column` is *not* among
430    /// the sub-query's values.
431    pub fn where_not_in_query<N: Model>(
432        self,
433        column: &str,
434        sub: Query<N>,
435        sub_column: &str,
436    ) -> Self {
437        let before = self.filters.len();
438        let mut query = self.where_in_query(column, sub, sub_column);
439        if query.filters.len() > before
440            && let Some(last) = query.filters.pop()
441        {
442            query.filters.push(Filter::Not(Box::new(last)));
443        }
444        query
445    }
446
447    /// Rows that have at least one related row in `children`, joined by the
448    /// children's `foreign_key` column to this model's `id` (`EXISTS`):
449    /// products with a 5-star review.
450    ///
451    /// ```
452    /// # use renox::prelude::*;
453    /// # #[derive(Model, serde::Serialize, Default)] struct Product { id: i64, name: String }
454    /// # #[derive(Model, serde::Serialize, Default)] struct Review { id: i64, product_id: i64, stars: i64 }
455    /// # async fn demo(db: Db) -> Result {
456    /// let loved = Product::query()
457    ///     .where_has(Review::where_eq("stars", 5), "product_id")
458    ///     .get(&db)
459    ///     .await?;
460    /// let unreviewed = Product::query().where_doesnt_have(Review::query(), "product_id").count(&db).await?;
461    /// # let _ = (loved, unreviewed); Ok(()) }
462    /// ```
463    pub fn where_has<N: Model>(self, children: Query<N>, foreign_key: &str) -> Self {
464        self.related(children, foreign_key, false)
465    }
466
467    /// Rows without any related row in `children` (`NOT EXISTS`).
468    pub fn where_doesnt_have<N: Model>(self, children: Query<N>, foreign_key: &str) -> Self {
469        self.related(children, foreign_key, true)
470    }
471
472    fn related<N: Model>(mut self, mut children: Query<N>, foreign_key: &str, not: bool) -> Self {
473        let foreign_key = children.column(foreign_key);
474        if let Some(error) = children.error.take() {
475            self.error.get_or_insert(error);
476        }
477        if let Some(foreign_key) = foreign_key {
478            let trashed = children.trashed_filter();
479            let mut filters = children.filters;
480            filters.extend(trashed);
481            self.filters.push(Filter::Exists {
482                not,
483                table: N::TABLE,
484                correlation: format!(
485                    "{}.{foreign_key} = {}.\"id\"",
486                    quote(N::TABLE),
487                    quote(M::TABLE)
488                ),
489                filters,
490            });
491            self.binds.extend(children.binds);
492        }
493        self
494    }
495
496    /// Any of the conditions `group` adds must hold (`OR`), in parentheses:
497    /// `.where_any(|q| q.where_eq("status", "new").where_op("total", ">", 100))`.
498    pub fn where_any(self, group: impl FnOnce(Self) -> Self) -> Self {
499        self.group(true, group)
500    }
501
502    /// All of the conditions `group` adds must hold, in parentheses; useful
503    /// inside `where_any`: `.where_any(|q| q.where_eq("a", 1).where_all(|q| …))`.
504    pub fn where_all(self, group: impl FnOnce(Self) -> Self) -> Self {
505        self.group(false, group)
506    }
507
508    fn group(mut self, any: bool, group: impl FnOnce(Self) -> Self) -> Self {
509        let mut inner = group(Self::new());
510        if let Some(error) = inner.error.take() {
511            self.error.get_or_insert(error);
512        }
513        self.filters.push(Filter::Group {
514            any,
515            filters: inner.filters,
516        });
517        self.binds.extend(inner.binds);
518        self
519    }
520
521    /// Applies `add` only when `condition` holds, e.g. an optional search:
522    /// `.when(!q.is_empty(), |query| query.where_like("name", format!("%{q}%")))`.
523    pub fn when(self, condition: bool, add: impl FnOnce(Self) -> Self) -> Self {
524        if condition { add(self) } else { self }
525    }
526
527    /// The rows matching a full-text search, best matches first (as
528    /// [`where_search`](Self::where_search) then
529    /// [`order_by_relevance`](Self::order_by_relevance)); more `order_by`
530    /// calls break ties. The model needs `#[model(search = "…")]` and its
531    /// index: see [`renox::db::search`](super::search). A text without any
532    /// word changes nothing.
533    ///
534    /// ```
535    /// # use renox::prelude::*;
536    /// #[derive(Model, serde::Serialize, Default)]
537    /// #[model(table = "posts", search = "title, body", soft_deletes)]
538    /// struct Post {
539    ///     id: i64,
540    ///     title: String,
541    ///     body: String,
542    ///     author_id: i64,
543    ///     deleted_at: Option<DateTime>,
544    /// }
545    ///
546    /// # async fn demo(db: Db, q: String) -> Result {
547    /// let page = Post::query()
548    ///     .where_eq("author_id", 7)
549    ///     .search(&q)
550    ///     .order_by_desc("id")
551    ///     .paginate(&db, 1, 20)
552    ///     .await?;
553    /// # let _ = page; Ok(()) }
554    /// ```
555    pub fn search(self, words: &str) -> Self {
556        self.where_search(words).order_by_relevance(words)
557    }
558
559    /// Keeps the rows matching a full-text search (every word, as a word
560    /// or the start of one), without ordering them. User input is safe
561    /// here: only its words reach the database, bound as a value.
562    pub fn where_search(mut self, words: &str) -> Self {
563        if let Some(problem) = super::search::unsearchable::<M>() {
564            self.error.get_or_insert(problem);
565            return self;
566        }
567        if let Some(terms) = super::search::bound(words) {
568            self.filters.push(Filter::Search {
569                sqlite: super::search::filter_sql::<M>(Dialect::Sqlite),
570                postgres: super::search::filter_sql::<M>(Dialect::Postgres),
571            });
572            self.binds.push(DbValue::Text(terms));
573        }
574        self
575    }
576
577    /// Sorts by how well rows match a full-text search, best first (rows
578    /// that don't match last); call `order_by` after it to break ties.
579    pub fn order_by_relevance(mut self, words: &str) -> Self {
580        if let Some(problem) = super::search::unsearchable::<M>() {
581            self.error.get_or_insert(problem);
582            return self;
583        }
584        if let Some(terms) = super::search::bound(words) {
585            self.order.push(Order::Relevance {
586                sqlite: super::search::rank_sql::<M>(Dialect::Sqlite),
587                postgres: super::search::rank_sql::<M>(Dialect::Postgres),
588            });
589            self.order_binds.push(DbValue::Text(terms));
590        }
591        self
592    }
593
594    /// Sorts by `column`, ascending; call again to add tie-breakers.
595    pub fn order_by(mut self, column: &str) -> Self {
596        if let Some(column) = self.column(column) {
597            self.order.push(Order::Sql(format!("{column} ASC")));
598        }
599        self
600    }
601
602    /// Sorts by `column`, descending; call again to add tie-breakers.
603    pub fn order_by_desc(mut self, column: &str) -> Self {
604        if let Some(column) = self.column(column) {
605            self.order.push(Order::Sql(format!("{column} DESC")));
606        }
607        self
608    }
609
610    /// Newest first, by `created_at` when the model has it, otherwise by `id`.
611    pub fn latest(self) -> Self {
612        let column = if M::COLUMNS.contains(&"created_at") {
613            "created_at"
614        } else {
615            "id"
616        };
617        self.order_by_desc(column).order_by_desc("id")
618    }
619
620    /// At most `limit` rows (a limit past `i64::MAX` means no limit).
621    pub fn limit(mut self, limit: u64) -> Self {
622        self.limit = Some(limit.min(i64::MAX as u64));
623        self
624    }
625
626    /// Skips the first `offset` rows.
627    pub fn offset(mut self, offset: u64) -> Self {
628        self.offset = Some(offset.min(i64::MAX as u64));
629        self
630    }
631
632    /// A condition in SQL, for what the builder doesn't cover (dates, JSON,
633    /// full-text search, …), with `?` for each value. Column names are yours
634    /// to quote; never build `sql` from user input.
635    ///
636    /// ```
637    /// # use renox::prelude::*;
638    /// # #[derive(Model, serde::Serialize, Default)] struct Order { id: i64, total: i64, created_at: Option<DateTime> }
639    /// # async fn demo(db: Db) -> Result {
640    /// let today = Order::query()
641    ///     .where_raw("DATE(created_at) = DATE(?)", [renox::db::now()])
642    ///     .order_by_raw("total DESC, id")
643    ///     .get(&db)
644    ///     .await?;
645    /// # let _ = today; Ok(()) }
646    /// ```
647    pub fn where_raw<V: ToDbValue>(
648        mut self,
649        sql: &str,
650        values: impl IntoIterator<Item = V>,
651    ) -> Self {
652        self.filters.push(Filter::Sql(format!("({sql})")));
653        self.binds
654            .extend(values.into_iter().map(|v| v.to_db_value()));
655        self
656    }
657
658    /// An `ORDER BY` term in SQL, e.g. `"total DESC, id"` (no values; never
659    /// from user input).
660    pub fn order_by_raw(mut self, sql: &str) -> Self {
661        self.order.push(Order::Sql(sql.to_owned()));
662        self
663    }
664
665    /// Groups rows by `column`, for `select_as` and `count`:
666    /// `.group_by("user_id").select_as::<(i64, i64), _>(&db, "user_id, COUNT(*)")`.
667    pub fn group_by(mut self, column: &str) -> Self {
668        if let Some(column) = self.column(column) {
669            self.group.push(column);
670        }
671        self
672    }
673
674    /// A condition on the groups, in SQL with `?` for each value:
675    /// `.having_raw("COUNT(*) > ?", [2])`.
676    pub fn having_raw<V: ToDbValue>(
677        mut self,
678        sql: &str,
679        values: impl IntoIterator<Item = V>,
680    ) -> Self {
681        self.having.push(format!("({sql})"));
682        self.having_binds
683            .extend(values.into_iter().map(|v| v.to_db_value()));
684        self
685    }
686
687    /// Locks the matching rows until the transaction ends (`FOR UPDATE`),
688    /// e.g. to read a balance and write it back safely. Use it on `&mut tx`.
689    /// PostgreSQL only: on SQLite a write transaction already holds the whole
690    /// database, so start it with `db.begin_immediate()` instead.
691    pub fn lock_for_update(mut self) -> Self {
692        self.lock = Some("FOR UPDATE");
693        self
694    }
695
696    /// Like `lock_for_update`, but others may still read-lock (`FOR SHARE`).
697    pub fn shared_lock(mut self) -> Self {
698        self.lock = Some("FOR SHARE");
699        self
700    }
701
702    /// Matches no rows at all, e.g. a default scope when no tenant is set.
703    pub fn none(mut self) -> Self {
704        self.filters.push(Filter::Sql("1 = 0".into()));
705        self
706    }
707
708    /// Include soft-deleted rows.
709    pub fn with_trashed(mut self) -> Self {
710        self.trashed = Trashed::With;
711        self
712    }
713
714    /// Only soft-deleted rows.
715    pub fn only_trashed(mut self) -> Self {
716        self.trashed = Trashed::Only;
717        self
718    }
719
720    fn check(&self) -> Result {
721        match &self.error {
722            Some(error) => Err(anyhow!("invalid query: {error}").into()),
723            None => Ok(()),
724        }
725    }
726
727    /// The soft-delete condition, if the model has soft deletes.
728    fn trashed_filter(&self) -> Option<Filter> {
729        if !M::SOFT_DELETES {
730            return None;
731        }
732        match self.trashed {
733            Trashed::Without => Some(Filter::Sql("\"deleted_at\" IS NULL".into())),
734            Trashed::Only => Some(Filter::Sql("\"deleted_at\" IS NOT NULL".into())),
735            Trashed::With => None,
736        }
737    }
738
739    fn where_sql(&self, dialect: Dialect) -> String {
740        let mut filters: Vec<String> = self.filters.iter().map(|f| f.render(dialect)).collect();
741        filters.extend(self.trashed_filter().map(|f| f.render(dialect)));
742        if filters.is_empty() {
743            String::new()
744        } else {
745            format!(" WHERE {}", filters.join(" AND "))
746        }
747    }
748
749    fn select_sql(&self, dialect: Dialect) -> String {
750        let columns: Vec<String> = if M::SELECT_ALL {
751            vec!["*".to_owned()]
752        } else {
753            M::COLUMNS.iter().map(|c| quote(c)).collect()
754        };
755        self.select_columns_sql(dialect, &columns.join(", "))
756    }
757
758    fn select_columns_sql(&self, dialect: Dialect, columns: &str) -> String {
759        let mut sql = format!(
760            "SELECT {columns} FROM {}{}",
761            quote(M::TABLE),
762            self.where_sql(dialect)
763        );
764        sql.push_str(&self.group_sql());
765        if !self.order.is_empty() {
766            let terms: Vec<&str> = self.order.iter().map(|o| o.render(dialect)).collect();
767            sql.push_str(&format!(" ORDER BY {}", terms.join(", ")));
768        }
769        match (self.limit, self.offset) {
770            (Some(limit), Some(offset)) => sql.push_str(&format!(" LIMIT {limit} OFFSET {offset}")),
771            (Some(limit), None) => sql.push_str(&format!(" LIMIT {limit}")),
772            // SQLite needs a LIMIT before OFFSET; -1 means none.
773            (None, Some(offset)) if dialect == Dialect::Sqlite => {
774                sql.push_str(&format!(" LIMIT -1 OFFSET {offset}"))
775            }
776            (None, Some(offset)) => sql.push_str(&format!(" OFFSET {offset}")),
777            (None, None) => {}
778        }
779        if let (Some(lock), Dialect::Postgres) = (self.lock, dialect) {
780            sql.push(' ');
781            sql.push_str(lock);
782        }
783        sql
784    }
785
786    fn group_sql(&self) -> String {
787        let mut sql = String::new();
788        if !self.group.is_empty() {
789            sql.push_str(&format!(" GROUP BY {}", self.group.join(", ")));
790        }
791        if !self.having.is_empty() {
792            sql.push_str(&format!(" HAVING {}", self.having.join(" AND ")));
793        }
794        sql
795    }
796
797    /// The query's values in the order its SQL uses them (without the
798    /// `ORDER BY` ones: see `select_binds`).
799    fn all_binds(&self) -> Vec<DbValue> {
800        let mut binds = self.binds.clone();
801        binds.extend(self.having_binds.iter().cloned());
802        binds
803    }
804
805    /// The values of a SELECT with its `ORDER BY` (`select_columns_sql`).
806    fn select_binds(&self) -> Vec<DbValue> {
807        let mut binds = self.all_binds();
808        binds.extend(self.order_binds.iter().cloned());
809        binds
810    }
811
812    fn clear_order(&mut self) {
813        self.order.clear();
814        self.order_binds.clear();
815    }
816
817    /// Records an error if `column` isn't one of the model's.
818    pub(crate) fn check_column(mut self, column: &str) -> Self {
819        let _ = self.column(column);
820        self
821    }
822
823    /// The SELECT this query runs and its values, e.g. to log or debug it.
824    pub fn to_sql(&self, dialect: Dialect) -> Result<(String, Vec<DbValue>)> {
825        self.check()?;
826        Ok((self.select_sql(dialect), self.select_binds()))
827    }
828
829    /// Selects `columns` (SQL, e.g. `"user_id, COUNT(*) AS orders"`) of the
830    /// matching rows, with `group_by`/`having_raw`, read into a
831    /// `#[derive(FromRow)]` struct or a tuple:
832    ///
833    /// ```
834    /// # use renox::prelude::*;
835    /// # #[derive(Model, serde::Serialize, Default)] struct Order { id: i64, user_id: i64, total: i64, status: String }
836    /// # async fn demo(db: Db) -> Result {
837    /// let big_spenders: Vec<(i64, i64)> = Order::where_eq("status", "paid")
838    ///     .group_by("user_id")
839    ///     .having_raw("SUM(total) > ?", [1_000_000])
840    ///     .order_by_raw("2 DESC")
841    ///     .select_as(&db, "user_id, CAST(SUM(total) AS BIGINT)")
842    ///     .await?;
843    /// # let _ = big_spenders; Ok(()) }
844    /// ```
845    pub async fn select_as<'c, T: super::FromRow, E: Executor<'c>>(
846        self,
847        db: E,
848        columns: &str,
849    ) -> Result<Vec<T>> {
850        self.check()?;
851        let db = db.into_conn();
852        let statement = self.select_columns_sql(db.dialect(), columns);
853        Ok(sql(statement)
854            .bind_all(self.select_binds())
855            .fetch_as(db)
856            .await?)
857    }
858
859    /// `bucket` and an aggregate per bucket (`chart::Trend`): the query's
860    /// conditions, grouped by the first column; its order, limit, groups
861    /// and lock don't apply.
862    pub(crate) async fn buckets<'c, E: Executor<'c>>(
863        mut self,
864        db: E,
865        bucket: &(dyn Fn(Dialect) -> String + Send + Sync),
866        aggregate: &str,
867    ) -> Result<Vec<(String, Option<f64>)>> {
868        self.check()?;
869        self.clear_order();
870        self.group = vec!["1".to_owned()];
871        self.having.clear();
872        self.having_binds.clear();
873        self.limit = None;
874        self.offset = None;
875        self.lock = None;
876        let db = db.into_conn();
877        let dialect = db.dialect();
878        let statement = self.select_columns_sql(
879            dialect,
880            &format!("{}, CAST({aggregate} AS DOUBLE PRECISION)", bucket(dialect)),
881        );
882        Ok(sql(statement)
883            .bind_all(self.all_binds())
884            .fetch_as(db)
885            .await?)
886    }
887
888    /// One aggregate over the matching rows (order and limit don't apply).
889    async fn aggregate<'c, T: FromDb, E: Executor<'c>>(
890        self,
891        db: E,
892        expression: String,
893    ) -> Result<T> {
894        self.check()?;
895        let db = db.into_conn();
896        let statement = format!(
897            "SELECT {expression} FROM {}{}",
898            quote(M::TABLE),
899            self.where_sql(db.dialect())
900        );
901        Ok(sql(statement).bind_all(self.binds).scalar(db).await?)
902    }
903
904    /// The sum of `column`, 0 without rows: `sum::<i64>(…)` for whole
905    /// numbers (money in its smallest unit), `sum::<f64>(…)` for measures.
906    pub async fn sum<'c, T: Number, E: Executor<'c>>(mut self, db: E, column: &str) -> Result<T> {
907        let Some(column) = self.column(column) else {
908            return Err(self.check().unwrap_err());
909        };
910        let expression = format!("CAST(COALESCE(SUM({column}), 0) AS {})", T::SQL_TYPE);
911        self.aggregate(db, expression).await
912    }
913
914    /// The average of `column`, or `None` without rows.
915    pub async fn avg<'c, E: Executor<'c>>(mut self, db: E, column: &str) -> Result<Option<f64>> {
916        let Some(column) = self.column(column) else {
917            return Err(self.check().unwrap_err());
918        };
919        self.aggregate(db, format!("CAST(AVG({column}) AS DOUBLE PRECISION)"))
920            .await
921    }
922
923    /// The smallest value of `column`, or `None` without rows.
924    pub async fn min<'c, T: FromDb, E: Executor<'c>>(
925        mut self,
926        db: E,
927        column: &str,
928    ) -> Result<Option<T>>
929    where
930        Option<T>: FromDb,
931    {
932        let Some(column) = self.column(column) else {
933            return Err(self.check().unwrap_err());
934        };
935        self.aggregate(db, format!("MIN({column})")).await
936    }
937
938    /// The largest value of `column`, or `None` without rows.
939    pub async fn max<'c, T: FromDb, E: Executor<'c>>(
940        mut self,
941        db: E,
942        column: &str,
943    ) -> Result<Option<T>>
944    where
945        Option<T>: FromDb,
946    {
947        let Some(column) = self.column(column) else {
948            return Err(self.check().unwrap_err());
949        };
950        self.aggregate(db, format!("MAX({column})")).await
951    }
952
953    /// One column of the matching rows, in the query's order:
954    /// `Product::query().order_by("name").pluck::<String, _>(&db, "name")`.
955    pub async fn pluck<'c, T: FromDb, E: Executor<'c>>(
956        mut self,
957        db: E,
958        column: &str,
959    ) -> Result<Vec<T>> {
960        let Some(column) = self.column(column) else {
961            return Err(self.check().unwrap_err());
962        };
963        self.check()?;
964        let db = db.into_conn();
965        let statement = self.select_columns_sql(db.dialect(), &column);
966        Ok(sql(statement)
967            .bind_all(self.select_binds())
968            .scalars(db)
969            .await?)
970    }
971
972    /// Sets columns on every matching row (and `updated_at` when the model
973    /// has it); returns how many changed.
974    ///
975    /// ```
976    /// # use renox::prelude::*;
977    /// # #[derive(Model, serde::Serialize, Default)] struct Order { id: i64, status: String, paid_at: Option<DateTime> }
978    /// # async fn demo(db: Db) -> Result {
979    /// Order::where_eq("status", "pending")
980    ///     .update(&db, &[("status", &"paid"), ("paid_at", &renox::db::now())])
981    ///     .await?;
982    /// # Ok(()) }
983    /// ```
984    pub async fn update<'c, E: Executor<'c>>(
985        mut self,
986        db: E,
987        values: &[(&str, &(dyn ToDbValue + Sync))],
988    ) -> Result<u64> {
989        let mut sets = Vec::new();
990        let mut binds = Vec::new();
991        for (column, value) in values {
992            if *column == "id" {
993                self.error
994                    .get_or_insert_with(|| "update can't change `id`".into());
995            }
996            if let Some(quoted) = self.column(column) {
997                sets.push(format!("{quoted} = ?"));
998                binds.push(value.to_db_value());
999            }
1000        }
1001        if M::COLUMNS.contains(&"updated_at") && !values.iter().any(|(c, _)| *c == "updated_at") {
1002            sets.push(format!("{} = ?", quote("updated_at")));
1003            binds.push(now().to_db_value());
1004        }
1005        if sets.is_empty() {
1006            self.check()?;
1007            return Ok(0);
1008        }
1009        self.set_rows(db, sets.join(", "), binds).await
1010    }
1011
1012    /// Adds `by` to `column` on every matching row (negative to subtract),
1013    /// in the database, so concurrent changes aren't lost:
1014    /// `Product::where_eq("id", id).increment(&db, "stock", -1)`.
1015    pub async fn increment<'c, E: Executor<'c>>(
1016        mut self,
1017        db: E,
1018        column: &str,
1019        by: i64,
1020    ) -> Result<u64> {
1021        let Some(quoted) = self.column(column) else {
1022            return Err(self.check().unwrap_err());
1023        };
1024        let mut sets = format!("{quoted} = {quoted} + ?");
1025        let mut binds = vec![DbValue::Integer(by)];
1026        if M::COLUMNS.contains(&"updated_at") {
1027            sets.push_str(&format!(", {} = ?", quote("updated_at")));
1028            binds.push(now().to_db_value());
1029        }
1030        self.set_rows(db, sets, binds).await
1031    }
1032
1033    async fn set_rows<'c, E: Executor<'c>>(
1034        self,
1035        db: E,
1036        sets: String,
1037        binds: Vec<DbValue>,
1038    ) -> Result<u64> {
1039        self.check()?;
1040        let db = db.into_conn();
1041        let statement = format!(
1042            "UPDATE {} SET {sets}{}",
1043            quote(M::TABLE),
1044            self.where_sql(db.dialect())
1045        );
1046        Ok(sql(statement)
1047            .bind_all(binds)
1048            .bind_all(self.binds)
1049            .execute(db)
1050            .await?)
1051    }
1052
1053    /// Like `first`, but no row becomes a 404 response.
1054    pub async fn first_or_404<'c, E: Executor<'c>>(self, db: E) -> Result<M> {
1055        self.first(db).await?.ok_or(crate::Error::NotFound)
1056    }
1057
1058    /// The first matching row, or `make()` saved as a new one. If another
1059    /// request creates it at the same moment (a unique index stops the
1060    /// second insert), the row it created is returned.
1061    //
1062    // Not an `async fn`: written that way, a handler awaiting it failed
1063    // axum's `Send` check (rustc issue #100013; see it/send_handlers.rs).
1064    #[allow(clippy::manual_async_fn)] // the `+ Send` in the signature is the point
1065    pub fn first_or_create<'a>(
1066        self,
1067        db: &'a Db,
1068        make: impl FnOnce() -> M + Send + 'a,
1069    ) -> impl Future<Output = Result<M>> + Send + 'a {
1070        async move {
1071            if let Some(found) = self.clone().first(db).await? {
1072                return Ok(found);
1073            }
1074            match M::create(db, make()).await {
1075                Ok(created) => Ok(created),
1076                Err(err) if err.is_unique_violation() => self.first(db).await?.ok_or(err),
1077                Err(err) => Err(err),
1078            }
1079        }
1080    }
1081
1082    /// Runs `each` on the matching rows `size` at a time, in id order, so a
1083    /// large table never sits in memory at once. (The query's own order and
1084    /// limit don't apply.)
1085    pub async fn chunk<F, Fut>(self, db: &Db, size: u64, mut each: F) -> Result<u64>
1086    where
1087        F: FnMut(Vec<M>) -> Fut,
1088        Fut: std::future::Future<Output = Result>,
1089    {
1090        let mut last: Option<M::Key> = None;
1091        let mut seen = 0_u64;
1092        loop {
1093            let mut page = self.clone();
1094            page.clear_order();
1095            page.limit = None;
1096            page.offset = None;
1097            if let Some(last) = last.take() {
1098                page = page.where_op("id", ">", last);
1099            }
1100            let rows = page.order_by("id").limit(size.max(1)).get(db).await?;
1101            let Some(tail) = rows.last() else { break };
1102            last = Some(tail.id());
1103            let full = rows.len() as u64 == size.max(1);
1104            seen += rows.len() as u64;
1105            each(rows).await?;
1106            if !full {
1107                break;
1108            }
1109        }
1110        Ok(seen)
1111    }
1112
1113    /// Runs the query and returns every matching row.
1114    pub async fn get<'c, E: Executor<'c>>(self, db: E) -> Result<Vec<M>> {
1115        self.check()?;
1116        let db = db.into_conn();
1117        let rows = sql(self.select_sql(db.dialect()))
1118            .bind_all(self.select_binds())
1119            .fetch_all(db)
1120            .await?;
1121        Ok(rows
1122            .iter()
1123            .map(M::from_row)
1124            .collect::<std::result::Result<_, _>>()?)
1125    }
1126
1127    /// The first matching row (in the query's order), or `None`.
1128    pub async fn first<'c, E: Executor<'c>>(self, db: E) -> Result<Option<M>> {
1129        Ok(self.limit(1).get(db).await?.into_iter().next())
1130    }
1131
1132    /// How many rows match (with `group_by`: how many groups).
1133    pub async fn count<'c, E: Executor<'c>>(self, db: E) -> Result<u64> {
1134        self.check()?;
1135        let db = db.into_conn();
1136        let statement = if self.group.is_empty() {
1137            format!(
1138                "SELECT COUNT(*) FROM {}{}",
1139                quote(M::TABLE),
1140                self.where_sql(db.dialect())
1141            )
1142        } else {
1143            format!(
1144                "SELECT COUNT(*) FROM (SELECT 1 AS one FROM {}{}{}) AS groups",
1145                quote(M::TABLE),
1146                self.where_sql(db.dialect()),
1147                self.group_sql()
1148            )
1149        };
1150        let count: i64 = sql(statement).bind_all(self.all_binds()).scalar(db).await?;
1151        Ok(count as u64)
1152    }
1153
1154    /// Whether any row matches.
1155    pub async fn exists<'c, E: Executor<'c>>(self, db: E) -> Result<bool> {
1156        Ok(self.count(db).await? > 0)
1157    }
1158
1159    /// One page of results plus the numbers needed to render page links.
1160    /// `page` starts at 1; `per_page` is capped at 1000.
1161    pub async fn paginate(self, db: &Db, page: u32, per_page: u32) -> Result<Paginated<M>> {
1162        let page = page.max(1);
1163        let per_page = per_page.clamp(1, 1000);
1164        let total = self.clone().count(db).await?;
1165        let items = self
1166            .limit(u64::from(per_page))
1167            .offset(u64::from(page - 1) * u64::from(per_page))
1168            .get(db)
1169            .await?;
1170        Ok(Paginated::new(items, page, per_page, total))
1171    }
1172
1173    /// One page without counting the rows (one query instead of two): for
1174    /// "previous / next" links on large tables. `page` starts at 1.
1175    pub async fn simple_paginate(self, db: &Db, page: u32, per_page: u32) -> Result<SimplePage<M>> {
1176        let page = page.max(1);
1177        let per_page = per_page.clamp(1, 1000);
1178        let mut items = self
1179            .limit(u64::from(per_page) + 1)
1180            .offset(u64::from(page - 1) * u64::from(per_page))
1181            .get(db)
1182            .await?;
1183        let has_next = items.len() > per_page as usize;
1184        items.truncate(per_page as usize);
1185        Ok(SimplePage {
1186            items,
1187            page,
1188            per_page,
1189            has_prev: page > 1,
1190            has_next,
1191        })
1192    }
1193
1194    /// The `per_page` newest rows after `cursor` (by id; the query's own
1195    /// order doesn't apply), for APIs and infinite scroll on big tables:
1196    /// unlike page numbers, rows added meanwhile don't shift the pages.
1197    /// "Newest" is the key's order: creation order for `i64`, ULIDs and
1198    /// UUID v7s; a `String` key pages in text order.
1199    ///
1200    /// ```
1201    /// # use renox::prelude::*;
1202    /// # #[derive(Model, serde::Serialize, Default)] struct Event { id: i64, name: String }
1203    /// #[derive(serde::Deserialize)]
1204    /// struct Params { cursor: Option<String> }
1205    ///
1206    /// async fn events(State(db): State<Db>, Query(p): Query<Params>) -> Result<Json<renox::db::CursorPage<Event>>> {
1207    ///     Ok(Json(Event::query().cursor_paginate(&db, p.cursor.as_deref(), 50).await?))
1208    /// }
1209    /// ```
1210    pub async fn cursor_paginate(
1211        mut self,
1212        db: &Db,
1213        cursor: Option<&str>,
1214        per_page: u32,
1215    ) -> Result<CursorPage<M>> {
1216        let per_page = per_page.clamp(1, 1000);
1217        if let Some(cursor) = cursor {
1218            let Ok(after) = cursor.parse::<M::Key>() else {
1219                return Err(crate::Error::BadRequest("invalid cursor".into()));
1220            };
1221            self = self.where_op("id", "<", after);
1222        }
1223        self.clear_order();
1224        let mut items = self
1225            .order_by_desc("id")
1226            .limit(u64::from(per_page) + 1)
1227            .get(db)
1228            .await?;
1229        let more = items.len() > per_page as usize;
1230        items.truncate(per_page as usize);
1231        let next_cursor = more
1232            .then(|| items.last().map(|last| last.id().to_string()))
1233            .flatten();
1234        Ok(CursorPage {
1235            items,
1236            per_page,
1237            next_cursor,
1238        })
1239    }
1240
1241    /// The first matching row, or `make()` unsaved (`firstOrNew`).
1242    pub async fn first_or_new(self, db: &Db, make: impl FnOnce() -> M + Send) -> Result<M> {
1243        Ok(match self.first(db).await? {
1244            Some(found) => found,
1245            None => make(),
1246        })
1247    }
1248
1249    /// Changes the first matching row with `change`, or creates `make()`
1250    /// with `change` applied (`updateOrCreate`); returns it saved.
1251    ///
1252    /// ```
1253    /// # use renox::prelude::*;
1254    /// # #[derive(Model, serde::Serialize, Default)] struct Setting { id: i64, user_id: i64, key: String, value: String }
1255    /// # async fn demo(db: Db) -> Result {
1256    /// let theme = Setting::where_eq("user_id", 7)
1257    ///     .where_eq("key", "theme")
1258    ///     .update_or_create(
1259    ///         &db,
1260    ///         || Setting { user_id: 7, key: "theme".into(), ..Default::default() },
1261    ///         |s| s.value = "dark".into(),
1262    ///     )
1263    ///     .await?;
1264    /// # let _ = theme; Ok(()) }
1265    /// ```
1266    #[allow(clippy::manual_async_fn)] // the `+ Send` in the signature is the point
1267    pub fn update_or_create<'a>(
1268        self,
1269        db: &'a Db,
1270        make: impl FnOnce() -> M + Send + 'a,
1271        change: impl FnOnce(&mut M) + Send + 'a,
1272    ) -> impl Future<Output = Result<M>> + Send + 'a {
1273        async move {
1274            let mut model = match self.first(db).await? {
1275                Some(found) => found,
1276                None => make(),
1277            };
1278            change(&mut model);
1279            model.save(db).await?;
1280            Ok(model)
1281        }
1282    }
1283
1284    /// Deletes every matching row (soft-deletes them for models with soft deletes).
1285    pub async fn delete<'c, E: Executor<'c>>(self, db: E) -> Result<u64> {
1286        if M::SOFT_DELETES {
1287            self.check()?;
1288            let db = db.into_conn();
1289            return Ok(sql(format!(
1290                "UPDATE {} SET \"deleted_at\" = ?{}",
1291                quote(M::TABLE),
1292                self.where_sql(db.dialect())
1293            ))
1294            .bind(now())
1295            .bind_all(self.binds)
1296            .execute(db)
1297            .await?);
1298        }
1299        self.force_delete(db).await
1300    }
1301
1302    /// Removes every matching row, even for models with soft deletes.
1303    pub async fn force_delete<'c, E: Executor<'c>>(self, db: E) -> Result<u64> {
1304        self.check()?;
1305        let db = db.into_conn();
1306        Ok(sql(format!(
1307            "DELETE FROM {}{}",
1308            quote(M::TABLE),
1309            self.where_sql(db.dialect())
1310        ))
1311        .bind_all(self.binds)
1312        .execute(db)
1313        .await?)
1314    }
1315}