#![cfg(feature = "redis")]
use dtmrs_core::{BranchResult, GlobalStatus, SagaStep, TransType};
use dtmrs_server::driver::Driver;
use dtmrs_server::saga_rows;
use dtmrs_store::Store;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
static REDIS_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
async fn store(prefix_hint: &str) -> Option<(tokio::sync::MutexGuard<'static, ()>, Store)> {
let guard = REDIS_LOCK.lock().await;
let Ok(url) = std::env::var("DTMRS_TEST_REDIS") else {
require_real_db("DTMRS_TEST_REDIS");
eprintln!(
"\n⚠ 跳过 Redis 测试({prefix_hint}):DTMRS_TEST_REDIS 没配。\n \
这不等于 Redis 后端通过 —— 它只有对着真 Redis 才能验。\n"
);
return None;
};
let s = Store::open(&url).await.expect("连不上 Redis");
s.as_redis().unwrap().flush_prefix().await.unwrap();
Some((guard, s))
}
#[derive(Default)]
struct Busi {
calls: Mutex<HashMap<String, usize>>,
fail_on: Mutex<Option<String>>,
}
impl Busi {
fn hits(&self, key: &str) -> usize {
*self.calls.lock().unwrap().get(key).unwrap_or(&0)
}
fn total(&self) -> usize {
self.calls.lock().unwrap().values().sum()
}
fn duplicates(&self) -> Vec<(String, usize)> {
self.calls
.lock()
.unwrap()
.iter()
.filter(|(_, &n)| n > 1)
.map(|(k, &n)| (k.clone(), n))
.collect()
}
}
fn registry(busi: Arc<Busi>, names: &[&str]) -> dtmrs_server::registry::Registry {
let mut r = dtmrs_server::registry::Registry::new();
for name in names {
let b = busi.clone();
let n = name.to_string();
r.register(name, move |ctx| {
let (b, n) = (b.clone(), n.clone());
async move {
let key = format!("{}|{}|{}", ctx.gid, ctx.branch_id, ctx.op.as_str());
*b.calls.lock().unwrap().entry(key).or_insert(0) += 1;
tokio::time::sleep(std::time::Duration::from_millis(30)).await;
let fail = b.fail_on.lock().unwrap().clone();
if fail.as_deref() == Some(n.as_str()) {
BranchResult::Failure
} else {
BranchResult::Success
}
}
});
}
r
}
#[tokio::test]
async fn redis_正向提交与逆序补偿() {
let Some((_guard, st)) = store("正向与补偿").await else {
return;
};
let busi = Arc::new(Busi::default());
let d = Driver::new(st.clone(), "tc-1".into()).with_registry(Arc::new(registry(
busi.clone(),
&["扣款", "退款", "发货", "退货"],
)));
let steps = vec![
SagaStep::new("local://扣款", "local://退款"),
SagaStep::new("local://发货", "local://退货"),
];
let (g, br) = saga_rows("r-ok", &steps);
st.create_global(&g, &br).await.unwrap();
d.process(&g).await.unwrap();
assert_eq!(
st.get_global("r-ok").await.unwrap().unwrap().status,
GlobalStatus::Succeed
);
assert_eq!(busi.hits("r-ok|01|action"), 1);
assert_eq!(busi.hits("r-ok|02|action"), 1);
assert_eq!(busi.hits("r-ok|01|compensate"), 0, "成功的事务不该有补偿");
*busi.fail_on.lock().unwrap() = Some("发货".into());
let (g2, br2) = saga_rows("r-fail", &steps);
st.create_global(&g2, &br2).await.unwrap();
d.process(&g2).await.unwrap();
assert_eq!(
st.get_global("r-fail").await.unwrap().unwrap().status,
GlobalStatus::Failed
);
assert_eq!(busi.hits("r-fail|01|compensate"), 1, "扣款要被补偿");
assert_eq!(busi.hits("r-fail|02|compensate"), 1, "失败分支也要补偿");
}
#[tokio::test]
async fn redis_超时不能触发回滚() {
let Some((_guard, st)) = store("超时不回滚").await else {
return;
};
let mut reg = dtmrs_server::registry::Registry::new();
reg.register("超时", |_| async { BranchResult::Unknown });
reg.register("补偿", |_| async { BranchResult::Success });
let d = Driver::new(st.clone(), "tc-1".into()).with_registry(Arc::new(reg));
let steps = vec![SagaStep::new("local://超时", "local://补偿")];
let (g, br) = saga_rows("r-timeout", &steps);
st.create_global(&g, &br).await.unwrap();
d.process(&g).await.unwrap();
assert_eq!(
st.get_global("r-timeout").await.unwrap().unwrap().status,
GlobalStatus::Submitted
);
}
#[tokio::test]
async fn redis_多实例并发不重复推进() {
let Some((_guard, st)) = store("多实例并发").await else {
return;
};
const N: usize = 20;
const INSTANCES: usize = 3;
let busi = Arc::new(Busi::default());
let names = ["扣款", "退款"];
let steps = vec![SagaStep::new("local://扣款", "local://退款")];
for i in 0..N {
let (g, br) = saga_rows(&format!("r-race-{i:02}"), &steps);
st.create_global(&g, &br).await.unwrap();
}
let mut handles = Vec::new();
for inst in 0..INSTANCES {
let st = st.clone();
let reg = Arc::new(registry(busi.clone(), &names));
handles.push(tokio::spawn(async move {
let d = Driver::new(st.clone(), format!("tc-{inst}")).with_registry(reg);
let mut done = 0;
for _ in 0..200 {
match d.store.lock_one_due(&d.owner, 30).await {
Ok(Some(g)) => {
let _ = d.process(&g).await;
done += 1;
}
Ok(None) => tokio::time::sleep(std::time::Duration::from_millis(10)).await,
Err(e) => panic!("抢活失败: {e}"),
}
}
(inst, done)
}));
}
let mut per_instance = Vec::new();
for h in handles {
per_instance.push(h.await.unwrap());
}
let mut succeeded = 0;
for i in 0..N {
let g = st
.get_global(&format!("r-race-{i:02}"))
.await
.unwrap()
.unwrap();
if g.status == GlobalStatus::Succeed {
succeeded += 1;
}
}
assert_eq!(succeeded, N, "{N} 笔都该成功");
let dups = busi.duplicates();
assert!(
dups.is_empty(),
"有分支被重复推进了(这就是重复扣款): {dups:?}"
);
assert_eq!(busi.total(), N, "{N} 笔 × 1 步 = {N} 次调用");
let working: Vec<_> = per_instance.iter().filter(|(_, n)| *n > 0).collect();
println!("各实例处理数: {per_instance:?}");
assert!(
working.len() >= 2,
"至少两个实例该抢到活,否则没真正并发: {per_instance:?}"
);
}
#[tokio::test]
async fn redis_终态会挂ttl() {
let Some((_guard, st)) = store("终态 TTL").await else {
return;
};
let mut reg = dtmrs_server::registry::Registry::new();
reg.register("好", |_| async { BranchResult::Success });
let d = Driver::new(st.clone(), "tc-1".into()).with_registry(Arc::new(reg));
let steps = vec![SagaStep::new("local://好", "local://好")];
let (g, br) = saga_rows("r-ttl", &steps);
st.create_global(&g, &br).await.unwrap();
let url = std::env::var("DTMRS_TEST_REDIS").unwrap();
let client = redis::Client::open(url).unwrap();
let mut c = client.get_multiplexed_async_connection().await.unwrap();
let ttl: i64 = redis::cmd("TTL")
.arg("dtmrs:g:r-ttl")
.query_async(&mut c)
.await
.unwrap();
assert_eq!(ttl, -1, "没终结的事务不该有 TTL,-1 表示永不过期");
d.process(&g).await.unwrap();
assert_eq!(
st.get_global("r-ttl").await.unwrap().unwrap().status,
GlobalStatus::Succeed
);
let ttl: i64 = redis::cmd("TTL")
.arg("dtmrs:g:r-ttl")
.query_async(&mut c)
.await
.unwrap();
assert!(ttl > 0, "终结之后该挂上 TTL,实际 {ttl}");
assert!(
ttl <= dtmrs_store::redis_store::DEFAULT_FINAL_TTL,
"TTL 不该超过默认值"
);
}
#[tokio::test]
async fn redis_终态不再被调度() {
let Some((_guard, st)) = store("终态不调度").await else {
return;
};
let steps = vec![SagaStep::new("local://x", "local://y")];
let (g, br) = saga_rows("r-final", &steps);
st.create_global(&g, &br).await.unwrap();
st.set_global_status("r-final", GlobalStatus::Succeed, TransType::Saga, "")
.await
.unwrap();
let got = st.lock_one_due("tc-1", 30).await.unwrap();
assert!(got.is_none(), "终态事务不该再被捞起来: {got:?}");
}
#[tokio::test]
async fn redis_租约到期后别的实例能接手() {
let Some((_guard, st)) = store("租约接手").await else {
return;
};
let steps = vec![SagaStep::new("local://x", "local://y")];
let (g, br) = saga_rows("r-lease", &steps);
st.create_global(&g, &br).await.unwrap();
let first = st.lock_one_due("tc-1", 1).await.unwrap();
assert_eq!(first.unwrap().gid, "r-lease");
assert!(
st.lock_one_due("tc-2", 30).await.unwrap().is_none(),
"租约期内不该被抢走"
);
tokio::time::sleep(std::time::Duration::from_millis(1200)).await;
let second = st.lock_one_due("tc-2", 30).await.unwrap();
assert_eq!(
second.expect("租约到期后该能接手").owner,
"tc-2",
"接手方要写上自己的 owner"
);
}
fn require_real_db(缺的变量: &str) {
if std::env::var("DTMRS_TEST_REQUIRE_REAL_DB").is_ok() {
panic!(
"设了 DTMRS_TEST_REQUIRE_REAL_DB,却没有 {缺的变量} —— \
这是 CI 配置坏了(容器没起来?变量名打错?),不是可以跳过的情况"
);
}
}