#[cfg(any(feature = "postgres", test))]
use std::borrow::Cow;
use std::fmt;
use std::sync::Arc;
use cookie::Key;
use sqlx::sqlite::{Sqlite, SqliteArguments, SqlitePool, SqliteRow};
use sqlx::{AssertSqlSafe, Column, Row as _};
use super::value::DbValue;
use super::{DbError, ToDbValue};
#[cfg(feature = "postgres")]
use sqlx::postgres::{PgArguments, PgPool, PgRow, Postgres};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Dialect {
Sqlite,
Postgres,
}
impl fmt::Display for Dialect {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
Dialect::Sqlite => "sqlite",
Dialect::Postgres => "postgres",
})
}
}
#[derive(Clone)]
pub struct Db {
pool: Pool,
schema: SchemaEpoch,
key: Option<Arc<Key>>,
}
#[derive(Clone, Default)]
pub(crate) struct SchemaEpoch(std::sync::Arc<std::sync::Mutex<Option<std::time::Instant>>>);
impl SchemaEpoch {
pub(crate) fn changed(&self) {
*self.0.lock().unwrap_or_else(|e| e.into_inner()) = Some(std::time::Instant::now());
}
pub(crate) fn is_stale(&self, age: std::time::Duration) -> bool {
self.0
.lock()
.unwrap_or_else(|e| e.into_inner())
.is_some_and(|at| age >= at.elapsed())
}
}
#[derive(Clone)]
enum Pool {
Sqlite(SqlitePool),
#[cfg(feature = "postgres")]
Postgres(PgPool),
}
impl fmt::Debug for Db {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Db")
.field("dialect", &self.dialect())
.finish_non_exhaustive()
}
}
impl From<SqlitePool> for Db {
fn from(pool: SqlitePool) -> Self {
Self::with_epoch(Pool::Sqlite(pool), SchemaEpoch::default())
}
}
#[cfg(feature = "postgres")]
impl From<PgPool> for Db {
fn from(pool: PgPool) -> Self {
Self::with_epoch(Pool::Postgres(pool), SchemaEpoch::default())
}
}
impl Db {
fn with_epoch(pool: Pool, schema: SchemaEpoch) -> Self {
Self {
pool,
schema,
key: None,
}
}
pub(crate) fn with_key(mut self, key: Key) -> Self {
self.key = Some(Arc::new(key));
self
}
pub(crate) fn from_sqlite(pool: SqlitePool, schema: SchemaEpoch) -> Self {
Self::with_epoch(Pool::Sqlite(pool), schema)
}
#[cfg(feature = "postgres")]
pub(crate) fn from_postgres(pool: PgPool, schema: SchemaEpoch) -> Self {
Self::with_epoch(Pool::Postgres(pool), schema)
}
pub(crate) fn schema_changed(&self) {
self.schema.changed();
}
pub fn dialect(&self) -> Dialect {
match self.pool {
Pool::Sqlite(_) => Dialect::Sqlite,
#[cfg(feature = "postgres")]
Pool::Postgres(_) => Dialect::Postgres,
}
}
pub fn sqlite(&self) -> Option<&SqlitePool> {
match &self.pool {
Pool::Sqlite(pool) => Some(pool),
#[cfg(feature = "postgres")]
Pool::Postgres(_) => None,
}
}
#[cfg(feature = "postgres")]
pub fn postgres(&self) -> Option<&PgPool> {
match &self.pool {
Pool::Postgres(pool) => Some(pool),
Pool::Sqlite(_) => None,
}
}
pub async fn begin(&self) -> Result<Transaction, DbError> {
let inner = match &self.pool {
Pool::Sqlite(pool) => TxInner::Sqlite(pool.begin().await?),
#[cfg(feature = "postgres")]
Pool::Postgres(pool) => TxInner::Postgres(pool.begin().await?),
};
Ok(Transaction {
inner,
key: self.key.clone(),
savepoints: 0,
})
}
pub async fn transaction<T, F>(&self, work: F) -> crate::Result<T>
where
F: for<'t> FnMut(
&'t mut Transaction,
) -> futures_util::future::BoxFuture<'t, crate::Result<T>>,
{
self.transaction_retrying(1, work).await
}
pub async fn transaction_retrying<T, F>(&self, attempts: u32, mut work: F) -> crate::Result<T>
where
F: for<'t> FnMut(
&'t mut Transaction,
) -> futures_util::future::BoxFuture<'t, crate::Result<T>>,
{
let attempts = attempts.max(1);
let mut attempt = 1;
loop {
let outcome = async {
let mut tx = self.begin().await?;
let value = work(&mut tx).await?;
tx.commit().await?;
Ok::<T, crate::Error>(value)
}
.await;
match outcome {
Err(err) if attempt < attempts && err.is_retryable() => {
let pause = std::time::Duration::from_millis(20 * u64::from(attempt));
tokio::time::sleep(pause).await;
attempt += 1;
}
other => return other,
}
}
}
pub fn retrying<'a, T, F, Fut>(
&'a self,
attempts: u32,
mut work: F,
) -> impl Future<Output = crate::Result<T>> + Send + 'a
where
T: Send + 'a,
F: FnMut() -> Fut + Send + 'a,
Fut: Future<Output = crate::Result<T>> + Send + 'a,
{
let attempts = attempts.max(1);
async move {
let mut attempt = 1;
loop {
match work().await {
Err(err) if attempt < attempts && err.is_retryable() => {
let pause = std::time::Duration::from_millis(20 * u64::from(attempt));
tokio::time::sleep(pause).await;
attempt += 1;
}
other => return other,
}
}
}
}
pub async fn begin_immediate(&self) -> Result<Transaction, DbError> {
let inner = match &self.pool {
Pool::Sqlite(pool) => TxInner::Sqlite(pool.begin_with("BEGIN IMMEDIATE").await?),
#[cfg(feature = "postgres")]
Pool::Postgres(pool) => TxInner::Postgres(pool.begin().await?),
};
Ok(Transaction {
inner,
key: self.key.clone(),
savepoints: 0,
})
}
pub async fn close(&self) {
match &self.pool {
Pool::Sqlite(pool) => pool.close().await,
#[cfg(feature = "postgres")]
Pool::Postgres(pool) => pool.close().await,
}
}
}
pub struct Transaction {
inner: TxInner,
key: Option<Arc<Key>>,
savepoints: u32,
}
enum TxInner {
Sqlite(sqlx::Transaction<'static, Sqlite>),
#[cfg(feature = "postgres")]
Postgres(sqlx::Transaction<'static, Postgres>),
}
impl Transaction {
pub fn dialect(&self) -> Dialect {
match self.inner {
TxInner::Sqlite(_) => Dialect::Sqlite,
#[cfg(feature = "postgres")]
TxInner::Postgres(_) => Dialect::Postgres,
}
}
pub async fn commit(self) -> Result<(), DbError> {
match self.inner {
TxInner::Sqlite(tx) => Ok(tx.commit().await?),
#[cfg(feature = "postgres")]
TxInner::Postgres(tx) => Ok(tx.commit().await?),
}
}
pub async fn rollback(self) -> Result<(), DbError> {
match self.inner {
TxInner::Sqlite(tx) => Ok(tx.rollback().await?),
#[cfg(feature = "postgres")]
TxInner::Postgres(tx) => Ok(tx.rollback().await?),
}
}
pub async fn savepoint<T, F>(&mut self, work: F) -> crate::Result<T>
where
F: for<'t> FnOnce(
&'t mut Transaction,
) -> futures_util::future::BoxFuture<'t, crate::Result<T>>,
{
let name = format!("renox_savepoint_{}", self.savepoints + 1);
sql(format!("SAVEPOINT {name}")).execute(&mut *self).await?;
self.savepoints += 1;
let outcome = work(self).await;
self.savepoints -= 1;
match outcome {
Ok(value) => {
sql(format!("RELEASE SAVEPOINT {name}"))
.execute(&mut *self)
.await?;
Ok(value)
}
Err(err) => {
sql(format!("ROLLBACK TO SAVEPOINT {name}"))
.execute(&mut *self)
.await?;
sql(format!("RELEASE SAVEPOINT {name}"))
.execute(&mut *self)
.await?;
Err(err)
}
}
}
}
impl fmt::Debug for Transaction {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Transaction")
.field("dialect", &self.dialect())
.finish_non_exhaustive()
}
}
pub trait Executor<'c>: Send + executor::Sealed {
#[doc(hidden)]
fn into_conn(self) -> Conn<'c>;
}
mod executor {
pub trait Sealed {}
impl Sealed for &super::Db {}
impl Sealed for &mut super::Transaction {}
impl Sealed for super::Conn<'_> {}
}
#[doc(hidden)]
pub enum Conn<'c> {
Pool(&'c Db),
Tx(&'c mut Transaction),
}
impl<'c> Executor<'c> for &'c Db {
fn into_conn(self) -> Conn<'c> {
Conn::Pool(self)
}
}
impl<'c> Executor<'c> for &'c mut Transaction {
fn into_conn(self) -> Conn<'c> {
Conn::Tx(self)
}
}
impl<'c> Executor<'c> for Conn<'c> {
fn into_conn(self) -> Conn<'c> {
self
}
}
impl Conn<'_> {
fn key(&self) -> Option<Arc<Key>> {
match self {
Conn::Pool(db) => db.key.clone(),
Conn::Tx(tx) => tx.key.clone(),
}
}
pub(crate) fn reborrow(&mut self) -> Conn<'_> {
match self {
Conn::Pool(db) => Conn::Pool(db),
Conn::Tx(tx) => Conn::Tx(tx),
}
}
pub(crate) fn dialect(&self) -> Dialect {
match self {
Conn::Pool(db) => db.dialect(),
Conn::Tx(tx) => tx.dialect(),
}
}
}
macro_rules! dispatch {
($conn:expr, |$build:ident, $exec:ident| $body:expr) => {
match $conn {
Conn::Pool(db) => match &db.pool {
Pool::Sqlite(pool) => {
#[allow(unused_variables)]
let $build = sqlite_query;
let $exec = pool;
($body).map_err(DbError::from)
}
#[cfg(feature = "postgres")]
Pool::Postgres(pool) => {
#[allow(unused_variables)]
let $build = postgres_query;
let $exec = pool;
($body).map_err(DbError::from)
}
},
Conn::Tx(tx) => match &mut tx.inner {
TxInner::Sqlite(tx) => {
#[allow(unused_variables)]
let $build = sqlite_query;
let $exec = &mut **tx;
($body).map_err(DbError::from)
}
#[cfg(feature = "postgres")]
TxInner::Postgres(tx) => {
#[allow(unused_variables)]
let $build = postgres_query;
let $exec = &mut **tx;
($body).map_err(DbError::from)
}
},
}
};
}
fn seal(key: Option<&Key>, args: Vec<DbValue>) -> Result<Vec<DbValue>, DbError> {
args.into_iter()
.map(|value| match value {
DbValue::Encrypted(plain) => match key {
Some(key) => Ok(DbValue::Text(super::encrypted::seal(key, &plain.0))),
None => Err(DbError::from(sqlx::Error::Encode(
"an Encrypted value needs the app's Db (it has the APP_KEY); \
this Db was made outside App"
.into(),
))),
},
other => Ok(other),
})
.collect()
}
fn sqlite_query(
sql: String,
args: Vec<DbValue>,
) -> sqlx::query::Query<'static, Sqlite, SqliteArguments> {
args.into_iter().fold(
sqlx::query(AssertSqlSafe(sql)),
|query, value| match value.for_sqlite() {
DbValue::Integer(v) => query.bind(v),
DbValue::Real(v) => query.bind(v),
DbValue::Text(v) => query.bind(v),
DbValue::Blob(v) => query.bind(v),
_ => query.bind(None::<i64>),
},
)
}
#[cfg(feature = "postgres")]
fn postgres_query(
sql: String,
args: Vec<DbValue>,
) -> sqlx::query::Query<'static, Postgres, PgArguments> {
let sql = numbered_placeholders(&sql).into_owned();
args.into_iter().fold(
sqlx::query(AssertSqlSafe(sql)),
|query, value| match value {
DbValue::Null => query.bind(UntypedNull),
DbValue::Integer(v) => query.bind(v),
DbValue::Real(v) => query.bind(v),
DbValue::Text(v) => query.bind(v),
DbValue::Blob(v) => query.bind(v),
DbValue::Bool(v) => query.bind(v),
DbValue::DateTime(v) => query.bind(v),
DbValue::NaiveDateTime(v) => query.bind(v),
DbValue::Date(v) => query.bind(v),
DbValue::Time(v) => query.bind(v),
DbValue::Json(v) => query.bind(sqlx::types::Json(v)),
#[cfg(feature = "uuid")]
DbValue::Uuid(v) => query.bind(v),
DbValue::Encrypted(_) => query.bind(UntypedNull),
},
)
}
#[cfg(feature = "postgres")]
struct UntypedNull;
#[cfg(feature = "postgres")]
impl sqlx::Type<Postgres> for UntypedNull {
fn type_info() -> sqlx::postgres::PgTypeInfo {
sqlx::postgres::PgTypeInfo::with_oid(sqlx::postgres::types::Oid(0))
}
}
#[cfg(feature = "postgres")]
impl sqlx::Encode<'_, Postgres> for UntypedNull {
fn encode_by_ref(
&self,
_buf: &mut sqlx::postgres::PgArgumentBuffer,
) -> Result<sqlx::encode::IsNull, sqlx::error::BoxDynError> {
Ok(sqlx::encode::IsNull::Yes)
}
}
#[cfg(any(feature = "postgres", test))]
pub(crate) fn numbered_placeholders(sql: &str) -> Cow<'_, str> {
if !sql.contains('?') {
return Cow::Borrowed(sql);
}
let bytes = sql.as_bytes();
let ident = |b: u8| b.is_ascii_alphanumeric() || b == b'_' || b == b'$';
let mut out = String::with_capacity(sql.len() + 8);
let mut n = 0;
let mut copied = 0; let mut i = 0;
while i < bytes.len() {
let at = i;
let end = match bytes[i] {
quote @ (b'\'' | b'"') => {
let escapes = quote == b'\''
&& i > 0
&& matches!(bytes[i - 1], b'E' | b'e')
&& (i < 2 || !ident(bytes[i - 2]));
let mut j = i + 1;
while j < bytes.len() && bytes[j] != quote {
j += if escapes && bytes[j] == b'\\' { 2 } else { 1 };
}
j + 1
}
b'-' if bytes.get(i + 1) == Some(&b'-') => {
sql[i..].find('\n').map_or(bytes.len(), |nl| i + nl + 1)
}
b'/' if bytes.get(i + 1) == Some(&b'*') => sql[i + 2..]
.find("*/")
.map_or(bytes.len(), |e| i + 2 + e + 2),
b'$' if i == 0 || !ident(bytes[i - 1]) => {
let tag_len = bytes[i + 1..]
.iter()
.take_while(|b| b.is_ascii_alphanumeric() || **b == b'_')
.count();
let starts_with_digit = bytes.get(i + 1).is_some_and(u8::is_ascii_digit);
if !starts_with_digit && bytes.get(i + 1 + tag_len) == Some(&b'$') {
let tag = &sql[i..i + tag_len + 2];
sql[i + tag.len()..]
.find(tag)
.map_or(bytes.len(), |e| i + tag.len() + e + tag.len())
} else {
i + 1
}
}
b'?' => {
n += 1;
out.push_str(&sql[copied..i]);
out.push_str(&format!("${n}"));
copied = i + 1;
i + 1
}
_ => i + 1,
};
i = end.min(bytes.len()).max(at + 1);
}
out.push_str(&sql[copied..]);
Cow::Owned(out)
}
pub fn sql(sql: impl Into<String>) -> Sql {
Sql {
sql: sql.into(),
args: Vec::new(),
}
}
#[derive(Debug, Clone)]
#[must_use = "a statement does nothing until it is run"]
pub struct Sql {
sql: String,
args: Vec<DbValue>,
}
impl Sql {
pub fn bind(mut self, value: impl ToDbValue) -> Self {
self.args.push(value.to_db_value());
self
}
pub fn bind_all(mut self, values: impl IntoIterator<Item = DbValue>) -> Self {
self.args.extend(values);
self
}
pub async fn fetch_all<'c>(self, db: impl Executor<'c>) -> Result<Vec<Row>, DbError> {
let Self { sql, args } = self;
super::query_log::record(&sql);
let conn = db.into_conn();
let key = conn.key();
let args = seal(key.as_deref(), args)?;
let rows: Vec<Row> = dispatch!(conn, |build, exec| build(sql, args)
.fetch_all(exec)
.await
.map(|rows| rows.into_iter().map(Row::from).collect()))?;
Ok(rows
.into_iter()
.map(|row| row.with_key(key.clone()))
.collect())
}
pub async fn fetch_optional<'c>(self, db: impl Executor<'c>) -> Result<Option<Row>, DbError> {
let Self { sql, args } = self;
super::query_log::record(&sql);
let conn = db.into_conn();
let key = conn.key();
let args = seal(key.as_deref(), args)?;
let row: Option<Row> = dispatch!(conn, |build, exec| build(sql, args)
.fetch_optional(exec)
.await
.map(|row| row.map(Row::from)))?;
Ok(row.map(|row| row.with_key(key)))
}
pub async fn fetch_one<'c>(self, db: impl Executor<'c>) -> Result<Row, DbError> {
self.fetch_optional(db)
.await?
.ok_or_else(|| DbError::from(sqlx::Error::RowNotFound))
}
pub async fn execute<'c>(self, db: impl Executor<'c>) -> Result<u64, DbError> {
let Self { sql, args } = self;
super::query_log::record(&sql);
let conn = db.into_conn();
let args = seal(conn.key().as_deref(), args)?;
dispatch!(conn, |build, exec| build(sql, args)
.execute(exec)
.await
.map(|done| done.rows_affected()))
}
pub async fn fetch_as<'c, T: super::FromRow>(
self,
db: impl Executor<'c>,
) -> Result<Vec<T>, DbError> {
self.fetch_all(db).await?.iter().map(T::from_row).collect()
}
pub async fn fetch_optional_as<'c, T: super::FromRow>(
self,
db: impl Executor<'c>,
) -> Result<Option<T>, DbError> {
self.fetch_optional(db)
.await?
.as_ref()
.map(T::from_row)
.transpose()
}
pub async fn fetch_one_as<'c, T: super::FromRow>(
self,
db: impl Executor<'c>,
) -> Result<T, DbError> {
T::from_row(&self.fetch_one(db).await?)
}
pub async fn scalar<'c, T: FromDb>(self, db: impl Executor<'c>) -> Result<T, DbError> {
self.fetch_one(db).await?.try_get(0)
}
pub async fn scalar_optional<'c, T: FromDb>(
self,
db: impl Executor<'c>,
) -> Result<Option<T>, DbError> {
self.fetch_optional(db)
.await?
.map(|row| row.try_get(0))
.transpose()
}
pub async fn scalars<'c, T: FromDb>(self, db: impl Executor<'c>) -> Result<Vec<T>, DbError> {
self.fetch_all(db)
.await?
.iter()
.map(|row| row.try_get(0))
.collect()
}
}
pub(crate) async fn script<'c>(db: impl Executor<'c>, sql: &str) -> Result<u64, DbError> {
let sql = sql.to_owned();
dispatch!(db.into_conn(), |build, exec| sqlx::raw_sql(AssertSqlSafe(
sql
))
.execute(exec)
.await
.map(|done| done.rows_affected()))
}
pub struct Row(pub(crate) RowInner, Option<Arc<Key>>);
pub(crate) enum RowInner {
Sqlite(SqliteRow),
#[cfg(feature = "postgres")]
Postgres(PgRow),
}
impl From<SqliteRow> for Row {
fn from(row: SqliteRow) -> Self {
Self(RowInner::Sqlite(row), None)
}
}
#[cfg(feature = "postgres")]
impl From<PgRow> for Row {
fn from(row: PgRow) -> Self {
Self(RowInner::Postgres(row), None)
}
}
impl Row {
fn with_key(mut self, key: Option<Arc<Key>>) -> Self {
self.1 = key;
self
}
pub fn try_get<T: FromDb>(&self, index: impl RowIndex) -> Result<T, DbError> {
super::encrypted::reading(self.1.as_ref(), || match &self.0 {
RowInner::Sqlite(row) => Ok(row.try_get(index)?),
#[cfg(feature = "postgres")]
RowInner::Postgres(row) => Ok(row.try_get(index)?),
})
}
pub(crate) fn json(&self, column: &str) -> serde_json::Value {
use serde_json::Value;
if let Ok(v) = self.try_get::<Option<i64>>(column) {
return v.map_or(Value::Null, Value::from);
}
if let Ok(v) = self.try_get::<Option<i32>>(column) {
return v.map_or(Value::Null, Value::from);
}
if let Ok(v) = self.try_get::<Option<i16>>(column) {
return v.map_or(Value::Null, Value::from);
}
if let Ok(v) = self.try_get::<Option<f64>>(column) {
return v.map_or(Value::Null, Value::from);
}
if let Ok(v) = self.try_get::<Option<f32>>(column) {
return v.map_or(Value::Null, |n| Value::from(f64::from(n)));
}
if let Ok(v) = self.try_get::<Option<bool>>(column) {
return v.map_or(Value::Null, Value::from);
}
if let Ok(v) = self.try_get::<Option<chrono::DateTime<chrono::Utc>>>(column) {
return v.map_or(Value::Null, |d| Value::from(d.to_rfc3339()));
}
if let Ok(v) = self.try_get::<Option<chrono::NaiveDateTime>>(column) {
return v.map_or(Value::Null, |d| {
Value::from(d.format("%Y-%m-%dT%H:%M:%S").to_string())
});
}
if let Ok(v) = self.try_get::<Option<chrono::NaiveDate>>(column) {
return v.map_or(Value::Null, |d| Value::from(d.to_string()));
}
if let Ok(v) = self.try_get::<Option<String>>(column) {
return v.map_or(Value::Null, Value::from);
}
if let Ok(v) = self.try_get::<Option<chrono::NaiveTime>>(column) {
return v.map_or(Value::Null, |t| Value::from(t.to_string()));
}
if let Ok(v) = self.try_get::<Option<serde_json::Value>>(column) {
return v.unwrap_or(Value::Null);
}
Value::Null
}
pub fn columns(&self) -> Vec<&str> {
match &self.0 {
RowInner::Sqlite(row) => row.columns().iter().map(Column::name).collect(),
#[cfg(feature = "postgres")]
RowInner::Postgres(row) => row.columns().iter().map(Column::name).collect(),
}
}
pub fn sqlite(&self) -> Option<&SqliteRow> {
match &self.0 {
RowInner::Sqlite(row) => Some(row),
#[cfg(feature = "postgres")]
RowInner::Postgres(_) => None,
}
}
#[cfg(feature = "postgres")]
pub fn postgres(&self) -> Option<&PgRow> {
match &self.0 {
RowInner::Postgres(row) => Some(row),
RowInner::Sqlite(_) => None,
}
}
}
impl fmt::Debug for Row {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Row")
.field("columns", &self.columns())
.finish_non_exhaustive()
}
}
pub trait FromDb: for<'r> sqlx::Decode<'r, Sqlite> + sqlx::Type<Sqlite> + bounds::Postgres {}
impl<T> FromDb for T where
T: for<'r> sqlx::Decode<'r, Sqlite> + sqlx::Type<Sqlite> + bounds::Postgres
{
}
pub trait RowIndex: sqlx::ColumnIndex<SqliteRow> + bounds::PostgresIndex {}
impl<T> RowIndex for T where T: sqlx::ColumnIndex<SqliteRow> + bounds::PostgresIndex {}
#[doc(hidden)]
pub mod bounds {
#[cfg(feature = "postgres")]
mod on {
use sqlx::postgres::PgRow;
pub trait Postgres:
for<'r> sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>
{
}
impl<T> Postgres for T where T: for<'r> sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>
{}
pub trait PostgresIndex: sqlx::ColumnIndex<PgRow> {}
impl<T: sqlx::ColumnIndex<PgRow>> PostgresIndex for T {}
}
#[cfg(not(feature = "postgres"))]
mod on {
pub trait Postgres {}
impl<T: ?Sized> Postgres for T {}
pub trait PostgresIndex {}
impl<T: ?Sized> PostgresIndex for T {}
}
pub use on::{Postgres, PostgresIndex};
}
#[cfg(test)]
mod tests {
use super::{numbered_placeholders, seal};
use crate::db::DbValue;
#[test]
fn numbers_placeholders_outside_quotes_and_comments() {
assert_eq!(
numbered_placeholders("SELECT * FROM t WHERE a = ? AND b IN (?, ?)"),
"SELECT * FROM t WHERE a = $1 AND b IN ($2, $3)"
);
assert_eq!(
numbered_placeholders("SELECT '?', \"a?\" -- why?\nFROM t /* ? */ WHERE x = ?"),
"SELECT '?', \"a?\" -- why?\nFROM t /* ? */ WHERE x = $1"
);
assert_eq!(
numbered_placeholders("SELECT 'it''s ?' WHERE y = ?"),
"SELECT 'it''s ?' WHERE y = $1"
);
assert_eq!(numbered_placeholders("SELECT 1"), "SELECT 1");
}
#[test]
fn skips_dollar_quoted_and_escaped_strings() {
assert_eq!(
numbered_placeholders("SELECT $$why?$$, $fn$ a ? b $fn$ WHERE x = ?"),
"SELECT $$why?$$, $fn$ a ? b $fn$ WHERE x = $1"
);
assert_eq!(
numbered_placeholders(r"SELECT E'it\'s ?', e'\\' WHERE y = ? AND z = ?"),
r"SELECT E'it\'s ?', e'\\' WHERE y = $1 AND z = $2"
);
assert_eq!(
numbered_placeholders(r"SELECT '\', a$b$ FROM t WHERE c = ?"),
r"SELECT '\', a$b$ FROM t WHERE c = $1"
);
assert_eq!(
numbered_placeholders("SELECT 'é?' WHERE ü = ?"),
"SELECT 'é?' WHERE ü = $1"
);
}
#[test]
fn placeholders_after_unfinished_quotes_stay() {
assert_eq!(numbered_placeholders("SELECT 'open ?"), "SELECT 'open ?");
assert_eq!(numbered_placeholders("SELECT 1 -- ?"), "SELECT 1 -- ?");
assert_eq!(numbered_placeholders("SELECT 1 /* ?"), "SELECT 1 /* ?");
assert_eq!(numbered_placeholders("SELECT $tag$ ?"), "SELECT $tag$ ?");
assert_eq!(numbered_placeholders("SELECT ? || $"), "SELECT $1 || $");
assert!(numbered_placeholders("SELECT $1, ?").ends_with(", $1"));
assert_eq!(numbered_placeholders("?"), "$1");
}
#[tokio::test]
async fn retrying_follows_the_error_kind() {
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use crate::db::DbError;
use crate::db::error::fake;
fn busy() -> crate::Error {
crate::Error::from(DbError::from(fake::coded(Some("5"))))
}
let db = super::super::connect(&crate::Config::default())
.await
.unwrap();
let tries = Arc::new(AtomicU32::new(0));
let counter = tries.clone();
let value = db
.retrying(3, move || {
let counter = counter.clone();
async move {
match counter.fetch_add(1, Ordering::SeqCst) {
0 | 1 => Err(busy()),
_ => Ok("done"),
}
}
})
.await
.unwrap();
assert_eq!((value, tries.load(Ordering::SeqCst)), ("done", 3));
let tries = Arc::new(AtomicU32::new(0));
let counter = tries.clone();
let err = db
.retrying(2, move || {
counter.fetch_add(1, Ordering::SeqCst);
async { Err::<(), _>(busy()) }
})
.await
.unwrap_err();
assert!(err.is_retryable());
assert_eq!(
tries.load(Ordering::SeqCst),
2,
"gives up at the last attempt"
);
let tries = Arc::new(AtomicU32::new(0));
let counter = tries.clone();
let err = db
.retrying(5, move || {
counter.fetch_add(1, Ordering::SeqCst);
async { Err::<(), _>(crate::Error::NotFound) }
})
.await
.unwrap_err();
assert!(matches!(err, crate::Error::NotFound));
assert_eq!(tries.load(Ordering::SeqCst), 1, "not retried");
let tries = Arc::new(AtomicU32::new(0));
let counter = tries.clone();
let value = db
.transaction_retrying(3, move |tx| {
let first = counter.fetch_add(1, Ordering::SeqCst) == 0;
Box::pin(async move {
let one: i64 = super::sql("SELECT CAST(1 AS BIGINT)").scalar(tx).await?;
if first { Err(busy()) } else { Ok(one) }
})
})
.await
.unwrap();
assert_eq!((value, tries.load(Ordering::SeqCst)), (1, 2));
}
#[test]
fn sealing_needs_the_apps_key() {
let err = seal(
None,
vec![DbValue::Encrypted(super::super::encrypted::Unsealed(
"x".into(),
))],
)
.unwrap_err();
assert!(
err.to_string()
.contains("an Encrypted value needs the app's Db"),
"{err}"
);
assert!(matches!(
seal(None, vec![DbValue::Integer(1)]).unwrap()[..],
[DbValue::Integer(1)]
));
}
}