Skip to main content

confium_coordinator/
saga.rs

1//! Saga pattern for multi-step ceremony management.
2
3use serde::{Deserialize, Serialize};
4use std::collections::HashMap;
5
6/// A saga step with forward action and compensation.
7#[derive(Debug, Clone)]
8pub struct SagaStep {
9    pub name: String,
10    pub completed: bool,
11    pub failed: bool,
12}
13
14/// Saga execution state.
15#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
16#[serde(rename_all = "snake_case")]
17pub enum SagaState {
18    Running,
19    Completed,
20    Compensating,
21    Compensated,
22    Failed,
23}
24
25/// A saga orchestrating multi-step ceremonies.
26pub struct Saga {
27    pub saga_id: String,
28    pub steps: Vec<SagaStep>,
29    pub state: SagaState,
30    pub current_step: usize,
31    pub results: HashMap<String, String>,
32}
33
34impl Saga {
35    pub fn new(saga_id: &str, step_names: &[&str]) -> Self {
36        let steps = step_names
37            .iter()
38            .map(|name| SagaStep {
39                name: name.to_string(),
40                completed: false,
41                failed: false,
42            })
43            .collect();
44        Self {
45            saga_id: saga_id.into(),
46            steps,
47            state: SagaState::Running,
48            current_step: 0,
49            results: HashMap::new(),
50        }
51    }
52
53    /// Execute the saga forward. Returns Ok when all steps complete,
54    /// Err when a step fails (triggers compensation).
55    pub fn execute<F>(&mut self, mut step_fn: F) -> Result<(), String>
56    where
57        F: FnMut(&str) -> Result<String, String>,
58    {
59        while self.current_step < self.steps.len() {
60            let step = &mut self.steps[self.current_step];
61            match step_fn(&step.name) {
62                Ok(result) => {
63                    self.results.insert(step.name.clone(), result);
64                    self.steps[self.current_step].completed = true;
65                    self.current_step += 1;
66                }
67                Err(e) => {
68                    self.steps[self.current_step].failed = true;
69                    self.state = SagaState::Compensating;
70                    return Err(e);
71                }
72            }
73        }
74        self.state = SagaState::Completed;
75        Ok(())
76    }
77
78    /// Compensate (roll back) completed steps in reverse order.
79    pub fn compensate<F>(&mut self, mut compensate_fn: F) -> Result<(), String>
80    where
81        F: FnMut(&str) -> Result<(), String>,
82    {
83        self.state = SagaState::Compensating;
84        let mut completed_indices: Vec<usize> = self
85            .steps
86            .iter()
87            .enumerate()
88            .filter(|(_, s)| s.completed)
89            .map(|(i, _)| i)
90            .collect();
91        completed_indices.reverse();
92
93        for idx in completed_indices {
94            let step_name = self.steps[idx].name.clone();
95            if let Err(e) = compensate_fn(&step_name) {
96                self.state = SagaState::Failed;
97                return Err(e);
98            }
99            self.steps[idx].completed = false;
100        }
101
102        self.state = SagaState::Compensated;
103        Ok(())
104    }
105
106    pub fn progress(&self) -> f64 {
107        let completed = self.steps.iter().filter(|s| s.completed).count();
108        completed as f64 / self.steps.len().max(1) as f64
109    }
110
111    pub fn completed_steps(&self) -> Vec<&str> {
112        self.steps
113            .iter()
114            .filter(|s| s.completed)
115            .map(|s| s.name.as_str())
116            .collect()
117    }
118
119    pub fn step_count(&self) -> usize {
120        self.steps.len()
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    use super::*;
127
128    #[test]
129    fn all_steps_succeed() {
130        let mut saga = Saga::new("s1", &["step1", "step2", "step3"]);
131        saga.execute(|name| Ok(format!("done-{name}"))).unwrap();
132        assert_eq!(saga.state, SagaState::Completed);
133        assert_eq!(saga.progress(), 1.0);
134    }
135
136    #[test]
137    fn step_failure_triggers_compensation() {
138        let mut saga = Saga::new("s1", &["step1", "step2", "step3"]);
139        let result = saga.execute(|name| {
140            if name == "step2" {
141                Err("step2 failed".into())
142            } else {
143                Ok("ok".into())
144            }
145        });
146        assert!(result.is_err());
147        assert_eq!(saga.state, SagaState::Compensating);
148        // step1 completed, step2 failed, step3 not reached
149        assert!(saga.steps[0].completed);
150        assert!(saga.steps[1].failed);
151        assert!(!saga.steps[2].completed);
152    }
153
154    #[test]
155    fn compensation_rolls_back() {
156        let mut saga = Saga::new("s1", &["step1", "step2", "step3"]);
157        saga.execute(|name| {
158            if name == "step3" {
159                Err("fail".into())
160            } else {
161                Ok("ok".into())
162            }
163        })
164        .unwrap_err();
165        saga.compensate(|_| Ok(())).unwrap();
166        assert_eq!(saga.state, SagaState::Compensated);
167        assert_eq!(saga.completed_steps().len(), 0);
168    }
169
170    #[test]
171    fn first_step_failure_no_compensation_needed() {
172        let mut saga = Saga::new("s1", &["step1", "step2"]);
173        saga.execute(|_| Err("immediate fail".into())).unwrap_err();
174        saga.compensate(|_| Ok(())).unwrap();
175        assert_eq!(saga.state, SagaState::Compensated);
176    }
177
178    #[test]
179    fn progress_tracking() {
180        let mut saga = Saga::new("s1", &["a", "b", "c", "d"]);
181        saga.execute(|name| {
182            if name == "c" {
183                Err("stop".into())
184            } else {
185                Ok("ok".into())
186            }
187        })
188        .unwrap_err();
189        assert!((saga.progress() - 0.5).abs() < 0.01);
190    }
191
192    #[test]
193    fn empty_saga_completes() {
194        let mut saga = Saga::new("empty", &[]);
195        saga.execute(|_| Ok("ok".into())).unwrap();
196        assert_eq!(saga.state, SagaState::Completed);
197    }
198
199    #[test]
200    fn results_recorded() {
201        let mut saga = Saga::new("s1", &["step1", "step2"]);
202        saga.execute(|name| Ok(format!("result-{name}"))).unwrap();
203        assert_eq!(saga.results.get("step1").unwrap(), "result-step1");
204        assert_eq!(saga.results.get("step2").unwrap(), "result-step2");
205    }
206
207    #[test]
208    fn compensation_failure_marks_failed() {
209        let mut saga = Saga::new("s1", &["step1", "step2"]);
210        saga.execute(|name| {
211            if name == "step2" {
212                Err("fail".into())
213            } else {
214                Ok("ok".into())
215            }
216        })
217        .unwrap_err();
218        saga.compensate(|name| {
219            if name == "step1" {
220                Err("compensation failed".into())
221            } else {
222                Ok(())
223            }
224        })
225        .unwrap_err();
226        assert_eq!(saga.state, SagaState::Failed);
227    }
228
229    #[test]
230    fn step_count() {
231        let saga = Saga::new("s1", &["a", "b", "c"]);
232        assert_eq!(saga.step_count(), 3);
233    }
234}