pub use dtmrs_core::Backend;
use dtmrs_core::dialect::check_len;
use dtmrs_core::{BranchOp, BranchStatus, GlobalStatus, TransType};
use sqlx::any::{AnyPoolOptions, AnyRow};
use sqlx::{AnyPool, Row};
use std::sync::Once;
pub type Result<T> = std::result::Result<T, sqlx::Error>;
pub const BIG: usize = 8192;
pub const MID: usize = 1024;
fn len_ok(col: &'static str, val: &str, max: usize) -> Result<()> {
check_len(col, val, max).map_err(|e| sqlx::Error::Encode(Box::new(e)))
}
pub fn now() -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.unwrap_or(0)
}
#[derive(Debug, Clone)]
pub enum SubmitOutcome {
Missing,
Advanced(Box<GlobalRow>),
Already,
}
#[derive(Debug, Clone)]
pub struct GlobalRow {
pub gid: String,
pub trans_type: TransType,
pub status: GlobalStatus,
pub payload: String,
pub next_cron_time: i64,
pub next_cron_interval: i64,
pub owner: String,
pub rollback_reason: String,
pub query_prepared: String,
pub create_time: i64,
pub finish_time: Option<i64>,
}
#[derive(Debug, Clone)]
pub struct BranchRow {
pub gid: String,
pub branch_id: String,
pub op: BranchOp,
pub url: String,
pub payload: String,
pub status: BranchStatus,
}
#[derive(Clone)]
pub struct SqlStore {
pool: AnyPool,
be: Backend,
}
static DRIVERS: Once = Once::new();
impl SqlStore {
pub async fn open(url: &str) -> Result<Self> {
DRIVERS.call_once(sqlx::any::install_default_drivers);
let mut url = url.to_string();
if url.starts_with("sqlite") && !url.contains("mode=") && !url.contains(":memory:") {
url.push_str(if url.contains('?') {
"&mode=rwc"
} else {
"?mode=rwc"
});
}
let max = if url.contains(":memory:") {
1
} else {
std::env::var("DTMRS_DB_POOL")
.ok()
.and_then(|v| v.parse::<u32>().ok())
.filter(|v| *v > 0)
.unwrap_or(32)
};
let be = Backend::from_url(&url);
let is_file_sqlite = be == Backend::Sqlite && !url.contains(":memory:");
let pool = AnyPoolOptions::new()
.max_connections(max)
.after_connect(move |conn, _| {
Box::pin(async move {
if is_file_sqlite {
for pragma in [
"PRAGMA journal_mode=WAL",
"PRAGMA synchronous=NORMAL",
"PRAGMA busy_timeout=5000",
] {
sqlx::query(pragma).execute(&mut *conn).await?;
}
}
Ok(())
})
})
.connect(&url)
.await?;
let s = Self { pool, be };
s.migrate_racy().await?;
Ok(s)
}
async fn migrate_racy(&self) -> Result<()> {
let mut last = None;
for attempt in 0..3 {
match self.migrate().await {
Ok(()) => return Ok(()),
Err(e) => {
last = Some(e);
tokio::time::sleep(std::time::Duration::from_millis(100 * (attempt + 1))).await;
}
}
}
Err(last.expect("循环至少失败一次"))
}
pub async fn migrate(&self) -> Result<()> {
let idt = self.be.id_text();
let ids = self.be.id_short();
let big = self.be.text(BIG);
let mid = self.be.text(MID);
let inline = self
.be
.inline_index("idx_status_cron", "status, next_cron_time");
sqlx::query(&format!(
"CREATE TABLE IF NOT EXISTS trans_global (
gid {idt} NOT NULL,
trans_type {ids} NOT NULL,
status {ids} NOT NULL,
payload {big} NOT NULL,
next_cron_time BIGINT NOT NULL DEFAULT 0,
next_cron_interval BIGINT NOT NULL DEFAULT 0,
owner {idt} NOT NULL,
rollback_reason {mid} NOT NULL,
query_prepared {mid} NOT NULL,
create_time BIGINT NOT NULL,
update_time BIGINT NOT NULL,
finish_time BIGINT,
PRIMARY KEY (gid){inline}
)"
))
.execute(&self.pool)
.await?;
if let Some(sql) =
self.be
.create_index("idx_status_cron", "trans_global", "status, next_cron_time")
{
sqlx::query(&sql).execute(&self.pool).await?;
}
sqlx::query(&format!(
"CREATE TABLE IF NOT EXISTS trans_branch_op (
gid {idt} NOT NULL,
branch_id {idt} NOT NULL,
op {ids} NOT NULL,
url {mid} NOT NULL,
payload {mid} NOT NULL,
status {ids} NOT NULL,
create_time BIGINT NOT NULL,
update_time BIGINT NOT NULL,
finish_time BIGINT,
PRIMARY KEY (gid, branch_id, op)
)"
))
.execute(&self.pool)
.await?;
Ok(())
}
pub fn backend(&self) -> Backend {
self.be
}
pub fn pool(&self) -> &AnyPool {
&self.pool
}
pub async fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> Result<bool> {
len_ok("gid", &g.gid, Backend::ID_MAX)?;
len_ok("payload", &g.payload, BIG)?;
len_ok("query_prepared", &g.query_prepared, MID)?;
for b in branches {
len_ok("branch_id", &b.branch_id, Backend::ID_MAX)?;
len_ok("url", &b.url, MID)?;
len_ok("payload", &b.payload, MID)?;
}
let mut tx = self.pool.begin().await?;
let t = now();
let n = sqlx::query(&self.be.q("{INS} trans_global
(gid,trans_type,status,payload,next_cron_time,next_cron_interval,
owner,rollback_reason,query_prepared,create_time,update_time)
VALUES (?,?,?,?,?,?,?,'',?,?,?)
{NOCONFLICT}"))
.bind(&g.gid)
.bind(g.trans_type.to_string())
.bind(g.status.as_str())
.bind(&g.payload)
.bind(g.next_cron_time)
.bind(g.next_cron_interval)
.bind(&g.owner)
.bind(&g.query_prepared)
.bind(t)
.bind(t)
.execute(&mut *tx)
.await?
.rows_affected();
if n == 0 {
tx.rollback().await?;
return Ok(false);
}
for b in branches {
sqlx::query(&self.be.q("{INS} trans_branch_op
(gid,branch_id,op,url,payload,status,create_time,update_time)
VALUES (?,?,?,?,?,?,?,?)
{NOCONFLICT}"))
.bind(&b.gid)
.bind(&b.branch_id)
.bind(b.op.as_str())
.bind(&b.url)
.bind(&b.payload)
.bind(b.status.as_str())
.bind(t)
.bind(t)
.execute(&mut *tx)
.await?;
}
tx.commit().await?;
Ok(true)
}
pub async fn get_global(&self, gid: &str) -> Result<Option<GlobalRow>> {
let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
.bind(gid)
.fetch_optional(&self.pool)
.await?;
Ok(row.map(global_from_row))
}
pub async fn list_branches(&self, gid: &str) -> Result<Vec<BranchRow>> {
let rows = sqlx::query(&self.be.q(
"SELECT gid,branch_id,op,url,payload,status FROM trans_branch_op
WHERE gid=? ORDER BY branch_id, op",
))
.bind(gid)
.fetch_all(&self.pool)
.await?;
Ok(rows
.into_iter()
.map(|r| BranchRow {
gid: r.get("gid"),
branch_id: r.get("branch_id"),
op: BranchOp::parse(r.get::<String, _>("op").as_str()).unwrap_or(BranchOp::Action),
url: r.get("url"),
payload: r.get("payload"),
status: BranchStatus::parse(r.get::<String, _>("status").as_str())
.unwrap_or(BranchStatus::Prepared),
})
.collect())
}
pub async fn set_global_status(
&self,
gid: &str,
status: GlobalStatus,
_trans_type: TransType,
reason: &str,
) -> Result<()> {
let t = now();
let fin = if status.is_final() { Some(t) } else { None };
let reason: String = reason.chars().take(MID).collect();
let reason = reason.as_str();
sqlx::query(&self.be.q(
"UPDATE trans_global SET status=?, update_time=?, finish_time=?,
rollback_reason = CASE WHEN ? <> '' THEN ? ELSE rollback_reason END
WHERE gid=?",
))
.bind(status.as_str())
.bind(t)
.bind(fin)
.bind(reason)
.bind(reason)
.bind(gid)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn set_branch_result(
&self,
gid: &str,
branch_id: &str,
op: BranchOp,
status: BranchStatus,
payload: &str,
) -> Result<()> {
len_ok("payload", payload, MID)?;
let t = now();
sqlx::query(&self.be.q(
"UPDATE trans_branch_op SET status=?, payload=?, update_time=?,
finish_time = CASE WHEN ? <> 'prepared' THEN ? ELSE finish_time END
WHERE gid=? AND branch_id=? AND op=?",
))
.bind(status.as_str())
.bind(payload)
.bind(t)
.bind(status.as_str())
.bind(t)
.bind(gid)
.bind(branch_id)
.bind(op.as_str())
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn set_branch_status(
&self,
gid: &str,
branch_id: &str,
op: BranchOp,
status: BranchStatus,
) -> Result<()> {
let t = now();
sqlx::query(
&self
.be
.q("UPDATE trans_branch_op SET status=?, update_time=?,
finish_time = CASE WHEN ? <> 'prepared' THEN ? ELSE finish_time END
WHERE gid=? AND branch_id=? AND op=?"),
)
.bind(status.as_str())
.bind(t)
.bind(status.as_str())
.bind(t)
.bind(gid)
.bind(branch_id)
.bind(op.as_str())
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn lock_one_due(&self, owner: &str, lease: i64) -> Result<Option<GlobalRow>> {
let mut tx = self.pool.begin().await?;
let t = now();
let gid: Option<String> = sqlx::query_scalar(&self.be.q(&format!(
"SELECT gid FROM trans_global
WHERE (status IN ('submitted','aborting')
OR (status = 'prepared' AND trans_type = 'msg'))
AND next_cron_time <= ?
LIMIT 1{}",
self.be.skip_locked()
)))
.bind(t)
.fetch_optional(&mut *tx)
.await?;
let Some(gid) = gid else {
tx.rollback().await?;
return Ok(None);
};
let n = sqlx::query(&self.be.q(
"UPDATE trans_global SET owner=?, next_cron_time=?, update_time=?
WHERE gid=? AND next_cron_time <= ?",
))
.bind(owner)
.bind(t + lease)
.bind(t)
.bind(&gid)
.bind(t)
.execute(&mut *tx)
.await?
.rows_affected();
if n == 0 {
tx.rollback().await?;
return Ok(None); }
let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
.bind(&gid)
.fetch_one(&mut *tx)
.await?;
tx.commit().await?;
Ok(Some(global_from_row(row)))
}
pub async fn schedule_retry(&self, gid: &str, interval: i64) -> Result<()> {
let t = now();
sqlx::query(&self.be.q(
"UPDATE trans_global SET next_cron_interval=?, next_cron_time=?, update_time=?
WHERE gid=?",
))
.bind(interval)
.bind(t + interval)
.bind(t)
.bind(gid)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn submit_prepared(
&self,
gid: &str,
owner: &str,
next_cron_time: i64,
) -> Result<SubmitOutcome> {
let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
.bind(gid)
.fetch_optional(&self.pool)
.await?;
let Some(row) = row else {
return Ok(SubmitOutcome::Missing);
};
let mut g = global_from_row(row);
if g.status != GlobalStatus::Prepared {
return Ok(SubmitOutcome::Already);
}
let t = now();
sqlx::query(&self.be.q("UPDATE trans_global SET status=?, update_time=?,
next_cron_time=?, next_cron_interval=0, owner=?
WHERE gid=? AND status=?"))
.bind(GlobalStatus::Submitted.as_str())
.bind(t)
.bind(next_cron_time)
.bind(owner)
.bind(gid)
.bind(GlobalStatus::Prepared.as_str())
.execute(&self.pool)
.await?;
g.status = GlobalStatus::Submitted;
g.next_cron_time = next_cron_time;
g.owner = owner.to_string();
Ok(SubmitOutcome::Advanced(Box::new(g)))
}
pub async fn schedule_now(&self, gid: &str) -> Result<()> {
sqlx::query(
&self
.be
.q("UPDATE trans_global SET next_cron_time=?, next_cron_interval=0 WHERE gid=?"),
)
.bind(now())
.bind(gid)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn register_branch(
&self,
gid: &str,
branch_id: &str,
ops: &[(BranchOp, String)],
) -> Result<()> {
len_ok("gid", gid, Backend::ID_MAX)?;
len_ok("branch_id", branch_id, Backend::ID_MAX)?;
for (_, url) in ops {
len_ok("url", url, MID)?;
}
let mut tx = self.pool.begin().await?;
let t = now();
for (op, url) in ops {
sqlx::query(&self.be.q("{INS} trans_branch_op
(gid,branch_id,op,url,payload,status,create_time,update_time)
VALUES (?,?,?,?,'',?,?,?)
{NOCONFLICT}"))
.bind(gid)
.bind(branch_id)
.bind(op.as_str())
.bind(url)
.bind(BranchStatus::Prepared.as_str())
.bind(t)
.bind(t)
.execute(&mut *tx)
.await?;
}
tx.commit().await?;
Ok(())
}
pub async fn list_recent(&self, limit: i64) -> Result<Vec<GlobalRow>> {
let rows = sqlx::query(&self.be.q(&format!(
"{SELECT_GLOBAL} ORDER BY create_time DESC LIMIT ?"
)))
.bind(limit)
.fetch_all(&self.pool)
.await?;
Ok(rows.into_iter().map(global_from_row).collect())
}
}
const SELECT_GLOBAL: &str = "SELECT gid,trans_type,status,payload,next_cron_time,
next_cron_interval,owner,rollback_reason,query_prepared,create_time,finish_time
FROM trans_global";
fn global_from_row(r: AnyRow) -> GlobalRow {
GlobalRow {
gid: r.get("gid"),
trans_type: TransType::parse(r.get::<String, _>("trans_type").as_str())
.unwrap_or(TransType::Saga),
status: GlobalStatus::parse(r.get::<String, _>("status").as_str())
.unwrap_or(GlobalStatus::Prepared),
payload: r.get("payload"),
next_cron_time: r.get("next_cron_time"),
next_cron_interval: r.get("next_cron_interval"),
owner: r.get("owner"),
rollback_reason: r.get("rollback_reason"),
query_prepared: r.get("query_prepared"),
create_time: r.get("create_time"),
finish_time: r.get("finish_time"),
}
}
#[cfg(test)]
mod tests {
use super::*;
static PG_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
fn require_real_db(缺的变量: &str) {
if std::env::var("DTMRS_TEST_REQUIRE_REAL_DB").is_ok() {
panic!(
"设了 DTMRS_TEST_REQUIRE_REAL_DB,却没有 {缺的变量} —— \
这是 CI 配置坏了(容器没起来?变量名打错?),不是可以跳过的情况"
);
}
}
async fn backends() -> (
tokio::sync::MutexGuard<'static, ()>,
Vec<(&'static str, Store)>,
) {
let guard = PG_LOCK.lock().await;
let mut v = vec![("sqlite", Store::open("sqlite::memory:").await.unwrap())];
for (name, env) in [("postgres", "DTMRS_TEST_PG"), ("mysql", "DTMRS_TEST_MYSQL")] {
if std::env::var(env).is_err() {
require_real_db(env);
continue;
}
if let Ok(url) = std::env::var(env) {
let s = Store::open(&url)
.await
.unwrap_or_else(|e| panic!("连不上 {env}: {e}"));
for t in ["trans_branch_op", "trans_global"] {
sqlx::query(&format!("DELETE FROM {t}"))
.execute(s.pool().expect("SQL 后端才有连接池"))
.await
.expect("清表");
}
v.push((name, s));
}
}
#[cfg(feature = "redis")]
if std::env::var("DTMRS_TEST_REDIS").is_err() {
require_real_db("DTMRS_TEST_REDIS");
}
#[cfg(feature = "redis")]
if let Ok(url) = std::env::var("DTMRS_TEST_REDIS") {
let s = Store::open(&url)
.await
.unwrap_or_else(|e| panic!("连不上 DTMRS_TEST_REDIS: {e}"));
s.as_redis()
.unwrap()
.flush_prefix()
.await
.expect("清 redis");
v.push(("redis", s));
}
(guard, v)
}
fn g(gid: &str) -> GlobalRow {
GlobalRow {
gid: gid.into(),
trans_type: TransType::Saga,
status: GlobalStatus::Submitted,
payload: "{}".into(),
next_cron_time: 0,
next_cron_interval: 0,
owner: String::new(),
rollback_reason: String::new(),
query_prepared: String::new(),
create_time: 0,
finish_time: None,
}
}
#[tokio::test]
async fn 重复提交同一个gid是幂等的() {
let (_g, bes) = backends().await;
for (name, s) in bes {
assert!(s.create_global(&g("t1"), &[]).await.unwrap(), "{name}");
assert!(!s.create_global(&g("t1"), &[]).await.unwrap(), "{name}");
assert_eq!(s.list_recent(10).await.unwrap().len(), 1, "{name}");
}
}
#[tokio::test]
async fn 租约只能被抢到一次() {
let (_g, bes) = backends().await;
for (name, s) in bes {
s.create_global(&g("t2"), &[]).await.unwrap();
let a = s.lock_one_due("worker-a", 60).await.unwrap();
assert!(a.is_some(), "{name}: 第一个实例应该抢到");
let b = s.lock_one_due("worker-b", 60).await.unwrap();
assert!(b.is_none(), "{name}: 租约期内不能被别人抢走");
}
}
#[tokio::test]
async fn 并发抢占要各拿各的不能全挤在同一笔上() {
const K: usize = 6;
let (_g, bes) = backends().await;
for (name, s) in bes {
for i in 0..K {
s.create_global(&g(&format!("par-{i}")), &[]).await.unwrap();
}
let mut hs = Vec::new();
for i in 0..K {
let s = s.clone();
hs.push(tokio::spawn(async move {
s.lock_one_due(&format!("w-{i}"), 60).await.unwrap()
}));
}
let mut got: Vec<String> = Vec::new();
for h in hs {
if let Some(row) = h.await.unwrap() {
got.push(row.gid);
}
}
let uniq: std::collections::HashSet<_> = got.iter().collect();
assert_eq!(uniq.len(), got.len(), "{name}: 同一笔被抢到了两次");
if name != "sqlite" {
assert_eq!(
got.len(),
K,
"{name}: 并发抢占退化成串行了(SKIP LOCKED 没生效?)"
);
}
}
}
#[tokio::test]
async fn 终态不再被调度() {
let (_g, bes) = backends().await;
for (name, s) in bes {
s.create_global(&g("t3"), &[]).await.unwrap();
s.set_global_status("t3", GlobalStatus::Succeed, TransType::Saga, "")
.await
.unwrap();
assert!(s.lock_one_due("w", 60).await.unwrap().is_none(), "{name}");
let got = s.get_global("t3").await.unwrap().unwrap();
assert_eq!(got.status, GlobalStatus::Succeed, "{name}");
assert!(got.finish_time.is_some(), "{name}: 终态要落 finish_time");
}
}
#[tokio::test]
async fn 分支状态可更新() {
let (_g, bes) = backends().await;
for (name, s) in bes {
let b = BranchRow {
gid: "t4".into(),
branch_id: "01".into(),
op: BranchOp::Action,
url: "http://x/a".into(),
payload: "{}".into(),
status: BranchStatus::Prepared,
};
s.create_global(&g("t4"), std::slice::from_ref(&b))
.await
.unwrap();
s.set_branch_status("t4", "01", BranchOp::Action, BranchStatus::Succeed)
.await
.unwrap();
let got = s.list_branches("t4").await.unwrap();
assert_eq!(got.len(), 1, "{name}");
assert_eq!(got[0].status, BranchStatus::Succeed, "{name}");
}
}
#[tokio::test]
async fn 回滚原因和回查地址能存取() {
let (_g, bes) = backends().await;
for (name, s) in bes {
let mut row = g("t5");
row.query_prepared = "http://busi/query".into();
s.create_global(&row, &[]).await.unwrap();
s.set_global_status(
"t5",
GlobalStatus::Aborting,
TransType::Saga,
"分支 02 返回 FAILURE",
)
.await
.unwrap();
let got = s.get_global("t5").await.unwrap().unwrap();
assert_eq!(got.query_prepared, "http://busi/query", "{name}");
assert_eq!(got.rollback_reason, "分支 02 返回 FAILURE", "{name}");
assert!(
got.finish_time.is_none(),
"{name}: 非终态不该有 finish_time"
);
s.set_global_status("t5", GlobalStatus::Failed, TransType::Saga, "")
.await
.unwrap();
let got = s.get_global("t5").await.unwrap().unwrap();
assert_eq!(
got.rollback_reason, "分支 02 返回 FAILURE",
"{name}: 空原因不能覆盖"
);
}
}
#[tokio::test]
async fn msg的prepared会被捞tcc的不会() {
let (_g, bes) = backends().await;
for (name, s) in bes {
let mut m = g("m1");
m.trans_type = TransType::Msg;
m.status = GlobalStatus::Prepared;
s.create_global(&m, &[]).await.unwrap();
let mut t = g("c1");
t.trans_type = TransType::Tcc;
t.status = GlobalStatus::Prepared;
s.create_global(&t, &[]).await.unwrap();
let got = s.lock_one_due("w", 60).await.unwrap();
assert_eq!(
got.map(|x| x.gid),
Some("m1".to_string()),
"{name}: 只该捞到 msg"
);
assert!(s.lock_one_due("w2", 60).await.unwrap().is_none(), "{name}");
}
}
#[tokio::test]
async fn 分支登记是幂等的() {
let (_g, bes) = backends().await;
for (name, s) in bes {
let mut t = g("c2");
t.trans_type = TransType::Tcc;
s.create_global(&t, &[]).await.unwrap();
let ops = [
(BranchOp::Confirm, "http://x/c".to_string()),
(BranchOp::Cancel, "http://x/n".to_string()),
];
s.register_branch("c2", "01", &ops).await.unwrap();
s.register_branch("c2", "01", &ops).await.unwrap(); assert_eq!(
s.list_branches("c2").await.unwrap().len(),
2,
"{name}: 不该重复插入"
);
}
}
}
#[cfg(feature = "redis")]
pub mod redis_store;
#[cfg(feature = "redis")]
pub use redis_store::RedisStore;
#[derive(Clone)]
enum Inner {
Sql(SqlStore),
#[cfg(feature = "redis")]
Redis(RedisStore),
}
#[derive(Clone)]
pub struct Store {
inner: Inner,
}
pub type StoreError = sqlx::Error;
#[cfg(feature = "redis")]
fn redis_err(e: redis::RedisError) -> sqlx::Error {
sqlx::Error::Configuration(Box::new(e))
}
pub fn is_redis_url(url: &str) -> bool {
let u = url.trim().to_ascii_lowercase();
u.starts_with("redis://") || u.starts_with("rediss://") || u.starts_with("redis+unix:")
}
impl Store {
pub async fn open(url: &str) -> Result<Self> {
if is_redis_url(url) {
#[cfg(feature = "redis")]
{
let r = RedisStore::open(url).await.map_err(redis_err)?;
return Ok(Self {
inner: Inner::Redis(r),
});
}
#[cfg(not(feature = "redis"))]
{
return Err(sqlx::Error::Configuration(
"这个 URL 要 Redis 后端,但构建时没开 dtmrs-store 的 `redis` feature".into(),
));
}
}
Ok(Self {
inner: Inner::Sql(SqlStore::open(url).await?),
})
}
pub fn is_redis(&self) -> bool {
match &self.inner {
Inner::Sql(_) => false,
#[cfg(feature = "redis")]
Inner::Redis(_) => true,
}
}
pub fn pool(&self) -> Option<&AnyPool> {
match &self.inner {
Inner::Sql(s) => Some(s.pool()),
#[cfg(feature = "redis")]
Inner::Redis(_) => None,
}
}
pub fn backend(&self) -> Option<Backend> {
match &self.inner {
Inner::Sql(s) => Some(s.backend()),
#[cfg(feature = "redis")]
Inner::Redis(_) => None,
}
}
#[cfg(feature = "redis")]
pub fn as_redis(&self) -> Option<&RedisStore> {
match &self.inner {
Inner::Redis(r) => Some(r),
_ => None,
}
}
}
macro_rules! dispatch {
($( $(#[$m:meta])* fn $name:ident (&self $(, $arg:ident : $ty:ty)* ) -> $ret:ty; )*) => {
impl Store {
$(
$(#[$m])*
pub async fn $name(&self $(, $arg: $ty)*) -> Result<$ret> {
match &self.inner {
Inner::Sql(s) => s.$name($($arg),*).await,
#[cfg(feature = "redis")]
Inner::Redis(r) => r.$name($($arg),*).await.map_err(redis_err),
}
}
)*
}
};
}
dispatch! {
fn migrate(&self) -> ();
fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> bool;
fn get_global(&self, gid: &str) -> Option<GlobalRow>;
fn list_branches(&self, gid: &str) -> Vec<BranchRow>;
fn lock_one_due(&self, owner: &str, lease: i64) -> Option<GlobalRow>;
fn set_global_status(&self, gid: &str, status: GlobalStatus, trans_type: TransType, reason: &str) -> ();
fn submit_prepared(&self, gid: &str, owner: &str, next_cron_time: i64) -> SubmitOutcome;
fn set_branch_result(&self, gid: &str, branch_id: &str, op: BranchOp, status: BranchStatus, payload: &str) -> ();
fn set_branch_status(&self, gid: &str, branch_id: &str, op: BranchOp, status: BranchStatus) -> ();
fn schedule_retry(&self, gid: &str, interval: i64) -> ();
fn schedule_now(&self, gid: &str) -> ();
fn register_branch(&self, gid: &str, branch_id: &str, ops: &[(BranchOp, String)]) -> ();
fn list_recent(&self, limit: i64) -> Vec<GlobalRow>;
}