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 #[cfg(feature = "grpc")]
109 pub fn with_grpc_ca_pem(mut self, pem: impl Into<Vec<u8>>) -> Self {
110 self.grpc = self.grpc.with_ca_pem(pem);
111 self
112 }
113
114 pub async fn run_forever(self, tick: Duration) {
142 let mut set = tokio::task::JoinSet::new();
143 for _ in 0..self.workers.max(1) {
144 let d = self.clone();
145 set.spawn(async move { d.worker_loop(tick).await });
146 }
147 set.join_next().await;
149 }
150
151 async fn worker_loop(&self, tick: Duration) {
153 loop {
154 match self.store.lock_one_due(&self.owner, self.lease).await {
155 Ok(Some(g)) => {
156 if let Err(e) = self.process(&g).await {
157 warn!(gid = %g.gid, error = %e, "推进出错,等下轮重试");
158 }
159 }
160 Ok(None) => tokio::time::sleep(tick).await,
161 Err(e) => {
162 warn!(error = %e, "取待办失败");
163 tokio::time::sleep(tick).await;
164 }
165 }
166 }
167 }
168
169 pub async fn process(&self, g: &GlobalRow) -> anyhow::Result<()> {
173 match g.trans_type {
174 TransType::Saga => self.process_saga(g).await,
175 TransType::Tcc => self.process_tcc(g).await,
176 TransType::Msg => self.process_msg(g).await,
177 TransType::Xa => self.process_xa(g).await,
178 TransType::Workflow => self.process_workflow(g).await,
179 }
180 }
181
182 async fn process_workflow(&self, g: &GlobalRow) -> anyhow::Result<()> {
189 let (name, input) = crate::workflow::decode_payload(&g.payload);
190 let mut status = g.status;
191
192 loop {
193 let rows = self.store.list_branches(&g.gid).await?;
194 let compensates = compensate_states(&rows);
195
196 match dtmrs_core::workflow_advance(status, &compensates) {
197 Advance::Finish(s) => {
198 info!(gid = %g.gid, status = s.as_str(), "workflow 事务终结");
199 self.store
200 .set_global_status(&g.gid, s, g.trans_type, "")
201 .await?;
202 return Ok(());
203 }
204 Advance::Wait => return Ok(()),
205
206 Advance::RunWorkflow => {
207 let Some(f) = self.workflows.get(&name) else {
208 warn!(gid = %g.gid, workflow = %name,
212 "workflow 未注册,按结果未知处理(会重试,不回滚)");
213 self.retry_later(g).await?;
214 return Ok(());
215 };
216 let ctx =
217 crate::workflow::WorkflowCtx::new(&g.gid, &input, self.store.clone(), rows);
218 match f(ctx).await {
219 Ok(()) => {
220 info!(gid = %g.gid, workflow = %name, "workflow 跑完");
221 self.store
222 .set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
223 .await?;
224 return Ok(());
225 }
226 Err(crate::workflow::WorkflowError::Rollback(reason)) => {
227 info!(gid = %g.gid, workflow = %name, %reason, "workflow 要求回滚");
228 self.store
229 .set_global_status(
230 &g.gid,
231 GlobalStatus::Aborting,
232 g.trans_type,
233 &reason,
234 )
235 .await?;
236 status = GlobalStatus::Aborting;
237 continue;
238 }
239 Err(crate::workflow::WorkflowError::Diverged {
240 branch_id: bid,
241 recorded,
242 got,
243 }) => {
244 warn!(gid = %g.gid, workflow = %name, branch = %bid,
248 %recorded, %got,
249 "workflow 重放走岔了,已停止推进,需要人工介入");
250 self.retry_later(g).await?;
251 return Ok(());
252 }
253 Err(e) => {
254 warn!(gid = %g.gid, workflow = %name, error = %e, "workflow 需要重试");
256 self.retry_later(g).await?;
257 return Ok(());
258 }
259 }
260 }
261
262 Advance::Call { index, op } => {
263 let bid = branch_id(index);
264 let Some(url) = url_of(&rows, &bid, op) else {
265 warn!(gid = %g.gid, branch = %bid, "workflow 补偿地址缺失");
267 self.retry_later(g).await?;
268 return Ok(());
269 };
270 let bp = payload_of(&rows, &bid, op);
271 match self.call_branch(g, &bid, op, &url, &bp).await {
272 BranchResult::Success => {
273 self.store
274 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
275 .await?;
276 }
277 _ => {
279 warn!(gid = %g.gid, branch = %bid, "workflow 补偿未成功,会重试");
280 self.retry_later(g).await?;
281 return Ok(());
282 }
283 }
284 }
285 }
286 }
287 }
288
289 async fn process_saga(&self, g: &GlobalRow) -> anyhow::Result<()> {
292 let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
293 if steps.is_empty() {
294 self.store
295 .set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
296 .await?;
297 return Ok(());
298 }
299 let mut status = g.status;
300
301 loop {
302 let (actions, compensates) = self.branch_states(&g.gid, steps.len()).await?;
303 match saga_advance(status, &actions, &compensates) {
304 Advance::Finish(s) => {
305 if s == GlobalStatus::Aborting {
306 status = s;
308 self.store
309 .set_global_status(&g.gid, s, g.trans_type, "分支已判失败")
310 .await?;
311 continue;
312 }
313 info!(gid = %g.gid, status = s.as_str(), "事务终结");
314 self.store
315 .set_global_status(&g.gid, s, g.trans_type, "")
316 .await?;
317 return Ok(());
318 }
319 Advance::Wait => return Ok(()),
320 Advance::RunWorkflow => {
323 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
324 return Ok(());
325 }
326 Advance::Call { index, op } => {
327 let branch_id = branch_id(index);
328 let url = match op {
329 BranchOp::Action => &steps[index].action,
330 _ => &steps[index].compensate,
331 };
332 match self
333 .call_branch(g, &branch_id, op, url, &steps[index].payload)
334 .await
335 {
336 BranchResult::Success => {
337 self.store
338 .set_branch_status(&g.gid, &branch_id, op, BranchStatus::Succeed)
339 .await?;
340 }
341 BranchResult::Failure => {
342 self.store
343 .set_branch_status(&g.gid, &branch_id, op, BranchStatus::Failed)
344 .await?;
345 if op == BranchOp::Action {
346 info!(gid = %g.gid, branch = %branch_id, "分支要求回滚");
348 status = GlobalStatus::Aborting;
349 self.store
350 .set_global_status(
351 &g.gid,
352 GlobalStatus::Aborting,
353 g.trans_type,
354 &format!("分支 {branch_id} 返回 FAILURE"),
355 )
356 .await?;
357 } else {
358 warn!(gid = %g.gid, branch = %branch_id, "补偿失败,需要人工介入");
360 self.retry_later(g).await?;
361 return Ok(());
362 }
363 }
364 BranchResult::Ongoing | BranchResult::Unknown => {
365 self.retry_later(g).await?;
367 return Ok(());
368 }
369 }
370 }
371 }
372 }
373 }
374
375 async fn process_tcc(&self, g: &GlobalRow) -> anyhow::Result<()> {
381 self.drive_two_phase(g, BranchOp::Confirm, BranchOp::Cancel, "TCC")
382 .await
383 }
384
385 async fn process_xa(&self, g: &GlobalRow) -> anyhow::Result<()> {
396 self.drive_two_phase(g, BranchOp::Commit, BranchOp::Rollback, "XA")
397 .await
398 }
399
400 async fn drive_two_phase(
403 &self,
404 g: &GlobalRow,
405 fwd: BranchOp,
406 bwd: BranchOp,
407 label: &str,
408 ) -> anyhow::Result<()> {
409 let rows = self.store.list_branches(&g.gid).await?;
410 let n = rows
411 .iter()
412 .filter_map(|r| index_of(&r.branch_id))
413 .max()
414 .map(|m| m + 1)
415 .unwrap_or(0);
416 if n == 0 {
417 if !rows.is_empty() {
428 error!(
429 gid = %g.gid,
430 branches = rows.len(),
431 ids = ?rows.iter().map(|r| r.branch_id.as_str()).take(5).collect::<Vec<_>>(),
432 "分支号全都无法解析成下标,无法推进。这笔事务需要人工介入 —— \
433 合法的分支号形如 01 / 02 / 100(见 is_canonical_branch_id)"
434 );
435 return Ok(());
436 }
437 let s = if g.status == GlobalStatus::Aborting {
438 GlobalStatus::Failed
439 } else {
440 GlobalStatus::Succeed
441 };
442 self.store
443 .set_global_status(&g.gid, s, g.trans_type, "")
444 .await?;
445 return Ok(());
446 }
447
448 let status = g.status;
449 loop {
450 let rows = self.store.list_branches(&g.gid).await?;
451 let (f, b) = split_by_op(&rows, n, fwd, bwd);
452 let adv = if fwd == BranchOp::Commit {
453 xa_advance(status, &f, &b)
454 } else {
455 tcc_advance(status, &f, &b)
456 };
457 match adv {
458 Advance::Finish(s) => {
459 info!(gid = %g.gid, status = s.as_str(), mode = label, "事务终结");
460 self.store
461 .set_global_status(&g.gid, s, g.trans_type, "")
462 .await?;
463 return Ok(());
464 }
465 Advance::Wait => return Ok(()),
466 Advance::RunWorkflow => {
469 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
470 return Ok(());
471 }
472 Advance::Call { index, op } => {
473 let bid = branch_id(index);
474 let Some(url) = url_of(&rows, &bid, op) else {
475 warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
476 "分支没登记这个操作的 URL,无法调用");
477 self.retry_later(g).await?;
478 return Ok(());
479 };
480 let bp = payload_of(&rows, &bid, op);
481 match self.call_branch(g, &bid, op, &url, &bp).await {
482 BranchResult::Success => {
483 self.store
484 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
485 .await?;
486 }
487 BranchResult::Failure => {
491 self.store
492 .set_branch_status(&g.gid, &bid, op, BranchStatus::Failed)
493 .await?;
494 warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
495 "二阶段失败,会持续重试,需要人工介入");
496 self.retry_later(g).await?;
497 return Ok(());
498 }
499 BranchResult::Ongoing | BranchResult::Unknown => {
500 self.retry_later(g).await?;
501 return Ok(());
502 }
503 }
504 }
505 }
506 }
507 }
508
509 async fn process_msg(&self, g: &GlobalRow) -> anyhow::Result<()> {
517 let mut status = g.status;
518
519 if status == GlobalStatus::Prepared {
520 if g.query_prepared.is_empty() {
521 warn!(gid = %g.gid, "msg 事务没提供 query_prepared,无法回查,等人处理");
523 self.retry_later(g).await?;
524 return Ok(());
525 }
526 match self
528 .call_branch(g, "00", BranchOp::Action, &g.query_prepared, "")
529 .await
530 {
531 BranchResult::Success => {
532 info!(gid = %g.gid, "回查:本地事务已提交 → 继续推进");
533 self.store
534 .set_global_status(&g.gid, GlobalStatus::Submitted, g.trans_type, "")
535 .await?;
536 status = GlobalStatus::Submitted;
537 }
538 BranchResult::Failure => {
539 info!(gid = %g.gid, "回查:本地事务未提交 → 整单作废");
542 self.store
543 .set_global_status(
544 &g.gid,
545 GlobalStatus::Failed,
546 g.trans_type,
547 "回查得到 FAILURE:本地事务未提交",
548 )
549 .await?;
550 return Ok(());
551 }
552 BranchResult::Ongoing | BranchResult::Unknown => {
553 self.retry_later(g).await?;
555 return Ok(());
556 }
557 }
558 }
559
560 let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
561 if steps.is_empty() {
562 self.store
563 .set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
564 .await?;
565 return Ok(());
566 }
567 loop {
568 let (actions, _) = self.branch_states(&g.gid, steps.len()).await?;
569 match msg_advance(status, &actions) {
570 Advance::Finish(s) => {
571 info!(gid = %g.gid, status = s.as_str(), "消息事务终结");
572 self.store
573 .set_global_status(&g.gid, s, g.trans_type, "")
574 .await?;
575 return Ok(());
576 }
577 Advance::Wait => return Ok(()),
578 Advance::RunWorkflow => {
581 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
582 return Ok(());
583 }
584 Advance::Call { index, op } => {
585 let bid = branch_id(index);
586 match self
587 .call_branch(g, &bid, op, &steps[index].action, &steps[index].payload)
588 .await
589 {
590 BranchResult::Success => {
591 self.store
592 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
593 .await?;
594 }
595 BranchResult::Failure | BranchResult::Ongoing | BranchResult::Unknown => {
597 self.retry_later(g).await?;
598 return Ok(());
599 }
600 }
601 }
602 }
603 }
604 }
605
606 async fn retry_later(&self, g: &GlobalRow) -> anyhow::Result<()> {
607 let iv = dtmrs_core::next_interval_with(g.next_cron_interval, self.retry);
608 self.store.schedule_retry(&g.gid, iv).await?;
609 Ok(())
610 }
611
612 async fn branch_states(
614 &self,
615 gid: &str,
616 n: usize,
617 ) -> anyhow::Result<(Vec<BranchStatus>, Vec<BranchStatus>)> {
618 let rows = self.store.list_branches(gid).await?;
619 let mut actions = vec![BranchStatus::Prepared; n];
620 let mut compensates = vec![BranchStatus::Prepared; n];
621 for r in rows {
622 let Some(i) = index_of(&r.branch_id) else {
623 continue;
624 };
625 if i >= n {
626 continue;
627 }
628 match r.op {
629 BranchOp::Action | BranchOp::Try => actions[i] = r.status,
630 BranchOp::Compensate | BranchOp::Cancel | BranchOp::Rollback => {
631 compensates[i] = r.status
632 }
633 _ => {}
634 }
635 }
636 Ok((actions, compensates))
637 }
638
639 async fn call_branch(
641 &self,
642 g: &GlobalRow,
643 branch_id: &str,
644 op: BranchOp,
645 url: &str,
646 payload: &str,
647 ) -> BranchResult {
648 match parse_target(url) {
649 Target::Local(name) => self.call_local(g, branch_id, op, &name).await,
650 Target::Http(u) => self.call_http(g, branch_id, op, &u, payload).await,
651 #[cfg(feature = "grpc")]
652 Target::Grpc(t) => {
653 self.grpc
654 .call(
655 &t,
656 &g.gid,
657 &g.trans_type.to_string(),
658 branch_id,
659 op.as_str(),
660 )
661 .await
662 }
663 #[cfg(not(feature = "grpc"))]
666 Target::Grpc(t) => {
667 warn!(gid = %g.gid, branch = %branch_id, endpoint = %t.endpoint,
668 "遇到 grpc:// 分支但本次构建关掉了 grpc feature,按结果未知处理");
669 BranchResult::Unknown
670 }
671 }
672 }
673
674 async fn call_local(
676 &self,
677 g: &GlobalRow,
678 branch_id: &str,
679 op: BranchOp,
680 name: &str,
681 ) -> BranchResult {
682 let Some(h) = self.registry.get(name) else {
683 warn!(gid = %g.gid, branch = %branch_id, handler = name,
686 "本地分支未注册,按结果未知处理(会重试,不回滚)");
687 return BranchResult::Unknown;
688 };
689 let ctx = BranchCtx {
690 gid: g.gid.clone(),
691 branch_id: branch_id.to_string(),
692 op,
693 trans_type: g.trans_type.to_string(),
694 };
695 let r = h(ctx).await;
696 info!(gid = %g.gid, branch = %branch_id, op = op.as_str(),
697 handler = name, result = ?r, "本地分支返回");
698 r
699 }
700
701 async fn call_http(
703 &self,
704 g: &GlobalRow,
705 branch_id: &str,
706 op: BranchOp,
707 url: &str,
708 payload: &str,
709 ) -> BranchResult {
710 let req = self
711 .http
712 .post(url)
713 .query(&[
714 ("gid", g.gid.as_str()),
715 ("trans_type", &g.trans_type.to_string()),
716 ("branch_id", branch_id),
717 ("op", op.as_str()),
718 ])
719 .header("content-type", "application/json")
720 .body(branch_payload(payload));
721 match req.send().await {
722 Ok(resp) => {
723 let code = resp.status().as_u16();
724 let body = resp.text().await.unwrap_or_default();
725 let r = BranchResult::from_http(code, &body);
726 info!(gid = %g.gid, branch = %branch_id, op = op.as_str(), code, result = ?r, "分支返回");
727 r
728 }
729 Err(e) => {
730 warn!(gid = %g.gid, branch = %branch_id, error = %e, "分支不可达,结果未知");
732 BranchResult::Unknown
733 }
734 }
735 }
736}
737
738pub fn branch_id(index: usize) -> String {
750 format!("{:02}", index + 1)
751}
752
753pub const MAX_BRANCH_INDEX: usize = 9999;
764
765fn index_of(branch_id: &str) -> Option<usize> {
771 branch_id
772 .parse::<usize>()
773 .ok()
774 .and_then(|v| v.checked_sub(1))
775 .filter(|&i| i <= MAX_BRANCH_INDEX)
776}
777
778pub fn is_canonical_branch_id(s: &str) -> bool {
794 index_of(s).map(branch_id).as_deref() == Some(s)
795}
796
797fn split_by_op(
799 rows: &[dtmrs_store::BranchRow],
800 n: usize,
801 fwd: BranchOp,
802 bwd: BranchOp,
803) -> (Vec<BranchStatus>, Vec<BranchStatus>) {
804 let mut a = vec![BranchStatus::Prepared; n];
805 let mut b = vec![BranchStatus::Prepared; n];
806 for r in rows {
807 let Some(i) = index_of(&r.branch_id) else {
808 continue;
809 };
810 if i >= n {
811 continue;
812 }
813 if r.op == fwd {
814 a[i] = r.status;
815 } else if r.op == bwd {
816 b[i] = r.status;
817 }
818 }
819 (a, b)
820}
821
822fn compensate_states(rows: &[dtmrs_store::BranchRow]) -> Vec<BranchStatus> {
827 let n = rows
828 .iter()
829 .filter(|r| r.op == BranchOp::Compensate)
830 .filter_map(|r| index_of(&r.branch_id))
831 .max()
832 .map(|m| m + 1)
833 .unwrap_or(0);
834 let mut v = vec![BranchStatus::Succeed; n];
837 for r in rows.iter().filter(|r| r.op == BranchOp::Compensate) {
838 if let Some(i) = index_of(&r.branch_id) {
839 if i < n {
840 v[i] = r.status;
841 }
842 }
843 }
844 v
845}
846
847fn payload_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> String {
850 rows.iter()
851 .find(|r| r.branch_id == branch_id && r.op == op)
852 .map(|r| r.payload.clone())
853 .unwrap_or_default()
854}
855
856fn url_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> Option<String> {
857 rows.iter()
858 .find(|r| r.branch_id == branch_id && r.op == op)
859 .map(|r| r.url.clone())
860}
861
862fn branch_payload(step_payload: &str) -> String {
867 if step_payload.trim().is_empty() {
868 "{}".to_string()
869 } else {
870 step_payload.to_string()
871 }
872}
873
874#[derive(Debug, Clone, Copy)]
876pub struct DriverConfig {
877 pub branch_timeout_secs: i64,
879 pub lease_secs: i64,
881 pub retry: dtmrs_core::RetryPolicy,
882 pub workers: usize,
884}
885
886impl Default for DriverConfig {
887 fn default() -> Self {
888 Self {
890 branch_timeout_secs: 10,
891 lease_secs: 30,
892 retry: dtmrs_core::RetryPolicy::default(),
893 workers: 16,
907 }
908 }
909}
910
911impl DriverConfig {
912 pub fn from_env() -> Self {
913 let d = Self::default();
914 let get = |k: &str, fallback: i64| {
915 std::env::var(k)
916 .ok()
917 .and_then(|v| v.parse::<i64>().ok())
918 .filter(|v| *v > 0)
919 .unwrap_or(fallback)
920 };
921 Self {
922 branch_timeout_secs: get("DTMRS_BRANCH_TIMEOUT", d.branch_timeout_secs),
923 lease_secs: get("DTMRS_LEASE", d.lease_secs),
924 retry: dtmrs_core::RetryPolicy::from_env(),
925 workers: get("DTMRS_WORKERS", d.workers as i64) as usize,
926 }
927 }
928}
929
930#[cfg(test)]
931mod tests {
932 use super::*;
933
934 #[test]
935 fn 分支号与下标互转() {
936 assert_eq!(branch_id(0), "01");
937 assert_eq!(branch_id(9), "10");
938 assert_eq!(index_of("01"), Some(0));
939 assert_eq!(index_of("10"), Some(9));
940 assert_eq!(index_of("00"), None);
941 assert_eq!(index_of("xx"), None);
942 }
943
944 #[test]
951 fn 分支号超过99后下标解析仍然正确() {
952 let ids: Vec<String> = (0..500).map(branch_id).collect();
953 for (i, id) in ids.iter().enumerate() {
954 assert_eq!(index_of(id), Some(i), "下标解析错了,执行顺序会乱");
955 }
956 let mut sorted = ids.clone();
958 sorted.sort();
959 assert_ne!(ids, sorted);
960 }
961}