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 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 Store {
pool: AnyPool,
be: Backend,
}
static DRIVERS: Once = Once::new();
impl Store {
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 { 8 };
let be = Backend::from_url(&url);
let pool = AnyPoolOptions::new()
.max_connections(max)
.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.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,
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("SELECT gid FROM trans_global
WHERE (status IN ('submitted','aborting')
OR (status = 'prepared' AND trans_type = 'msg'))
AND next_cron_time <= ?
ORDER BY next_cron_time LIMIT 1"))
.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 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(());
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 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())
.await
.expect("清表");
}
v.push((name, 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 终态不再被调度() {
let (_g, bes) = backends().await;
for (name, s) in bes {
s.create_global(&g("t3"), &[]).await.unwrap();
s.set_global_status("t3", GlobalStatus::Succeed, "")
.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, "分支 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, "")
.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}: 不该重复插入"
);
}
}
}