use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use crate::manager::{TransactionDefinition, TransactionManager};
use crate::status::TransactionStatus;
use crate::{TransactionError, TransactionResult};
#[derive(Clone, Default)]
pub struct TransactionManagerRegistry {
managers: HashMap<String, Arc<dyn TransactionManager>>,
default_name: Option<String>,
}
impl std::fmt::Debug for TransactionManagerRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TransactionManagerRegistry")
.field("managers", &self.managers.keys().collect::<Vec<_>>())
.field("default_name", &self.default_name)
.finish()
}
}
impl TransactionManagerRegistry {
pub fn new() -> Self {
Self {
managers: HashMap::new(),
default_name: None,
}
}
pub fn register(&mut self, name: impl Into<String>, manager: Arc<dyn TransactionManager>) {
let key = name.into();
if self.managers.is_empty() || self.default_name.is_none() {
self.default_name = Some(key.clone());
}
self.managers.insert(key, manager);
}
pub fn register_default(
&mut self,
name: impl Into<String>,
manager: Arc<dyn TransactionManager>,
) {
let key = name.into();
self.default_name = Some(key.clone());
self.managers.insert(key, manager);
}
pub fn get(&self, name: &str) -> Option<Arc<dyn TransactionManager>> {
self.managers.get(name).cloned()
}
pub fn default_manager(&self) -> Option<Arc<dyn TransactionManager>> {
self.default_name
.as_ref()
.and_then(|n| self.managers.get(n).cloned())
}
pub fn default_name(&self) -> Option<&str> {
self.default_name.as_deref()
}
pub fn manager_names(&self) -> Vec<&str> {
self.managers.keys().map(String::as_str).collect()
}
pub fn contains(&self, name: &str) -> bool {
self.managers.contains_key(name)
}
pub fn len(&self) -> usize {
self.managers.len()
}
pub fn is_empty(&self) -> bool {
self.managers.is_empty()
}
pub fn remove(&mut self, name: &str) -> Option<Arc<dyn TransactionManager>> {
let removed = self.managers.remove(name);
if removed.is_some() && self.default_name.as_deref() == Some(name) {
self.default_name = None;
}
removed
}
pub fn into_delegate(self) -> TransactionResult<DelegatingTransactionManager> {
let default = self.default_manager().ok_or_else(|| {
TransactionError::InvalidState("No default transaction manager registered".into())
})?;
Ok(DelegatingTransactionManager {
registry: self,
fallback: default,
})
}
}
#[derive(Clone)]
pub struct DelegatingTransactionManager {
registry: TransactionManagerRegistry,
fallback: Arc<dyn TransactionManager>,
}
impl std::fmt::Debug for DelegatingTransactionManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DelegatingTransactionManager")
.field("registry", &self.registry)
.field("fallback", &"dyn TransactionManager")
.finish()
}
}
impl DelegatingTransactionManager {
pub fn new(registry: &TransactionManagerRegistry) -> TransactionResult<Self> {
registry.clone().into_delegate()
}
pub fn registry(&self) -> &TransactionManagerRegistry {
&self.registry
}
pub async fn begin_for(
&self,
name: &str,
definition: &TransactionDefinition,
) -> TransactionResult<TransactionStatus> {
self.registry
.get(name)
.ok_or_else(|| {
TransactionError::NotFound(format!("Transaction manager '{}' not found", name))
})?
.begin(definition)
.await
}
pub async fn commit_for(&self, name: &str, status: TransactionStatus) -> TransactionResult<()> {
self.registry
.get(name)
.ok_or_else(|| {
TransactionError::NotFound(format!("Transaction manager '{}' not found", name))
})?
.commit(status)
.await
}
pub async fn rollback_for(
&self,
name: &str,
status: TransactionStatus,
) -> TransactionResult<()> {
self.registry
.get(name)
.ok_or_else(|| {
TransactionError::NotFound(format!("Transaction manager '{}' not found", name))
})?
.rollback(status)
.await
}
}
#[async_trait]
impl TransactionManager for DelegatingTransactionManager {
async fn begin(
&self,
definition: &TransactionDefinition,
) -> TransactionResult<TransactionStatus> {
self.fallback.begin(definition).await
}
async fn commit(&self, status: TransactionStatus) -> TransactionResult<()> {
self.fallback.commit(status).await
}
async fn rollback(&self, status: TransactionStatus) -> TransactionResult<()> {
self.fallback.rollback(status).await
}
fn name(&self) -> &'static str {
"delegating"
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::manager::NoopTransactionManager;
fn noop(name: &str) -> Arc<dyn TransactionManager> {
Arc::new(crate::manager::NoopTransactionManager)
}
#[derive(Debug, Default)]
struct RecordingManager {
name: String,
}
impl RecordingManager {
fn new(name: impl Into<String>) -> Self {
Self { name: name.into() }
}
}
#[async_trait]
impl TransactionManager for RecordingManager {
async fn begin(
&self,
definition: &TransactionDefinition,
) -> TransactionResult<TransactionStatus> {
Ok(TransactionStatus::new(&definition.name))
}
async fn commit(&self, status: TransactionStatus) -> TransactionResult<()> {
status.mark_completed();
Ok(())
}
async fn rollback(&self, status: TransactionStatus) -> TransactionResult<()> {
status.mark_completed();
Ok(())
}
fn name(&self) -> &str {
&self.name
}
}
#[test]
fn test_registry_new_is_empty() {
let registry = TransactionManagerRegistry::new();
assert!(registry.is_empty());
assert_eq!(registry.len(), 0);
assert!(registry.default_manager().is_none());
}
#[test]
fn test_register_sets_first_as_default() {
let mut registry = TransactionManagerRegistry::new();
registry.register("primary", Arc::new(RecordingManager::new("primary")));
registry.register("secondary", Arc::new(RecordingManager::new("secondary")));
assert_eq!(registry.len(), 2);
assert_eq!(registry.default_name(), Some("primary"));
assert!(registry.contains("primary"));
assert!(registry.contains("secondary"));
}
#[test]
fn test_register_default_overrides() {
let mut registry = TransactionManagerRegistry::new();
registry.register("a", Arc::new(RecordingManager::new("a")));
registry.register_default("b", Arc::new(RecordingManager::new("b")));
assert_eq!(registry.default_name(), Some("b"));
}
#[test]
fn test_get_by_name() {
let mut registry = TransactionManagerRegistry::new();
registry.register("orders", Arc::new(RecordingManager::new("orders")));
assert!(registry.get("orders").is_some());
assert!(registry.get("unknown").is_none());
}
#[test]
fn test_remove_manager() {
let mut registry = TransactionManagerRegistry::new();
registry.register("primary", Arc::new(RecordingManager::new("primary")));
registry.register("secondary", Arc::new(RecordingManager::new("secondary")));
let removed = registry.remove("primary");
assert!(removed.is_some());
assert!(!registry.contains("primary"));
assert!(registry.default_name().is_none());
}
#[test]
fn test_manager_names() {
let mut registry = TransactionManagerRegistry::new();
registry.register("alpha", Arc::new(RecordingManager::new("alpha")));
registry.register("beta", Arc::new(RecordingManager::new("beta")));
let mut names = registry.manager_names();
names.sort();
assert_eq!(names, vec!["alpha", "beta"]);
}
#[tokio::test]
async fn test_delegate_uses_default_manager() {
let mut registry = TransactionManagerRegistry::new();
registry.register("primary", Arc::new(RecordingManager::new("primary")));
let delegate = registry.into_delegate().unwrap();
let def = TransactionDefinition::new("test-tx");
let status = delegate.begin(&def).await.unwrap();
assert!(status.is_new_transaction());
delegate.commit(status).await.unwrap();
}
#[tokio::test]
async fn test_delegate_begin_for_named_source() {
let mut registry = TransactionManagerRegistry::new();
registry.register("orders", Arc::new(RecordingManager::new("orders")));
registry.register("inventory", Arc::new(RecordingManager::new("inventory")));
let delegate = registry.into_delegate().unwrap();
let def = TransactionDefinition::new("order-tx");
let status = delegate.begin_for("orders", &def).await.unwrap();
delegate.commit_for("orders", status).await.unwrap();
let status2 = delegate.begin_for("inventory", &def).await.unwrap();
delegate.rollback_for("inventory", status2).await.unwrap();
}
#[tokio::test]
async fn test_delegate_begin_for_unknown_fails() {
let mut registry = TransactionManagerRegistry::new();
registry.register("primary", Arc::new(RecordingManager::new("primary")));
let delegate = registry.into_delegate().unwrap();
let def = TransactionDefinition::new("test");
let result = delegate.begin_for("nonexistent", &def).await;
assert!(result.is_err());
}
#[test]
fn test_into_delegate_empty_registry_fails() {
let registry = TransactionManagerRegistry::new();
let result = registry.into_delegate();
assert!(result.is_err());
}
}