use crate::driver::branch_id;
use dtmrs_core::{BranchOp, BranchResult, BranchStatus};
use dtmrs_store::{BranchRow, Store};
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum WorkflowError {
Rollback(String),
Retry(String),
Diverged {
branch_id: String,
recorded: String,
got: String,
},
Internal(String),
}
impl std::fmt::Display for WorkflowError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Rollback(m) => write!(f, "业务要求回滚: {m}"),
Self::Retry(m) => write!(f, "需要重试: {m}"),
Self::Diverged {
branch_id,
recorded,
got,
} => write!(
f,
"重放走岔了:分支 {branch_id} 上次记录的是「{recorded}」,这次却是「{got}」。\
函数必须是确定性的,副作用都要放进 branch().run() 里"
),
Self::Internal(m) => write!(f, "内部错误: {m}"),
}
}
}
impl std::error::Error for WorkflowError {}
pub type WorkflowResult<T> = Result<T, WorkflowError>;
pub struct WorkflowCtx {
pub gid: String,
pub input: String,
store: Store,
seq: usize,
recorded: HashMap<String, BranchRow>,
}
impl WorkflowCtx {
pub(crate) fn new(gid: &str, input: &str, store: Store, rows: Vec<BranchRow>) -> Self {
let recorded = rows
.into_iter()
.filter(|r| r.op == BranchOp::Action)
.map(|r| (r.branch_id.clone(), r))
.collect();
Self {
gid: gid.to_string(),
input: input.to_string(),
store,
seq: 0,
recorded,
}
}
pub fn branch(&mut self, name: &str) -> BranchBuilder<'_> {
BranchBuilder {
ctx: self,
name: name.to_string(),
compensate: None,
}
}
pub fn branch_count(&self) -> usize {
self.seq
}
}
pub struct BranchBuilder<'a> {
ctx: &'a mut WorkflowCtx,
name: String,
compensate: Option<String>,
}
impl BranchBuilder<'_> {
pub fn on_rollback(mut self, compensate: &str) -> Self {
self.compensate = Some(compensate.to_string());
self
}
pub async fn run<F, Fut>(self, f: F) -> WorkflowResult<()>
where
F: FnOnce() -> Fut,
Fut: Future<Output = BranchResult>,
{
self.run_with(|| async move { (f().await, String::new()) })
.await
.map(|_| ())
}
pub async fn run_with<F, Fut>(self, f: F) -> WorkflowResult<String>
where
F: FnOnce() -> Fut,
Fut: Future<Output = (BranchResult, String)>,
{
let bid = branch_id(self.ctx.seq);
self.ctx.seq += 1;
let gid = self.ctx.gid.clone();
if let Some(row) = self.ctx.recorded.get(&bid) {
if row.url != self.name {
return Err(WorkflowError::Diverged {
branch_id: bid,
recorded: row.url.clone(),
got: self.name,
});
}
if row.status == BranchStatus::Succeed {
return Ok(row.payload.clone());
}
}
let mut ops = Vec::with_capacity(2);
if let Some(c) = &self.compensate {
ops.push((BranchOp::Compensate, c.clone()));
}
ops.push((BranchOp::Action, self.name.clone()));
self.ctx
.store
.register_branch(&gid, &bid, &ops)
.await
.map_err(|e| WorkflowError::Internal(format!("登记分支 {bid} 失败: {e}")))?;
let (result, data) = f().await;
match result {
BranchResult::Success => {
self.ctx
.store
.set_branch_result(&gid, &bid, BranchOp::Action, BranchStatus::Succeed, &data)
.await
.map_err(|e| WorkflowError::Internal(format!("存分支结果失败: {e}")))?;
Ok(data)
}
BranchResult::Failure => {
let _ = self
.ctx
.store
.set_branch_status(&gid, &bid, BranchOp::Action, BranchStatus::Failed)
.await;
Err(WorkflowError::Rollback(format!(
"分支 {bid}({})返回 FAILURE",
self.name
)))
}
BranchResult::Ongoing | BranchResult::Unknown => Err(WorkflowError::Retry(format!(
"分支 {bid}({})结果未知或处理中",
self.name
))),
}
}
}
type BoxFut = Pin<Box<dyn Future<Output = WorkflowResult<()>> + Send>>;
type WorkflowFn = Arc<dyn Fn(WorkflowCtx) -> BoxFut + Send + Sync>;
#[derive(Default)]
pub struct WorkflowRegistry {
fns: HashMap<String, WorkflowFn>,
}
impl WorkflowRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register<F, Fut>(&mut self, name: &str, f: F) -> &mut Self
where
F: Fn(WorkflowCtx) -> Fut + Send + Sync + 'static,
Fut: Future<Output = WorkflowResult<()>> + Send + 'static,
{
let h: WorkflowFn = Arc::new(move |ctx| Box::pin(f(ctx)));
self.fns.insert(name.to_string(), h);
self
}
pub fn get(&self, name: &str) -> Option<WorkflowFn> {
self.fns.get(name).cloned()
}
pub fn contains(&self, name: &str) -> bool {
self.fns.contains_key(name)
}
pub fn names(&self) -> Vec<&str> {
self.fns.keys().map(String::as_str).collect()
}
}
impl std::fmt::Debug for WorkflowRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WorkflowRegistry")
.field("workflows", &self.names())
.finish()
}
}
pub(crate) fn encode_payload(name: &str, input: &str) -> String {
serde_json::json!({ "name": name, "input": input }).to_string()
}
pub(crate) fn decode_payload(payload: &str) -> (String, String) {
let v: serde_json::Value = serde_json::from_str(payload).unwrap_or_default();
(
v.get("name")
.and_then(|x| x.as_str())
.unwrap_or("")
.to_string(),
v.get("input")
.and_then(|x| x.as_str())
.unwrap_or("")
.to_string(),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn payload编解码能往返() {
let p = encode_payload("下单", r#"{"oid":1}"#);
let (n, i) = decode_payload(&p);
assert_eq!(n, "下单");
assert_eq!(i, r#"{"oid":1}"#);
}
#[test]
fn 坏payload不会panic() {
assert_eq!(decode_payload("不是 json"), (String::new(), String::new()));
assert_eq!(decode_payload(""), (String::new(), String::new()));
assert_eq!(decode_payload("{}"), (String::new(), String::new()));
}
#[test]
fn 注册表按名字查函数() {
let mut r = WorkflowRegistry::new();
r.register("a", |_ctx| async { Ok(()) });
assert!(r.contains("a"));
assert!(r.get("a").is_some());
assert!(r.get("没这个").is_none());
}
#[test]
fn 分岔错误的说明要能指出问题所在() {
let e = WorkflowError::Diverged {
branch_id: "02".into(),
recorded: "扣款".into(),
got: "发货".into(),
};
let s = e.to_string();
assert!(s.contains("02") && s.contains("扣款") && s.contains("发货"));
assert!(s.contains("确定性"), "得告诉用户根因是函数不确定");
}
}