use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use crate::database::Session;
use crate::database::sharding::ShardRouter;
use super::store::{InMemorySagaLog, SagaLog, SagaLogStore};
use super::types::*;
#[derive(Debug)]
pub struct SagaExecutionResult {
pub saga_id: String,
pub success: bool,
pub status: SagaStatus,
pub completed_steps: Vec<String>,
pub compensated_steps: Vec<String>,
pub failure: Option<SagaFailure>,
}
#[derive(Debug)]
pub struct SagaFailure {
pub step_name: String,
pub error: String,
}
fn mark_compensation_failed(log: &mut SagaLog, step_name: &str, error: &str) {
if let Some(entry) = log
.steps
.iter_mut()
.find(|s| s.name == step_name && s.action_success)
{
entry.compensation_success = Some(false);
entry.error = Some(error.to_string());
}
}
pub struct SagaRecovery {
store: Arc<dyn SagaLogStore>,
}
impl SagaRecovery {
pub fn new(store: Arc<dyn SagaLogStore>) -> Self {
Self { store }
}
pub async fn list_pending(&self) -> Result<Vec<SagaLog>, String> {
self.store.load_pending().await
}
}
pub struct SagaOrchestrator {
router: Arc<ShardRouter>,
saga_log: Arc<dyn SagaLogStore>,
}
impl SagaOrchestrator {
pub fn new(router: Arc<ShardRouter>) -> Self {
Self {
router,
saga_log: Arc::new(InMemorySagaLog::new()),
}
}
pub fn new_with_log_store(router: Arc<ShardRouter>, saga_log: Arc<dyn SagaLogStore>) -> Self {
Self { router, saga_log }
}
pub async fn compensate_recovered(
&self,
saga_id: &str,
steps: &[SagaStep],
) -> SagaExecutionResult {
let stored = match self.saga_log.get(saga_id).await {
Ok(Some(log)) => log,
_ => {
return SagaExecutionResult {
saga_id: saga_id.to_string(),
success: false,
status: SagaStatus::Failed,
completed_steps: Vec::new(),
compensated_steps: Vec::new(),
failure: Some(SagaFailure {
step_name: saga_id.to_string(),
error: "recovered saga log not found".to_string(),
}),
};
}
};
let mut log = stored.clone();
let mut compensated: Vec<String> = Vec::new();
let mut replay_failed = false;
log.status = SagaStatus::Compensating;
let _ = self.saga_log.persist(&log).await;
let step_index_map: HashMap<&str, usize> = steps
.iter()
.enumerate()
.map(|(i, s)| (s.name.as_str(), i))
.collect();
let replay_list: Vec<SagaStepLog> = log
.steps
.iter()
.filter(|s| s.action_success)
.rev()
.cloned()
.collect();
for step_log in replay_list {
let Some(&idx) = step_index_map.get(step_log.name.as_str()) else {
replay_failed = true;
continue;
};
if let Ok(Some(session)) = self.router.get_session(step_log.shard_id).await {
match steps[idx].compensation.execute(&session).await {
Ok(()) => compensated.push(step_log.name.clone()),
Err(comp_err) => {
replay_failed = true;
mark_compensation_failed(
&mut log,
&step_log.name,
&format!("compensation replay failed: {comp_err}"),
);
}
}
}
}
let final_status = if replay_failed {
SagaStatus::CompensationFailed
} else {
SagaStatus::Failed
};
log.status = final_status;
let _ = self.saga_log.persist(&log).await;
SagaExecutionResult {
saga_id: saga_id.to_string(),
success: false,
status: final_status,
completed_steps: Vec::new(),
compensated_steps: compensated,
failure: None,
}
}
pub async fn execute_saga(&self, steps: Vec<SagaStep>) -> SagaExecutionResult {
let saga_id = uuid::Uuid::new_v4().to_string();
let mut log = SagaLog {
saga_id: saga_id.clone(),
status: SagaStatus::Running,
steps: Vec::new(),
};
let _ = self.saga_log.persist(&log).await;
let mut completed_steps: Vec<(String, u32, Box<dyn SagaAction>)> = Vec::new();
let mut completed_names: Vec<String> = Vec::new();
let step_index_map: HashMap<&str, usize> = steps
.iter()
.enumerate()
.map(|(i, s)| (s.name.as_str(), i))
.collect();
for step in &steps {
let session_result = self.router.get_session(step.shard_id).await;
match session_result {
Ok(Some(session)) => match step.action.execute(&session).await {
Ok(()) => {
log.steps.push(SagaStepLog {
name: step.name.clone(),
shard_id: step.shard_id,
action_success: true,
compensation_success: None,
error: None,
});
completed_names.push(step.name.clone());
let _ = self.saga_log.persist(&log).await;
}
Err(e) => {
log.steps.push(SagaStepLog {
name: step.name.clone(),
shard_id: step.shard_id,
action_success: false,
compensation_success: None,
error: Some(e.to_string()),
});
let mut compensated: Vec<String> = Vec::new();
let mut compensation_failed = false;
log.status = SagaStatus::Compensating;
let _ = self.saga_log.persist(&log).await;
for (completed_name, completed_shard_id, _) in completed_steps.iter().rev()
{
if let Ok(Some(session)) =
self.router.get_session(*completed_shard_id).await
{
if let Some(&idx) = step_index_map.get(completed_name.as_str()) {
match steps[idx].compensation.execute(&session).await {
Ok(()) => {
compensated.push(completed_name.clone());
}
Err(comp_err) => {
compensation_failed = true;
mark_compensation_failed(
&mut log,
completed_name,
&format!("compensation failed: {comp_err}"),
);
}
}
}
}
}
let final_status = if compensation_failed {
SagaStatus::CompensationFailed
} else {
SagaStatus::Failed
};
log.status = final_status;
let _ = self.saga_log.persist(&log).await;
return SagaExecutionResult {
saga_id,
success: false,
status: final_status,
completed_steps: completed_names,
compensated_steps: compensated,
failure: Some(SagaFailure {
step_name: step.name.clone(),
error: e.to_string(),
}),
};
}
},
Err(e) => {
log.status = SagaStatus::Failed;
let _ = self.saga_log.persist(&log).await;
return SagaExecutionResult {
saga_id,
success: false,
status: SagaStatus::Failed,
completed_steps: completed_names,
compensated_steps: Vec::new(),
failure: Some(SagaFailure {
step_name: step.name.clone(),
error: e.to_string(),
}),
};
}
Ok(None) => {
log.status = SagaStatus::Failed;
let _ = self.saga_log.persist(&log).await;
return SagaExecutionResult {
saga_id,
success: false,
status: SagaStatus::Failed,
completed_steps: completed_names,
compensated_steps: Vec::new(),
failure: Some(SagaFailure {
step_name: step.name.clone(),
error: format!("No session available for shard {}", step.shard_id),
}),
};
}
}
completed_steps.push((step.name.clone(), step.shard_id, {
struct NoopAction;
#[async_trait]
impl SagaAction for NoopAction {
async fn execute(&self, _session: &Session) -> Result<(), SagaError> {
Ok(())
}
fn name(&self) -> &str {
"noop"
}
}
Box::new(NoopAction) as Box<dyn SagaAction>
}));
}
log.status = SagaStatus::Completed;
let _ = self.saga_log.persist(&log).await;
SagaExecutionResult {
saga_id,
success: true,
status: SagaStatus::Completed,
completed_steps: completed_names,
compensated_steps: Vec::new(),
failure: None,
}
}
pub async fn get_saga_log(&self, saga_id: &str) -> Option<SagaLog> {
self.saga_log.get(saga_id).await.ok().flatten()
}
}