Skip to main content

dtmrs_server/
workflow.rs

1//! workflow 模式:把事务流程写成一个**普通函数**,崩溃后从断点续跑。
2//!
3//! # 跟前四种模式的差别
4//!
5//! SAGA / TCC / msg / XA 都要求提交时就把步骤声明清楚。真实业务经常不满足:
6//! 第三步做不做取决于第二步返回了什么,中间还有 `if`、有循环。
7//!
8//! workflow 模式让你直接写:
9//!
10//! ```ignore
11//! tc.workflow("下单", |mut wf| async move {
12//!     let oid = wf.branch("建订单").on_rollback("local://取消订单")
13//!         .run_with(|| async { (BranchResult::Success, new_order_id()) }).await?;
14//!
15//!     wf.branch("扣款").on_rollback("local://退款")
16//!         .run(|| async { deduct(&oid).await }).await?;
17//!
18//!     // 控制流是真的控制流
19//!     if need_ship(&oid) {
20//!         wf.branch("发货").on_rollback("local://退货")
21//!             .run(|| async { ship(&oid).await }).await?;
22//!     }
23//!     Ok(())
24//! })
25//! ```
26//!
27//! # 崩溃恢复靠重放 + 结果记忆化
28//!
29//! 进程崩了重启,TC 会把这个函数**从头再跑一遍**。已经成功过的分支不重新执行,
30//! 而是把上次存的返回值原样还给你 —— 所以函数会沿着上次的路径走到断点,
31//! 然后继续往下。
32//!
33//! ```text
34//! 第一次:  建订单(真跑,存 oid) → 扣款(真跑) → 崩溃
35//! 重启后:  建订单(记忆化,还回 oid) → 扣款(记忆化) → 发货(真跑) → 完成
36//!                    ↑ 副作用不会重做
37//! ```
38//!
39//! # ⚠ 你的函数必须是确定性的
40//!
41//! 重放是**从头再跑**,所以分支之间的那些代码会被执行多次。它们必须在相同的
42//! 分支返回值下走相同的路径:
43//!
44//! - ❌ `if rand() > 0.5`、`if now().hour() < 12`、读一个会变的全局状态
45//! - ❌ 在分支**外面**直接写数据库 —— 那部分不会被记忆化,重放时会重复执行
46//! - ✅ 所有副作用都放进 `branch(...).run(...)` 里
47//!
48//! 写岔了会怎样?本模块会**当场发现并拒绝继续**(见 [`WorkflowError::Diverged`]),
49//! 而不是静默补偿错对象。这是刻意的:静默走错比停下来严重得多。
50//!
51//! # 为什么这个模式只在嵌入式形态下提供
52//!
53//! 因为「步骤」是**代码**,没法表示成一个 URL 存进数据库。DTM 那边也是同理:
54//! workflow 的函数体在客户端进程里,TC 只存状态。
55//! 我们把 TC 也放在同一个进程里,所以这件事反而更自然。
56
57use 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/// 函数跑不下去的原因。用 `?` 往外抛。
66#[derive(Debug, Clone, PartialEq, Eq)]
67pub enum WorkflowError {
68    /// 业务**明确**要求回滚 —— 会触发逆序补偿。
69    /// 只有这一种会回滚,其它都是重试
70    Rollback(String),
71    /// 结果未知或还在处理中 —— 退避重试,**绝不回滚**。
72    /// 分支返回 `Ongoing` / `Unknown` 时自动变成这个
73    Retry(String),
74    /// **重放走岔了**:这次跑到第 N 个分支时的名字,跟上次记录的对不上。
75    ///
76    /// 说明函数不是确定性的(或者代码改了步骤顺序又碰上老事务)。
77    /// 这时候**不能继续**:按位置记忆化会把 A 的结果当成 B 的,
78    /// 回滚时也会补偿错对象。停下来等人比静默走错安全得多。
79    Diverged {
80        branch_id: String,
81        recorded: String,
82        got: String,
83    },
84    /// 存储层出错之类
85    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
111/// 传给 workflow 函数的上下文。分支都从这里开。
112pub struct WorkflowCtx {
113    pub gid: String,
114    /// 提交时带的输入数据,原样透传
115    pub input: String,
116    store: Store,
117    /// 下一个分支的序号
118    seq: usize,
119    /// 本次运行开始时库里已有的分支行(上次运行留下的),key 是 `branch_id`
120    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    /// 开一个分支。
140    ///
141    /// `name` 是这个分支的逻辑名字,用来做**重放分岔检测** ——
142    /// 重放时第 N 个分支的名字必须跟上次一致。取个稳定的名字,别用时间戳之类。
143    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    /// 已经登记过的分支数(含本次运行新增的)
152    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    /// 登记这个分支的补偿。地址跟 saga 的补偿一样,可以是
165    /// `local://名字` / `http://...` / `grpc://...`。
166    ///
167    /// 不登记补偿的分支在回滚时**不会被补偿** —— 只适合本来就没副作用的步骤
168    /// (比如纯查询)。有副作用就一定要给。
169    pub fn on_rollback(mut self, compensate: &str) -> Self {
170        self.compensate = Some(compensate.to_string());
171        self
172    }
173
174    /// 跑这个分支(不带返回数据)。
175    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    /// 跑这个分支,并记住它的返回数据。**重放时原样还回来,不会重新执行。**
186    ///
187    /// 数据大小受 `trans_branch_op.payload` 列限制(MySQL 上是 VARCHAR),
188    /// 放个 id 或一小段 JSON 就好,别塞大对象。
189    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        // ---- 1. 重放:这个位置上次是什么?----
199        if let Some(row) = self.ctx.recorded.get(&bid) {
200            // 分岔检测先做,不管上次成没成 —— 名字对不上就说明函数走了另一条路,
201            // 再往下按位置记忆化就是张冠李戴
202            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                // 记忆化命中:**不重新执行**,把上次的结果还回去
211                return Ok(row.payload.clone());
212            }
213            // 上次没成(崩在中间 / 超时)—— 往下走,重新执行一遍。
214            // 重复执行的安全性由业务侧的子事务屏障保证
215        }
216
217        // ---- 2. 先登记补偿,再执行动作 ----
218        // 顺序不能反:反过来的话,动作执行完、补偿还没登记时崩溃,
219        // 副作用就永远没人收拾了。跟 TCC「先 registerBranch 再调 try」同一条教训。
220        let mut ops = Vec::with_capacity(2);
221        if let Some(c) = &self.compensate {
222            ops.push((BranchOp::Compensate, c.clone()));
223        }
224        // action 行的 url 存**分支名**,供下次重放做分岔检测
225        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        // ---- 3. 真正执行 ----
233        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            // **绝不能当失败**:对方可能已经成功了,回滚会造成不一致
255            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/// 按名字存 workflow 函数。
267///
268/// 跟 `local://` 分支同一个道理:**函数没法持久化**,库里存的是名字,
269/// 重启后靠这张表把名字解析回函数。所以重启后必须注册同名的 workflow。
270#[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
311/// `trans_global.payload` 里存的东西:workflow 名字 + 输入数据
312pub(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        // 老版本留下的数据、或者被截断的 payload —— 解不出来也不能崩,
345        // 崩了整个推进器就停了
346        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}