1pub use dtmrs_core::Backend;
21
22use dtmrs_core::dialect::check_len;
23use dtmrs_core::{BranchOp, BranchStatus, GlobalStatus, TransType};
24use sqlx::any::{AnyPoolOptions, AnyRow};
25use sqlx::{AnyPool, Row};
26use std::sync::Once;
27
28pub type Result<T> = std::result::Result<T, sqlx::Error>;
29
30pub const BIG: usize = 8192;
32pub const MID: usize = 1024;
34
35fn len_ok(col: &'static str, val: &str, max: usize) -> Result<()> {
41 check_len(col, val, max).map_err(|e| sqlx::Error::Encode(Box::new(e)))
42}
43
44pub fn now() -> i64 {
45 std::time::SystemTime::now()
46 .duration_since(std::time::UNIX_EPOCH)
47 .map(|d| d.as_secs() as i64)
48 .unwrap_or(0)
49}
50
51#[derive(Debug, Clone)]
52pub struct GlobalRow {
53 pub gid: String,
54 pub trans_type: TransType,
55 pub status: GlobalStatus,
56 pub payload: String,
57 pub next_cron_time: i64,
58 pub next_cron_interval: i64,
59 pub owner: String,
60 pub rollback_reason: String,
61 pub query_prepared: String,
64 pub create_time: i64,
65 pub finish_time: Option<i64>,
66}
67
68#[derive(Debug, Clone)]
69pub struct BranchRow {
70 pub gid: String,
71 pub branch_id: String,
72 pub op: BranchOp,
73 pub url: String,
74 pub payload: String,
75 pub status: BranchStatus,
76}
77
78#[derive(Clone)]
79pub struct SqlStore {
80 pool: AnyPool,
81 be: Backend,
82}
83
84static DRIVERS: Once = Once::new();
85
86impl SqlStore {
87 pub async fn open(url: &str) -> Result<Self> {
91 DRIVERS.call_once(sqlx::any::install_default_drivers);
92
93 let mut url = url.to_string();
96 if url.starts_with("sqlite") && !url.contains("mode=") && !url.contains(":memory:") {
97 url.push_str(if url.contains('?') {
98 "&mode=rwc"
99 } else {
100 "?mode=rwc"
101 });
102 }
103 let max = if url.contains(":memory:") { 1 } else { 8 };
105 let be = Backend::from_url(&url);
106 let pool = AnyPoolOptions::new()
107 .max_connections(max)
108 .connect(&url)
109 .await?;
110 let s = Self { pool, be };
111 s.migrate_racy().await?;
112 Ok(s)
113 }
114
115 async fn migrate_racy(&self) -> Result<()> {
125 let mut last = None;
126 for attempt in 0..3 {
127 match self.migrate().await {
128 Ok(()) => return Ok(()),
129 Err(e) => {
130 last = Some(e);
131 tokio::time::sleep(std::time::Duration::from_millis(100 * (attempt + 1))).await;
133 }
134 }
135 }
136 Err(last.expect("循环至少失败一次"))
137 }
138
139 pub async fn migrate(&self) -> Result<()> {
140 let idt = self.be.id_text();
141 let ids = self.be.id_short();
142 let big = self.be.text(BIG);
145 let mid = self.be.text(MID);
146 let inline = self
149 .be
150 .inline_index("idx_status_cron", "status, next_cron_time");
151
152 sqlx::query(&format!(
153 "CREATE TABLE IF NOT EXISTS trans_global (
154 gid {idt} NOT NULL,
155 trans_type {ids} NOT NULL,
156 status {ids} NOT NULL,
157 payload {big} NOT NULL,
158 next_cron_time BIGINT NOT NULL DEFAULT 0,
159 next_cron_interval BIGINT NOT NULL DEFAULT 0,
160 owner {idt} NOT NULL,
161 rollback_reason {mid} NOT NULL,
162 query_prepared {mid} NOT NULL,
163 create_time BIGINT NOT NULL,
164 update_time BIGINT NOT NULL,
165 finish_time BIGINT,
166 PRIMARY KEY (gid){inline}
167 )"
168 ))
169 .execute(&self.pool)
170 .await?;
171 if let Some(sql) =
173 self.be
174 .create_index("idx_status_cron", "trans_global", "status, next_cron_time")
175 {
176 sqlx::query(&sql).execute(&self.pool).await?;
177 }
178 sqlx::query(&format!(
179 "CREATE TABLE IF NOT EXISTS trans_branch_op (
180 gid {idt} NOT NULL,
181 branch_id {idt} NOT NULL,
182 op {ids} NOT NULL,
183 url {mid} NOT NULL,
184 payload {mid} NOT NULL,
185 status {ids} NOT NULL,
186 create_time BIGINT NOT NULL,
187 update_time BIGINT NOT NULL,
188 finish_time BIGINT,
189 PRIMARY KEY (gid, branch_id, op)
190 )"
191 ))
192 .execute(&self.pool)
193 .await?;
194 Ok(())
195 }
196
197 pub fn backend(&self) -> Backend {
198 self.be
199 }
200
201 pub fn pool(&self) -> &AnyPool {
202 &self.pool
203 }
204
205 pub async fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> Result<bool> {
210 len_ok("gid", &g.gid, Backend::ID_MAX)?;
212 len_ok("payload", &g.payload, BIG)?;
213 len_ok("query_prepared", &g.query_prepared, MID)?;
214 for b in branches {
215 len_ok("branch_id", &b.branch_id, Backend::ID_MAX)?;
216 len_ok("url", &b.url, MID)?;
217 len_ok("payload", &b.payload, MID)?;
218 }
219 let mut tx = self.pool.begin().await?;
220 let t = now();
221 let n = sqlx::query(&self.be.q("{INS} trans_global
222 (gid,trans_type,status,payload,next_cron_time,next_cron_interval,
223 owner,rollback_reason,query_prepared,create_time,update_time)
224 VALUES (?,?,?,?,?,?,'','',?,?,?)
225 {NOCONFLICT}"))
226 .bind(&g.gid)
227 .bind(g.trans_type.to_string())
228 .bind(g.status.as_str())
229 .bind(&g.payload)
230 .bind(g.next_cron_time)
231 .bind(g.next_cron_interval)
232 .bind(&g.query_prepared)
233 .bind(t)
234 .bind(t)
235 .execute(&mut *tx)
236 .await?
237 .rows_affected();
238 if n == 0 {
239 tx.rollback().await?;
240 return Ok(false);
241 }
242 for b in branches {
243 sqlx::query(&self.be.q("{INS} trans_branch_op
244 (gid,branch_id,op,url,payload,status,create_time,update_time)
245 VALUES (?,?,?,?,?,?,?,?)
246 {NOCONFLICT}"))
247 .bind(&b.gid)
248 .bind(&b.branch_id)
249 .bind(b.op.as_str())
250 .bind(&b.url)
251 .bind(&b.payload)
252 .bind(b.status.as_str())
253 .bind(t)
254 .bind(t)
255 .execute(&mut *tx)
256 .await?;
257 }
258 tx.commit().await?;
259 Ok(true)
260 }
261
262 pub async fn get_global(&self, gid: &str) -> Result<Option<GlobalRow>> {
263 let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
264 .bind(gid)
265 .fetch_optional(&self.pool)
266 .await?;
267 Ok(row.map(global_from_row))
268 }
269
270 pub async fn list_branches(&self, gid: &str) -> Result<Vec<BranchRow>> {
271 let rows = sqlx::query(&self.be.q(
272 "SELECT gid,branch_id,op,url,payload,status FROM trans_branch_op
273 WHERE gid=? ORDER BY branch_id, op",
274 ))
275 .bind(gid)
276 .fetch_all(&self.pool)
277 .await?;
278 Ok(rows
279 .into_iter()
280 .map(|r| BranchRow {
281 gid: r.get("gid"),
282 branch_id: r.get("branch_id"),
283 op: BranchOp::parse(r.get::<String, _>("op").as_str()).unwrap_or(BranchOp::Action),
284 url: r.get("url"),
285 payload: r.get("payload"),
286 status: BranchStatus::parse(r.get::<String, _>("status").as_str())
287 .unwrap_or(BranchStatus::Prepared),
288 })
289 .collect())
290 }
291
292 pub async fn set_global_status(
293 &self,
294 gid: &str,
295 status: GlobalStatus,
296 reason: &str,
297 ) -> Result<()> {
298 let t = now();
299 let fin = if status.is_final() { Some(t) } else { None };
300 let reason: String = reason.chars().take(MID).collect();
304 let reason = reason.as_str();
305 sqlx::query(&self.be.q(
307 "UPDATE trans_global SET status=?, update_time=?, finish_time=?,
308 rollback_reason = CASE WHEN ? <> '' THEN ? ELSE rollback_reason END
309 WHERE gid=?",
310 ))
311 .bind(status.as_str())
312 .bind(t)
313 .bind(fin)
314 .bind(reason)
315 .bind(reason)
316 .bind(gid)
317 .execute(&self.pool)
318 .await?;
319 Ok(())
320 }
321
322 pub async fn set_branch_result(
329 &self,
330 gid: &str,
331 branch_id: &str,
332 op: BranchOp,
333 status: BranchStatus,
334 payload: &str,
335 ) -> Result<()> {
336 len_ok("payload", payload, MID)?;
337 let t = now();
338 sqlx::query(&self.be.q(
339 "UPDATE trans_branch_op SET status=?, payload=?, update_time=?,
340 finish_time = CASE WHEN ? <> 'prepared' THEN ? ELSE finish_time END
341 WHERE gid=? AND branch_id=? AND op=?",
342 ))
343 .bind(status.as_str())
344 .bind(payload)
345 .bind(t)
346 .bind(status.as_str())
347 .bind(t)
348 .bind(gid)
349 .bind(branch_id)
350 .bind(op.as_str())
351 .execute(&self.pool)
352 .await?;
353 Ok(())
354 }
355
356 pub async fn set_branch_status(
357 &self,
358 gid: &str,
359 branch_id: &str,
360 op: BranchOp,
361 status: BranchStatus,
362 ) -> Result<()> {
363 let t = now();
364 sqlx::query(
365 &self
366 .be
367 .q("UPDATE trans_branch_op SET status=?, update_time=?,
368 finish_time = CASE WHEN ? <> 'prepared' THEN ? ELSE finish_time END
369 WHERE gid=? AND branch_id=? AND op=?"),
370 )
371 .bind(status.as_str())
372 .bind(t)
373 .bind(status.as_str())
374 .bind(t)
375 .bind(gid)
376 .bind(branch_id)
377 .bind(op.as_str())
378 .execute(&self.pool)
379 .await?;
380 Ok(())
381 }
382
383 pub async fn lock_one_due(&self, owner: &str, lease: i64) -> Result<Option<GlobalRow>> {
388 let mut tx = self.pool.begin().await?;
389 let t = now();
390 let gid: Option<String> = sqlx::query_scalar(&self.be.q("SELECT gid FROM trans_global
391 WHERE (status IN ('submitted','aborting')
392 OR (status = 'prepared' AND trans_type = 'msg'))
393 AND next_cron_time <= ?
394 ORDER BY next_cron_time LIMIT 1"))
395 .bind(t)
396 .fetch_optional(&mut *tx)
397 .await?;
398 let Some(gid) = gid else {
399 tx.rollback().await?;
400 return Ok(None);
401 };
402 let n = sqlx::query(&self.be.q(
404 "UPDATE trans_global SET owner=?, next_cron_time=?, update_time=?
405 WHERE gid=? AND next_cron_time <= ?",
406 ))
407 .bind(owner)
408 .bind(t + lease)
409 .bind(t)
410 .bind(&gid)
411 .bind(t)
412 .execute(&mut *tx)
413 .await?
414 .rows_affected();
415 if n == 0 {
416 tx.rollback().await?;
417 return Ok(None); }
419 let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
420 .bind(&gid)
421 .fetch_one(&mut *tx)
422 .await?;
423 tx.commit().await?;
424 Ok(Some(global_from_row(row)))
425 }
426
427 pub async fn schedule_retry(&self, gid: &str, interval: i64) -> Result<()> {
429 let t = now();
430 sqlx::query(&self.be.q(
431 "UPDATE trans_global SET next_cron_interval=?, next_cron_time=?, update_time=?
432 WHERE gid=?",
433 ))
434 .bind(interval)
435 .bind(t + interval)
436 .bind(t)
437 .bind(gid)
438 .execute(&self.pool)
439 .await?;
440 Ok(())
441 }
442
443 pub async fn schedule_now(&self, gid: &str) -> Result<()> {
445 sqlx::query(
446 &self
447 .be
448 .q("UPDATE trans_global SET next_cron_time=?, next_cron_interval=0 WHERE gid=?"),
449 )
450 .bind(now())
451 .bind(gid)
452 .execute(&self.pool)
453 .await?;
454 Ok(())
455 }
456
457 pub async fn register_branch(
464 &self,
465 gid: &str,
466 branch_id: &str,
467 ops: &[(BranchOp, String)],
468 ) -> Result<()> {
469 len_ok("gid", gid, Backend::ID_MAX)?;
470 len_ok("branch_id", branch_id, Backend::ID_MAX)?;
471 for (_, url) in ops {
472 len_ok("url", url, MID)?;
473 }
474 let mut tx = self.pool.begin().await?;
475 let t = now();
476 for (op, url) in ops {
477 sqlx::query(&self.be.q("{INS} trans_branch_op
478 (gid,branch_id,op,url,payload,status,create_time,update_time)
479 VALUES (?,?,?,?,'',?,?,?)
480 {NOCONFLICT}"))
481 .bind(gid)
482 .bind(branch_id)
483 .bind(op.as_str())
484 .bind(url)
485 .bind(BranchStatus::Prepared.as_str())
486 .bind(t)
487 .bind(t)
488 .execute(&mut *tx)
489 .await?;
490 }
491 tx.commit().await?;
492 Ok(())
493 }
494
495 pub async fn list_recent(&self, limit: i64) -> Result<Vec<GlobalRow>> {
496 let rows = sqlx::query(&self.be.q(&format!(
497 "{SELECT_GLOBAL} ORDER BY create_time DESC LIMIT ?"
498 )))
499 .bind(limit)
500 .fetch_all(&self.pool)
501 .await?;
502 Ok(rows.into_iter().map(global_from_row).collect())
503 }
504}
505
506const SELECT_GLOBAL: &str = "SELECT gid,trans_type,status,payload,next_cron_time,
508 next_cron_interval,owner,rollback_reason,query_prepared,create_time,finish_time
509 FROM trans_global";
510
511fn global_from_row(r: AnyRow) -> GlobalRow {
512 GlobalRow {
513 gid: r.get("gid"),
514 trans_type: TransType::parse(r.get::<String, _>("trans_type").as_str())
515 .unwrap_or(TransType::Saga),
516 status: GlobalStatus::parse(r.get::<String, _>("status").as_str())
517 .unwrap_or(GlobalStatus::Prepared),
518 payload: r.get("payload"),
519 next_cron_time: r.get("next_cron_time"),
520 next_cron_interval: r.get("next_cron_interval"),
521 owner: r.get("owner"),
522 rollback_reason: r.get("rollback_reason"),
523 query_prepared: r.get("query_prepared"),
524 create_time: r.get("create_time"),
525 finish_time: r.get("finish_time"),
526 }
527}
528
529#[cfg(test)]
530mod tests {
531 use super::*;
532
533 static PG_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
545
546 async fn backends() -> (
551 tokio::sync::MutexGuard<'static, ()>,
552 Vec<(&'static str, Store)>,
553 ) {
554 let guard = PG_LOCK.lock().await;
555 let mut v = vec![("sqlite", Store::open("sqlite::memory:").await.unwrap())];
556 for (name, env) in [("postgres", "DTMRS_TEST_PG"), ("mysql", "DTMRS_TEST_MYSQL")] {
558 if let Ok(url) = std::env::var(env) {
559 let s = Store::open(&url)
560 .await
561 .unwrap_or_else(|e| panic!("连不上 {env}: {e}"));
562 for t in ["trans_branch_op", "trans_global"] {
563 sqlx::query(&format!("DELETE FROM {t}"))
564 .execute(s.pool().expect("SQL 后端才有连接池"))
565 .await
566 .expect("清表");
567 }
568 v.push((name, s));
569 }
570 }
571 #[cfg(feature = "redis")]
572 if let Ok(url) = std::env::var("DTMRS_TEST_REDIS") {
573 let s = Store::open(&url)
574 .await
575 .unwrap_or_else(|e| panic!("连不上 DTMRS_TEST_REDIS: {e}"));
576 s.as_redis()
578 .unwrap()
579 .flush_prefix()
580 .await
581 .expect("清 redis");
582 v.push(("redis", s));
583 }
584 (guard, v)
585 }
586
587 fn g(gid: &str) -> GlobalRow {
588 GlobalRow {
589 gid: gid.into(),
590 trans_type: TransType::Saga,
591 status: GlobalStatus::Submitted,
592 payload: "{}".into(),
593 next_cron_time: 0,
594 next_cron_interval: 0,
595 owner: String::new(),
596 rollback_reason: String::new(),
597 query_prepared: String::new(),
598 create_time: 0,
599 finish_time: None,
600 }
601 }
602
603 #[tokio::test]
604 async fn 重复提交同一个gid是幂等的() {
605 let (_g, bes) = backends().await;
606 for (name, s) in bes {
607 assert!(s.create_global(&g("t1"), &[]).await.unwrap(), "{name}");
608 assert!(!s.create_global(&g("t1"), &[]).await.unwrap(), "{name}");
610 assert_eq!(s.list_recent(10).await.unwrap().len(), 1, "{name}");
611 }
612 }
613
614 #[tokio::test]
615 async fn 租约只能被抢到一次() {
616 let (_g, bes) = backends().await;
617 for (name, s) in bes {
618 s.create_global(&g("t2"), &[]).await.unwrap();
619 let a = s.lock_one_due("worker-a", 60).await.unwrap();
620 assert!(a.is_some(), "{name}: 第一个实例应该抢到");
621 let b = s.lock_one_due("worker-b", 60).await.unwrap();
623 assert!(b.is_none(), "{name}: 租约期内不能被别人抢走");
624 }
625 }
626
627 #[tokio::test]
628 async fn 终态不再被调度() {
629 let (_g, bes) = backends().await;
630 for (name, s) in bes {
631 s.create_global(&g("t3"), &[]).await.unwrap();
632 s.set_global_status("t3", GlobalStatus::Succeed, "")
633 .await
634 .unwrap();
635 assert!(s.lock_one_due("w", 60).await.unwrap().is_none(), "{name}");
636 let got = s.get_global("t3").await.unwrap().unwrap();
637 assert_eq!(got.status, GlobalStatus::Succeed, "{name}");
638 assert!(got.finish_time.is_some(), "{name}: 终态要落 finish_time");
639 }
640 }
641
642 #[tokio::test]
643 async fn 分支状态可更新() {
644 let (_g, bes) = backends().await;
645 for (name, s) in bes {
646 let b = BranchRow {
647 gid: "t4".into(),
648 branch_id: "01".into(),
649 op: BranchOp::Action,
650 url: "http://x/a".into(),
651 payload: "{}".into(),
652 status: BranchStatus::Prepared,
653 };
654 s.create_global(&g("t4"), std::slice::from_ref(&b))
655 .await
656 .unwrap();
657 s.set_branch_status("t4", "01", BranchOp::Action, BranchStatus::Succeed)
658 .await
659 .unwrap();
660 let got = s.list_branches("t4").await.unwrap();
661 assert_eq!(got.len(), 1, "{name}");
662 assert_eq!(got[0].status, BranchStatus::Succeed, "{name}");
663 }
664 }
665
666 #[tokio::test]
667 async fn 回滚原因和回查地址能存取() {
668 let (_g, bes) = backends().await;
670 for (name, s) in bes {
671 let mut row = g("t5");
672 row.query_prepared = "http://busi/query".into();
673 s.create_global(&row, &[]).await.unwrap();
674 s.set_global_status("t5", GlobalStatus::Aborting, "分支 02 返回 FAILURE")
675 .await
676 .unwrap();
677 let got = s.get_global("t5").await.unwrap().unwrap();
678 assert_eq!(got.query_prepared, "http://busi/query", "{name}");
679 assert_eq!(got.rollback_reason, "分支 02 返回 FAILURE", "{name}");
680 assert!(
681 got.finish_time.is_none(),
682 "{name}: 非终态不该有 finish_time"
683 );
684
685 s.set_global_status("t5", GlobalStatus::Failed, "")
687 .await
688 .unwrap();
689 let got = s.get_global("t5").await.unwrap().unwrap();
690 assert_eq!(
691 got.rollback_reason, "分支 02 返回 FAILURE",
692 "{name}: 空原因不能覆盖"
693 );
694 }
695 }
696
697 #[tokio::test]
698 async fn msg的prepared会被捞tcc的不会() {
699 let (_g, bes) = backends().await;
700 for (name, s) in bes {
701 let mut m = g("m1");
702 m.trans_type = TransType::Msg;
703 m.status = GlobalStatus::Prepared;
704 s.create_global(&m, &[]).await.unwrap();
705 let mut t = g("c1");
706 t.trans_type = TransType::Tcc;
707 t.status = GlobalStatus::Prepared;
708 s.create_global(&t, &[]).await.unwrap();
709
710 let got = s.lock_one_due("w", 60).await.unwrap();
711 assert_eq!(
712 got.map(|x| x.gid),
713 Some("m1".to_string()),
714 "{name}: 只该捞到 msg"
715 );
716 assert!(s.lock_one_due("w2", 60).await.unwrap().is_none(), "{name}");
718 }
719 }
720
721 #[tokio::test]
722 async fn 分支登记是幂等的() {
723 let (_g, bes) = backends().await;
724 for (name, s) in bes {
725 let mut t = g("c2");
726 t.trans_type = TransType::Tcc;
727 s.create_global(&t, &[]).await.unwrap();
728 let ops = [
729 (BranchOp::Confirm, "http://x/c".to_string()),
730 (BranchOp::Cancel, "http://x/n".to_string()),
731 ];
732 s.register_branch("c2", "01", &ops).await.unwrap();
733 s.register_branch("c2", "01", &ops).await.unwrap(); assert_eq!(
735 s.list_branches("c2").await.unwrap().len(),
736 2,
737 "{name}: 不该重复插入"
738 );
739 }
740 }
741}
742
743#[cfg(feature = "redis")]
746pub mod redis_store;
747#[cfg(feature = "redis")]
748pub use redis_store::RedisStore;
749
750#[derive(Clone)]
764enum Inner {
765 Sql(SqlStore),
766 #[cfg(feature = "redis")]
767 Redis(RedisStore),
768}
769
770#[derive(Clone)]
780pub struct Store {
781 inner: Inner,
782}
783
784pub type StoreError = sqlx::Error;
789
790#[cfg(feature = "redis")]
791fn redis_err(e: redis::RedisError) -> sqlx::Error {
792 sqlx::Error::Configuration(Box::new(e))
793}
794
795pub fn is_redis_url(url: &str) -> bool {
797 let u = url.trim().to_ascii_lowercase();
798 u.starts_with("redis://") || u.starts_with("rediss://") || u.starts_with("redis+unix:")
799}
800
801impl Store {
802 pub async fn open(url: &str) -> Result<Self> {
804 if is_redis_url(url) {
805 #[cfg(feature = "redis")]
806 {
807 let r = RedisStore::open(url).await.map_err(redis_err)?;
808 return Ok(Self {
809 inner: Inner::Redis(r),
810 });
811 }
812 #[cfg(not(feature = "redis"))]
813 {
814 return Err(sqlx::Error::Configuration(
817 "这个 URL 要 Redis 后端,但构建时没开 dtmrs-store 的 `redis` feature".into(),
818 ));
819 }
820 }
821 Ok(Self {
822 inner: Inner::Sql(SqlStore::open(url).await?),
823 })
824 }
825
826 pub fn is_redis(&self) -> bool {
828 match &self.inner {
829 Inner::Sql(_) => false,
830 #[cfg(feature = "redis")]
831 Inner::Redis(_) => true,
832 }
833 }
834
835 pub fn pool(&self) -> Option<&AnyPool> {
838 match &self.inner {
839 Inner::Sql(s) => Some(s.pool()),
840 #[cfg(feature = "redis")]
841 Inner::Redis(_) => None,
842 }
843 }
844
845 pub fn backend(&self) -> Option<Backend> {
847 match &self.inner {
848 Inner::Sql(s) => Some(s.backend()),
849 #[cfg(feature = "redis")]
850 Inner::Redis(_) => None,
851 }
852 }
853
854 #[cfg(feature = "redis")]
856 pub fn as_redis(&self) -> Option<&RedisStore> {
857 match &self.inner {
858 Inner::Redis(r) => Some(r),
859 _ => None,
860 }
861 }
862}
863
864macro_rules! dispatch {
867 ($( $(#[$m:meta])* fn $name:ident (&self $(, $arg:ident : $ty:ty)* ) -> $ret:ty; )*) => {
868 impl Store {
869 $(
870 $(#[$m])*
871 pub async fn $name(&self $(, $arg: $ty)*) -> Result<$ret> {
872 match &self.inner {
873 Inner::Sql(s) => s.$name($($arg),*).await,
874 #[cfg(feature = "redis")]
875 Inner::Redis(r) => r.$name($($arg),*).await.map_err(redis_err),
876 }
877 }
878 )*
879 }
880 };
881}
882
883dispatch! {
884 fn migrate(&self) -> ();
886 fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> bool;
888 fn get_global(&self, gid: &str) -> Option<GlobalRow>;
889 fn list_branches(&self, gid: &str) -> Vec<BranchRow>;
890 fn lock_one_due(&self, owner: &str, lease: i64) -> Option<GlobalRow>;
892 fn set_global_status(&self, gid: &str, status: GlobalStatus, reason: &str) -> ();
893 fn set_branch_result(&self, gid: &str, branch_id: &str, op: BranchOp, status: BranchStatus, payload: &str) -> ();
894 fn set_branch_status(&self, gid: &str, branch_id: &str, op: BranchOp, status: BranchStatus) -> ();
895 fn schedule_retry(&self, gid: &str, interval: i64) -> ();
896 fn schedule_now(&self, gid: &str) -> ();
897 fn register_branch(&self, gid: &str, branch_id: &str, ops: &[(BranchOp, String)]) -> ();
898 fn list_recent(&self, limit: i64) -> Vec<GlobalRow>;
899}