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::{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 let s = if g.status == GlobalStatus::Aborting {
406 GlobalStatus::Failed
407 } else {
408 GlobalStatus::Succeed
409 };
410 self.store
411 .set_global_status(&g.gid, s, g.trans_type, "")
412 .await?;
413 return Ok(());
414 }
415
416 let status = g.status;
417 loop {
418 let rows = self.store.list_branches(&g.gid).await?;
419 let (f, b) = split_by_op(&rows, n, fwd, bwd);
420 let adv = if fwd == BranchOp::Commit {
421 xa_advance(status, &f, &b)
422 } else {
423 tcc_advance(status, &f, &b)
424 };
425 match adv {
426 Advance::Finish(s) => {
427 info!(gid = %g.gid, status = s.as_str(), mode = label, "事务终结");
428 self.store
429 .set_global_status(&g.gid, s, g.trans_type, "")
430 .await?;
431 return Ok(());
432 }
433 Advance::Wait => return Ok(()),
434 Advance::RunWorkflow => {
437 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
438 return Ok(());
439 }
440 Advance::Call { index, op } => {
441 let bid = branch_id(index);
442 let Some(url) = url_of(&rows, &bid, op) else {
443 warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
444 "分支没登记这个操作的 URL,无法调用");
445 self.retry_later(g).await?;
446 return Ok(());
447 };
448 let bp = payload_of(&rows, &bid, op);
449 match self.call_branch(g, &bid, op, &url, &bp).await {
450 BranchResult::Success => {
451 self.store
452 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
453 .await?;
454 }
455 BranchResult::Failure => {
459 self.store
460 .set_branch_status(&g.gid, &bid, op, BranchStatus::Failed)
461 .await?;
462 warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
463 "二阶段失败,会持续重试,需要人工介入");
464 self.retry_later(g).await?;
465 return Ok(());
466 }
467 BranchResult::Ongoing | BranchResult::Unknown => {
468 self.retry_later(g).await?;
469 return Ok(());
470 }
471 }
472 }
473 }
474 }
475 }
476
477 async fn process_msg(&self, g: &GlobalRow) -> anyhow::Result<()> {
485 let mut status = g.status;
486
487 if status == GlobalStatus::Prepared {
488 if g.query_prepared.is_empty() {
489 warn!(gid = %g.gid, "msg 事务没提供 query_prepared,无法回查,等人处理");
491 self.retry_later(g).await?;
492 return Ok(());
493 }
494 match self
496 .call_branch(g, "00", BranchOp::Action, &g.query_prepared, "")
497 .await
498 {
499 BranchResult::Success => {
500 info!(gid = %g.gid, "回查:本地事务已提交 → 继续推进");
501 self.store
502 .set_global_status(&g.gid, GlobalStatus::Submitted, g.trans_type, "")
503 .await?;
504 status = GlobalStatus::Submitted;
505 }
506 BranchResult::Failure => {
507 info!(gid = %g.gid, "回查:本地事务未提交 → 整单作废");
510 self.store
511 .set_global_status(
512 &g.gid,
513 GlobalStatus::Failed,
514 g.trans_type,
515 "回查得到 FAILURE:本地事务未提交",
516 )
517 .await?;
518 return Ok(());
519 }
520 BranchResult::Ongoing | BranchResult::Unknown => {
521 self.retry_later(g).await?;
523 return Ok(());
524 }
525 }
526 }
527
528 let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
529 if steps.is_empty() {
530 self.store
531 .set_global_status(&g.gid, GlobalStatus::Succeed, g.trans_type, "")
532 .await?;
533 return Ok(());
534 }
535 loop {
536 let (actions, _) = self.branch_states(&g.gid, steps.len()).await?;
537 match msg_advance(status, &actions) {
538 Advance::Finish(s) => {
539 info!(gid = %g.gid, status = s.as_str(), "消息事务终结");
540 self.store
541 .set_global_status(&g.gid, s, g.trans_type, "")
542 .await?;
543 return Ok(());
544 }
545 Advance::Wait => return Ok(()),
546 Advance::RunWorkflow => {
549 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
550 return Ok(());
551 }
552 Advance::Call { index, op } => {
553 let bid = branch_id(index);
554 match self
555 .call_branch(g, &bid, op, &steps[index].action, &steps[index].payload)
556 .await
557 {
558 BranchResult::Success => {
559 self.store
560 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
561 .await?;
562 }
563 BranchResult::Failure | BranchResult::Ongoing | BranchResult::Unknown => {
565 self.retry_later(g).await?;
566 return Ok(());
567 }
568 }
569 }
570 }
571 }
572 }
573
574 async fn retry_later(&self, g: &GlobalRow) -> anyhow::Result<()> {
575 let iv = dtmrs_core::next_interval_with(g.next_cron_interval, self.retry);
576 self.store.schedule_retry(&g.gid, iv).await?;
577 Ok(())
578 }
579
580 async fn branch_states(
582 &self,
583 gid: &str,
584 n: usize,
585 ) -> anyhow::Result<(Vec<BranchStatus>, Vec<BranchStatus>)> {
586 let rows = self.store.list_branches(gid).await?;
587 let mut actions = vec![BranchStatus::Prepared; n];
588 let mut compensates = vec![BranchStatus::Prepared; n];
589 for r in rows {
590 let Some(i) = index_of(&r.branch_id) else {
591 continue;
592 };
593 if i >= n {
594 continue;
595 }
596 match r.op {
597 BranchOp::Action | BranchOp::Try => actions[i] = r.status,
598 BranchOp::Compensate | BranchOp::Cancel | BranchOp::Rollback => {
599 compensates[i] = r.status
600 }
601 _ => {}
602 }
603 }
604 Ok((actions, compensates))
605 }
606
607 async fn call_branch(
609 &self,
610 g: &GlobalRow,
611 branch_id: &str,
612 op: BranchOp,
613 url: &str,
614 payload: &str,
615 ) -> BranchResult {
616 match parse_target(url) {
617 Target::Local(name) => self.call_local(g, branch_id, op, &name).await,
618 Target::Http(u) => self.call_http(g, branch_id, op, &u, payload).await,
619 #[cfg(feature = "grpc")]
620 Target::Grpc(t) => {
621 self.grpc
622 .call(
623 &t,
624 &g.gid,
625 &g.trans_type.to_string(),
626 branch_id,
627 op.as_str(),
628 )
629 .await
630 }
631 #[cfg(not(feature = "grpc"))]
634 Target::Grpc(t) => {
635 warn!(gid = %g.gid, branch = %branch_id, endpoint = %t.endpoint,
636 "遇到 grpc:// 分支但本次构建关掉了 grpc feature,按结果未知处理");
637 BranchResult::Unknown
638 }
639 }
640 }
641
642 async fn call_local(
644 &self,
645 g: &GlobalRow,
646 branch_id: &str,
647 op: BranchOp,
648 name: &str,
649 ) -> BranchResult {
650 let Some(h) = self.registry.get(name) else {
651 warn!(gid = %g.gid, branch = %branch_id, handler = name,
654 "本地分支未注册,按结果未知处理(会重试,不回滚)");
655 return BranchResult::Unknown;
656 };
657 let ctx = BranchCtx {
658 gid: g.gid.clone(),
659 branch_id: branch_id.to_string(),
660 op,
661 trans_type: g.trans_type.to_string(),
662 };
663 let r = h(ctx).await;
664 info!(gid = %g.gid, branch = %branch_id, op = op.as_str(),
665 handler = name, result = ?r, "本地分支返回");
666 r
667 }
668
669 async fn call_http(
671 &self,
672 g: &GlobalRow,
673 branch_id: &str,
674 op: BranchOp,
675 url: &str,
676 payload: &str,
677 ) -> BranchResult {
678 let req = self
679 .http
680 .post(url)
681 .query(&[
682 ("gid", g.gid.as_str()),
683 ("trans_type", &g.trans_type.to_string()),
684 ("branch_id", branch_id),
685 ("op", op.as_str()),
686 ])
687 .header("content-type", "application/json")
688 .body(branch_payload(payload));
689 match req.send().await {
690 Ok(resp) => {
691 let code = resp.status().as_u16();
692 let body = resp.text().await.unwrap_or_default();
693 let r = BranchResult::from_http(code, &body);
694 info!(gid = %g.gid, branch = %branch_id, op = op.as_str(), code, result = ?r, "分支返回");
695 r
696 }
697 Err(e) => {
698 warn!(gid = %g.gid, branch = %branch_id, error = %e, "分支不可达,结果未知");
700 BranchResult::Unknown
701 }
702 }
703 }
704}
705
706pub fn branch_id(index: usize) -> String {
708 format!("{:02}", index + 1)
709}
710
711fn index_of(branch_id: &str) -> Option<usize> {
712 branch_id
713 .parse::<usize>()
714 .ok()
715 .and_then(|v| v.checked_sub(1))
716}
717
718fn split_by_op(
720 rows: &[dtmrs_store::BranchRow],
721 n: usize,
722 fwd: BranchOp,
723 bwd: BranchOp,
724) -> (Vec<BranchStatus>, Vec<BranchStatus>) {
725 let mut a = vec![BranchStatus::Prepared; n];
726 let mut b = vec![BranchStatus::Prepared; n];
727 for r in rows {
728 let Some(i) = index_of(&r.branch_id) else {
729 continue;
730 };
731 if i >= n {
732 continue;
733 }
734 if r.op == fwd {
735 a[i] = r.status;
736 } else if r.op == bwd {
737 b[i] = r.status;
738 }
739 }
740 (a, b)
741}
742
743fn compensate_states(rows: &[dtmrs_store::BranchRow]) -> Vec<BranchStatus> {
748 let n = rows
749 .iter()
750 .filter(|r| r.op == BranchOp::Compensate)
751 .filter_map(|r| index_of(&r.branch_id))
752 .max()
753 .map(|m| m + 1)
754 .unwrap_or(0);
755 let mut v = vec![BranchStatus::Succeed; n];
758 for r in rows.iter().filter(|r| r.op == BranchOp::Compensate) {
759 if let Some(i) = index_of(&r.branch_id) {
760 if i < n {
761 v[i] = r.status;
762 }
763 }
764 }
765 v
766}
767
768fn payload_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> String {
771 rows.iter()
772 .find(|r| r.branch_id == branch_id && r.op == op)
773 .map(|r| r.payload.clone())
774 .unwrap_or_default()
775}
776
777fn url_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> Option<String> {
778 rows.iter()
779 .find(|r| r.branch_id == branch_id && r.op == op)
780 .map(|r| r.url.clone())
781}
782
783fn branch_payload(step_payload: &str) -> String {
788 if step_payload.trim().is_empty() {
789 "{}".to_string()
790 } else {
791 step_payload.to_string()
792 }
793}
794
795#[derive(Debug, Clone, Copy)]
797pub struct DriverConfig {
798 pub branch_timeout_secs: i64,
800 pub lease_secs: i64,
802 pub retry: dtmrs_core::RetryPolicy,
803 pub workers: usize,
805}
806
807impl Default for DriverConfig {
808 fn default() -> Self {
809 Self {
811 branch_timeout_secs: 10,
812 lease_secs: 30,
813 retry: dtmrs_core::RetryPolicy::default(),
814 workers: 16,
828 }
829 }
830}
831
832impl DriverConfig {
833 pub fn from_env() -> Self {
834 let d = Self::default();
835 let get = |k: &str, fallback: i64| {
836 std::env::var(k)
837 .ok()
838 .and_then(|v| v.parse::<i64>().ok())
839 .filter(|v| *v > 0)
840 .unwrap_or(fallback)
841 };
842 Self {
843 branch_timeout_secs: get("DTMRS_BRANCH_TIMEOUT", d.branch_timeout_secs),
844 lease_secs: get("DTMRS_LEASE", d.lease_secs),
845 retry: dtmrs_core::RetryPolicy::from_env(),
846 workers: get("DTMRS_WORKERS", d.workers as i64) as usize,
847 }
848 }
849}
850
851#[cfg(test)]
852mod tests {
853 use super::*;
854
855 #[test]
856 fn 分支号与下标互转() {
857 assert_eq!(branch_id(0), "01");
858 assert_eq!(branch_id(9), "10");
859 assert_eq!(index_of("01"), Some(0));
860 assert_eq!(index_of("10"), Some(9));
861 assert_eq!(index_of("00"), None);
862 assert_eq!(index_of("xx"), None);
863 }
864}