use crate::{BranchRow, GlobalRow, 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 bfield(branch_id: &str, op: BranchOp) -> String {
format!("{}\x1f{}", branch_id, op.as_str())
}
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,
"status": b.status.as_str(),
})
.to_string()
}
fn branch_from(gid: &str, field: &str, val: &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: v
.get("status")
.and_then(|x| x.as_str())
.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', '', 'rollback_reason', '',
'query_prepared', ARGV[7], 'create_time', ARGV[8], 'update_time', ARGV[8])
-- 分支从第 10 个 ARGV 开始,每两个一组(field, value)
for i = 10, #ARGV, 2 do
redis.call('HSETNX', KEYS[2], ARGV[i], ARGV[i + 1])
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])
-- 管理视图只看最近的,索引留个上限免得内存无限涨
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"
});
for b in branches {
inv.arg(Self::bfield(&b.branch_id, b.op))
.arg(Self::branch_value(b));
}
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_map(|(f, val)| Self::branch_from(gid, f, val))
.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
local st = redis.call('HGET', gkey, 'status')
local tt = redis.call('HGET', gkey, 'trans_type')
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,
reason: &str,
) -> Result<()> {
let t = crate::now();
let reason: String = reason.chars().take(MID).collect();
let script = redis::Script::new(&format!(
r#"
{sched}
if redis.call('EXISTS', KEYS[1]) == 0 then return 0 end
local tt = redis.call('HGET', KEYS[1], 'trans_type')
redis.call('HSET', KEYS[1], 'status', ARGV[1], 'update_time', ARGV[2])
if ARGV[3] == '1' then
redis.call('HSET', KEYS[1], 'finish_time', ARGV[2])
end
if ARGV[4] ~= '' then
redis.call('HSET', KEYS[1], 'rollback_reason', ARGV[4])
end
if schedulable(ARGV[1], tt) then
local nct = redis.call('HGET', KEYS[1], 'next_cron_time')
redis.call('ZADD', KEYS[2], nct, ARGV[5])
else
redis.call('ZREM', KEYS[2], ARGV[5])
end
-- 终态挂 TTL:秒杀跑几千万笔之后,不回收内存会撑爆
if ARGV[3] == '1' and tonumber(ARGV[6]) > 0 then
redis.call('EXPIRE', KEYS[1], ARGV[6])
redis.call('EXPIRE', KEYS[3], ARGV[6])
end
return 1
"#,
sched = LUA_SCHEDULABLE
));
let mut c = self.conn.clone();
let _: i64 = script
.key(self.gkey(gid))
.key(self.ikey())
.key(self.bkey(gid))
.arg(status.as_str())
.arg(t)
.arg(if status.is_final() { "1" } else { "0" })
.arg(reason)
.arg(gid)
.arg(self.final_ttl)
.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['status'] = ARGV[2]
if ARGV[4] == '1' then v['payload'] = ARGV[3] end
redis.call('HSET', KEYS[1], ARGV[1], cjson.encode(v))
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("1")
.invoke_async(&mut c)
.await?;
Ok(())
}
pub async fn set_branch_status(
&self,
gid: &str,
branch_id: &str,
op: BranchOp,
status: BranchStatus,
) -> Result<()> {
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['status'] = ARGV[2]
redis.call('HSET', KEYS[1], ARGV[1], cjson.encode(v))
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())
.invoke_async(&mut c)
.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 schedule_now(&self, gid: &str) -> Result<()> {
let t = crate::now();
let script = redis::Script::new(&format!(
r#"
{sched}
if redis.call('EXISTS', KEYS[1]) == 0 then return 0 end
redis.call('HSET', KEYS[1], 'next_cron_time', ARGV[1], 'next_cron_interval', 0)
local st = redis.call('HGET', KEYS[1], 'status')
local tt = redis.call('HGET', KEYS[1], 'trans_type')
-- 这里跟 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 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());
}
}