1use crate::registry::{parse_target, BranchCtx, Registry, Target};
7use dtmrs_core::{
8 msg_advance, saga_advance, tcc_advance, xa_advance, Advance, BranchOp, BranchResult,
9 BranchStatus, GlobalStatus, SagaStep, TransType,
10};
11use dtmrs_store::{GlobalRow, Store};
12use std::sync::Arc;
13use std::time::Duration;
14use tracing::{error, info, warn};
15
16#[derive(Clone)]
17pub struct Driver {
18 pub store: Store,
19 pub http: reqwest::Client,
20 pub owner: String,
21 pub lease: i64,
23 pub retry: dtmrs_core::RetryPolicy,
25 branch_timeout_secs: u64,
27 pub workers: usize,
29 pub registry: Arc<Registry>,
31 #[cfg(feature = "grpc")]
33 pub grpc: crate::grpc::client::GrpcCaller,
34 pub workflows: Arc<crate::workflow::WorkflowRegistry>,
36}
37
38impl Driver {
39 pub fn new(store: Store, owner: String) -> Self {
42 Self::with_config(store, owner, DriverConfig::default())
43 }
44
45 pub fn from_env(store: Store, owner: String) -> Self {
60 Self::with_config(store, owner, DriverConfig::from_env())
61 }
62
63 pub fn with_config(store: Store, owner: String, cfg: DriverConfig) -> Self {
64 Self {
65 store,
66 http: reqwest::Client::builder()
67 .timeout(Duration::from_secs(cfg.branch_timeout_secs.max(1) as u64))
68 .build()
69 .expect("build http client"),
70 owner,
71 lease: cfg.lease_secs,
72 retry: cfg.retry,
73 branch_timeout_secs: cfg.branch_timeout_secs.max(1) as u64,
74 workers: cfg.workers.max(1),
75 registry: Arc::new(Registry::new()),
76 #[cfg(feature = "grpc")]
77 grpc: crate::grpc::client::GrpcCaller::new(Duration::from_secs(
78 cfg.branch_timeout_secs.max(1) as u64,
79 )),
80 workflows: Arc::new(crate::workflow::WorkflowRegistry::new()),
81 }
82 }
83
84 pub fn http_timeout_secs(&self) -> u64 {
86 self.branch_timeout_secs
87 }
88
89 pub fn with_registry(mut self, r: Arc<Registry>) -> Self {
91 self.registry = r;
92 self
93 }
94
95 pub fn with_workflows(mut self, w: Arc<crate::workflow::WorkflowRegistry>) -> Self {
97 self.workflows = w;
98 self
99 }
100
101 pub async fn run_forever(self, tick: Duration) {
129 let mut set = tokio::task::JoinSet::new();
130 for _ in 0..self.workers.max(1) {
131 let d = self.clone();
132 set.spawn(async move { d.worker_loop(tick).await });
133 }
134 set.join_next().await;
136 }
137
138 async fn worker_loop(&self, tick: Duration) {
140 loop {
141 match self.store.lock_one_due(&self.owner, self.lease).await {
142 Ok(Some(g)) => {
143 if let Err(e) = self.process(&g).await {
144 warn!(gid = %g.gid, error = %e, "推进出错,等下轮重试");
145 }
146 }
147 Ok(None) => tokio::time::sleep(tick).await,
148 Err(e) => {
149 warn!(error = %e, "取待办失败");
150 tokio::time::sleep(tick).await;
151 }
152 }
153 }
154 }
155
156 pub async fn process(&self, g: &GlobalRow) -> anyhow::Result<()> {
160 match g.trans_type {
161 TransType::Saga => self.process_saga(g).await,
162 TransType::Tcc => self.process_tcc(g).await,
163 TransType::Msg => self.process_msg(g).await,
164 TransType::Xa => self.process_xa(g).await,
165 TransType::Workflow => self.process_workflow(g).await,
166 }
167 }
168
169 async fn process_workflow(&self, g: &GlobalRow) -> anyhow::Result<()> {
176 let (name, input) = crate::workflow::decode_payload(&g.payload);
177 let mut status = g.status;
178
179 loop {
180 let rows = self.store.list_branches(&g.gid).await?;
181 let compensates = compensate_states(&rows);
182
183 match dtmrs_core::workflow_advance(status, &compensates) {
184 Advance::Finish(s) => {
185 info!(gid = %g.gid, status = s.as_str(), "workflow 事务终结");
186 self.store
187 .set_global_status(&g.gid, s, g.trans_type, "")
188 .await?;
189 return Ok(());
190 }
191 Advance::Wait => return Ok(()),
192
193 Advance::RunWorkflow => {
194 let Some(f) = self.workflows.get(&name) else {
195 warn!(gid = %g.gid, workflow = %name,
199 "workflow 未注册,按结果未知处理(会重试,不回滚)");
200 self.retry_later(g).await?;
201 return Ok(());
202 };
203 let ctx =
204 crate::workflow::WorkflowCtx::new(&g.gid, &input, self.store.clone(), rows);
205 match f(ctx).await {
206 Ok(()) => {
207 info!(gid = %g.gid, workflow = %name, "workflow 跑完");
208 self.store
209 .set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
210 .await?;
211 return Ok(());
212 }
213 Err(crate::workflow::WorkflowError::Rollback(reason)) => {
214 info!(gid = %g.gid, workflow = %name, %reason, "workflow 要求回滚");
215 self.store
216 .set_global_status(
217 &g.gid,
218 GlobalStatus::Aborting,
219 g.trans_type,
220 &reason,
221 )
222 .await?;
223 status = GlobalStatus::Aborting;
224 continue;
225 }
226 Err(crate::workflow::WorkflowError::Diverged {
227 branch_id: bid,
228 recorded,
229 got,
230 }) => {
231 warn!(gid = %g.gid, workflow = %name, branch = %bid,
235 %recorded, %got,
236 "workflow 重放走岔了,已停止推进,需要人工介入");
237 self.retry_later(g).await?;
238 return Ok(());
239 }
240 Err(e) => {
241 warn!(gid = %g.gid, workflow = %name, error = %e, "workflow 需要重试");
243 self.retry_later(g).await?;
244 return Ok(());
245 }
246 }
247 }
248
249 Advance::Call { index, op } => {
250 let bid = branch_id(index);
251 let Some(url) = url_of(&rows, &bid, op) else {
252 warn!(gid = %g.gid, branch = %bid, "workflow 补偿地址缺失");
254 self.retry_later(g).await?;
255 return Ok(());
256 };
257 let bp = payload_of(&rows, &bid, op);
258 match self.call_branch(g, &bid, op, &url, &bp).await {
259 BranchResult::Success => {
260 self.store
261 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
262 .await?;
263 }
264 _ => {
266 warn!(gid = %g.gid, branch = %bid, "workflow 补偿未成功,会重试");
267 self.retry_later(g).await?;
268 return Ok(());
269 }
270 }
271 }
272 }
273 }
274 }
275
276 async fn process_saga(&self, g: &GlobalRow) -> anyhow::Result<()> {
279 let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
280 if steps.is_empty() {
281 self.store
282 .set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
283 .await?;
284 return Ok(());
285 }
286 let mut status = g.status;
287
288 loop {
289 let (actions, compensates) = self.branch_states(&g.gid, steps.len()).await?;
290 match saga_advance(status, &actions, &compensates) {
291 Advance::Finish(s) => {
292 if s == GlobalStatus::Aborting {
293 status = s;
295 self.store
296 .set_global_status(&g.gid, s, g.trans_type, "分支已判失败")
297 .await?;
298 continue;
299 }
300 info!(gid = %g.gid, status = s.as_str(), "事务终结");
301 self.store
302 .set_global_status(&g.gid, s, g.trans_type, "")
303 .await?;
304 return Ok(());
305 }
306 Advance::Wait => return Ok(()),
307 Advance::RunWorkflow => {
310 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
311 return Ok(());
312 }
313 Advance::Call { index, op } => {
314 let branch_id = branch_id(index);
315 let url = match op {
316 BranchOp::Action => &steps[index].action,
317 _ => &steps[index].compensate,
318 };
319 match self
320 .call_branch(g, &branch_id, op, url, &steps[index].payload)
321 .await
322 {
323 BranchResult::Success => {
324 self.store
325 .set_branch_status(&g.gid, &branch_id, op, BranchStatus::Succeed)
326 .await?;
327 }
328 BranchResult::Failure => {
329 self.store
330 .set_branch_status(&g.gid, &branch_id, op, BranchStatus::Failed)
331 .await?;
332 if op == BranchOp::Action {
333 info!(gid = %g.gid, branch = %branch_id, "分支要求回滚");
335 status = GlobalStatus::Aborting;
336 self.store
337 .set_global_status(
338 &g.gid,
339 GlobalStatus::Aborting,
340 g.trans_type,
341 &format!("分支 {branch_id} 返回 FAILURE"),
342 )
343 .await?;
344 } else {
345 warn!(gid = %g.gid, branch = %branch_id, "补偿失败,需要人工介入");
347 self.retry_later(g).await?;
348 return Ok(());
349 }
350 }
351 BranchResult::Ongoing | BranchResult::Unknown => {
352 self.retry_later(g).await?;
354 return Ok(());
355 }
356 }
357 }
358 }
359 }
360 }
361
362 async fn process_tcc(&self, g: &GlobalRow) -> anyhow::Result<()> {
368 self.drive_two_phase(g, BranchOp::Confirm, BranchOp::Cancel, "TCC")
369 .await
370 }
371
372 async fn process_xa(&self, g: &GlobalRow) -> anyhow::Result<()> {
383 self.drive_two_phase(g, BranchOp::Commit, BranchOp::Rollback, "XA")
384 .await
385 }
386
387 async fn drive_two_phase(
390 &self,
391 g: &GlobalRow,
392 fwd: BranchOp,
393 bwd: BranchOp,
394 label: &str,
395 ) -> anyhow::Result<()> {
396 let rows = self.store.list_branches(&g.gid).await?;
397 let n = rows
398 .iter()
399 .filter_map(|r| index_of(&r.branch_id))
400 .max()
401 .map(|m| m + 1)
402 .unwrap_or(0);
403 if n == 0 {
404 if !rows.is_empty() {
415 error!(
416 gid = %g.gid,
417 branches = rows.len(),
418 ids = ?rows.iter().map(|r| r.branch_id.as_str()).take(5).collect::<Vec<_>>(),
419 "分支号全都无法解析成下标,无法推进。这笔事务需要人工介入 —— \
420 合法的分支号形如 01 / 02 / 100(见 is_canonical_branch_id)"
421 );
422 return Ok(());
423 }
424 let s = if g.status == GlobalStatus::Aborting {
425 GlobalStatus::Failed
426 } else {
427 GlobalStatus::Succeed
428 };
429 self.store
430 .set_global_status(&g.gid, s, g.trans_type, "")
431 .await?;
432 return Ok(());
433 }
434
435 let status = g.status;
436 loop {
437 let rows = self.store.list_branches(&g.gid).await?;
438 let (f, b) = split_by_op(&rows, n, fwd, bwd);
439 let adv = if fwd == BranchOp::Commit {
440 xa_advance(status, &f, &b)
441 } else {
442 tcc_advance(status, &f, &b)
443 };
444 match adv {
445 Advance::Finish(s) => {
446 info!(gid = %g.gid, status = s.as_str(), mode = label, "事务终结");
447 self.store
448 .set_global_status(&g.gid, s, g.trans_type, "")
449 .await?;
450 return Ok(());
451 }
452 Advance::Wait => return Ok(()),
453 Advance::RunWorkflow => {
456 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
457 return Ok(());
458 }
459 Advance::Call { index, op } => {
460 let bid = branch_id(index);
461 let Some(url) = url_of(&rows, &bid, op) else {
462 warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
463 "分支没登记这个操作的 URL,无法调用");
464 self.retry_later(g).await?;
465 return Ok(());
466 };
467 let bp = payload_of(&rows, &bid, op);
468 match self.call_branch(g, &bid, op, &url, &bp).await {
469 BranchResult::Success => {
470 self.store
471 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
472 .await?;
473 }
474 BranchResult::Failure => {
478 self.store
479 .set_branch_status(&g.gid, &bid, op, BranchStatus::Failed)
480 .await?;
481 warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
482 "二阶段失败,会持续重试,需要人工介入");
483 self.retry_later(g).await?;
484 return Ok(());
485 }
486 BranchResult::Ongoing | BranchResult::Unknown => {
487 self.retry_later(g).await?;
488 return Ok(());
489 }
490 }
491 }
492 }
493 }
494 }
495
496 async fn process_msg(&self, g: &GlobalRow) -> anyhow::Result<()> {
504 let mut status = g.status;
505
506 if status == GlobalStatus::Prepared {
507 if g.query_prepared.is_empty() {
508 warn!(gid = %g.gid, "msg 事务没提供 query_prepared,无法回查,等人处理");
510 self.retry_later(g).await?;
511 return Ok(());
512 }
513 match self
515 .call_branch(g, "00", BranchOp::Action, &g.query_prepared, "")
516 .await
517 {
518 BranchResult::Success => {
519 info!(gid = %g.gid, "回查:本地事务已提交 → 继续推进");
520 self.store
521 .set_global_status(&g.gid, GlobalStatus::Submitted, g.trans_type, "")
522 .await?;
523 status = GlobalStatus::Submitted;
524 }
525 BranchResult::Failure => {
526 info!(gid = %g.gid, "回查:本地事务未提交 → 整单作废");
529 self.store
530 .set_global_status(
531 &g.gid,
532 GlobalStatus::Failed,
533 g.trans_type,
534 "回查得到 FAILURE:本地事务未提交",
535 )
536 .await?;
537 return Ok(());
538 }
539 BranchResult::Ongoing | BranchResult::Unknown => {
540 self.retry_later(g).await?;
542 return Ok(());
543 }
544 }
545 }
546
547 let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
548 if steps.is_empty() {
549 self.store
550 .set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
551 .await?;
552 return Ok(());
553 }
554 loop {
555 let (actions, _) = self.branch_states(&g.gid, steps.len()).await?;
556 match msg_advance(status, &actions) {
557 Advance::Finish(s) => {
558 info!(gid = %g.gid, status = s.as_str(), "消息事务终结");
559 self.store
560 .set_global_status(&g.gid, s, g.trans_type, "")
561 .await?;
562 return Ok(());
563 }
564 Advance::Wait => return Ok(()),
565 Advance::RunWorkflow => {
568 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
569 return Ok(());
570 }
571 Advance::Call { index, op } => {
572 let bid = branch_id(index);
573 match self
574 .call_branch(g, &bid, op, &steps[index].action, &steps[index].payload)
575 .await
576 {
577 BranchResult::Success => {
578 self.store
579 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
580 .await?;
581 }
582 BranchResult::Failure | BranchResult::Ongoing | BranchResult::Unknown => {
584 self.retry_later(g).await?;
585 return Ok(());
586 }
587 }
588 }
589 }
590 }
591 }
592
593 async fn retry_later(&self, g: &GlobalRow) -> anyhow::Result<()> {
594 let iv = dtmrs_core::next_interval_with(g.next_cron_interval, self.retry);
595 self.store.schedule_retry(&g.gid, iv).await?;
596 Ok(())
597 }
598
599 async fn branch_states(
601 &self,
602 gid: &str,
603 n: usize,
604 ) -> anyhow::Result<(Vec<BranchStatus>, Vec<BranchStatus>)> {
605 let rows = self.store.list_branches(gid).await?;
606 let mut actions = vec![BranchStatus::Prepared; n];
607 let mut compensates = vec![BranchStatus::Prepared; n];
608 for r in rows {
609 let Some(i) = index_of(&r.branch_id) else {
610 continue;
611 };
612 if i >= n {
613 continue;
614 }
615 match r.op {
616 BranchOp::Action | BranchOp::Try => actions[i] = r.status,
617 BranchOp::Compensate | BranchOp::Cancel | BranchOp::Rollback => {
618 compensates[i] = r.status
619 }
620 _ => {}
621 }
622 }
623 Ok((actions, compensates))
624 }
625
626 async fn call_branch(
628 &self,
629 g: &GlobalRow,
630 branch_id: &str,
631 op: BranchOp,
632 url: &str,
633 payload: &str,
634 ) -> BranchResult {
635 match parse_target(url) {
636 Target::Local(name) => self.call_local(g, branch_id, op, &name).await,
637 Target::Http(u) => self.call_http(g, branch_id, op, &u, payload).await,
638 #[cfg(feature = "grpc")]
639 Target::Grpc(t) => {
640 self.grpc
641 .call(
642 &t,
643 &g.gid,
644 &g.trans_type.to_string(),
645 branch_id,
646 op.as_str(),
647 )
648 .await
649 }
650 #[cfg(not(feature = "grpc"))]
653 Target::Grpc(t) => {
654 warn!(gid = %g.gid, branch = %branch_id, endpoint = %t.endpoint,
655 "遇到 grpc:// 分支但本次构建关掉了 grpc feature,按结果未知处理");
656 BranchResult::Unknown
657 }
658 }
659 }
660
661 async fn call_local(
663 &self,
664 g: &GlobalRow,
665 branch_id: &str,
666 op: BranchOp,
667 name: &str,
668 ) -> BranchResult {
669 let Some(h) = self.registry.get(name) else {
670 warn!(gid = %g.gid, branch = %branch_id, handler = name,
673 "本地分支未注册,按结果未知处理(会重试,不回滚)");
674 return BranchResult::Unknown;
675 };
676 let ctx = BranchCtx {
677 gid: g.gid.clone(),
678 branch_id: branch_id.to_string(),
679 op,
680 trans_type: g.trans_type.to_string(),
681 };
682 let r = h(ctx).await;
683 info!(gid = %g.gid, branch = %branch_id, op = op.as_str(),
684 handler = name, result = ?r, "本地分支返回");
685 r
686 }
687
688 async fn call_http(
690 &self,
691 g: &GlobalRow,
692 branch_id: &str,
693 op: BranchOp,
694 url: &str,
695 payload: &str,
696 ) -> BranchResult {
697 let req = self
698 .http
699 .post(url)
700 .query(&[
701 ("gid", g.gid.as_str()),
702 ("trans_type", &g.trans_type.to_string()),
703 ("branch_id", branch_id),
704 ("op", op.as_str()),
705 ])
706 .header("content-type", "application/json")
707 .body(branch_payload(payload));
708 match req.send().await {
709 Ok(resp) => {
710 let code = resp.status().as_u16();
711 let body = resp.text().await.unwrap_or_default();
712 let r = BranchResult::from_http(code, &body);
713 info!(gid = %g.gid, branch = %branch_id, op = op.as_str(), code, result = ?r, "分支返回");
714 r
715 }
716 Err(e) => {
717 warn!(gid = %g.gid, branch = %branch_id, error = %e, "分支不可达,结果未知");
719 BranchResult::Unknown
720 }
721 }
722 }
723}
724
725pub fn branch_id(index: usize) -> String {
737 format!("{:02}", index + 1)
738}
739
740pub const MAX_BRANCH_INDEX: usize = 9999;
751
752fn index_of(branch_id: &str) -> Option<usize> {
758 branch_id
759 .parse::<usize>()
760 .ok()
761 .and_then(|v| v.checked_sub(1))
762 .filter(|&i| i <= MAX_BRANCH_INDEX)
763}
764
765pub fn is_canonical_branch_id(s: &str) -> bool {
781 index_of(s).map(branch_id).as_deref() == Some(s)
782}
783
784fn split_by_op(
786 rows: &[dtmrs_store::BranchRow],
787 n: usize,
788 fwd: BranchOp,
789 bwd: BranchOp,
790) -> (Vec<BranchStatus>, Vec<BranchStatus>) {
791 let mut a = vec![BranchStatus::Prepared; n];
792 let mut b = vec![BranchStatus::Prepared; n];
793 for r in rows {
794 let Some(i) = index_of(&r.branch_id) else {
795 continue;
796 };
797 if i >= n {
798 continue;
799 }
800 if r.op == fwd {
801 a[i] = r.status;
802 } else if r.op == bwd {
803 b[i] = r.status;
804 }
805 }
806 (a, b)
807}
808
809fn compensate_states(rows: &[dtmrs_store::BranchRow]) -> Vec<BranchStatus> {
814 let n = rows
815 .iter()
816 .filter(|r| r.op == BranchOp::Compensate)
817 .filter_map(|r| index_of(&r.branch_id))
818 .max()
819 .map(|m| m + 1)
820 .unwrap_or(0);
821 let mut v = vec![BranchStatus::Succeed; n];
824 for r in rows.iter().filter(|r| r.op == BranchOp::Compensate) {
825 if let Some(i) = index_of(&r.branch_id) {
826 if i < n {
827 v[i] = r.status;
828 }
829 }
830 }
831 v
832}
833
834fn payload_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> String {
837 rows.iter()
838 .find(|r| r.branch_id == branch_id && r.op == op)
839 .map(|r| r.payload.clone())
840 .unwrap_or_default()
841}
842
843fn url_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> Option<String> {
844 rows.iter()
845 .find(|r| r.branch_id == branch_id && r.op == op)
846 .map(|r| r.url.clone())
847}
848
849fn branch_payload(step_payload: &str) -> String {
854 if step_payload.trim().is_empty() {
855 "{}".to_string()
856 } else {
857 step_payload.to_string()
858 }
859}
860
861#[derive(Debug, Clone, Copy)]
863pub struct DriverConfig {
864 pub branch_timeout_secs: i64,
866 pub lease_secs: i64,
868 pub retry: dtmrs_core::RetryPolicy,
869 pub workers: usize,
871}
872
873impl Default for DriverConfig {
874 fn default() -> Self {
875 Self {
877 branch_timeout_secs: 10,
878 lease_secs: 30,
879 retry: dtmrs_core::RetryPolicy::default(),
880 workers: 16,
894 }
895 }
896}
897
898impl DriverConfig {
899 pub fn from_env() -> Self {
900 let d = Self::default();
901 let get = |k: &str, fallback: i64| {
902 std::env::var(k)
903 .ok()
904 .and_then(|v| v.parse::<i64>().ok())
905 .filter(|v| *v > 0)
906 .unwrap_or(fallback)
907 };
908 Self {
909 branch_timeout_secs: get("DTMRS_BRANCH_TIMEOUT", d.branch_timeout_secs),
910 lease_secs: get("DTMRS_LEASE", d.lease_secs),
911 retry: dtmrs_core::RetryPolicy::from_env(),
912 workers: get("DTMRS_WORKERS", d.workers as i64) as usize,
913 }
914 }
915}
916
917#[cfg(test)]
918mod tests {
919 use super::*;
920
921 #[test]
922 fn 分支号与下标互转() {
923 assert_eq!(branch_id(0), "01");
924 assert_eq!(branch_id(9), "10");
925 assert_eq!(index_of("01"), Some(0));
926 assert_eq!(index_of("10"), Some(9));
927 assert_eq!(index_of("00"), None);
928 assert_eq!(index_of("xx"), None);
929 }
930
931 #[test]
938 fn 分支号超过99后下标解析仍然正确() {
939 let ids: Vec<String> = (0..500).map(branch_id).collect();
940 for (i, id) in ids.iter().enumerate() {
941 assert_eq!(index_of(id), Some(i), "下标解析错了,执行顺序会乱");
942 }
943 let mut sorted = ids.clone();
945 sorted.sort();
946 assert_ne!(ids, sorted);
947 }
948}