use std::sync::{Arc, Mutex};
use saddle_observability::ActiveCall;
use tokio::{sync::mpsc, task::JoinHandle};
use crate::{
Result,
error::{cleanup_failed, transaction_cancelled, transaction_rollback_failed},
};
pub(crate) struct RollbackJob {
pub(crate) inner: sqlx::Transaction<'static, sqlx::MySql>,
pub(crate) boundary: ActiveCall,
pub(crate) rollback: ActiveCall,
}
impl RollbackJob {
async fn run(self) {
match self.inner.rollback().await {
Ok(()) => {
let error = transaction_cancelled();
self.rollback.succeed();
self.boundary.fail(&error);
}
Err(_) => {
let error = transaction_rollback_failed();
self.rollback.fail(&error);
self.boundary.fail(&error);
}
}
}
pub(crate) fn fail_without_worker(self) {
let error = transaction_rollback_failed();
self.rollback.fail(&error);
self.boundary.fail(&error);
drop(self.inner);
}
}
pub(crate) struct CleanupCoordinator {
sender: Mutex<Option<mpsc::UnboundedSender<RollbackJob>>>,
worker: Mutex<Option<JoinHandle<()>>>,
}
impl CleanupCoordinator {
pub(crate) fn start() -> Arc<Self> {
let (sender, mut receiver) = mpsc::unbounded_channel::<RollbackJob>();
let worker = tokio::spawn(async move {
while let Some(job) = receiver.recv().await {
job.run().await;
}
});
Arc::new(Self {
sender: Mutex::new(Some(sender)),
worker: Mutex::new(Some(worker)),
})
}
pub(crate) fn transaction_sender(&self) -> Option<mpsc::UnboundedSender<RollbackJob>> {
self.sender.lock().unwrap().as_ref().cloned()
}
pub(crate) async fn shutdown(&self) -> Result<()> {
self.sender.lock().unwrap().take();
let worker = self.worker.lock().unwrap().take();
if let Some(worker) = worker {
worker.await.map_err(|_| cleanup_failed())?;
}
Ok(())
}
}