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.set_global_status(&g.gid, s, "").await?;
187 return Ok(());
188 }
189 Advance::Wait => return Ok(()),
190
191 Advance::RunWorkflow => {
192 let Some(f) = self.workflows.get(&name) else {
193 warn!(gid = %g.gid, workflow = %name,
197 "workflow 未注册,按结果未知处理(会重试,不回滚)");
198 self.retry_later(g).await?;
199 return Ok(());
200 };
201 let ctx =
202 crate::workflow::WorkflowCtx::new(&g.gid, &input, self.store.clone(), rows);
203 match f(ctx).await {
204 Ok(()) => {
205 info!(gid = %g.gid, workflow = %name, "workflow 跑完");
206 self.store
207 .set_global_status(&g.gid, GlobalStatus::Succeed, "")
208 .await?;
209 return Ok(());
210 }
211 Err(crate::workflow::WorkflowError::Rollback(reason)) => {
212 info!(gid = %g.gid, workflow = %name, %reason, "workflow 要求回滚");
213 self.store
214 .set_global_status(&g.gid, GlobalStatus::Aborting, &reason)
215 .await?;
216 status = GlobalStatus::Aborting;
217 continue;
218 }
219 Err(crate::workflow::WorkflowError::Diverged {
220 branch_id: bid,
221 recorded,
222 got,
223 }) => {
224 warn!(gid = %g.gid, workflow = %name, branch = %bid,
228 %recorded, %got,
229 "workflow 重放走岔了,已停止推进,需要人工介入");
230 self.retry_later(g).await?;
231 return Ok(());
232 }
233 Err(e) => {
234 warn!(gid = %g.gid, workflow = %name, error = %e, "workflow 需要重试");
236 self.retry_later(g).await?;
237 return Ok(());
238 }
239 }
240 }
241
242 Advance::Call { index, op } => {
243 let bid = branch_id(index);
244 let Some(url) = url_of(&rows, &bid, op) else {
245 warn!(gid = %g.gid, branch = %bid, "workflow 补偿地址缺失");
247 self.retry_later(g).await?;
248 return Ok(());
249 };
250 let bp = payload_of(&rows, &bid, op);
251 match self.call_branch(g, &bid, op, &url, &bp).await {
252 BranchResult::Success => {
253 self.store
254 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
255 .await?;
256 }
257 _ => {
259 warn!(gid = %g.gid, branch = %bid, "workflow 补偿未成功,会重试");
260 self.retry_later(g).await?;
261 return Ok(());
262 }
263 }
264 }
265 }
266 }
267 }
268
269 async fn process_saga(&self, g: &GlobalRow) -> anyhow::Result<()> {
272 let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
273 if steps.is_empty() {
274 self.store
275 .set_global_status(&g.gid, GlobalStatus::Succeed, "")
276 .await?;
277 return Ok(());
278 }
279 let mut status = g.status;
280
281 loop {
282 let (actions, compensates) = self.branch_states(&g.gid, steps.len()).await?;
283 match saga_advance(status, &actions, &compensates) {
284 Advance::Finish(s) => {
285 if s == GlobalStatus::Aborting {
286 status = s;
288 self.store
289 .set_global_status(&g.gid, s, "分支已判失败")
290 .await?;
291 continue;
292 }
293 info!(gid = %g.gid, status = s.as_str(), "事务终结");
294 self.store.set_global_status(&g.gid, s, "").await?;
295 return Ok(());
296 }
297 Advance::Wait => return Ok(()),
298 Advance::RunWorkflow => {
301 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
302 return Ok(());
303 }
304 Advance::Call { index, op } => {
305 let branch_id = branch_id(index);
306 let url = match op {
307 BranchOp::Action => &steps[index].action,
308 _ => &steps[index].compensate,
309 };
310 match self
311 .call_branch(g, &branch_id, op, url, &steps[index].payload)
312 .await
313 {
314 BranchResult::Success => {
315 self.store
316 .set_branch_status(&g.gid, &branch_id, op, BranchStatus::Succeed)
317 .await?;
318 }
319 BranchResult::Failure => {
320 self.store
321 .set_branch_status(&g.gid, &branch_id, op, BranchStatus::Failed)
322 .await?;
323 if op == BranchOp::Action {
324 info!(gid = %g.gid, branch = %branch_id, "分支要求回滚");
326 status = GlobalStatus::Aborting;
327 self.store
328 .set_global_status(
329 &g.gid,
330 GlobalStatus::Aborting,
331 &format!("分支 {branch_id} 返回 FAILURE"),
332 )
333 .await?;
334 } else {
335 warn!(gid = %g.gid, branch = %branch_id, "补偿失败,需要人工介入");
337 self.retry_later(g).await?;
338 return Ok(());
339 }
340 }
341 BranchResult::Ongoing | BranchResult::Unknown => {
342 self.retry_later(g).await?;
344 return Ok(());
345 }
346 }
347 }
348 }
349 }
350 }
351
352 async fn process_tcc(&self, g: &GlobalRow) -> anyhow::Result<()> {
358 self.drive_two_phase(g, BranchOp::Confirm, BranchOp::Cancel, "TCC")
359 .await
360 }
361
362 async fn process_xa(&self, g: &GlobalRow) -> anyhow::Result<()> {
373 self.drive_two_phase(g, BranchOp::Commit, BranchOp::Rollback, "XA")
374 .await
375 }
376
377 async fn drive_two_phase(
380 &self,
381 g: &GlobalRow,
382 fwd: BranchOp,
383 bwd: BranchOp,
384 label: &str,
385 ) -> anyhow::Result<()> {
386 let rows = self.store.list_branches(&g.gid).await?;
387 let n = rows
388 .iter()
389 .filter_map(|r| index_of(&r.branch_id))
390 .max()
391 .map(|m| m + 1)
392 .unwrap_or(0);
393 if n == 0 {
394 let s = if g.status == GlobalStatus::Aborting {
396 GlobalStatus::Failed
397 } else {
398 GlobalStatus::Succeed
399 };
400 self.store.set_global_status(&g.gid, s, "").await?;
401 return Ok(());
402 }
403
404 let status = g.status;
405 loop {
406 let rows = self.store.list_branches(&g.gid).await?;
407 let (f, b) = split_by_op(&rows, n, fwd, bwd);
408 let adv = if fwd == BranchOp::Commit {
409 xa_advance(status, &f, &b)
410 } else {
411 tcc_advance(status, &f, &b)
412 };
413 match adv {
414 Advance::Finish(s) => {
415 info!(gid = %g.gid, status = s.as_str(), mode = label, "事务终结");
416 self.store.set_global_status(&g.gid, s, "").await?;
417 return Ok(());
418 }
419 Advance::Wait => return Ok(()),
420 Advance::RunWorkflow => {
423 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
424 return Ok(());
425 }
426 Advance::Call { index, op } => {
427 let bid = branch_id(index);
428 let Some(url) = url_of(&rows, &bid, op) else {
429 warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
430 "分支没登记这个操作的 URL,无法调用");
431 self.retry_later(g).await?;
432 return Ok(());
433 };
434 let bp = payload_of(&rows, &bid, op);
435 match self.call_branch(g, &bid, op, &url, &bp).await {
436 BranchResult::Success => {
437 self.store
438 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
439 .await?;
440 }
441 BranchResult::Failure => {
445 self.store
446 .set_branch_status(&g.gid, &bid, op, BranchStatus::Failed)
447 .await?;
448 warn!(gid = %g.gid, branch = %bid, op = op.as_str(), mode = label,
449 "二阶段失败,会持续重试,需要人工介入");
450 self.retry_later(g).await?;
451 return Ok(());
452 }
453 BranchResult::Ongoing | BranchResult::Unknown => {
454 self.retry_later(g).await?;
455 return Ok(());
456 }
457 }
458 }
459 }
460 }
461 }
462
463 async fn process_msg(&self, g: &GlobalRow) -> anyhow::Result<()> {
471 let mut status = g.status;
472
473 if status == GlobalStatus::Prepared {
474 if g.query_prepared.is_empty() {
475 warn!(gid = %g.gid, "msg 事务没提供 query_prepared,无法回查,等人处理");
477 self.retry_later(g).await?;
478 return Ok(());
479 }
480 match self
482 .call_branch(g, "00", BranchOp::Action, &g.query_prepared, "")
483 .await
484 {
485 BranchResult::Success => {
486 info!(gid = %g.gid, "回查:本地事务已提交 → 继续推进");
487 self.store
488 .set_global_status(&g.gid, GlobalStatus::Submitted, "")
489 .await?;
490 status = GlobalStatus::Submitted;
491 }
492 BranchResult::Failure => {
493 info!(gid = %g.gid, "回查:本地事务未提交 → 整单作废");
496 self.store
497 .set_global_status(
498 &g.gid,
499 GlobalStatus::Failed,
500 "回查得到 FAILURE:本地事务未提交",
501 )
502 .await?;
503 return Ok(());
504 }
505 BranchResult::Ongoing | BranchResult::Unknown => {
506 self.retry_later(g).await?;
508 return Ok(());
509 }
510 }
511 }
512
513 let steps: Vec<SagaStep> = serde_json::from_str(&g.payload).unwrap_or_default();
514 if steps.is_empty() {
515 self.store
516 .set_global_status(&g.gid, GlobalStatus::Succeed, "")
517 .await?;
518 return Ok(());
519 }
520 loop {
521 let (actions, _) = self.branch_states(&g.gid, steps.len()).await?;
522 match msg_advance(status, &actions) {
523 Advance::Finish(s) => {
524 info!(gid = %g.gid, status = s.as_str(), "消息事务终结");
525 self.store.set_global_status(&g.gid, s, "").await?;
526 return Ok(());
527 }
528 Advance::Wait => return Ok(()),
529 Advance::RunWorkflow => {
532 warn!(gid = %g.gid, "非 workflow 事务收到 RunWorkflow 决策,跳过");
533 return Ok(());
534 }
535 Advance::Call { index, op } => {
536 let bid = branch_id(index);
537 match self
538 .call_branch(g, &bid, op, &steps[index].action, &steps[index].payload)
539 .await
540 {
541 BranchResult::Success => {
542 self.store
543 .set_branch_status(&g.gid, &bid, op, BranchStatus::Succeed)
544 .await?;
545 }
546 BranchResult::Failure | BranchResult::Ongoing | BranchResult::Unknown => {
548 self.retry_later(g).await?;
549 return Ok(());
550 }
551 }
552 }
553 }
554 }
555 }
556
557 async fn retry_later(&self, g: &GlobalRow) -> anyhow::Result<()> {
558 let iv = dtmrs_core::next_interval_with(g.next_cron_interval, self.retry);
559 self.store.schedule_retry(&g.gid, iv).await?;
560 Ok(())
561 }
562
563 async fn branch_states(
565 &self,
566 gid: &str,
567 n: usize,
568 ) -> anyhow::Result<(Vec<BranchStatus>, Vec<BranchStatus>)> {
569 let rows = self.store.list_branches(gid).await?;
570 let mut actions = vec![BranchStatus::Prepared; n];
571 let mut compensates = vec![BranchStatus::Prepared; n];
572 for r in rows {
573 let Some(i) = index_of(&r.branch_id) else {
574 continue;
575 };
576 if i >= n {
577 continue;
578 }
579 match r.op {
580 BranchOp::Action | BranchOp::Try => actions[i] = r.status,
581 BranchOp::Compensate | BranchOp::Cancel | BranchOp::Rollback => {
582 compensates[i] = r.status
583 }
584 _ => {}
585 }
586 }
587 Ok((actions, compensates))
588 }
589
590 async fn call_branch(
592 &self,
593 g: &GlobalRow,
594 branch_id: &str,
595 op: BranchOp,
596 url: &str,
597 payload: &str,
598 ) -> BranchResult {
599 match parse_target(url) {
600 Target::Local(name) => self.call_local(g, branch_id, op, &name).await,
601 Target::Http(u) => self.call_http(g, branch_id, op, &u, payload).await,
602 #[cfg(feature = "grpc")]
603 Target::Grpc(t) => {
604 self.grpc
605 .call(
606 &t,
607 &g.gid,
608 &g.trans_type.to_string(),
609 branch_id,
610 op.as_str(),
611 )
612 .await
613 }
614 #[cfg(not(feature = "grpc"))]
617 Target::Grpc(t) => {
618 warn!(gid = %g.gid, branch = %branch_id, endpoint = %t.endpoint,
619 "遇到 grpc:// 分支但本次构建关掉了 grpc feature,按结果未知处理");
620 BranchResult::Unknown
621 }
622 }
623 }
624
625 async fn call_local(
627 &self,
628 g: &GlobalRow,
629 branch_id: &str,
630 op: BranchOp,
631 name: &str,
632 ) -> BranchResult {
633 let Some(h) = self.registry.get(name) else {
634 warn!(gid = %g.gid, branch = %branch_id, handler = name,
637 "本地分支未注册,按结果未知处理(会重试,不回滚)");
638 return BranchResult::Unknown;
639 };
640 let ctx = BranchCtx {
641 gid: g.gid.clone(),
642 branch_id: branch_id.to_string(),
643 op,
644 trans_type: g.trans_type.to_string(),
645 };
646 let r = h(ctx).await;
647 info!(gid = %g.gid, branch = %branch_id, op = op.as_str(),
648 handler = name, result = ?r, "本地分支返回");
649 r
650 }
651
652 async fn call_http(
654 &self,
655 g: &GlobalRow,
656 branch_id: &str,
657 op: BranchOp,
658 url: &str,
659 payload: &str,
660 ) -> BranchResult {
661 let req = self
662 .http
663 .post(url)
664 .query(&[
665 ("gid", g.gid.as_str()),
666 ("trans_type", &g.trans_type.to_string()),
667 ("branch_id", branch_id),
668 ("op", op.as_str()),
669 ])
670 .header("content-type", "application/json")
671 .body(branch_payload(payload));
672 match req.send().await {
673 Ok(resp) => {
674 let code = resp.status().as_u16();
675 let body = resp.text().await.unwrap_or_default();
676 let r = BranchResult::from_http(code, &body);
677 info!(gid = %g.gid, branch = %branch_id, op = op.as_str(), code, result = ?r, "分支返回");
678 r
679 }
680 Err(e) => {
681 warn!(gid = %g.gid, branch = %branch_id, error = %e, "分支不可达,结果未知");
683 BranchResult::Unknown
684 }
685 }
686 }
687}
688
689pub fn branch_id(index: usize) -> String {
691 format!("{:02}", index + 1)
692}
693
694fn index_of(branch_id: &str) -> Option<usize> {
695 branch_id
696 .parse::<usize>()
697 .ok()
698 .and_then(|v| v.checked_sub(1))
699}
700
701fn split_by_op(
703 rows: &[dtmrs_store::BranchRow],
704 n: usize,
705 fwd: BranchOp,
706 bwd: BranchOp,
707) -> (Vec<BranchStatus>, Vec<BranchStatus>) {
708 let mut a = vec![BranchStatus::Prepared; n];
709 let mut b = vec![BranchStatus::Prepared; n];
710 for r in rows {
711 let Some(i) = index_of(&r.branch_id) else {
712 continue;
713 };
714 if i >= n {
715 continue;
716 }
717 if r.op == fwd {
718 a[i] = r.status;
719 } else if r.op == bwd {
720 b[i] = r.status;
721 }
722 }
723 (a, b)
724}
725
726fn compensate_states(rows: &[dtmrs_store::BranchRow]) -> Vec<BranchStatus> {
731 let n = rows
732 .iter()
733 .filter(|r| r.op == BranchOp::Compensate)
734 .filter_map(|r| index_of(&r.branch_id))
735 .max()
736 .map(|m| m + 1)
737 .unwrap_or(0);
738 let mut v = vec![BranchStatus::Succeed; n];
741 for r in rows.iter().filter(|r| r.op == BranchOp::Compensate) {
742 if let Some(i) = index_of(&r.branch_id) {
743 if i < n {
744 v[i] = r.status;
745 }
746 }
747 }
748 v
749}
750
751fn payload_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> String {
754 rows.iter()
755 .find(|r| r.branch_id == branch_id && r.op == op)
756 .map(|r| r.payload.clone())
757 .unwrap_or_default()
758}
759
760fn url_of(rows: &[dtmrs_store::BranchRow], branch_id: &str, op: BranchOp) -> Option<String> {
761 rows.iter()
762 .find(|r| r.branch_id == branch_id && r.op == op)
763 .map(|r| r.url.clone())
764}
765
766fn branch_payload(step_payload: &str) -> String {
771 if step_payload.trim().is_empty() {
772 "{}".to_string()
773 } else {
774 step_payload.to_string()
775 }
776}
777
778#[derive(Debug, Clone, Copy)]
780pub struct DriverConfig {
781 pub branch_timeout_secs: i64,
783 pub lease_secs: i64,
785 pub retry: dtmrs_core::RetryPolicy,
786 pub workers: usize,
788}
789
790impl Default for DriverConfig {
791 fn default() -> Self {
792 Self {
794 branch_timeout_secs: 10,
795 lease_secs: 30,
796 retry: dtmrs_core::RetryPolicy::default(),
797 workers: 16,
811 }
812 }
813}
814
815impl DriverConfig {
816 pub fn from_env() -> Self {
817 let d = Self::default();
818 let get = |k: &str, fallback: i64| {
819 std::env::var(k)
820 .ok()
821 .and_then(|v| v.parse::<i64>().ok())
822 .filter(|v| *v > 0)
823 .unwrap_or(fallback)
824 };
825 Self {
826 branch_timeout_secs: get("DTMRS_BRANCH_TIMEOUT", d.branch_timeout_secs),
827 lease_secs: get("DTMRS_LEASE", d.lease_secs),
828 retry: dtmrs_core::RetryPolicy::from_env(),
829 workers: get("DTMRS_WORKERS", d.workers as i64) as usize,
830 }
831 }
832}
833
834#[cfg(test)]
835mod tests {
836 use super::*;
837
838 #[test]
839 fn 分支号与下标互转() {
840 assert_eq!(branch_id(0), "01");
841 assert_eq!(branch_id(9), "10");
842 assert_eq!(index_of("01"), Some(0));
843 assert_eq!(index_of("10"), Some(9));
844 assert_eq!(index_of("00"), None);
845 assert_eq!(index_of("xx"), None);
846 }
847}