use crate::registry::{parse_target, BranchCtx, Registry, Target};
use dtmrs_core::{
msg_advance, saga_advance, tcc_advance, xa_advance, Advance, BranchOp, BranchResult,
BranchStatus, GlobalStatus, SagaStep, TransType,
};
use dtmrs_store::{GlobalRow, Store};
use std::sync::Arc;
use std::time::Duration;
use tracing::{error, info, warn};
#[derive(Clone)]
pub struct Driver {
pub store: Store,
pub http: reqwest::Client,
pub owner: String,
pub lease: i64,
pub retry: dtmrs_core::RetryPolicy,
branch_timeout_secs: u64,
pub workers: usize,
pub registry: Arc<Registry>,
#[cfg(feature = "grpc")]
pub grpc: crate::grpc::client::GrpcCaller,
pub workflows: Arc<crate::workflow::WorkflowRegistry>,
}
impl Driver {
pub fn new(store: Store, owner: String) -> Self {
Self::with_config(store, owner, DriverConfig::default())
}
pub fn from_env(store: Store, owner: String) -> Self {
Self::with_config(store, owner, DriverConfig::from_env())
}
pub fn with_config(store: Store, owner: String, cfg: DriverConfig) -> Self {
Self {
store,
http: reqwest::Client::builder()
.timeout(Duration::from_secs(cfg.branch_timeout_secs.max(1) as u64))
.build()
.expect("build http client"),
owner,
lease: cfg.lease_secs,
retry: cfg.retry,
branch_timeout_secs: cfg.branch_timeout_secs.max(1) as u64,
workers: cfg.workers.max(1),
registry: Arc::new(Registry::new()),
#[cfg(feature = "grpc")]
grpc: crate::grpc::client::GrpcCaller::new(Duration::from_secs(
cfg.branch_timeout_secs.max(1) as u64,
)),
workflows: Arc::new(crate::workflow::WorkflowRegistry::new()),
}
}
pub fn http_timeout_secs(&self) -> u64 {
self.branch_timeout_secs
}
pub fn with_registry(mut self, r: Arc<Registry>) -> Self {
self.registry = r;
self
}
pub fn with_workflows(mut self, w: Arc<crate::workflow::WorkflowRegistry>) -> Self {
self.workflows = w;
self
}
#[cfg(feature = "grpc")]
pub fn with_grpc_ca_pem(mut self, pem: impl Into<Vec<u8>>) -> Self {
self.grpc = self.grpc.with_ca_pem(pem);
self
}
pub async fn run_forever(self, tick: Duration) {
let mut set = tokio::task::JoinSet::new();
for _ in 0..self.workers.max(1) {
let d = self.clone();
set.spawn(async move { d.worker_loop(tick).await });
}
set.join_next().await;
}
async fn worker_loop(&self, tick: Duration) {
loop {
match self.store.lock_one_due(&self.owner, self.lease).await {
Ok(Some(g)) => {
if let Err(e) = self.process(&g).await {
warn!(gid = %g.gid, error = %e, "推进出错,等下轮重试");
}
}
Ok(None) => tokio::time::sleep(tick).await,
Err(e) => {
warn!(error = %e, "取待办失败");
tokio::time::sleep(tick).await;
}
}
}
}
pub async fn process(&self, g: &GlobalRow) -> anyhow::Result<()> {
match g.trans_type {
TransType::Saga => self.process_saga(g).await,
TransType::Tcc => self.process_tcc(g).await,
TransType::Msg => self.process_msg(g).await,
TransType::Xa => self.process_xa(g).await,
TransType::Workflow => self.process_workflow(g).await,
}
}
async fn process_workflow(&self, g: &GlobalRow) -> anyhow::Result<()> {
let (name, input) = crate::workflow::decode_payload(&g.payload);
let mut status = g.status;
loop {
let rows = self.store.list_branches(&g.gid).await?;
let compensates = compensate_states(&rows);
match dtmrs_core::workflow_advance(status, &compensates) {
Advance::Finish(s) => {
info!(gid = %g.gid, status = s.as_str(), "workflow 事务终结");
self.store
.set_global_status(&g.gid, s, g.trans_type, "")
.await?;
return Ok(());
}
Advance::Wait => return Ok(()),
Advance::RunWorkflow => {
let Some(f) = self.workflows.get(&name) else {
warn!(gid = %g.gid, workflow = %name,
"workflow 未注册,按结果未知处理(会重试,不回滚)");
self.retry_later(g).await?;
return Ok(());
};
let ctx =
crate::workflow::WorkflowCtx::new(&g.gid, &input, self.store.clone(), rows);
match f(ctx).await {
Ok(()) => {
info!(gid = %g.gid, workflow = %name, "workflow 跑完");
self.store
.set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
.await?;
return Ok(());
}
Err(crate::workflow::WorkflowError::Rollback(reason)) => {
info!(gid = %g.gid, workflow = %name, %reason, "workflow 要求回滚");
self.store
.set_global_status(
&g.gid,
GlobalStatus::Aborting,
g.trans_type,
&reason,
)
.await?;
status = GlobalStatus::Aborting;
continue;
}
Err(crate::workflow::WorkflowError::Diverged {
branch_id: bid,
recorded,
got,
}) => {
warn!(gid = %g.gid, workflow = %name, branch = %bid,
%recorded, %got,
"workflow 重放走岔了,已停止推进,需要人工介入");
self.retry_later(g).await?;
return Ok(());
}
Err(e) => {
warn!(gid = %g.gid, workflow = %name, error = %e, "workflow 需要重试");
self.retry_later(g).await?;
return Ok(());
}
}
}
Advance::Call { index, op } => {
let bid = branch_id(index);
let Some(url) = url_of(&rows, &bid, op) else {
warn!(gid = %g.gid, branch = %bid, "workflow 补偿地址缺失");
self.retry_later(g).await?;
return Ok(());
};
let bp = payload_of(&rows, &bid, op);
match self.call_branch(g, &bid, op, &url, &bp).await {
BranchResult::Success => {
self.store
.set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
.await?;
}
_ => {
warn!(gid = %g.gid, branch = %bid, "workflow 补偿未成功,会重试");
self.retry_later(g).await?;
return Ok(());
}
}
}
}
}
}
async fn process_saga(&self, g: &GlobalRow) -> anyhow::Result<()> {
let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
if steps.is_empty() {
self.store
.set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
.await?;
return Ok(());
}
let mut status = g.status;
loop {
let (actions, compensates) = self.branch_states(&g.gid, steps.len()).await?;
match saga_advance(status, &actions, &compensates) {
Advance::Finish(s) => {
if s == GlobalStatus::Aborting {
status = s;
self.store
.set_global_status(&g.gid, s, g.trans_type, "分支已判失败")
.await?;
continue;
}
info!(gid = %g.gid, status = s.as_str(), "事务终结");
self.store
.set_global_status(&g.gid, s, g.trans_type, "")
.await?;
return Ok(());
}
Advance::Wait => return Ok(()),
Advance::RunWorkflow => {
warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
return Ok(());
}
Advance::Call { index, op } => {
let branch_id = branch_id(index);
let url = match op {
BranchOp::Action => &steps[index].action,
_ => &steps[index].compensate,
};
match self
.call_branch(g, &branch_id, op, url, &steps[index].payload)
.await
{
BranchResult::Success => {
self.store
.set_branch_status(&g.gid, &branch_id, op, BranchStatus::Succeed)
.await?;
}
BranchResult::Failure => {
self.store
.set_branch_status(&g.gid, &branch_id, op, BranchStatus::Failed)
.await?;
if op == BranchOp::Action {
info!(gid = %g.gid, branch = %branch_id, "分支要求回滚");
status = GlobalStatus::Aborting;
self.store
.set_global_status(
&g.gid,
GlobalStatus::Aborting,
g.trans_type,
&format!("分支 {branch_id} 返回 FAILURE"),
)
.await?;
} else {
warn!(gid = %g.gid, branch = %branch_id, "补偿失败,需要人工介入");
self.retry_later(g).await?;
return Ok(());
}
}
BranchResult::Ongoing | BranchResult::Unknown => {
self.retry_later(g).await?;
return Ok(());
}
}
}
}
}
}
async fn process_tcc(&self, g: &GlobalRow) -> anyhow::Result<()> {
self.drive_two_phase(g, BranchOp::Confirm, BranchOp::Cancel, "TCC")
.await
}
async fn process_xa(&self, g: &GlobalRow) -> anyhow::Result<()> {
self.drive_two_phase(g, BranchOp::Commit, BranchOp::Rollback, "XA")
.await
}
async fn drive_two_phase(
&self,
g: &GlobalRow,
fwd: BranchOp,
bwd: BranchOp,
label: &str,
) -> anyhow::Result<()> {
let rows = self.store.list_branches(&g.gid).await?;
let n = rows
.iter()
.filter_map(|r| index_of(&r.branch_id))
.max()
.map(|m| m + 1)
.unwrap_or(0);
if n == 0 {
if !rows.is_empty() {
error!(
gid = %g.gid,
branches = rows.len(),
ids = ?rows.iter().map(|r| r.branch_id.as_str()).take(5).collect::<Vec<_>>(),
"分支号全都无法解析成下标,无法推进。这笔事务需要人工介入 —— \
合法的分支号形如 01 / 02 / 100(见 is_canonical_branch_id)"
);
return Ok(());
}
let s = if g.status == GlobalStatus::Aborting {
GlobalStatus::Failed
} else {
GlobalStatus::Succeed
};
self.store
.set_global_status(&g.gid, s, g.trans_type, "")
.await?;
return Ok(());
}
let status = g.status;
loop {
let rows = self.store.list_branches(&g.gid).await?;
let (f, b) = split_by_op(&rows, n, fwd, bwd);
let adv = if fwd == BranchOp::Commit {
xa_advance(status, &f, &b)
} else {
tcc_advance(status, &f, &b)
};
match adv {
Advance::Finish(s) => {
info!(gid = %g.gid, status = s.as_str(), mode = label, "事务终结");
self.store
.set_global_status(&g.gid, s, g.trans_type, "")
.await?;
return Ok(());
}
Advance::Wait => return Ok(()),
Advance::RunWorkflow => {
warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
return Ok(());
}
Advance::Call { index, op } => {
let bid = branch_id(index);
let Some(url) = url_of(&rows, &bid, op) else {
warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
"分支没登记这个操作的 URL,无法调用");
self.retry_later(g).await?;
return Ok(());
};
let bp = payload_of(&rows, &bid, op);
match self.call_branch(g, &bid, op, &url, &bp).await {
BranchResult::Success => {
self.store
.set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
.await?;
}
BranchResult::Failure => {
self.store
.set_branch_status(&g.gid, &bid, op, BranchStatus::Failed)
.await?;
warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
"二阶段失败,会持续重试,需要人工介入");
self.retry_later(g).await?;
return Ok(());
}
BranchResult::Ongoing | BranchResult::Unknown => {
self.retry_later(g).await?;
return Ok(());
}
}
}
}
}
}
async fn process_msg(&self, g: &GlobalRow) -> anyhow::Result<()> {
let mut status = g.status;
if status == GlobalStatus::Prepared {
if g.query_prepared.is_empty() {
warn!(gid = %g.gid, "msg 事务没提供 query_prepared,无法回查,等人处理");
self.retry_later(g).await?;
return Ok(());
}
match self
.call_branch(g, "00", BranchOp::Action, &g.query_prepared, "")
.await
{
BranchResult::Success => {
info!(gid = %g.gid, "回查:本地事务已提交 → 继续推进");
self.store
.set_global_status(&g.gid, GlobalStatus::Submitted, g.trans_type, "")
.await?;
status = GlobalStatus::Submitted;
}
BranchResult::Failure => {
info!(gid = %g.gid, "回查:本地事务未提交 → 整单作废");
self.store
.set_global_status(
&g.gid,
GlobalStatus::Failed,
g.trans_type,
"回查得到 FAILURE:本地事务未提交",
)
.await?;
return Ok(());
}
BranchResult::Ongoing | BranchResult::Unknown => {
self.retry_later(g).await?;
return Ok(());
}
}
}
let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
if steps.is_empty() {
self.store
.set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
.await?;
return Ok(());
}
loop {
let (actions, _) = self.branch_states(&g.gid, steps.len()).await?;
match msg_advance(status, &actions) {
Advance::Finish(s) => {
info!(gid = %g.gid, status = s.as_str(), "消息事务终结");
self.store
.set_global_status(&g.gid, s, g.trans_type, "")
.await?;
return Ok(());
}
Advance::Wait => return Ok(()),
Advance::RunWorkflow => {
warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
return Ok(());
}
Advance::Call { index, op } => {
let bid = branch_id(index);
match self
.call_branch(g, &bid, op, &steps[index].action, &steps[index].payload)
.await
{
BranchResult::Success => {
self.store
.set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
.await?;
}
BranchResult::Failure | BranchResult::Ongoing | BranchResult::Unknown => {
self.retry_later(g).await?;
return Ok(());
}
}
}
}
}
}
async fn retry_later(&self, g: &GlobalRow) -> anyhow::Result<()> {
let iv = dtmrs_core::next_interval_with(g.next_cron_interval, self.retry);
self.store.schedule_retry(&g.gid, iv).await?;
Ok(())
}
async fn branch_states(
&self,
gid: &str,
n: usize,
) -> anyhow::Result<(Vec<BranchStatus>, Vec<BranchStatus>)> {
let rows = self.store.list_branches(gid).await?;
let mut actions = vec![BranchStatus::Prepared; n];
let mut compensates = vec![BranchStatus::Prepared; n];
for r in rows {
let Some(i) = index_of(&r.branch_id) else {
continue;
};
if i >= n {
continue;
}
match r.op {
BranchOp::Action | BranchOp::Try => actions[i] = r.status,
BranchOp::Compensate | BranchOp::Cancel | BranchOp::Rollback => {
compensates[i] = r.status
}
_ => {}
}
}
Ok((actions, compensates))
}
async fn call_branch(
&self,
g: &GlobalRow,
branch_id: &str,
op: BranchOp,
url: &str,
payload: &str,
) -> BranchResult {
match parse_target(url) {
Target::Local(name) => self.call_local(g, branch_id, op, &name).await,
Target::Http(u) => self.call_http(g, branch_id, op, &u, payload).await,
#[cfg(feature = "grpc")]
Target::Grpc(t) => {
self.grpc
.call(
&t,
&g.gid,
&g.trans_type.to_string(),
branch_id,
op.as_str(),
)
.await
}
#[cfg(not(feature = "grpc"))]
Target::Grpc(t) => {
warn!(gid = %g.gid, branch = %branch_id, endpoint = %t.endpoint,
"遇到 grpc:// 分支但本次构建关掉了 grpc feature,按结果未知处理");
BranchResult::Unknown
}
}
}
async fn call_local(
&self,
g: &GlobalRow,
branch_id: &str,
op: BranchOp,
name: &str,
) -> BranchResult {
let Some(h) = self.registry.get(name) else {
warn!(gid = %g.gid, branch = %branch_id, handler = name,
"本地分支未注册,按结果未知处理(会重试,不回滚)");
return BranchResult::Unknown;
};
let ctx = BranchCtx {
gid: g.gid.clone(),
branch_id: branch_id.to_string(),
op,
trans_type: g.trans_type.to_string(),
};
let r = h(ctx).await;
info!(gid = %g.gid, branch = %branch_id, op = op.as_str(),
handler = name, result = ?r, "本地分支返回");
r
}
async fn call_http(
&self,
g: &GlobalRow,
branch_id: &str,
op: BranchOp,
url: &str,
payload: &str,
) -> BranchResult {
let req = self
.http
.post(url)
.query(&[
("gid", g.gid.as_str()),
("trans_type", &g.trans_type.to_string()),
("branch_id", branch_id),
("op", op.as_str()),
])
.header("content-type", "application/json")
.body(branch_payload(payload));
match req.send().await {
Ok(resp) => {
let code = resp.status().as_u16();
let body = resp.text().await.unwrap_or_default();
let r = BranchResult::from_http(code, &body);
info!(gid = %g.gid, branch = %branch_id, op = op.as_str(), code, result = ?r, "分支返回");
r
}
Err(e) => {
warn!(gid = %g.gid, branch = %branch_id, error = %e, "分支不可达,结果未知");
BranchResult::Unknown
}
}
}
}
pub fn branch_id(index: usize) -> String {
format!("{:02}", index + 1)
}
pub const MAX_BRANCH_INDEX: usize = 9999;
fn index_of(branch_id: &str) -> Option<usize> {
branch_id
.parse::<usize>()
.ok()
.and_then(|v| v.checked_sub(1))
.filter(|&i| i <= MAX_BRANCH_INDEX)
}
pub fn is_canonical_branch_id(s: &str) -> bool {
index_of(s).map(branch_id).as_deref() == Some(s)
}
fn split_by_op(
rows: &[dtmrs_store::BranchRow],
n: usize,
fwd: BranchOp,
bwd: BranchOp,
) -> (Vec<BranchStatus>, Vec<BranchStatus>) {
let mut a = vec![BranchStatus::Prepared; n];
let mut b = vec![BranchStatus::Prepared; n];
for r in rows {
let Some(i) = index_of(&r.branch_id) else {
continue;
};
if i >= n {
continue;
}
if r.op == fwd {
a[i] = r.status;
} else if r.op == bwd {
b[i] = r.status;
}
}
(a, b)
}
fn compensate_states(rows: &[dtmrs_store::BranchRow]) -> Vec<BranchStatus> {
let n = rows
.iter()
.filter(|r| r.op == BranchOp::Compensate)
.filter_map(|r| index_of(&r.branch_id))
.max()
.map(|m| m + 1)
.unwrap_or(0);
let mut v = vec![BranchStatus::Succeed; n];
for r in rows.iter().filter(|r| r.op == BranchOp::Compensate) {
if let Some(i) = index_of(&r.branch_id) {
if i < n {
v[i] = r.status;
}
}
}
v
}
fn payload_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> String {
rows.iter()
.find(|r| r.branch_id == branch_id && r.op == op)
.map(|r| r.payload.clone())
.unwrap_or_default()
}
fn url_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> Option<String> {
rows.iter()
.find(|r| r.branch_id == branch_id && r.op == op)
.map(|r| r.url.clone())
}
fn branch_payload(step_payload: &str) -> String {
if step_payload.trim().is_empty() {
"{}".to_string()
} else {
step_payload.to_string()
}
}
#[derive(Debug, Clone, Copy)]
pub struct DriverConfig {
pub branch_timeout_secs: i64,
pub lease_secs: i64,
pub retry: dtmrs_core::RetryPolicy,
pub workers: usize,
}
impl Default for DriverConfig {
fn default() -> Self {
Self {
branch_timeout_secs: 10,
lease_secs: 30,
retry: dtmrs_core::RetryPolicy::default(),
workers: 16,
}
}
}
impl DriverConfig {
pub fn from_env() -> Self {
let d = Self::default();
let get = |k: &str, fallback: i64| {
std::env::var(k)
.ok()
.and_then(|v| v.parse::<i64>().ok())
.filter(|v| *v > 0)
.unwrap_or(fallback)
};
Self {
branch_timeout_secs: get("DTMRS_BRANCH_TIMEOUT", d.branch_timeout_secs),
lease_secs: get("DTMRS_LEASE", d.lease_secs),
retry: dtmrs_core::RetryPolicy::from_env(),
workers: get("DTMRS_WORKERS", d.workers as i64) as usize,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn 分支号与下标互转() {
assert_eq!(branch_id(0), "01");
assert_eq!(branch_id(9), "10");
assert_eq!(index_of("01"), Some(0));
assert_eq!(index_of("10"), Some(9));
assert_eq!(index_of("00"), None);
assert_eq!(index_of("xx"), None);
}
#[test]
fn 分支号超过99后下标解析仍然正确() {
let ids: Vec<String> = (0..500).map(branch_id).collect();
for (i, id) in ids.iter().enumerate() {
assert_eq!(index_of(id), Some(i), "下标解析错了,执行顺序会乱");
}
let mut sorted = ids.clone();
sorted.sort();
assert_ne!(ids, sorted);
}
}