use crate::{ConnAcquireErr, ConnectOptions, DbErr, RuntimeErr};
use std::{sync::Arc, time::Duration};
pub(crate) type BeforeAcquireFn<DB> = Arc<
dyn for<'c> Fn(
&'c mut <DB as sqlx::Database>::Connection,
sqlx::pool::PoolConnectionMetadata,
) -> futures_util::future::BoxFuture<'c, Result<bool, sqlx::Error>>
+ Send
+ Sync,
>;
pub fn sqlx_error_to_exec_err(err: sqlx::Error) -> DbErr {
DbErr::Exec(RuntimeErr::SqlxError(err.into()))
}
pub fn sqlx_error_to_query_err(err: sqlx::Error) -> DbErr {
DbErr::Query(RuntimeErr::SqlxError(err.into()))
}
pub fn sqlx_error_to_conn_err(err: sqlx::Error) -> DbErr {
DbErr::Conn(RuntimeErr::SqlxError(err.into()))
}
pub fn sqlx_map_err_ignore_not_found<T: std::fmt::Debug>(
err: Result<Option<T>, sqlx::Error>,
) -> Result<Option<T>, DbErr> {
if let Err(sqlx::Error::RowNotFound) = err {
Ok(None)
} else {
err.map_err(sqlx_error_to_query_err)
}
}
pub fn sqlx_conn_acquire_err(sqlx_err: sqlx::Error) -> DbErr {
match sqlx_err {
sqlx::Error::PoolTimedOut => DbErr::ConnectionAcquire(ConnAcquireErr::Timeout),
sqlx::Error::PoolClosed => DbErr::ConnectionAcquire(ConnAcquireErr::ConnectionClosed),
_ => DbErr::Conn(RuntimeErr::SqlxError(sqlx_err.into())),
}
}
impl ConnectOptions {
pub fn sqlx_pool_options<DB>(self) -> sqlx::pool::PoolOptions<DB>
where
DB: sqlx::Database,
{
let mut opt = sqlx::pool::PoolOptions::new();
if let Some(max_connections) = self.max_connections {
opt = opt.max_connections(max_connections);
}
if let Some(min_connections) = self.min_connections {
opt = opt.min_connections(min_connections);
}
if let Some(connect_timeout) = self.connect_timeout {
opt = opt.acquire_timeout(connect_timeout);
}
if let Some(idle_timeout) = self.idle_timeout {
opt = opt.idle_timeout(idle_timeout);
}
if let Some(acquire_timeout) = self.acquire_timeout {
opt = opt.acquire_timeout(acquire_timeout);
}
if let Some(max_lifetime) = self.max_lifetime {
opt = opt.max_lifetime(max_lifetime);
}
opt = opt.test_before_acquire(self.test_before_acquire);
opt
}
pub(crate) fn apply_before_acquire<DB>(
mut opt: sqlx::pool::PoolOptions<DB>,
ping_after_idle: Option<Duration>,
user_cb: Option<BeforeAcquireFn<DB>>,
) -> sqlx::pool::PoolOptions<DB>
where
DB: sqlx::Database,
{
use sqlx::Connection;
if ping_after_idle.is_none() && user_cb.is_none() {
return opt;
}
if ping_after_idle.is_some() {
opt = opt.test_before_acquire(false);
}
opt.before_acquire(move |conn, meta| {
let user_cb = user_cb.clone();
Box::pin(async move {
if let Some(threshold) = ping_after_idle {
if meta.idle_for >= threshold {
conn.ping().await?;
}
}
match user_cb {
Some(user_cb) => user_cb(conn, meta).await,
None => Ok(true),
}
})
})
}
}
#[cfg(all(test, feature = "sqlx-postgres"))]
mod tests {
use crate::ConnectOptions;
use sqlx::Connection;
use std::time::Duration;
#[test]
fn idle_shorthand_disables_test_before_acquire() {
let mut opt = ConnectOptions::new("postgres://localhost/db");
assert!(opt.get_test_before_acquire());
assert_eq!(opt.get_test_before_acquire_if_idle_for(), None);
opt.test_before_acquire_if_idle_for(Duration::from_secs(30));
assert!(!opt.get_test_before_acquire());
assert_eq!(
opt.get_test_before_acquire_if_idle_for(),
Some(Duration::from_secs(30))
);
}
#[test]
fn compose_shorthand_and_user_callback() {
let mut opt = ConnectOptions::new("postgres://localhost/db");
opt.test_before_acquire_if_idle_for(Duration::from_secs(30))
.map_sqlx_postgres_before_acquire(|conn, _meta| {
Box::pin(async move {
conn.ping().await?;
Ok(true)
})
});
let pool_opts = ConnectOptions::apply_before_acquire::<sqlx::Postgres>(
sqlx::pool::PoolOptions::new(),
opt.get_test_before_acquire_if_idle_for(),
opt.pg_before_acquire_fn.clone(),
);
let _ = pool_opts;
}
#[test]
fn apply_before_acquire_noop_when_unset() {
let opts = ConnectOptions::apply_before_acquire::<sqlx::Postgres>(
sqlx::pool::PoolOptions::new(),
None,
None,
);
let _ = opts;
}
}