use std::time::Duration;
use sqlx::{Connection, PgConnection, PgPool};
use crate::error::FromDatabaseError;
#[derive(Clone, Debug)]
pub(crate) struct Session {
pool: PgPool,
statement_timeout: Duration,
}
impl Session {
pub(crate) fn new(pool: PgPool, statement_timeout: Duration) -> Self {
Self {
pool,
statement_timeout,
}
}
pub(crate) async fn run<T, F>(&self, op: F) -> Result<T, sqlx::Error>
where
F: AsyncFnOnce(&mut PgConnection) -> Result<T, sqlx::Error>,
{
let mut conn = self.pool.acquire().await?;
if self.statement_timeout.is_zero() {
return op(&mut conn).await;
}
let timeout_ms = i64::try_from(self.statement_timeout.as_millis())
.unwrap_or(i64::MAX)
.to_string();
let mut tx = conn.begin().await?;
sqlx::query_scalar!(
"SELECT set_config('statement_timeout', $1, true)",
timeout_ms
)
.fetch_one(&mut *tx)
.await?;
let result = op(&mut tx).await;
match result {
Ok(value) => {
tx.commit().await?;
Ok(value)
}
Err(err) => {
let _ = tx.rollback().await;
Err(err)
}
}
}
#[allow(
clippy::unused_self,
reason = "kept as a method for call-site symmetry with `run`"
)]
pub(crate) fn map_err<E: FromDatabaseError>(&self, err: sqlx::Error) -> E {
E::from_database_error(err)
}
}