1use crate::driver::branch_id;
58use dtmrs_core::{BranchOp, BranchResult, BranchStatus};
59use dtmrs_store::{BranchRow, Store};
60use std::collections::HashMap;
61use std::future::Future;
62use std::pin::Pin;
63use std::sync::Arc;
64
65#[derive(Debug, Clone, PartialEq, Eq)]
67pub enum WorkflowError {
68 Rollback(String),
71 Retry(String),
74 Diverged {
80 branch_id: String,
81 recorded: String,
82 got: String,
83 },
84 Internal(String),
86}
87
88impl std::fmt::Display for WorkflowError {
89 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
90 match self {
91 Self::Rollback(m) => write!(f, "业务要求回滚: {m}"),
92 Self::Retry(m) => write!(f, "需要重试: {m}"),
93 Self::Diverged {
94 branch_id,
95 recorded,
96 got,
97 } => write!(
98 f,
99 "重放走岔了:分支 {branch_id} 上次记录的是「{recorded}」,这次却是「{got}」。\
100 函数必须是确定性的,副作用都要放进 branch().run() 里"
101 ),
102 Self::Internal(m) => write!(f, "内部错误: {m}"),
103 }
104 }
105}
106
107impl std::error::Error for WorkflowError {}
108
109pub type WorkflowResult<T> = Result<T, WorkflowError>;
110
111pub struct WorkflowCtx {
113 pub gid: String,
114 pub input: String,
116 store: Store,
117 seq: usize,
119 recorded: HashMap<String, BranchRow>,
121}
122
123impl WorkflowCtx {
124 pub(crate) fn new(gid: &str, input: &str, store: Store, rows: Vec<BranchRow>) -> Self {
125 let recorded = rows
126 .into_iter()
127 .filter(|r| r.op == BranchOp::Action)
128 .map(|r| (r.branch_id.clone(), r))
129 .collect();
130 Self {
131 gid: gid.to_string(),
132 input: input.to_string(),
133 store,
134 seq: 0,
135 recorded,
136 }
137 }
138
139 pub fn branch(&mut self, name: &str) -> BranchBuilder<'_> {
144 BranchBuilder {
145 ctx: self,
146 name: name.to_string(),
147 compensate: None,
148 }
149 }
150
151 pub fn branch_count(&self) -> usize {
153 self.seq
154 }
155}
156
157pub struct BranchBuilder<'a> {
158 ctx: &'a mut WorkflowCtx,
159 name: String,
160 compensate: Option<String>,
161}
162
163impl BranchBuilder<'_> {
164 pub fn on_rollback(mut self, compensate: &str) -> Self {
170 self.compensate = Some(compensate.to_string());
171 self
172 }
173
174 pub async fn run<F, Fut>(self, f: F) -> WorkflowResult<()>
176 where
177 F: FnOnce() -> Fut,
178 Fut: Future<Output = BranchResult>,
179 {
180 self.run_with(|| async move { (f().await, String::new()) })
181 .await
182 .map(|_| ())
183 }
184
185 pub async fn run_with<F, Fut>(self, f: F) -> WorkflowResult<String>
190 where
191 F: FnOnce() -> Fut,
192 Fut: Future<Output = (BranchResult, String)>,
193 {
194 let bid = branch_id(self.ctx.seq);
195 self.ctx.seq += 1;
196 let gid = self.ctx.gid.clone();
197
198 if let Some(row) = self.ctx.recorded.get(&bid) {
200 if row.url != self.name {
203 return Err(WorkflowError::Diverged {
204 branch_id: bid,
205 recorded: row.url.clone(),
206 got: self.name,
207 });
208 }
209 if row.status == BranchStatus::Succeed {
210 return Ok(row.payload.clone());
212 }
213 }
216
217 let mut ops = Vec::with_capacity(2);
221 if let Some(c) = &self.compensate {
222 ops.push((BranchOp::Compensate, c.clone()));
223 }
224 ops.push((BranchOp::Action, self.name.clone()));
226 self.ctx
227 .store
228 .register_branch(&gid, &bid, &ops)
229 .await
230 .map_err(|e| WorkflowError::Internal(format!("登记分支 {bid} 失败: {e}")))?;
231
232 let (result, data) = f().await;
234 match result {
235 BranchResult::Success => {
236 self.ctx
237 .store
238 .set_branch_result(&gid, &bid, BranchOp::Action, BranchStatus::Succeed, &data)
239 .await
240 .map_err(|e| WorkflowError::Internal(format!("存分支结果失败: {e}")))?;
241 Ok(data)
242 }
243 BranchResult::Failure => {
244 let _ = self
245 .ctx
246 .store
247 .set_branch_status(&gid, &bid, BranchOp::Action, BranchStatus::Failed)
248 .await;
249 Err(WorkflowError::Rollback(format!(
250 "分支 {bid}({})返回 FAILURE",
251 self.name
252 )))
253 }
254 BranchResult::Ongoing | BranchResult::Unknown => Err(WorkflowError::Retry(format!(
256 "分支 {bid}({})结果未知或处理中",
257 self.name
258 ))),
259 }
260 }
261}
262
263type BoxFut = Pin<Box<dyn Future<Output = WorkflowResult<()>> + Send>>;
264type WorkflowFn = Arc<dyn Fn(WorkflowCtx) -> BoxFut + Send + Sync>;
265
266#[derive(Default)]
271pub struct WorkflowRegistry {
272 fns: HashMap<String, WorkflowFn>,
273}
274
275impl WorkflowRegistry {
276 pub fn new() -> Self {
277 Self::default()
278 }
279
280 pub fn register<F, Fut>(&mut self, name: &str, f: F) -> &mut Self
281 where
282 F: Fn(WorkflowCtx) -> Fut + Send + Sync + 'static,
283 Fut: Future<Output = WorkflowResult<()>> + Send + 'static,
284 {
285 let h: WorkflowFn = Arc::new(move |ctx| Box::pin(f(ctx)));
286 self.fns.insert(name.to_string(), h);
287 self
288 }
289
290 pub fn get(&self, name: &str) -> Option<WorkflowFn> {
291 self.fns.get(name).cloned()
292 }
293
294 pub fn contains(&self, name: &str) -> bool {
295 self.fns.contains_key(name)
296 }
297
298 pub fn names(&self) -> Vec<&str> {
299 self.fns.keys().map(String::as_str).collect()
300 }
301}
302
303impl std::fmt::Debug for WorkflowRegistry {
304 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
305 f.debug_struct("WorkflowRegistry")
306 .field("workflows", &self.names())
307 .finish()
308 }
309}
310
311pub(crate) fn encode_payload(name: &str, input: &str) -> String {
313 serde_json::json!({ "name": name, "input": input }).to_string()
314}
315
316pub(crate) fn decode_payload(payload: &str) -> (String, String) {
317 let v: serde_json::Value = serde_json::from_str(payload).unwrap_or_default();
318 (
319 v.get("name")
320 .and_then(|x| x.as_str())
321 .unwrap_or("")
322 .to_string(),
323 v.get("input")
324 .and_then(|x| x.as_str())
325 .unwrap_or("")
326 .to_string(),
327 )
328}
329
330#[cfg(test)]
331mod tests {
332 use super::*;
333
334 #[test]
335 fn payload编解码能往返() {
336 let p = encode_payload("下单", r#"{"oid":1}"#);
337 let (n, i) = decode_payload(&p);
338 assert_eq!(n, "下单");
339 assert_eq!(i, r#"{"oid":1}"#);
340 }
341
342 #[test]
343 fn 坏payload不会panic() {
344 assert_eq!(decode_payload("不是 json"), (String::new(), String::new()));
347 assert_eq!(decode_payload(""), (String::new(), String::new()));
348 assert_eq!(decode_payload("{}"), (String::new(), String::new()));
349 }
350
351 #[test]
352 fn 注册表按名字查函数() {
353 let mut r = WorkflowRegistry::new();
354 r.register("a", |_ctx| async { Ok(()) });
355 assert!(r.contains("a"));
356 assert!(r.get("a").is_some());
357 assert!(r.get("没这个").is_none());
358 }
359
360 #[test]
361 fn 分岔错误的说明要能指出问题所在() {
362 let e = WorkflowError::Diverged {
363 branch_id: "02".into(),
364 recorded: "扣款".into(),
365 got: "发货".into(),
366 };
367 let s = e.to_string();
368 assert!(s.contains("02") && s.contains("扣款") && s.contains("发货"));
369 assert!(s.contains("确定性"), "得告诉用户根因是函数不确定");
370 }
371}