confium_coordinator/
saga.rs1use serde::{Deserialize, Serialize};
4use std::collections::HashMap;
5
6#[derive(Debug, Clone)]
8pub struct SagaStep {
9 pub name: String,
10 pub completed: bool,
11 pub failed: bool,
12}
13
14#[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
25pub 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 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 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 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}