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
17const LARGE_IN: usize = 1000;
19
20pub 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#[derive(Clone)]
42enum Filter {
43 Sql(String),
45 Like { column: String, not: bool },
47 JsonIn { kind: &'static str, column: String },
49 Group { any: bool, filters: Vec<Filter> },
51 Not(Box<Filter>),
53 Exists {
55 not: bool,
56 table: &'static str,
57 correlation: String,
58 filters: Vec<Filter>,
59 },
60 InQuery {
62 column: String,
63 table: &'static str,
64 sub_column: String,
65 filters: Vec<Filter>,
66 },
67 Search { sqlite: String, postgres: String },
70}
71
72#[derive(Clone)]
74enum Order {
75 Sql(String),
77 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 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
156fn 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
193pub 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 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 pub fn where_eq(self, column: &str, value: impl ToDbValue) -> Self {
284 self.where_op(column, "=", value)
285 }
286
287 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 pub fn where_like(self, column: &str, pattern: impl ToDbValue) -> Self {
310 self.where_op(column, "like", pattern)
311 }
312
313 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 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 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 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 Filter::Sql("0 = 1".into())
355 } else if let Some((kind, list)) = large_list(&values) {
356 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 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 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 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 pub fn where_has<N: Model>(self, children: Query<N>, foreign_key: &str) -> Self {
464 self.related(children, foreign_key, false)
465 }
466
467 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 pub fn where_any(self, group: impl FnOnce(Self) -> Self) -> Self {
499 self.group(true, group)
500 }
501
502 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 pub fn when(self, condition: bool, add: impl FnOnce(Self) -> Self) -> Self {
524 if condition { add(self) } else { self }
525 }
526
527 pub fn search(self, words: &str) -> Self {
556 self.where_search(words).order_by_relevance(words)
557 }
558
559 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 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 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 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 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 pub fn limit(mut self, limit: u64) -> Self {
622 self.limit = Some(limit.min(i64::MAX as u64));
623 self
624 }
625
626 pub fn offset(mut self, offset: u64) -> Self {
628 self.offset = Some(offset.min(i64::MAX as u64));
629 self
630 }
631
632 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 pub fn order_by_raw(mut self, sql: &str) -> Self {
661 self.order.push(Order::Sql(sql.to_owned()));
662 self
663 }
664
665 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 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 pub fn lock_for_update(mut self) -> Self {
692 self.lock = Some("FOR UPDATE");
693 self
694 }
695
696 pub fn shared_lock(mut self) -> Self {
698 self.lock = Some("FOR SHARE");
699 self
700 }
701
702 pub fn none(mut self) -> Self {
704 self.filters.push(Filter::Sql("1 = 0".into()));
705 self
706 }
707
708 pub fn with_trashed(mut self) -> Self {
710 self.trashed = Trashed::With;
711 self
712 }
713
714 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 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 (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 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 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 pub(crate) fn check_column(mut self, column: &str) -> Self {
819 let _ = self.column(column);
820 self
821 }
822
823 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 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 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 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 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 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 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 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 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 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 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 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 #[allow(clippy::manual_async_fn)] 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 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 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 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 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 pub async fn exists<'c, E: Executor<'c>>(self, db: E) -> Result<bool> {
1156 Ok(self.count(db).await? > 0)
1157 }
1158
1159 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 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 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 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 #[allow(clippy::manual_async_fn)] 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 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 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}