1#[cfg(any(feature = "postgres", test))]
10use std::borrow::Cow;
11use std::fmt;
12use std::sync::Arc;
13
14use cookie::Key;
15
16use sqlx::sqlite::{Sqlite, SqliteArguments, SqlitePool, SqliteRow};
17use sqlx::{AssertSqlSafe, Column, Row as _};
18
19use super::value::DbValue;
20use super::{DbError, ToDbValue};
21
22#[cfg(feature = "postgres")]
23use sqlx::postgres::{PgArguments, PgPool, PgRow, Postgres};
24
25#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub enum Dialect {
28 Sqlite,
30 Postgres,
32}
33
34impl fmt::Display for Dialect {
35 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
36 f.write_str(match self {
37 Dialect::Sqlite => "sqlite",
38 Dialect::Postgres => "postgres",
39 })
40 }
41}
42
43#[derive(Clone)]
49pub struct Db {
50 pool: Pool,
51 schema: SchemaEpoch,
52 key: Option<Arc<Key>>,
54}
55
56#[derive(Clone, Default)]
66pub(crate) struct SchemaEpoch(std::sync::Arc<std::sync::Mutex<Option<std::time::Instant>>>);
67
68impl SchemaEpoch {
69 pub(crate) fn changed(&self) {
70 *self.0.lock().unwrap_or_else(|e| e.into_inner()) = Some(std::time::Instant::now());
71 }
72
73 pub(crate) fn is_stale(&self, age: std::time::Duration) -> bool {
75 self.0
76 .lock()
77 .unwrap_or_else(|e| e.into_inner())
78 .is_some_and(|at| age >= at.elapsed())
79 }
80}
81
82#[derive(Clone)]
83enum Pool {
84 Sqlite(SqlitePool),
85 #[cfg(feature = "postgres")]
86 Postgres(PgPool),
87}
88
89impl fmt::Debug for Db {
90 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
91 f.debug_struct("Db")
92 .field("dialect", &self.dialect())
93 .finish_non_exhaustive()
94 }
95}
96
97impl From<SqlitePool> for Db {
98 fn from(pool: SqlitePool) -> Self {
99 Self::with_epoch(Pool::Sqlite(pool), SchemaEpoch::default())
100 }
101}
102
103#[cfg(feature = "postgres")]
104impl From<PgPool> for Db {
105 fn from(pool: PgPool) -> Self {
106 Self::with_epoch(Pool::Postgres(pool), SchemaEpoch::default())
107 }
108}
109
110impl Db {
111 fn with_epoch(pool: Pool, schema: SchemaEpoch) -> Self {
112 Self {
113 pool,
114 schema,
115 key: None,
116 }
117 }
118
119 pub(crate) fn with_key(mut self, key: Key) -> Self {
122 self.key = Some(Arc::new(key));
123 self
124 }
125
126 pub(crate) fn from_sqlite(pool: SqlitePool, schema: SchemaEpoch) -> Self {
127 Self::with_epoch(Pool::Sqlite(pool), schema)
128 }
129
130 #[cfg(feature = "postgres")]
131 pub(crate) fn from_postgres(pool: PgPool, schema: SchemaEpoch) -> Self {
132 Self::with_epoch(Pool::Postgres(pool), schema)
133 }
134
135 pub(crate) fn schema_changed(&self) {
138 self.schema.changed();
139 }
140
141 pub fn dialect(&self) -> Dialect {
143 match self.pool {
144 Pool::Sqlite(_) => Dialect::Sqlite,
145 #[cfg(feature = "postgres")]
146 Pool::Postgres(_) => Dialect::Postgres,
147 }
148 }
149
150 pub fn sqlite(&self) -> Option<&SqlitePool> {
152 match &self.pool {
153 Pool::Sqlite(pool) => Some(pool),
154 #[cfg(feature = "postgres")]
155 Pool::Postgres(_) => None,
156 }
157 }
158
159 #[cfg(feature = "postgres")]
161 pub fn postgres(&self) -> Option<&PgPool> {
162 match &self.pool {
163 Pool::Postgres(pool) => Some(pool),
164 Pool::Sqlite(_) => None,
165 }
166 }
167
168 pub async fn begin(&self) -> Result<Transaction, DbError> {
171 let inner = match &self.pool {
172 Pool::Sqlite(pool) => TxInner::Sqlite(pool.begin().await?),
173 #[cfg(feature = "postgres")]
174 Pool::Postgres(pool) => TxInner::Postgres(pool.begin().await?),
175 };
176 Ok(Transaction {
177 inner,
178 key: self.key.clone(),
179 savepoints: 0,
180 })
181 }
182
183 pub async fn transaction<T, F>(&self, work: F) -> crate::Result<T>
205 where
206 F: for<'t> FnMut(
207 &'t mut Transaction,
208 ) -> futures_util::future::BoxFuture<'t, crate::Result<T>>,
209 {
210 self.transaction_retrying(1, work).await
211 }
212
213 pub async fn transaction_retrying<T, F>(&self, attempts: u32, mut work: F) -> crate::Result<T>
218 where
219 F: for<'t> FnMut(
220 &'t mut Transaction,
221 ) -> futures_util::future::BoxFuture<'t, crate::Result<T>>,
222 {
223 let attempts = attempts.max(1);
224 let mut attempt = 1;
225 loop {
226 let outcome = async {
227 let mut tx = self.begin().await?;
228 let value = work(&mut tx).await?;
229 tx.commit().await?;
230 Ok::<T, crate::Error>(value)
231 }
232 .await;
233 match outcome {
234 Err(err) if attempt < attempts && err.is_retryable() => {
235 let pause = std::time::Duration::from_millis(20 * u64::from(attempt));
236 tokio::time::sleep(pause).await;
237 attempt += 1;
238 }
239 other => return other,
240 }
241 }
242 }
243
244 pub fn retrying<'a, T, F, Fut>(
284 &'a self,
285 attempts: u32,
286 mut work: F,
287 ) -> impl Future<Output = crate::Result<T>> + Send + 'a
288 where
289 T: Send + 'a,
290 F: FnMut() -> Fut + Send + 'a,
291 Fut: Future<Output = crate::Result<T>> + Send + 'a,
292 {
293 let attempts = attempts.max(1);
294 async move {
295 let mut attempt = 1;
296 loop {
297 match work().await {
298 Err(err) if attempt < attempts && err.is_retryable() => {
299 let pause = std::time::Duration::from_millis(20 * u64::from(attempt));
300 tokio::time::sleep(pause).await;
301 attempt += 1;
302 }
303 other => return other,
304 }
305 }
306 }
307 }
308
309 pub async fn begin_immediate(&self) -> Result<Transaction, DbError> {
313 let inner = match &self.pool {
314 Pool::Sqlite(pool) => TxInner::Sqlite(pool.begin_with("BEGIN IMMEDIATE").await?),
315 #[cfg(feature = "postgres")]
316 Pool::Postgres(pool) => TxInner::Postgres(pool.begin().await?),
317 };
318 Ok(Transaction {
319 inner,
320 key: self.key.clone(),
321 savepoints: 0,
322 })
323 }
324
325 pub async fn close(&self) {
327 match &self.pool {
328 Pool::Sqlite(pool) => pool.close().await,
329 #[cfg(feature = "postgres")]
330 Pool::Postgres(pool) => pool.close().await,
331 }
332 }
333}
334
335pub struct Transaction {
337 inner: TxInner,
338 key: Option<Arc<Key>>,
339 savepoints: u32,
341}
342
343enum TxInner {
344 Sqlite(sqlx::Transaction<'static, Sqlite>),
345 #[cfg(feature = "postgres")]
346 Postgres(sqlx::Transaction<'static, Postgres>),
347}
348
349impl Transaction {
350 pub fn dialect(&self) -> Dialect {
352 match self.inner {
353 TxInner::Sqlite(_) => Dialect::Sqlite,
354 #[cfg(feature = "postgres")]
355 TxInner::Postgres(_) => Dialect::Postgres,
356 }
357 }
358
359 pub async fn commit(self) -> Result<(), DbError> {
361 match self.inner {
362 TxInner::Sqlite(tx) => Ok(tx.commit().await?),
363 #[cfg(feature = "postgres")]
364 TxInner::Postgres(tx) => Ok(tx.commit().await?),
365 }
366 }
367
368 pub async fn rollback(self) -> Result<(), DbError> {
370 match self.inner {
371 TxInner::Sqlite(tx) => Ok(tx.rollback().await?),
372 #[cfg(feature = "postgres")]
373 TxInner::Postgres(tx) => Ok(tx.rollback().await?),
374 }
375 }
376
377 pub async fn savepoint<T, F>(&mut self, work: F) -> crate::Result<T>
409 where
410 F: for<'t> FnOnce(
411 &'t mut Transaction,
412 ) -> futures_util::future::BoxFuture<'t, crate::Result<T>>,
413 {
414 let name = format!("renox_savepoint_{}", self.savepoints + 1);
415 sql(format!("SAVEPOINT {name}")).execute(&mut *self).await?;
416 self.savepoints += 1;
417 let outcome = work(self).await;
418 self.savepoints -= 1;
419 match outcome {
420 Ok(value) => {
421 sql(format!("RELEASE SAVEPOINT {name}"))
422 .execute(&mut *self)
423 .await?;
424 Ok(value)
425 }
426 Err(err) => {
427 sql(format!("ROLLBACK TO SAVEPOINT {name}"))
428 .execute(&mut *self)
429 .await?;
430 sql(format!("RELEASE SAVEPOINT {name}"))
431 .execute(&mut *self)
432 .await?;
433 Err(err)
434 }
435 }
436 }
437}
438
439impl fmt::Debug for Transaction {
440 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
441 f.debug_struct("Transaction")
442 .field("dialect", &self.dialect())
443 .finish_non_exhaustive()
444 }
445}
446
447pub trait Executor<'c>: Send + executor::Sealed {
449 #[doc(hidden)]
450 fn into_conn(self) -> Conn<'c>;
451}
452
453mod executor {
454 pub trait Sealed {}
455 impl Sealed for &super::Db {}
456 impl Sealed for &mut super::Transaction {}
457 impl Sealed for super::Conn<'_> {}
458}
459
460#[doc(hidden)]
461pub enum Conn<'c> {
462 Pool(&'c Db),
463 Tx(&'c mut Transaction),
464}
465
466impl<'c> Executor<'c> for &'c Db {
467 fn into_conn(self) -> Conn<'c> {
468 Conn::Pool(self)
469 }
470}
471
472impl<'c> Executor<'c> for &'c mut Transaction {
473 fn into_conn(self) -> Conn<'c> {
474 Conn::Tx(self)
475 }
476}
477
478impl<'c> Executor<'c> for Conn<'c> {
479 fn into_conn(self) -> Conn<'c> {
480 self
481 }
482}
483
484impl Conn<'_> {
485 fn key(&self) -> Option<Arc<Key>> {
487 match self {
488 Conn::Pool(db) => db.key.clone(),
489 Conn::Tx(tx) => tx.key.clone(),
490 }
491 }
492
493 pub(crate) fn reborrow(&mut self) -> Conn<'_> {
495 match self {
496 Conn::Pool(db) => Conn::Pool(db),
497 Conn::Tx(tx) => Conn::Tx(tx),
498 }
499 }
500
501 pub(crate) fn dialect(&self) -> Dialect {
503 match self {
504 Conn::Pool(db) => db.dialect(),
505 Conn::Tx(tx) => tx.dialect(),
506 }
507 }
508}
509
510macro_rules! dispatch {
513 ($conn:expr, |$build:ident, $exec:ident| $body:expr) => {
514 match $conn {
515 Conn::Pool(db) => match &db.pool {
516 Pool::Sqlite(pool) => {
517 #[allow(unused_variables)]
518 let $build = sqlite_query;
519 let $exec = pool;
520 ($body).map_err(DbError::from)
521 }
522 #[cfg(feature = "postgres")]
523 Pool::Postgres(pool) => {
524 #[allow(unused_variables)]
525 let $build = postgres_query;
526 let $exec = pool;
527 ($body).map_err(DbError::from)
528 }
529 },
530 Conn::Tx(tx) => match &mut tx.inner {
531 TxInner::Sqlite(tx) => {
532 #[allow(unused_variables)]
533 let $build = sqlite_query;
534 let $exec = &mut **tx;
535 ($body).map_err(DbError::from)
536 }
537 #[cfg(feature = "postgres")]
538 TxInner::Postgres(tx) => {
539 #[allow(unused_variables)]
540 let $build = postgres_query;
541 let $exec = &mut **tx;
542 ($body).map_err(DbError::from)
543 }
544 },
545 }
546 };
547}
548
549fn seal(key: Option<&Key>, args: Vec<DbValue>) -> Result<Vec<DbValue>, DbError> {
551 args.into_iter()
552 .map(|value| match value {
553 DbValue::Encrypted(plain) => match key {
554 Some(key) => Ok(DbValue::Text(super::encrypted::seal(key, &plain.0))),
555 None => Err(DbError::from(sqlx::Error::Encode(
556 "an Encrypted value needs the app's Db (it has the APP_KEY); \
557 this Db was made outside App"
558 .into(),
559 ))),
560 },
561 other => Ok(other),
562 })
563 .collect()
564}
565
566fn sqlite_query(
567 sql: String,
568 args: Vec<DbValue>,
569) -> sqlx::query::Query<'static, Sqlite, SqliteArguments> {
570 args.into_iter().fold(
571 sqlx::query(AssertSqlSafe(sql)),
572 |query, value| match value.for_sqlite() {
573 DbValue::Integer(v) => query.bind(v),
574 DbValue::Real(v) => query.bind(v),
575 DbValue::Text(v) => query.bind(v),
576 DbValue::Blob(v) => query.bind(v),
577 _ => query.bind(None::<i64>),
578 },
579 )
580}
581
582#[cfg(feature = "postgres")]
583fn postgres_query(
584 sql: String,
585 args: Vec<DbValue>,
586) -> sqlx::query::Query<'static, Postgres, PgArguments> {
587 let sql = numbered_placeholders(&sql).into_owned();
588 args.into_iter().fold(
589 sqlx::query(AssertSqlSafe(sql)),
590 |query, value| match value {
591 DbValue::Null => query.bind(UntypedNull),
592 DbValue::Integer(v) => query.bind(v),
593 DbValue::Real(v) => query.bind(v),
594 DbValue::Text(v) => query.bind(v),
595 DbValue::Blob(v) => query.bind(v),
596 DbValue::Bool(v) => query.bind(v),
597 DbValue::DateTime(v) => query.bind(v),
598 DbValue::NaiveDateTime(v) => query.bind(v),
599 DbValue::Date(v) => query.bind(v),
600 DbValue::Time(v) => query.bind(v),
601 DbValue::Json(v) => query.bind(sqlx::types::Json(v)),
602 #[cfg(feature = "uuid")]
603 DbValue::Uuid(v) => query.bind(v),
604 DbValue::Encrypted(_) => query.bind(UntypedNull),
606 },
607 )
608}
609
610#[cfg(feature = "postgres")]
613struct UntypedNull;
614
615#[cfg(feature = "postgres")]
616impl sqlx::Type<Postgres> for UntypedNull {
617 fn type_info() -> sqlx::postgres::PgTypeInfo {
618 sqlx::postgres::PgTypeInfo::with_oid(sqlx::postgres::types::Oid(0))
620 }
621}
622
623#[cfg(feature = "postgres")]
624impl sqlx::Encode<'_, Postgres> for UntypedNull {
625 fn encode_by_ref(
626 &self,
627 _buf: &mut sqlx::postgres::PgArgumentBuffer,
628 ) -> Result<sqlx::encode::IsNull, sqlx::error::BoxDynError> {
629 Ok(sqlx::encode::IsNull::Yes)
630 }
631}
632
633#[cfg(any(feature = "postgres", test))]
637pub(crate) fn numbered_placeholders(sql: &str) -> Cow<'_, str> {
638 if !sql.contains('?') {
639 return Cow::Borrowed(sql);
640 }
641 let bytes = sql.as_bytes();
643 let ident = |b: u8| b.is_ascii_alphanumeric() || b == b'_' || b == b'$';
644 let mut out = String::with_capacity(sql.len() + 8);
645 let mut n = 0;
646 let mut copied = 0; let mut i = 0;
648 while i < bytes.len() {
649 let at = i;
650 let end = match bytes[i] {
651 quote @ (b'\'' | b'"') => {
652 let escapes = quote == b'\''
654 && i > 0
655 && matches!(bytes[i - 1], b'E' | b'e')
656 && (i < 2 || !ident(bytes[i - 2]));
657 let mut j = i + 1;
658 while j < bytes.len() && bytes[j] != quote {
661 j += if escapes && bytes[j] == b'\\' { 2 } else { 1 };
662 }
663 j + 1
664 }
665 b'-' if bytes.get(i + 1) == Some(&b'-') => {
666 sql[i..].find('\n').map_or(bytes.len(), |nl| i + nl + 1)
667 }
668 b'/' if bytes.get(i + 1) == Some(&b'*') => sql[i + 2..]
669 .find("*/")
670 .map_or(bytes.len(), |e| i + 2 + e + 2),
671 b'$' if i == 0 || !ident(bytes[i - 1]) => {
672 let tag_len = bytes[i + 1..]
676 .iter()
677 .take_while(|b| b.is_ascii_alphanumeric() || **b == b'_')
678 .count();
679 let starts_with_digit = bytes.get(i + 1).is_some_and(u8::is_ascii_digit);
680 if !starts_with_digit && bytes.get(i + 1 + tag_len) == Some(&b'$') {
681 let tag = &sql[i..i + tag_len + 2];
682 sql[i + tag.len()..]
683 .find(tag)
684 .map_or(bytes.len(), |e| i + tag.len() + e + tag.len())
685 } else {
686 i + 1
687 }
688 }
689 b'?' => {
690 n += 1;
691 out.push_str(&sql[copied..i]);
692 out.push_str(&format!("${n}"));
693 copied = i + 1;
694 i + 1
695 }
696 _ => i + 1,
697 };
698 i = end.min(bytes.len()).max(at + 1);
699 }
700 out.push_str(&sql[copied..]);
701 Cow::Owned(out)
702}
703
704pub fn sql(sql: impl Into<String>) -> Sql {
719 Sql {
720 sql: sql.into(),
721 args: Vec::new(),
722 }
723}
724
725#[derive(Debug, Clone)]
727#[must_use = "a statement does nothing until it is run"]
728pub struct Sql {
729 sql: String,
730 args: Vec<DbValue>,
731}
732
733impl Sql {
734 pub fn bind(mut self, value: impl ToDbValue) -> Self {
736 self.args.push(value.to_db_value());
737 self
738 }
739
740 pub fn bind_all(mut self, values: impl IntoIterator<Item = DbValue>) -> Self {
742 self.args.extend(values);
743 self
744 }
745
746 pub async fn fetch_all<'c>(self, db: impl Executor<'c>) -> Result<Vec<Row>, DbError> {
748 let Self { sql, args } = self;
749 super::query_log::record(&sql);
750 let conn = db.into_conn();
751 let key = conn.key();
752 let args = seal(key.as_deref(), args)?;
753 let rows: Vec<Row> = dispatch!(conn, |build, exec| build(sql, args)
754 .fetch_all(exec)
755 .await
756 .map(|rows| rows.into_iter().map(Row::from).collect()))?;
757 Ok(rows
758 .into_iter()
759 .map(|row| row.with_key(key.clone()))
760 .collect())
761 }
762
763 pub async fn fetch_optional<'c>(self, db: impl Executor<'c>) -> Result<Option<Row>, DbError> {
765 let Self { sql, args } = self;
766 super::query_log::record(&sql);
767 let conn = db.into_conn();
768 let key = conn.key();
769 let args = seal(key.as_deref(), args)?;
770 let row: Option<Row> = dispatch!(conn, |build, exec| build(sql, args)
771 .fetch_optional(exec)
772 .await
773 .map(|row| row.map(Row::from)))?;
774 Ok(row.map(|row| row.with_key(key)))
775 }
776
777 pub async fn fetch_one<'c>(self, db: impl Executor<'c>) -> Result<Row, DbError> {
779 self.fetch_optional(db)
780 .await?
781 .ok_or_else(|| DbError::from(sqlx::Error::RowNotFound))
782 }
783
784 pub async fn execute<'c>(self, db: impl Executor<'c>) -> Result<u64, DbError> {
786 let Self { sql, args } = self;
787 super::query_log::record(&sql);
788 let conn = db.into_conn();
789 let args = seal(conn.key().as_deref(), args)?;
790 dispatch!(conn, |build, exec| build(sql, args)
791 .execute(exec)
792 .await
793 .map(|done| done.rows_affected()))
794 }
795
796 pub async fn fetch_as<'c, T: super::FromRow>(
799 self,
800 db: impl Executor<'c>,
801 ) -> Result<Vec<T>, DbError> {
802 self.fetch_all(db).await?.iter().map(T::from_row).collect()
803 }
804
805 pub async fn fetch_optional_as<'c, T: super::FromRow>(
807 self,
808 db: impl Executor<'c>,
809 ) -> Result<Option<T>, DbError> {
810 self.fetch_optional(db)
811 .await?
812 .as_ref()
813 .map(T::from_row)
814 .transpose()
815 }
816
817 pub async fn fetch_one_as<'c, T: super::FromRow>(
819 self,
820 db: impl Executor<'c>,
821 ) -> Result<T, DbError> {
822 T::from_row(&self.fetch_one(db).await?)
823 }
824
825 pub async fn scalar<'c, T: FromDb>(self, db: impl Executor<'c>) -> Result<T, DbError> {
827 self.fetch_one(db).await?.try_get(0)
828 }
829
830 pub async fn scalar_optional<'c, T: FromDb>(
832 self,
833 db: impl Executor<'c>,
834 ) -> Result<Option<T>, DbError> {
835 self.fetch_optional(db)
836 .await?
837 .map(|row| row.try_get(0))
838 .transpose()
839 }
840
841 pub async fn scalars<'c, T: FromDb>(self, db: impl Executor<'c>) -> Result<Vec<T>, DbError> {
843 self.fetch_all(db)
844 .await?
845 .iter()
846 .map(|row| row.try_get(0))
847 .collect()
848 }
849}
850
851pub(crate) async fn script<'c>(db: impl Executor<'c>, sql: &str) -> Result<u64, DbError> {
854 let sql = sql.to_owned();
855 dispatch!(db.into_conn(), |build, exec| sqlx::raw_sql(AssertSqlSafe(
856 sql
857 ))
858 .execute(exec)
859 .await
860 .map(|done| done.rows_affected()))
861}
862
863pub struct Row(pub(crate) RowInner, Option<Arc<Key>>);
865
866pub(crate) enum RowInner {
867 Sqlite(SqliteRow),
868 #[cfg(feature = "postgres")]
869 Postgres(PgRow),
870}
871
872impl From<SqliteRow> for Row {
873 fn from(row: SqliteRow) -> Self {
874 Self(RowInner::Sqlite(row), None)
875 }
876}
877
878#[cfg(feature = "postgres")]
879impl From<PgRow> for Row {
880 fn from(row: PgRow) -> Self {
881 Self(RowInner::Postgres(row), None)
882 }
883}
884
885impl Row {
886 fn with_key(mut self, key: Option<Arc<Key>>) -> Self {
887 self.1 = key;
888 self
889 }
890
891 pub fn try_get<T: FromDb>(&self, index: impl RowIndex) -> Result<T, DbError> {
893 super::encrypted::reading(self.1.as_ref(), || match &self.0 {
894 RowInner::Sqlite(row) => Ok(row.try_get(index)?),
895 #[cfg(feature = "postgres")]
896 RowInner::Postgres(row) => Ok(row.try_get(index)?),
897 })
898 }
899
900 pub(crate) fn json(&self, column: &str) -> serde_json::Value {
903 use serde_json::Value;
904 if let Ok(v) = self.try_get::<Option<i64>>(column) {
905 return v.map_or(Value::Null, Value::from);
906 }
907 if let Ok(v) = self.try_get::<Option<i32>>(column) {
908 return v.map_or(Value::Null, Value::from);
909 }
910 if let Ok(v) = self.try_get::<Option<i16>>(column) {
912 return v.map_or(Value::Null, Value::from);
913 }
914 if let Ok(v) = self.try_get::<Option<f64>>(column) {
915 return v.map_or(Value::Null, Value::from);
916 }
917 if let Ok(v) = self.try_get::<Option<f32>>(column) {
918 return v.map_or(Value::Null, |n| Value::from(f64::from(n)));
919 }
920 if let Ok(v) = self.try_get::<Option<bool>>(column) {
921 return v.map_or(Value::Null, Value::from);
922 }
923 if let Ok(v) = self.try_get::<Option<chrono::DateTime<chrono::Utc>>>(column) {
924 return v.map_or(Value::Null, |d| Value::from(d.to_rfc3339()));
925 }
926 if let Ok(v) = self.try_get::<Option<chrono::NaiveDateTime>>(column) {
927 return v.map_or(Value::Null, |d| {
928 Value::from(d.format("%Y-%m-%dT%H:%M:%S").to_string())
929 });
930 }
931 if let Ok(v) = self.try_get::<Option<chrono::NaiveDate>>(column) {
932 return v.map_or(Value::Null, |d| Value::from(d.to_string()));
933 }
934 if let Ok(v) = self.try_get::<Option<String>>(column) {
935 return v.map_or(Value::Null, Value::from);
936 }
937 if let Ok(v) = self.try_get::<Option<chrono::NaiveTime>>(column) {
940 return v.map_or(Value::Null, |t| Value::from(t.to_string()));
941 }
942 if let Ok(v) = self.try_get::<Option<serde_json::Value>>(column) {
943 return v.unwrap_or(Value::Null);
944 }
945 Value::Null
946 }
947
948 pub fn columns(&self) -> Vec<&str> {
950 match &self.0 {
951 RowInner::Sqlite(row) => row.columns().iter().map(Column::name).collect(),
952 #[cfg(feature = "postgres")]
953 RowInner::Postgres(row) => row.columns().iter().map(Column::name).collect(),
954 }
955 }
956
957 pub fn sqlite(&self) -> Option<&SqliteRow> {
959 match &self.0 {
960 RowInner::Sqlite(row) => Some(row),
961 #[cfg(feature = "postgres")]
962 RowInner::Postgres(_) => None,
963 }
964 }
965
966 #[cfg(feature = "postgres")]
968 pub fn postgres(&self) -> Option<&PgRow> {
969 match &self.0 {
970 RowInner::Postgres(row) => Some(row),
971 RowInner::Sqlite(_) => None,
972 }
973 }
974}
975
976impl fmt::Debug for Row {
977 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
978 f.debug_struct("Row")
979 .field("columns", &self.columns())
980 .finish_non_exhaustive()
981 }
982}
983
984pub trait FromDb: for<'r> sqlx::Decode<'r, Sqlite> + sqlx::Type<Sqlite> + bounds::Postgres {}
992
993impl<T> FromDb for T where
994 T: for<'r> sqlx::Decode<'r, Sqlite> + sqlx::Type<Sqlite> + bounds::Postgres
995{
996}
997
998pub trait RowIndex: sqlx::ColumnIndex<SqliteRow> + bounds::PostgresIndex {}
1000
1001impl<T> RowIndex for T where T: sqlx::ColumnIndex<SqliteRow> + bounds::PostgresIndex {}
1002
1003#[doc(hidden)]
1005pub mod bounds {
1006 #[cfg(feature = "postgres")]
1007 mod on {
1008 use sqlx::postgres::PgRow;
1009
1010 pub trait Postgres:
1011 for<'r> sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>
1012 {
1013 }
1014 impl<T> Postgres for T where T: for<'r> sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>
1015 {}
1016
1017 pub trait PostgresIndex: sqlx::ColumnIndex<PgRow> {}
1018 impl<T: sqlx::ColumnIndex<PgRow>> PostgresIndex for T {}
1019 }
1020
1021 #[cfg(not(feature = "postgres"))]
1022 mod on {
1023 pub trait Postgres {}
1024 impl<T: ?Sized> Postgres for T {}
1025
1026 pub trait PostgresIndex {}
1027 impl<T: ?Sized> PostgresIndex for T {}
1028 }
1029
1030 pub use on::{Postgres, PostgresIndex};
1031}
1032
1033#[cfg(test)]
1034mod tests {
1035 use super::{numbered_placeholders, seal};
1036 use crate::db::DbValue;
1037
1038 #[test]
1039 fn numbers_placeholders_outside_quotes_and_comments() {
1040 assert_eq!(
1041 numbered_placeholders("SELECT * FROM t WHERE a = ? AND b IN (?, ?)"),
1042 "SELECT * FROM t WHERE a = $1 AND b IN ($2, $3)"
1043 );
1044 assert_eq!(
1045 numbered_placeholders("SELECT '?', \"a?\" -- why?\nFROM t /* ? */ WHERE x = ?"),
1046 "SELECT '?', \"a?\" -- why?\nFROM t /* ? */ WHERE x = $1"
1047 );
1048 assert_eq!(
1049 numbered_placeholders("SELECT 'it''s ?' WHERE y = ?"),
1050 "SELECT 'it''s ?' WHERE y = $1"
1051 );
1052 assert_eq!(numbered_placeholders("SELECT 1"), "SELECT 1");
1053 }
1054
1055 #[test]
1057 fn skips_dollar_quoted_and_escaped_strings() {
1058 assert_eq!(
1059 numbered_placeholders("SELECT $$why?$$, $fn$ a ? b $fn$ WHERE x = ?"),
1060 "SELECT $$why?$$, $fn$ a ? b $fn$ WHERE x = $1"
1061 );
1062 assert_eq!(
1063 numbered_placeholders(r"SELECT E'it\'s ?', e'\\' WHERE y = ? AND z = ?"),
1064 r"SELECT E'it\'s ?', e'\\' WHERE y = $1 AND z = $2"
1065 );
1066 assert_eq!(
1068 numbered_placeholders(r"SELECT '\', a$b$ FROM t WHERE c = ?"),
1069 r"SELECT '\', a$b$ FROM t WHERE c = $1"
1070 );
1071 assert_eq!(
1072 numbered_placeholders("SELECT 'é?' WHERE ü = ?"),
1073 "SELECT 'é?' WHERE ü = $1"
1074 );
1075 }
1076
1077 #[test]
1080 fn placeholders_after_unfinished_quotes_stay() {
1081 assert_eq!(numbered_placeholders("SELECT 'open ?"), "SELECT 'open ?");
1082 assert_eq!(numbered_placeholders("SELECT 1 -- ?"), "SELECT 1 -- ?");
1083 assert_eq!(numbered_placeholders("SELECT 1 /* ?"), "SELECT 1 /* ?");
1084 assert_eq!(numbered_placeholders("SELECT $tag$ ?"), "SELECT $tag$ ?");
1085 assert_eq!(numbered_placeholders("SELECT ? || $"), "SELECT $1 || $");
1086 assert!(numbered_placeholders("SELECT $1, ?").ends_with(", $1"));
1088 assert_eq!(numbered_placeholders("?"), "$1");
1089 }
1090
1091 #[tokio::test]
1094 async fn retrying_follows_the_error_kind() {
1095 use std::sync::Arc;
1096 use std::sync::atomic::{AtomicU32, Ordering};
1097
1098 use crate::db::DbError;
1099 use crate::db::error::fake;
1100
1101 fn busy() -> crate::Error {
1102 crate::Error::from(DbError::from(fake::coded(Some("5"))))
1103 }
1104 let db = super::super::connect(&crate::Config::default())
1105 .await
1106 .unwrap();
1107
1108 let tries = Arc::new(AtomicU32::new(0));
1109 let counter = tries.clone();
1110 let value = db
1111 .retrying(3, move || {
1112 let counter = counter.clone();
1113 async move {
1114 match counter.fetch_add(1, Ordering::SeqCst) {
1115 0 | 1 => Err(busy()),
1116 _ => Ok("done"),
1117 }
1118 }
1119 })
1120 .await
1121 .unwrap();
1122 assert_eq!((value, tries.load(Ordering::SeqCst)), ("done", 3));
1123
1124 let tries = Arc::new(AtomicU32::new(0));
1125 let counter = tries.clone();
1126 let err = db
1127 .retrying(2, move || {
1128 counter.fetch_add(1, Ordering::SeqCst);
1129 async { Err::<(), _>(busy()) }
1130 })
1131 .await
1132 .unwrap_err();
1133 assert!(err.is_retryable());
1134 assert_eq!(
1135 tries.load(Ordering::SeqCst),
1136 2,
1137 "gives up at the last attempt"
1138 );
1139
1140 let tries = Arc::new(AtomicU32::new(0));
1141 let counter = tries.clone();
1142 let err = db
1143 .retrying(5, move || {
1144 counter.fetch_add(1, Ordering::SeqCst);
1145 async { Err::<(), _>(crate::Error::NotFound) }
1146 })
1147 .await
1148 .unwrap_err();
1149 assert!(matches!(err, crate::Error::NotFound));
1150 assert_eq!(tries.load(Ordering::SeqCst), 1, "not retried");
1151
1152 let tries = Arc::new(AtomicU32::new(0));
1153 let counter = tries.clone();
1154 let value = db
1155 .transaction_retrying(3, move |tx| {
1156 let first = counter.fetch_add(1, Ordering::SeqCst) == 0;
1157 Box::pin(async move {
1158 let one: i64 = super::sql("SELECT CAST(1 AS BIGINT)").scalar(tx).await?;
1159 if first { Err(busy()) } else { Ok(one) }
1160 })
1161 })
1162 .await
1163 .unwrap();
1164 assert_eq!((value, tries.load(Ordering::SeqCst)), (1, 2));
1165 }
1166
1167 #[test]
1169 fn sealing_needs_the_apps_key() {
1170 let err = seal(
1171 None,
1172 vec![DbValue::Encrypted(super::super::encrypted::Unsealed(
1173 "x".into(),
1174 ))],
1175 )
1176 .unwrap_err();
1177 assert!(
1178 err.to_string()
1179 .contains("an Encrypted value needs the app's Db"),
1180 "{err}"
1181 );
1182 assert!(matches!(
1184 seal(None, vec![DbValue::Integer(1)]).unwrap()[..],
1185 [DbValue::Integer(1)]
1186 ));
1187 }
1188}