use crate::{TokenRow, BranchRow, GlobalRow, SubmitOutcome, MID};
use dtmrs_core::dialect::check_len;
use dtmrs_core::{Backend, BranchOp, BranchStatus, GlobalStatus, TransType};
use redis::aio::MultiplexedConnection;
use redis::AsyncCommands;
pub const RECENT_CAP: isize = 1000;
pub const DEFAULT_FINAL_TTL: i64 = 7 * 24 * 3600;
const SCAN_LIMIT: usize = 20;
type Result<T> = std::result::Result<T, redis::RedisError>;
fn err(msg: &str) -> redis::RedisError {
redis::RedisError::from((redis::ErrorKind::Client, "dtmrs", msg.to_string()))
}
fn schedulable(status: GlobalStatus, tt: TransType) -> bool {
matches!(status, GlobalStatus::Submitted | GlobalStatus::Aborting)
|| (status == GlobalStatus::Prepared && tt == TransType::Msg)
}
const LUA_SCHEDULABLE: &str = r#"
local function schedulable(status, tt)
if status == 'submitted' or status == 'aborting' then return true end
if status == 'prepared' and tt == 'msg' then return true end
return false
end
"#;
#[derive(Clone)]
pub struct RedisStore {
conn: MultiplexedConnection,
prefix: String,
final_ttl: i64,
}
impl RedisStore {
pub async fn open(url: &str) -> Result<Self> {
let client = redis::Client::open(url)?;
let conn = client.get_multiplexed_async_connection().await?;
Ok(Self {
conn,
prefix: "dtmrs:".to_string(),
final_ttl: DEFAULT_FINAL_TTL,
})
}
pub fn with_prefix(mut self, p: &str) -> Self {
self.prefix = p.to_string();
self
}
pub fn with_ttl(mut self, secs: i64) -> Self {
self.final_ttl = secs;
self
}
fn gkey(&self, gid: &str) -> String {
format!("{}g:{}", self.prefix, gid)
}
fn bkey(&self, gid: &str) -> String {
format!("{}b:{}", self.prefix, gid)
}
fn ikey(&self) -> String {
format!("{}idx", self.prefix)
}
fn akey(&self) -> String {
format!("{}all", self.prefix)
}
fn tkey(&self, hash: &str) -> String {
format!("{}tok:{}", self.prefix, hash)
}
fn tidx(&self) -> String {
format!("{}tokens", self.prefix)
}
fn bfield(branch_id: &str, op: BranchOp) -> String {
format!("{}\x1f{}", branch_id, op.as_str())
}
fn sfield(branch_id: &str, op: BranchOp) -> String {
format!("{}\x1f{}\x1fs", branch_id, op.as_str())
}
const SUFFIX: &'static str = "\x1fs";
fn parse_bfield(f: &str) -> Option<(String, BranchOp)> {
let (b, o) = f.split_once('\x1f')?;
Some((b.to_string(), BranchOp::parse(o)?))
}
fn check_global(g: &GlobalRow) -> Result<()> {
check_len("gid", &g.gid, Backend::ID_MAX).map_err(|e| err(&e.to_string()))?;
check_len("payload", &g.payload, crate::BIG).map_err(|e| err(&e.to_string()))?;
check_len("query_prepared", &g.query_prepared, MID).map_err(|e| err(&e.to_string()))?;
Ok(())
}
fn check_branch(b: &BranchRow) -> Result<()> {
check_len("branch_id", &b.branch_id, Backend::ID_MAX).map_err(|e| err(&e.to_string()))?;
check_len("url", &b.url, MID).map_err(|e| err(&e.to_string()))?;
check_len("payload", &b.payload, MID).map_err(|e| err(&e.to_string()))?;
Ok(())
}
fn branch_value(b: &BranchRow) -> String {
serde_json::json!({
"url": b.url,
"payload": b.payload,
})
.to_string()
}
fn branch_from(gid: &str, field: &str, val: &str, status: Option<&str>) -> Option<BranchRow> {
let (branch_id, op) = Self::parse_bfield(field)?;
let v: serde_json::Value = serde_json::from_str(val).ok()?;
Some(BranchRow {
gid: gid.to_string(),
branch_id,
op,
url: v
.get("url")
.and_then(|x| x.as_str())
.unwrap_or("")
.to_string(),
payload: v
.get("payload")
.and_then(|x| x.as_str())
.unwrap_or("")
.to_string(),
status: status
.and_then(BranchStatus::parse)
.unwrap_or(BranchStatus::Prepared),
})
}
fn global_from(map: &std::collections::HashMap<String, String>) -> Option<GlobalRow> {
let get = |k: &str| map.get(k).cloned().unwrap_or_default();
let num = |k: &str| map.get(k).and_then(|v| v.parse::<i64>().ok()).unwrap_or(0);
if map.is_empty() {
return None;
}
Some(GlobalRow {
gid: get("gid"),
trans_type: TransType::parse(&get("trans_type")).unwrap_or(TransType::Saga),
status: GlobalStatus::parse(&get("status")).unwrap_or(GlobalStatus::Prepared),
payload: get("payload"),
next_cron_time: num("next_cron_time"),
next_cron_interval: num("next_cron_interval"),
owner: get("owner"),
rollback_reason: get("rollback_reason"),
query_prepared: get("query_prepared"),
create_time: num("create_time"),
finish_time: map.get("finish_time").and_then(|v| v.parse::<i64>().ok()),
})
}
pub async fn migrate(&self) -> Result<()> {
Ok(())
}
pub async fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> Result<bool> {
Self::check_global(g)?;
for b in branches {
Self::check_branch(b)?;
}
let t = crate::now();
let script = redis::Script::new(&format!(
r#"
if redis.call('EXISTS', KEYS[1]) == 1 then return 0 end
redis.call('HSET', KEYS[1],
'gid', ARGV[1], 'trans_type', ARGV[2], 'status', ARGV[3],
'payload', ARGV[4], 'next_cron_time', ARGV[5],
'next_cron_interval', ARGV[6], 'owner', ARGV[10], 'rollback_reason', '',
'query_prepared', ARGV[7], 'create_time', ARGV[8], 'update_time', ARGV[8])
-- 分支从第 11 个 ARGV 开始,每两个一组(field, value),
-- 定义和状态各占一组。**一条 HSET 全写完** —— 原来是每个字段
-- 一条 HSETNX,分支多的时候命令数线性涨。
-- 上面已经确认过全局键不存在,所以不需要 NX 语义
if #ARGV >= 11 then
redis.call('HSET', KEYS[2], unpack(ARGV, 11))
end
if ARGV[9] == '1' then
redis.call('ZADD', KEYS[3], ARGV[5], ARGV[1])
end
redis.call('ZADD', KEYS[4], ARGV[8], ARGV[1])
-- ⚠ 裁剪**必须留在写路径上**。
--
-- 曾经为了省一条命令把它挪去 list_recent,结果是:没人打开管理台
-- 就永远不裁剪。实测跑 5000 笔之后索引里就是 5000 个成员(上限本该
-- 是 1000)—— 每笔事务留一个成员且永不回收,跑久了必然撑爆内存。
-- 而它本身极便宜:索引已经在上限内时删 0 个成员,只是一次 O(log N)。
-- 省 4% 的命令数换一个无界增长,不划算
redis.call('ZREMRANGEBYRANK', KEYS[4], 0, -{cap})
return 1
"#,
cap = RECENT_CAP + 1
));
let mut inv = script.prepare_invoke();
inv.key(self.gkey(&g.gid))
.key(self.bkey(&g.gid))
.key(self.ikey())
.key(self.akey())
.arg(&g.gid)
.arg(g.trans_type.to_string())
.arg(g.status.as_str())
.arg(&g.payload)
.arg(g.next_cron_time)
.arg(g.next_cron_interval)
.arg(&g.query_prepared)
.arg(t)
.arg(if schedulable(g.status, g.trans_type) {
"1"
} else {
"0"
})
.arg(&g.owner);
for b in branches {
inv.arg(Self::bfield(&b.branch_id, b.op))
.arg(Self::branch_value(b));
if b.status != BranchStatus::Prepared {
inv.arg(Self::sfield(&b.branch_id, b.op))
.arg(b.status.as_str());
}
}
let mut c = self.conn.clone();
let created: i64 = inv.invoke_async(&mut c).await?;
Ok(created == 1)
}
pub async fn get_global(&self, gid: &str) -> Result<Option<GlobalRow>> {
let mut c = self.conn.clone();
let map: std::collections::HashMap<String, String> = c.hgetall(self.gkey(gid)).await?;
Ok(Self::global_from(&map))
}
pub async fn list_branches(&self, gid: &str) -> Result<Vec<BranchRow>> {
let mut c = self.conn.clone();
let map: std::collections::HashMap<String, String> = c.hgetall(self.bkey(gid)).await?;
let mut v: Vec<BranchRow> = map
.iter()
.filter(|(f, _)| !f.ends_with(Self::SUFFIX))
.filter_map(|(f, val)| {
Self::branch_from(
gid,
f,
val,
map.get(&format!("{f}{}", Self::SUFFIX)).map(|s| s.as_str()),
)
})
.collect();
v.sort_by(|a, b| {
a.branch_id
.cmp(&b.branch_id)
.then(a.op.as_str().cmp(b.op.as_str()))
});
Ok(v)
}
pub async fn lock_one_due(&self, owner: &str, lease: i64) -> Result<Option<GlobalRow>> {
let t = crate::now();
let script = redis::Script::new(&format!(
r#"
{sched}
local cands = redis.call('ZRANGEBYSCORE', KEYS[1], '-inf', ARGV[1],
'LIMIT', 0, {scan})
for i = 1, #cands do
local gid = cands[i]
local gkey = ARGV[4] .. gid
-- 一条 HMGET 拿两个字段。索引自愈保留着(下面两个 ZREM),
-- 只是把两次 HGET 合成一次
local v = redis.call('HMGET', gkey, 'status', 'trans_type')
local st, tt = v[1], v[2]
if not st then
-- 事务本体没了(过期了),索引里的残留清掉
redis.call('ZREM', KEYS[1], gid)
elseif not schedulable(st, tt) then
-- 状态已经不该被调度,从索引摘掉
redis.call('ZREM', KEYS[1], gid)
else
-- 抢到了:立刻把下次调度时间推到租约之后,等于占坑
redis.call('HSET', gkey, 'owner', ARGV[2],
'next_cron_time', ARGV[3], 'update_time', ARGV[1])
redis.call('ZADD', KEYS[1], ARGV[3], gid)
return redis.call('HGETALL', gkey)
end
end
return nil
"#,
sched = LUA_SCHEDULABLE,
scan = SCAN_LIMIT
));
let mut c = self.conn.clone();
let flat: Option<Vec<String>> = script
.key(self.ikey())
.arg(t)
.arg(owner)
.arg(t + lease)
.arg(format!("{}g:", self.prefix))
.invoke_async(&mut c)
.await?;
let Some(flat) = flat else { return Ok(None) };
let map: std::collections::HashMap<String, String> = flat
.chunks(2)
.filter(|c| c.len() == 2)
.map(|c| (c[0].clone(), c[1].clone()))
.collect();
Ok(Self::global_from(&map))
}
pub async fn set_global_status(
&self,
gid: &str,
status: GlobalStatus,
trans_type: TransType,
reason: &str,
) -> Result<()> {
let t = crate::now();
let reason: String = reason.chars().take(MID).collect();
if !schedulable(status, trans_type) {
let mut c = self.conn.clone();
let mut pipe = redis::pipe();
pipe.atomic();
let mut fields: Vec<(&str, String)> = vec![
("status", status.as_str().to_string()),
("update_time", t.to_string()),
];
if status.is_final() {
fields.push(("finish_time", t.to_string()));
}
if !reason.is_empty() {
fields.push(("rollback_reason", reason.clone()));
}
pipe.hset_multiple(self.gkey(gid), &fields).ignore();
pipe.zrem(self.ikey(), gid).ignore();
if status.is_final() && self.final_ttl > 0 {
pipe.expire(self.gkey(gid), self.final_ttl).ignore();
pipe.expire(self.bkey(gid), self.final_ttl).ignore();
}
let _: () = pipe.query_async(&mut c).await?;
return Ok(());
}
let script = redis::Script::new(
r#"
-- 走到这里说明新状态是**可调度**的(submitted / aborting /
-- msg 的 prepared),得把它按原来的到期时间放回索引,
-- 所以还是要读一次 next_cron_time。
-- trans_type 由调用方传进来了,不用再读(见函数头注释)
local nct = redis.call('HGET', KEYS[1], 'next_cron_time')
if not nct then return 0 end
-- 要写的字段先攒起来,最后一条 HSET 落完 —— 原来状态、finish_time、
-- rollback_reason 是分三条写的。
-- ⚠ 这里是普通字符串字面量不是 format!,花括号**不要写成 {{}}**
local f = {'status', ARGV[1], 'update_time', ARGV[2]}
if ARGV[3] == '1' then
f[#f + 1] = 'finish_time'
f[#f + 1] = ARGV[2]
end
if ARGV[4] ~= '' then
f[#f + 1] = 'rollback_reason'
f[#f + 1] = ARGV[4]
end
redis.call('HSET', KEYS[1], unpack(f))
redis.call('ZADD', KEYS[2], nct, ARGV[5])
return 1
"#,
);
let mut c = self.conn.clone();
let _: i64 = script
.key(self.gkey(gid))
.key(self.ikey())
.arg(status.as_str())
.arg(t)
.arg(if status.is_final() { "1" } else { "0" })
.arg(reason)
.arg(gid)
.invoke_async(&mut c)
.await?;
Ok(())
}
pub async fn set_branch_result(
&self,
gid: &str,
branch_id: &str,
op: BranchOp,
status: BranchStatus,
payload: &str,
) -> Result<()> {
check_len("payload", payload, MID).map_err(|e| err(&e.to_string()))?;
let script = redis::Script::new(
r#"
local cur = redis.call('HGET', KEYS[1], ARGV[1])
if not cur then return 0 end
local v = cjson.decode(cur)
v['payload'] = ARGV[3]
redis.call('HSET', KEYS[1], ARGV[1], cjson.encode(v), ARGV[4], ARGV[2])
return 1
"#,
);
let mut c = self.conn.clone();
let _: i64 = script
.key(self.bkey(gid))
.arg(Self::bfield(branch_id, op))
.arg(status.as_str())
.arg(payload)
.arg(Self::sfield(branch_id, op))
.invoke_async(&mut c)
.await?;
Ok(())
}
pub async fn set_branch_status(
&self,
gid: &str,
branch_id: &str,
op: BranchOp,
status: BranchStatus,
) -> Result<()> {
let mut c = self.conn.clone();
let _: () = c
.hset(self.bkey(gid), Self::sfield(branch_id, op), status.as_str())
.await?;
Ok(())
}
pub async fn schedule_retry(&self, gid: &str, interval: i64) -> Result<()> {
let t = crate::now();
let script = redis::Script::new(
r#"
if redis.call('EXISTS', KEYS[1]) == 0 then return 0 end
redis.call('HSET', KEYS[1], 'next_cron_interval', ARGV[1],
'next_cron_time', ARGV[2], 'update_time', ARGV[3])
-- 只更新已经在索引里的,别把不该调度的塞回去
if redis.call('ZSCORE', KEYS[2], ARGV[4]) then
redis.call('ZADD', KEYS[2], ARGV[2], ARGV[4])
end
return 1
"#,
);
let mut c = self.conn.clone();
let _: i64 = script
.key(self.gkey(gid))
.key(self.ikey())
.arg(interval)
.arg(t + interval)
.arg(t)
.arg(gid)
.invoke_async(&mut c)
.await?;
Ok(())
}
pub async fn submit_prepared(
&self,
gid: &str,
owner: &str,
next_cron_time: i64,
) -> Result<SubmitOutcome> {
let t = crate::now();
let script = redis::Script::new(&format!(
r#"
{sched}
local v = redis.call('HMGET', KEYS[1], 'status', 'trans_type')
if not v[1] then return {{'MISSING'}} end
if v[1] ~= ARGV[3] then return {{'ALREADY'}} end
redis.call('HSET', KEYS[1], 'status', ARGV[2], 'update_time', ARGV[1],
'next_cron_time', ARGV[5], 'next_cron_interval', 0,
'owner', ARGV[6])
if schedulable(ARGV[2], v[2]) then
redis.call('ZADD', KEYS[2], ARGV[5], ARGV[4])
end
-- 尾巴上带回事务体,调用方占了租约就能直接开推,不用再读一次。
-- 首元素是标记,后面是 HGETALL 的 field/value 对
local r = redis.call('HGETALL', KEYS[1])
table.insert(r, 1, 'ADVANCED')
return r
"#,
sched = LUA_SCHEDULABLE
));
let mut c = self.conn.clone();
let r: Vec<String> = script
.key(self.gkey(gid))
.key(self.ikey())
.arg(t)
.arg(GlobalStatus::Submitted.as_str())
.arg(GlobalStatus::Prepared.as_str())
.arg(gid)
.arg(next_cron_time)
.arg(owner)
.invoke_async(&mut c)
.await?;
Ok(match r.first().map(String::as_str) {
Some("ADVANCED") => {
let map: std::collections::HashMap<String, String> = r[1..]
.chunks(2)
.filter(|c| c.len() == 2)
.map(|c| (c[0].clone(), c[1].clone()))
.collect();
match Self::global_from(&map) {
Some(g) => SubmitOutcome::Advanced(Box::new(g)),
None => SubmitOutcome::Already,
}
}
Some("MISSING") => SubmitOutcome::Missing,
_ => SubmitOutcome::Already,
})
}
pub async fn schedule_now(&self, gid: &str) -> Result<()> {
let t = crate::now();
let script = redis::Script::new(&format!(
r#"
{sched}
-- 一条 HMGET 顶掉原来的 EXISTS + HGET(status) + HGET(trans_type)
local v = redis.call('HMGET', KEYS[1], 'status', 'trans_type')
local st, tt = v[1], v[2]
if not st then return 0 end
redis.call('HSET', KEYS[1], 'next_cron_time', ARGV[1], 'next_cron_interval', 0)
-- 这里跟 schedule_retry 不同:submit / abort 之后事务**变得**可调度了,
-- 索引里没有就得加进去
if schedulable(st, tt) then
redis.call('ZADD', KEYS[2], ARGV[1], ARGV[2])
end
return 1
"#,
sched = LUA_SCHEDULABLE
));
let mut c = self.conn.clone();
let _: i64 = script
.key(self.gkey(gid))
.key(self.ikey())
.arg(t)
.arg(gid)
.invoke_async(&mut c)
.await?;
Ok(())
}
pub async fn register_branch(
&self,
gid: &str,
branch_id: &str,
ops: &[(BranchOp, String)],
) -> Result<()> {
check_len("gid", gid, Backend::ID_MAX).map_err(|e| err(&e.to_string()))?;
check_len("branch_id", branch_id, Backend::ID_MAX).map_err(|e| err(&e.to_string()))?;
for (_, url) in ops {
check_len("url", url, MID).map_err(|e| err(&e.to_string()))?;
}
let mut c = self.conn.clone();
for (op, url) in ops {
let row = BranchRow {
gid: gid.to_string(),
branch_id: branch_id.to_string(),
op: *op,
url: url.clone(),
payload: String::new(),
status: BranchStatus::Prepared,
};
let _: i64 = c
.hset_nx(
self.bkey(gid),
Self::bfield(branch_id, *op),
Self::branch_value(&row),
)
.await?;
}
Ok(())
}
pub async fn list_recent(&self, limit: i64) -> Result<Vec<GlobalRow>> {
let mut c = self.conn.clone();
let gids: Vec<String> = c
.zrevrange(self.akey(), 0, (limit.max(1) - 1) as isize)
.await?;
let mut out = Vec::with_capacity(gids.len());
for gid in gids {
let map: std::collections::HashMap<String, String> = c.hgetall(self.gkey(&gid)).await?;
if let Some(g) = Self::global_from(&map) {
out.push(g);
}
}
Ok(out)
}
pub async fn create_token(&self, hash: &str, name: &str, secret: &str) -> Result<()> {
let mut c = self.conn.clone();
let mut pipe = redis::pipe();
pipe.atomic()
.hset_multiple(
self.tkey(hash),
&[
("name", name.to_string()),
("create_time", crate::now().to_string()),
("last_used", "0".into()),
("use_count", "0".into()),
("last_ip", String::new()),
("revoked", "0".into()),
("secret", secret.to_string()),
],
)
.ignore()
.sadd(self.tidx(), hash)
.ignore();
pipe.query_async::<()>(&mut c).await?;
Ok(())
}
pub async fn list_tokens(&self) -> Result<Vec<TokenRow>> {
let mut c = self.conn.clone();
let hashes: Vec<String> = redis::cmd("SMEMBERS")
.arg(self.tidx())
.query_async(&mut c)
.await?;
let mut out = Vec::with_capacity(hashes.len());
for h in hashes {
let m: std::collections::HashMap<String, String> = redis::cmd("HGETALL")
.arg(self.tkey(&h))
.query_async(&mut c)
.await?;
if m.is_empty() {
continue;
}
let g = |k: &str| m.get(k).cloned().unwrap_or_default();
let n = |k: &str| g(k).parse::<i64>().unwrap_or(0);
out.push(TokenRow {
token_hash: h,
name: g("name"),
create_time: n("create_time"),
last_used: n("last_used"),
use_count: n("use_count"),
last_ip: g("last_ip"),
revoked: n("revoked"),
secret: g("secret"),
});
}
out.sort_by(|a, b| b.create_time.cmp(&a.create_time));
Ok(out)
}
pub async fn revoke_token(&self, hash: &str) -> Result<bool> {
let mut c = self.conn.clone();
let script = redis::Script::new(
r#"
if redis.call('EXISTS', KEYS[1]) == 0 then return 0 end
if redis.call('HGET', KEYS[1], 'revoked') ~= '0' then return 0 end
redis.call('HSET', KEYS[1], 'revoked', ARGV[1])
return 1
"#,
);
let n: i64 = script
.key(self.tkey(hash))
.arg(crate::now())
.invoke_async(&mut c)
.await?;
Ok(n == 1)
}
pub async fn active_token_hashes(&self) -> Result<Vec<String>> {
Ok(self
.list_tokens()
.await?
.into_iter()
.filter(|t| t.revoked == 0)
.map(|t| t.token_hash)
.collect())
}
pub async fn touch_token(&self, hash: &str, ip: &str) -> Result<()> {
let mut c = self.conn.clone();
let mut pipe = redis::pipe();
pipe.atomic()
.hset(self.tkey(hash), "last_used", crate::now())
.ignore()
.hincr(self.tkey(hash), "use_count", 1)
.ignore()
.hset(self.tkey(hash), "last_ip", ip)
.ignore();
pipe.query_async::<()>(&mut c).await?;
Ok(())
}
pub async fn flush_prefix(&self) -> Result<()> {
let mut c = self.conn.clone();
let keys: Vec<String> = c.keys(format!("{}*", self.prefix)).await?;
if !keys.is_empty() {
let _: i64 = c.del(keys).await?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn 可调度判断跟sql后端的where一致() {
assert!(schedulable(GlobalStatus::Submitted, TransType::Saga));
assert!(schedulable(GlobalStatus::Aborting, TransType::Saga));
assert!(schedulable(GlobalStatus::Prepared, TransType::Msg));
assert!(!schedulable(GlobalStatus::Prepared, TransType::Tcc));
assert!(!schedulable(GlobalStatus::Prepared, TransType::Xa));
for tt in [TransType::Saga, TransType::Msg, TransType::Workflow] {
assert!(!schedulable(GlobalStatus::Succeed, tt));
assert!(!schedulable(GlobalStatus::Failed, tt));
}
}
#[test]
fn 分支字段名能往返() {
let f = RedisStore::bfield("01", BranchOp::Compensate);
assert_eq!(
RedisStore::parse_bfield(&f),
Some(("01".to_string(), BranchOp::Compensate))
);
assert!(f.contains('\x1f'));
assert_eq!(RedisStore::parse_bfield("没有分隔符"), None);
}
#[test]
fn 没终结的事务不该有完成时间() {
let mut m = std::collections::HashMap::new();
m.insert("gid".to_string(), "g1".to_string());
m.insert("status".to_string(), "submitted".to_string());
let g = RedisStore::global_from(&m).unwrap();
assert_eq!(g.finish_time, None);
assert_eq!(g.status, GlobalStatus::Submitted);
m.insert("finish_time".to_string(), "123".to_string());
assert_eq!(RedisStore::global_from(&m).unwrap().finish_time, Some(123));
}
#[test]
fn 空哈希解成none() {
assert!(RedisStore::global_from(&std::collections::HashMap::new()).is_none());
}
}