use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TransactionPhase {
BeforeCommit,
AfterCommit,
AfterRollback,
AfterCompletion,
}
#[async_trait::async_trait]
pub trait TransactionSynchronization: Send + Sync {
async fn before_commit(&self, _tx_name: &str) -> Result<(), ()> {
Ok(())
}
async fn before_completion(&self, _tx_name: &str) {}
async fn after_commit(&self, _tx_name: &str) {}
async fn after_rollback(&self, _tx_name: &str) {}
async fn after_completion(&self, _tx_name: &str, _committed: bool) {}
fn name(&self) -> &'static str {
std::any::type_name::<Self>()
}
}
pub struct PhaseListener {
phase: TransactionPhase,
callback: Arc<dyn Fn(&str) + Send + Sync>,
}
impl PhaseListener {
pub fn new(phase: TransactionPhase, callback: impl Fn(&str) + Send + Sync + 'static) -> Self {
Self {
phase,
callback: Arc::new(callback),
}
}
}
#[async_trait::async_trait]
impl TransactionSynchronization for PhaseListener {
async fn before_commit(&self, tx_name: &str) -> Result<(), ()> {
if self.phase == TransactionPhase::BeforeCommit {
(self.callback)(tx_name);
}
Ok(())
}
async fn after_commit(&self, tx_name: &str) {
if self.phase == TransactionPhase::AfterCommit {
(self.callback)(tx_name);
}
}
async fn after_rollback(&self, tx_name: &str) {
if self.phase == TransactionPhase::AfterRollback {
(self.callback)(tx_name);
}
}
async fn after_completion(&self, tx_name: &str, _committed: bool) {
if self.phase == TransactionPhase::AfterCompletion {
(self.callback)(tx_name);
}
}
fn name(&self) -> &'static str {
match self.phase {
TransactionPhase::BeforeCommit => "PhaseListener::BeforeCommit",
TransactionPhase::AfterCommit => "PhaseListener::AfterCommit",
TransactionPhase::AfterRollback => "PhaseListener::AfterRollback",
TransactionPhase::AfterCompletion => "PhaseListener::AfterCompletion",
}
}
}
pub struct SynchronizationRegistry {
synchronizations: Arc<RwLock<Vec<Arc<dyn TransactionSynchronization>>>>,
}
impl SynchronizationRegistry {
pub fn new() -> Self {
Self {
synchronizations: Arc::new(RwLock::new(Vec::new())),
}
}
pub fn register(&self, sync: impl TransactionSynchronization + 'static) {
if let Ok(mut guard) = self.synchronizations.try_write() {
guard.push(Arc::new(sync));
}
}
pub async fn count(&self) -> usize {
self.synchronizations.read().await.len()
}
pub async fn clear(&self) {
self.synchronizations.write().await.clear();
}
pub async fn trigger_before_commit(&self, tx_name: &str) -> Result<(), ()> {
let syncs = self.synchronizations.read().await;
for sync in syncs.iter() {
sync.before_commit(tx_name).await?;
}
Ok(())
}
pub async fn trigger_before_completion(&self, tx_name: &str) {
let syncs = self.synchronizations.read().await;
for sync in syncs.iter() {
sync.before_completion(tx_name).await;
}
}
pub async fn trigger_after_commit(&self, tx_name: &str) {
let syncs = self.synchronizations.read().await;
for sync in syncs.iter() {
sync.after_commit(tx_name).await;
}
}
pub async fn trigger_after_rollback(&self, tx_name: &str) {
let syncs = self.synchronizations.read().await;
for sync in syncs.iter() {
sync.after_rollback(tx_name).await;
}
}
pub async fn trigger_after_completion(&self, tx_name: &str, committed: bool) {
let syncs = self.synchronizations.read().await;
for sync in syncs.iter() {
sync.after_completion(tx_name, committed).await;
}
}
}
impl Default for SynchronizationRegistry {
fn default() -> Self {
Self::new()
}
}
impl Clone for SynchronizationRegistry {
fn clone(&self) -> Self {
Self {
synchronizations: self.synchronizations.clone(),
}
}
}
pub struct LoggingSynchronization;
#[async_trait::async_trait]
impl TransactionSynchronization for LoggingSynchronization {
async fn before_commit(&self, tx_name: &str) -> Result<(), ()> {
println!("[TxSync] before_commit: {}", tx_name);
Ok(())
}
async fn after_commit(&self, tx_name: &str) {
println!("[TxSync] after_commit: {}", tx_name);
}
async fn after_rollback(&self, tx_name: &str) {
println!("[TxSync] after_rollback: {}", tx_name);
}
async fn after_completion(&self, tx_name: &str, committed: bool) {
println!(
"[TxSync] after_completion: {} ({})",
tx_name,
if committed {
"committed"
} else {
"rolled back"
}
);
}
fn name(&self) -> &'static str {
"LoggingSynchronization"
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
struct CountingSync {
before_commit: AtomicUsize,
after_commit: AtomicUsize,
after_rollback: AtomicUsize,
after_completion: AtomicUsize,
}
impl CountingSync {
fn new() -> Self {
Self {
before_commit: AtomicUsize::new(0),
after_commit: AtomicUsize::new(0),
after_rollback: AtomicUsize::new(0),
after_completion: AtomicUsize::new(0),
}
}
}
#[async_trait::async_trait]
impl TransactionSynchronization for CountingSync {
async fn before_commit(&self, _: &str) -> Result<(), ()> {
self.before_commit.fetch_add(1, Ordering::SeqCst);
Ok(())
}
async fn after_commit(&self, _: &str) {
self.after_commit.fetch_add(1, Ordering::SeqCst);
}
async fn after_rollback(&self, _: &str) {
self.after_rollback.fetch_add(1, Ordering::SeqCst);
}
async fn after_completion(&self, _: &str, _: bool) {
self.after_completion.fetch_add(1, Ordering::SeqCst);
}
}
struct W(Arc<CountingSync>);
#[async_trait::async_trait]
impl TransactionSynchronization for W {
async fn before_commit(&self, tx: &str) -> Result<(), ()> {
self.0.before_commit(tx).await
}
async fn after_commit(&self, tx: &str) {
self.0.after_commit(tx).await
}
async fn after_rollback(&self, tx: &str) {
self.0.after_rollback(tx).await
}
async fn after_completion(&self, tx: &str, c: bool) {
self.0.after_completion(tx, c).await
}
}
#[tokio::test]
async fn test_commit_lifecycle() {
let registry = SynchronizationRegistry::new();
let counter = Arc::new(CountingSync::new());
registry.register(W(counter.clone()));
registry.trigger_before_commit("tx1").await.unwrap();
registry.trigger_before_completion("tx1").await;
registry.trigger_after_commit("tx1").await;
registry.trigger_after_completion("tx1", true).await;
assert_eq!(counter.before_commit.load(Ordering::SeqCst), 1);
assert_eq!(counter.after_commit.load(Ordering::SeqCst), 1);
assert_eq!(counter.after_rollback.load(Ordering::SeqCst), 0);
assert_eq!(counter.after_completion.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_rollback_lifecycle() {
let registry = SynchronizationRegistry::new();
let counter = Arc::new(CountingSync::new());
registry.register(W(counter.clone()));
registry.trigger_before_completion("tx2").await;
registry.trigger_after_rollback("tx2").await;
registry.trigger_after_completion("tx2", false).await;
assert_eq!(counter.before_commit.load(Ordering::SeqCst), 0);
assert_eq!(counter.after_commit.load(Ordering::SeqCst), 0);
assert_eq!(counter.after_rollback.load(Ordering::SeqCst), 1);
assert_eq!(counter.after_completion.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_phase_listener() {
let fired = Arc::new(AtomicUsize::new(0));
let f = fired.clone();
let listener = PhaseListener::new(TransactionPhase::AfterCommit, move |_| {
f.fetch_add(1, Ordering::SeqCst);
});
let registry = SynchronizationRegistry::new();
registry.register(listener);
registry.trigger_after_commit("tx3").await;
registry.trigger_after_rollback("tx3").await;
assert_eq!(fired.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_before_commit_veto() {
struct VetoSync;
#[async_trait::async_trait]
impl TransactionSynchronization for VetoSync {
async fn before_commit(&self, _: &str) -> Result<(), ()> {
Err(())
}
}
let registry = SynchronizationRegistry::new();
registry.register(VetoSync);
assert!(registry.trigger_before_commit("tx4").await.is_err());
}
#[tokio::test]
async fn test_clear() {
let registry = SynchronizationRegistry::new();
registry.register(LoggingSynchronization);
assert_eq!(registry.count().await, 1);
registry.clear().await;
assert_eq!(registry.count().await, 0);
}
}