Skip to main content

renox_core/db/
conn.rs

1//! Renox's own handle on the database, so the same code runs on SQLite and
2//! PostgreSQL: [`Db`] (the pool), [`Transaction`], [`Row`] and raw SQL
3//! through [`sql()`].
4//!
5//! SQL written for Renox uses `?` placeholders. On PostgreSQL they become
6//! `$1`, `$2`, … before the query is sent (a `?` inside quotes or comments is
7//! left alone).
8
9#[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/// Which database engine a [`Db`] talks to.
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub enum Dialect {
28    /// SQLite.
29    Sqlite,
30    /// PostgreSQL (the `postgres` feature).
31    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/// The database pool. Take it in a handler with `State(db): State<Db>`.
44///
45/// Models, queries and [`sql()`] take `&db` (or `&mut tx` inside a
46/// [`Transaction`]). For something only sqlx can do, reach the underlying
47/// pool with [`Db::sqlite`] or `Db::postgres` (with the `postgres` feature).
48#[derive(Clone)]
49pub struct Db {
50    pool: Pool,
51    schema: SchemaEpoch,
52    /// Seals and opens `Encrypted` columns (`APP_KEY`, set at boot).
53    key: Option<Arc<Key>>,
54}
55
56/// When this process last changed the schema (migrations ran).
57///
58/// A pooled connection keeps the schema it read, and prepares statements
59/// against it until it next steps one. sqlx takes a statement's columns from
60/// that first prepare, so after `ALTER TABLE users ADD COLUMN …` a
61/// `SELECT *` on an older connection gets one column more than sqlx expects
62/// and panics (sqlx-sqlite `row.rs`), or PostgreSQL refuses a cached plan.
63/// The pools made by [`super::connect`] drop connections opened before the
64/// last change instead of reusing them.
65#[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    /// Whether a connection of this age was opened before the last change.
74    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    /// The key `Encrypted` columns are sealed and opened with (the app's
120    /// `APP_KEY`, set at boot).
121    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    /// Records that the schema changed, so older pooled connections aren't
136    /// reused (see `SchemaEpoch`).
137    pub(crate) fn schema_changed(&self) {
138        self.schema.changed();
139    }
140
141    /// The engine this pool talks to.
142    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    /// The sqlx pool, when the database is SQLite.
151    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    /// The sqlx pool, when the database is PostgreSQL.
160    #[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    /// Starts a transaction. Pass `&mut tx` wherever a `&db` goes, then
169    /// `tx.commit()`; dropping it without committing rolls it back.
170    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    /// Runs `work` in a transaction: committed when it returns `Ok`, rolled
184    /// back when it returns `Err`.
185    ///
186    /// ```
187    /// # use renox::prelude::*;
188    /// # async fn demo(db: Db) -> Result {
189    /// let moved = db
190    ///     .transaction_retrying(3, |tx| {
191    ///         Box::pin(async move {
192    ///             renox::db::sql("UPDATE accounts SET balance = balance - 100 WHERE id = 1")
193    ///                 .execute(&mut *tx)
194    ///                 .await?;
195    ///             renox::db::sql("UPDATE accounts SET balance = balance + 100 WHERE id = 2")
196    ///                 .execute(&mut *tx)
197    ///                 .await?;
198    ///             Ok(100)
199    ///         })
200    ///     })
201    ///     .await?;
202    /// # let _: i64 = moved; Ok(()) }
203    /// ```
204    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    /// Like [`Db::transaction`], trying up to `attempts` times when another
214    /// transaction got in the way (SQLite busy, PostgreSQL serialization
215    /// failure or deadlock), with a short, growing pause in between. `work`
216    /// must be safe to run again.
217    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    /// Runs `work` again (up to `attempts` times, with a short, growing
245    /// pause) while it fails because another transaction got in the way
246    /// (SQLite busy, PostgreSQL serialization failure or deadlock). `work`
247    /// opens and commits its own transaction, so unlike
248    /// [`Db::transaction_retrying`] it can borrow from the caller, and it can
249    /// roll back and still return a value (drop or `rollback` the
250    /// transaction instead of committing).
251    ///
252    /// ```
253    /// # use renox::prelude::*;
254    /// enum Transfer { Done, Short }
255    ///
256    /// # async fn demo(db: Db, amount: i64, from: i64, to: i64) -> Result {
257    /// let outcome = db
258    ///     .retrying(3, || async {
259    ///         let mut tx = db.begin().await?;
260    ///         let taken = renox::db::sql(
261    ///             "UPDATE accounts SET balance = balance - ? WHERE id = ? AND balance >= ?",
262    ///         )
263    ///         .bind(amount) // borrowed from the caller
264    ///         .bind(from)
265    ///         .bind(amount)
266    ///         .execute(&mut tx)
267    ///         .await?;
268    ///         if taken == 0 {
269    ///             tx.rollback().await?;
270    ///             return Ok(Transfer::Short); // rolled back, with an answer
271    ///         }
272    ///         renox::db::sql("UPDATE accounts SET balance = balance + ? WHERE id = ?")
273    ///             .bind(amount)
274    ///             .bind(to)
275    ///             .execute(&mut tx)
276    ///             .await?;
277    ///         tx.commit().await?;
278    ///         Ok(Transfer::Done)
279    ///     })
280    ///     .await?;
281    /// # let _ = outcome; Ok(()) }
282    /// ```
283    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    /// A transaction that takes SQLite's write lock at once (`BEGIN
310    /// IMMEDIATE`), so what it reads can't change before it writes: the
311    /// SQLite counterpart of `lock_for_update`. Same as `begin` on PostgreSQL.
312    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    /// Closes every connection; later queries fail.
326    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
335/// A database transaction, from [`Db::begin`].
336pub struct Transaction {
337    inner: TxInner,
338    key: Option<Arc<Key>>,
339    /// Savepoints open inside it (for their names).
340    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    /// The engine this transaction runs on.
351    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    /// Commits the transaction.
360    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    /// Rolls the transaction back (dropping it uncommitted does the same).
369    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    /// Runs `work` inside a savepoint (a transaction inside the
378    /// transaction): when it returns `Err`, only what it did is undone and
379    /// the transaction goes on; when it returns `Ok`, its changes stay, to
380    /// be committed (or rolled back) with the rest. Savepoints nest.
381    ///
382    /// ```
383    /// # use renox::prelude::*;
384    /// # async fn demo(db: Db) -> Result {
385    /// let mut tx = db.begin().await?;
386    /// renox::db::sql("INSERT INTO orders (total) VALUES (?)").bind(18_000).execute(&mut tx).await?;
387    /// // Optional: a voucher that may be used up already. If it is, the order stays.
388    /// let voucher = tx
389    ///     .savepoint(|tx| {
390    ///         Box::pin(async move {
391    ///             let used = renox::db::sql("UPDATE vouchers SET used = used + 1 WHERE code = ? AND used < max_uses")
392    ///                 .bind("KOPI10")
393    ///                 .execute(&mut *tx)
394    ///                 .await?;
395    ///             renox::abort_if(used == 0, StatusCode::CONFLICT, "used up")?;
396    ///             Ok(())
397    ///         })
398    ///     })
399    ///     .await;
400    /// tx.commit().await?;
401    /// # let _ = voucher; Ok(()) }
402    /// ```
403    ///
404    /// On PostgreSQL a failed statement stops the whole transaction until
405    /// it rolls back; inside a savepoint only the savepoint rolls back, so
406    /// the transaction can go on after an expected failure (a unique
407    /// violation you handle, say).
408    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
447/// Where a query runs: `&db` or `&mut tx`. Sealed: only those.
448pub 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    /// The key for `Encrypted` columns of the database this runs on.
486    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    /// The same connection for one more statement.
494    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    /// Which engine the query will run on, for SQL that differs.
502    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
510/// Runs `$body` against whichever backend `$conn` is: `$build` makes the
511/// backend's query from `(sql, args)` and `$exec` is its sqlx executor.
512macro_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
549/// Seals the values of `Encrypted` fields with the database's key.
550fn 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            // Sealed into `Text` before the statement is built.
605            DbValue::Encrypted(_) => query.bind(UntypedNull),
606        },
607    )
608}
609
610/// A `NULL` parameter without a type, so PostgreSQL takes the type from
611/// where it's used (a `NULL` sent as `BIGINT` can't go into a `TIMESTAMPTZ`).
612#[cfg(feature = "postgres")]
613struct UntypedNull;
614
615#[cfg(feature = "postgres")]
616impl sqlx::Type<Postgres> for UntypedNull {
617    fn type_info() -> sqlx::postgres::PgTypeInfo {
618        // OID 0 means "unspecified" in the protocol.
619        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/// Rewrites `?` placeholders as PostgreSQL's `$1`, `$2`, …, skipping quoted
634/// strings (`'…'`, `E'…'` with backslash escapes, dollar-quoted `$$…$$` and
635/// `$tag$…$tag$`), quoted identifiers and comments.
636#[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    // Every delimiter is ASCII, so byte offsets always fall between characters.
642    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; // `sql[..copied]` is in `out`
647    let mut i = 0;
648    while i < bytes.len() {
649        let at = i;
650        let end = match bytes[i] {
651            quote @ (b'\'' | b'"') => {
652                // `E'…'` (or `e'…'`): a backslash escapes the next character.
653                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                // A doubled quote inside is an escaped quote: the loop ends at
659                // the first one and the next one opens a new run.
660                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                // `$tag$` opens a dollar-quoted string (the tag may be empty
673                // and doesn't start with a digit, unlike `$1`), closed by the
674                // same `$tag$`.
675                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
704/// A raw SQL statement with `?` placeholders.
705///
706/// ```
707/// # use renox::prelude::*;
708/// # async fn demo(db: Db) -> Result {
709/// let rows = renox::db::sql("SELECT name FROM products WHERE price < ?")
710///     .bind(20_000)
711///     .fetch_all(&db)
712///     .await?;
713/// let name: String = rows[0].try_get("name")?;
714///
715/// let total: i64 = renox::db::sql("SELECT COUNT(*) FROM products").scalar(&db).await?;
716/// # let _ = (name, total); Ok(()) }
717/// ```
718pub fn sql(sql: impl Into<String>) -> Sql {
719    Sql {
720        sql: sql.into(),
721        args: Vec::new(),
722    }
723}
724
725/// A statement built by [`sql()`].
726#[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    /// Binds the next `?`.
735    pub fn bind(mut self, value: impl ToDbValue) -> Self {
736        self.args.push(value.to_db_value());
737        self
738    }
739
740    /// Binds several values in order.
741    pub fn bind_all(mut self, values: impl IntoIterator<Item = DbValue>) -> Self {
742        self.args.extend(values);
743        self
744    }
745
746    /// Runs the statement and returns every row.
747    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    /// Runs the statement and returns the first row, if any.
764    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    /// The first row; an error if there is none.
778    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    /// Runs the statement and returns the number of rows it changed.
785    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    /// Every row, read into `T`: a model, a `#[derive(FromRow)]` struct or a
797    /// tuple; see [`super::FromRow`].
798    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    /// The first row read into `T`, or `None` if there is no row.
806    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    /// The first row read into `T`; an error if there is none.
818    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    /// The first column of the first row; an error if there is no row.
826    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    /// The first column of the first row, or `None` if there is no row.
831    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    /// The first column of every row.
842    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
851/// Runs SQL that may hold several statements and no parameters, e.g. a
852/// migration file. Returns the number of rows changed.
853pub(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
863/// One result row. Read columns by name or position with [`Row::try_get`].
864pub 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    /// A column's value, by name (`"name"`) or position (`0`).
892    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    /// A column's value as JSON, whatever its type; `null` when it can't be
901    /// read as a number, boolean, text, date, time, timestamp or JSON.
902    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        // PostgreSQL's SMALLINT and REAL decode only as these.
911        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        // After text, so SQLite text that looks like a time stays as written;
938        // PostgreSQL's TIME decodes only as this.
939        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    /// The column names, in order.
949    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    /// The sqlx row, when the database is SQLite.
958    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    /// The sqlx row, when the database is PostgreSQL.
967    #[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
984/// A Rust type that can be read from a column on every enabled database:
985/// the integer and float types, `bool`, `String`, `Vec<u8>`, the chrono date
986/// and time types, `Option<T>` of those, `Json<T>`, `Encrypted<T>`, `Ulid`,
987/// `Uuid` (the `uuid` feature) and `#[derive(DbEnum)]` enums. Those are the
988/// types Renox promises; the trait is implemented through sqlx's own
989/// traits, so other types sqlx decodes work too, without that promise (see
990/// docs/stability.md).
991pub 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
998/// A column name (`&str`) or position (`usize`).
999pub trait RowIndex: sqlx::ColumnIndex<SqliteRow> + bounds::PostgresIndex {}
1000
1001impl<T> RowIndex for T where T: sqlx::ColumnIndex<SqliteRow> + bounds::PostgresIndex {}
1002
1003/// Extra bounds that only apply when the `postgres` feature is on.
1004#[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    /// PostgreSQL's other string forms keep their `?` (#219).
1056    #[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        // A plain string keeps its backslash; `$1` and `a$b` aren't quotes.
1067        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    /// Unterminated quotes, comments and dollar quotes swallow the rest; a
1078    /// `$` at the very end or before a digit is no quote (#254).
1079    #[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        // A `$1` doesn't open a quote, so the `?` after it is still seen.
1087        assert!(numbered_placeholders("SELECT $1, ?").ends_with(", $1"));
1088        assert_eq!(numbered_placeholders("?"), "$1");
1089    }
1090
1091    /// `retrying` and `transaction_retrying` try again after a busy error,
1092    /// stop at the last attempt, and return other errors at once.
1093    #[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    // #254: an Encrypted value on a Db made outside App has no key to seal with.
1168    #[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        // Plain values pass through untouched.
1183        assert!(matches!(
1184            seal(None, vec![DbValue::Integer(1)]).unwrap()[..],
1185            [DbValue::Integer(1)]
1186        ));
1187    }
1188}