use std::future::Future;
use std::sync::Arc;
use aro_core::error::RepoError;
use aro_core::repository::{BoxFuture, TransactionManager};
use aro_web::dep::Dep;
use crate::error::from_fletch_err;
use crate::repo::{FletchDatabase, FletchEntity, FletchTransactionalRepository};
pub(crate) struct FletchTransactionState<DB: FletchDatabase> {
pub(crate) tx: tokio::sync::Mutex<Option<fletch_orm::Transaction<'static, DB>>>,
}
pub struct FletchTransactionContext<DB: FletchDatabase> {
state: Arc<FletchTransactionState<DB>>,
}
impl<DB: FletchDatabase> Clone for FletchTransactionContext<DB> {
fn clone(&self) -> Self {
Self {
state: Arc::clone(&self.state),
}
}
}
impl<DB: FletchDatabase> FletchTransactionContext<DB> {
pub fn repo<T>(&self) -> FletchTransactionalRepository<T, DB>
where
T: FletchEntity<DB>,
{
FletchTransactionalRepository::new(Arc::clone(&self.state))
}
}
async fn begin_transaction<DB: FletchDatabase>(
pool: &fletch_orm::Pool<DB>,
) -> Result<FletchTransactionContext<DB>, RepoError> {
let tx = pool.begin().await.map_err(from_fletch_err)?;
Ok(FletchTransactionContext {
state: Arc::new(FletchTransactionState {
tx: tokio::sync::Mutex::new(Some(tx)),
}),
})
}
async fn finish_transaction<DB: FletchDatabase, R: Send>(
state: Arc<FletchTransactionState<DB>>,
result: Result<R, RepoError>,
) -> Result<R, RepoError> {
let mut guard = state.tx.lock().await;
let tx = guard
.take()
.ok_or_else(|| RepoError::Unknown("transaction is no longer active".into()))?;
drop(guard);
match result {
Ok(value) => {
tx.commit().await.map_err(from_fletch_err)?;
Ok(value)
}
Err(err) => {
if let Err(rollback_err) = tx.rollback().await {
return Err(from_fletch_err(rollback_err));
}
Err(err)
}
}
}
pub struct FletchTransactionManager<DB: FletchDatabase> {
pool: fletch_orm::Pool<DB>,
}
impl<DB: FletchDatabase> FletchTransactionManager<DB> {
pub fn new(pool: fletch_orm::Pool<DB>) -> Self {
Self { pool }
}
pub fn from_dep(dep: &Dep<fletch_orm::Pool<DB>>) -> Self {
Self::new((**dep).clone())
}
pub fn pool(&self) -> &fletch_orm::Pool<DB> {
&self.pool
}
pub async fn transaction<F, Fut, R>(&self, f: F) -> Result<R, RepoError>
where
F: FnOnce(FletchTransactionContext<DB>) -> Fut,
Fut: Future<Output = Result<R, RepoError>>,
R: Send,
{
let ctx = begin_transaction(&self.pool).await?;
let state = Arc::clone(&ctx.state);
let result = f(ctx).await;
finish_transaction(state, result).await
}
pub fn transaction_boxed<F, Fut, R>(&self, f: F) -> BoxFuture<'_, Result<R, RepoError>>
where
F: FnOnce(FletchTransactionContext<DB>) -> Fut + Send + 'static,
Fut: Future<Output = Result<R, RepoError>> + Send + 'static,
R: Send + 'static,
{
let pool = self.pool.clone();
Box::pin(async move {
let manager = FletchTransactionManager { pool };
manager.transaction(f).await
})
}
}
impl<DB: FletchDatabase> TransactionManager for FletchTransactionManager<DB> {
type Context = FletchTransactionContext<DB>;
fn transaction<F, Fut, R>(&self, f: F) -> BoxFuture<'_, Result<R, RepoError>>
where
F: FnOnce(FletchTransactionContext<DB>) -> Fut + Send + 'static,
Fut: Future<Output = Result<R, RepoError>> + Send + 'static,
R: Send + 'static,
{
self.transaction_boxed(f)
}
}