use std::marker::PhantomData;
use actix_web::{dev::Extensions, FromRequest, HttpMessage, ResponseError};
use futures_core::future::LocalBoxFuture;
use sqlx::Transaction;
use crate::{
error::Error,
slot::{Lease, Slot},
};
#[derive(Debug)]
pub struct Tx<DB: sqlx::Database, E = Error>(Lease<sqlx::Transaction<'static, DB>>, PhantomData<E>);
impl<DB: sqlx::Database, E> Tx<DB, E> {
pub async fn commit(self) -> Result<(), sqlx::Error> {
self.0.steal().commit().await
}
}
impl<DB: sqlx::Database, E> AsRef<sqlx::Transaction<'static, DB>> for Tx<DB, E> {
fn as_ref(&self) -> &sqlx::Transaction<'static, DB> {
&self.0
}
}
impl<DB: sqlx::Database, E> AsMut<sqlx::Transaction<'static, DB>> for Tx<DB, E> {
fn as_mut(&mut self) -> &mut sqlx::Transaction<'static, DB> {
&mut self.0
}
}
impl<DB: sqlx::Database, E> std::ops::Deref for Tx<DB, E> {
type Target = sqlx::Transaction<'static, DB>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<DB: sqlx::Database, E> std::ops::DerefMut for Tx<DB, E> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.0
}
}
impl<DB: sqlx::Database, E> FromRequest for Tx<DB, E>
where
E: From<Error> + ResponseError + 'static,
{
type Error = E;
type Future = LocalBoxFuture<'static, Result<Self, Self::Error>>;
#[inline]
fn from_request(req: &actix_web::HttpRequest, _: &mut actix_web::dev::Payload) -> Self::Future {
let req = req.clone();
Box::pin(async move {
let mut ext = req
.extensions_mut()
.remove::<Lazy<DB>>()
.ok_or(Error::MissingExtension)?;
let tx = ext.get_or_begin().await?;
Ok(Self(tx, PhantomData))
})
}
}
pub(crate) struct TxSlot<DB: sqlx::Database>(Slot<Option<Slot<Transaction<'static, DB>>>>);
impl<DB: sqlx::Database> TxSlot<DB> {
pub(crate) fn bind(extensions: &mut Extensions, pool: &sqlx::Pool<DB>) -> Self {
let (slot, tx) = Slot::new_leased(None);
extensions.insert(Lazy {
pool: pool.clone(),
tx,
});
Self(slot)
}
pub(crate) async fn commit(self) -> Result<(), sqlx::Error> {
if let Some(tx) = self.0.into_inner().flatten().and_then(Slot::into_inner) {
tx.commit().await?;
}
Ok(())
}
}
struct Lazy<DB: sqlx::Database> {
pool: sqlx::Pool<DB>,
tx: Lease<Option<Slot<Transaction<'static, DB>>>>,
}
impl<DB: sqlx::Database> Lazy<DB> {
async fn get_or_begin(&mut self) -> Result<Lease<Transaction<'static, DB>>, Error> {
let tx = if let Some(tx) = self.tx.as_mut() {
tx
} else {
let tx = self.pool.begin().await?;
self.tx.insert(Slot::new(tx))
};
tx.lease().ok_or(Error::OverlappingExtractors)
}
}
#[cfg(any(
feature = "any",
feature = "mssql",
feature = "mysql",
feature = "postgres",
feature = "sqlite"
))]
mod sqlx_impls {
use std::fmt::Debug;
use futures_core::{future::BoxFuture, stream::BoxStream};
macro_rules! impl_executor {
($db:path) => {
impl<'c, E: Debug + Send> sqlx::Executor<'c> for &'c mut super::Tx<$db, E> {
type Database = $db;
#[allow(clippy::type_complexity)]
fn fetch_many<'e, 'q: 'e, Q: 'q>(
self,
query: Q,
) -> BoxStream<
'e,
Result<
sqlx::Either<
<Self::Database as sqlx::Database>::QueryResult,
<Self::Database as sqlx::Database>::Row,
>,
sqlx::Error,
>,
>
where
'c: 'e,
Q: sqlx::Execute<'q, Self::Database>,
{
(&mut **self).fetch_many(query)
}
fn fetch_optional<'e, 'q: 'e, Q: 'q>(
self,
query: Q,
) -> BoxFuture<
'e,
Result<Option<<Self::Database as sqlx::Database>::Row>, sqlx::Error>,
>
where
'c: 'e,
Q: sqlx::Execute<'q, Self::Database>,
{
(&mut **self).fetch_optional(query)
}
fn prepare_with<'e, 'q: 'e>(
self,
sql: &'q str,
parameters: &'e [<Self::Database as sqlx::Database>::TypeInfo],
) -> BoxFuture<
'e,
Result<
<Self::Database as sqlx::database::HasStatement<'q>>::Statement,
sqlx::Error,
>,
>
where
'c: 'e,
{
(&mut **self).prepare_with(sql, parameters)
}
fn describe<'e, 'q: 'e>(
self,
sql: &'q str,
) -> BoxFuture<'e, Result<sqlx::Describe<Self::Database>, sqlx::Error>>
where
'c: 'e,
{
(&mut **self).describe(sql)
}
}
};
}
#[cfg(feature = "any")]
impl_executor!(sqlx::Any);
#[cfg(feature = "mssql")]
impl_executor!(sqlx::Mssql);
#[cfg(feature = "mysql")]
impl_executor!(sqlx::MySql);
#[cfg(feature = "postgres")]
impl_executor!(sqlx::Postgres);
#[cfg(feature = "sqlite")]
impl_executor!(sqlx::Sqlite);
}