use dtmrs_core::{Backend, BranchOp, GlobalStatus, TransType};
use dtmrs_server::driver::Driver;
use dtmrs_server::tcc_rows;
use dtmrs_store::Store;
use dtmrs_xa::{PreparedXact, Resolved, Xa};
use sqlx::any::AnyPoolOptions;
use sqlx::AnyPool;
struct Target {
name: &'static str,
pool: AnyPool,
be: Backend,
xa: Xa,
}
async fn targets(ids: &[(i32, i64)]) -> Vec<Target> {
sqlx::any::install_default_drivers();
let mut out = Vec::new();
for (name, env) in [
("postgres", "DTMRS_TEST_XA_PG"),
("mysql", "DTMRS_TEST_XA_MYSQL"),
] {
let Ok(url) = std::env::var(env) else {
continue;
};
let be = Backend::from_url(&url);
let xa = Xa::from_url(&url).expect("这两种库都支持 XA");
let pool = open(&url, be).await;
reset(&pool, be, &xa, ids).await;
out.push(Target { name, pool, be, xa });
}
if out.is_empty() {
require_real_db("DTMRS_TEST_XA_PG / DTMRS_TEST_XA_MYSQL");
eprintln!(
"\n⚠ 跳过 XA 测试:DTMRS_TEST_XA_PG / DTMRS_TEST_XA_MYSQL 都没配。\n \
这不等于 XA 通过 —— XA 只有对着真数据库才能验。\n"
);
}
out
}
async fn open(url: &str, be: Backend) -> AnyPool {
let lock_sql = match be {
Backend::Postgres => "SET lock_timeout = '5s'",
Backend::MySql => "SET SESSION innodb_lock_wait_timeout = 5",
_ => "SELECT 1",
};
AnyPoolOptions::new()
.max_connections(6)
.after_connect(move |conn, _| {
Box::pin(async move {
sqlx::query(lock_sql).execute(&mut *conn).await?;
Ok(())
})
})
.connect(url)
.await
.unwrap()
}
async fn create_acct_racy(pool: &AnyPool) {
const SQL: &str = "CREATE TABLE IF NOT EXISTS xa_acct(id INT PRIMARY KEY, bal BIGINT)";
let mut last = None;
for _ in 0..5 {
match sqlx::query(SQL).execute(pool).await {
Ok(_) => return,
Err(e) => {
last = Some(e);
tokio::time::sleep(std::time::Duration::from_millis(80)).await;
}
}
}
panic!("建 xa_acct 失败: {}", last.unwrap());
}
async fn reset(pool: &AnyPool, be: Backend, xa: &Xa, ids: &[(i32, i64)]) {
for x in xa.list_prepared(pool).await.unwrap() {
let _ = xa.rollback_prepared(pool, &x.xid).await;
}
create_acct_racy(pool).await;
for (id, bal) in ids {
let sql = match be {
Backend::MySql => {
"INSERT INTO xa_acct (id,bal) VALUES (?,?) ON DUPLICATE KEY UPDATE bal=VALUES(bal)"
}
_ => "INSERT INTO xa_acct (id,bal) VALUES ($1,$2) ON CONFLICT (id) DO UPDATE SET bal = EXCLUDED.bal",
};
sqlx::query(sql)
.bind(id)
.bind(bal)
.execute(pool)
.await
.unwrap();
}
}
async fn bal(t: &Target, ids: &[i32]) -> Vec<i64> {
let mut v = Vec::new();
for id in ids {
let b: i64 = sqlx::query_scalar(&t.be.q("SELECT bal FROM xa_acct WHERE id=?"))
.bind(id)
.fetch_one(&t.pool)
.await
.unwrap();
v.push(b);
}
v
}
async fn hanging(t: &Target, gid: &str) -> Vec<PreparedXact> {
let key = format!("_{gid}_");
t.xa.list_prepared(&t.pool)
.await
.unwrap()
.into_iter()
.filter(|x| x.xid.contains(&key))
.collect()
}
async fn spawn_rm(pool: AnyPool, xa: Xa) -> String {
use axum::extract::{Query, State};
use axum::routing::post;
use axum::Router;
use std::collections::HashMap;
async fn handler(
State((pool, xa, commit)): State<(AnyPool, Xa, bool)>,
Query(q): Query<HashMap<String, String>>,
) -> &'static str {
let gid = q.get("gid").cloned().unwrap_or_default();
let bid = q.get("branch_id").cloned().unwrap_or_default();
let x = xa.xid(&gid, &bid);
let r = if commit {
xa.commit_prepared(&pool, &x).await
} else {
xa.rollback_prepared(&pool, &x).await
};
match r {
Ok(_) => "SUCCESS",
Err(e) => {
eprintln!("[RM] 解决 {x} 失败: {e}");
"ONGOING"
}
}
}
let app = Router::new()
.route(
"/commit",
post(handler).with_state((pool.clone(), xa, true)),
)
.route("/rollback", post(handler).with_state((pool, xa, false)));
let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = l.local_addr().unwrap();
tokio::spawn(async move { axum::serve(l, app).await.unwrap() });
format!("http://{addr}")
}
async fn new_xa_trans(store: &Store, gid: &str) {
let mut g = tcc_rows(gid);
g.trans_type = TransType::Xa;
store.create_global(&g, &[]).await.unwrap();
}
async fn prepare_branch(
t: &Target,
store: &Store,
base: &str,
gid: &str,
bid: &str,
id: i32,
delta: i64,
) {
store
.register_branch(
gid,
bid,
&[
(BranchOp::Commit, format!("{base}/commit")),
(BranchOp::Rollback, format!("{base}/rollback")),
],
)
.await
.unwrap();
let mut br = t.xa.begin(&t.pool, gid, bid).await.unwrap();
let sql = t.be.q("UPDATE xa_acct SET bal = bal + ? WHERE id = ?");
sqlx::query(&sql)
.bind(delta)
.bind(id)
.execute(br.conn())
.await
.unwrap();
br.prepare().await.unwrap();
}
#[tokio::test]
async fn xa_提交前改动不可见_提交后一起生效() {
for t in targets(&[(11, 1000), (12, 0)]).await {
let base = spawn_rm(t.pool.clone(), t.xa).await;
let store = Store::open("sqlite::memory:").await.unwrap();
let d = Driver::new(store.clone(), "tc".into());
let gid = "xaok";
new_xa_trans(&store, gid).await;
prepare_branch(&t, &store, &base, gid, "01", 11, -100).await;
prepare_branch(&t, &store, &base, gid, "02", 12, 100).await;
assert_eq!(
bal(&t, &[11, 12]).await,
vec![1000, 0],
"{}: prepare 后不该可见",
t.name
);
assert_eq!(hanging(&t, gid).await.len(), 2, "{}: 应有 2 个挂着", t.name);
store
.set_global_status(gid, GlobalStatus::Submitted, TransType::Xa, "")
.await
.unwrap();
let g = store.get_global(gid).await.unwrap().unwrap();
d.process(&g).await.unwrap();
assert_eq!(
store.get_global(gid).await.unwrap().unwrap().status,
GlobalStatus::Succeed,
"{}",
t.name
);
assert_eq!(
bal(&t, &[11, 12]).await,
vec![900, 100],
"{}: 两个分支的改动一起生效",
t.name
);
assert!(
hanging(&t, gid).await.is_empty(),
"{}: 必须全部解决",
t.name
);
}
}
#[tokio::test]
async fn xa_中止则全部回滚_余额不动() {
for t in targets(&[(21, 1000), (22, 0)]).await {
let base = spawn_rm(t.pool.clone(), t.xa).await;
let store = Store::open("sqlite::memory:").await.unwrap();
let d = Driver::new(store.clone(), "tc".into());
let gid = "xarb";
new_xa_trans(&store, gid).await;
prepare_branch(&t, &store, &base, gid, "01", 21, -100).await;
prepare_branch(&t, &store, &base, gid, "02", 22, 100).await;
store
.set_global_status(
gid,
GlobalStatus::Aborting,
TransType::Xa,
"某分支一阶段失败",
)
.await
.unwrap();
let g = store.get_global(gid).await.unwrap().unwrap();
d.process(&g).await.unwrap();
assert_eq!(
store.get_global(gid).await.unwrap().unwrap().status,
GlobalStatus::Failed,
"{}",
t.name
);
assert_eq!(
bal(&t, &[21, 22]).await,
vec![1000, 0],
"{}: 余额一点没动",
t.name
);
assert!(hanging(&t, gid).await.is_empty(), "{}", t.name);
}
}
#[tokio::test]
async fn xa_二阶段幂等_重复提交不报错() {
for t in targets(&[(31, 1000)]).await {
let base = spawn_rm(t.pool.clone(), t.xa).await;
let store = Store::open("sqlite::memory:").await.unwrap();
let gid = "xaidem";
new_xa_trans(&store, gid).await;
prepare_branch(&t, &store, &base, gid, "01", 31, -10).await;
let x = t.xa.xid(gid, "01");
assert_eq!(
t.xa.commit_prepared(&t.pool, &x).await.unwrap(),
Resolved::Done,
"{}",
t.name
);
assert_eq!(
t.xa.commit_prepared(&t.pool, &x).await.unwrap(),
Resolved::AlreadyResolved,
"{}: 重复提交必须是 AlreadyResolved 而不是报错",
t.name
);
assert_eq!(
t.xa.rollback_prepared(&t.pool, &x).await.unwrap(),
Resolved::AlreadyResolved,
"{}",
t.name
);
assert_eq!(bal(&t, &[31]).await, vec![990], "{}: 只生效一次", t.name);
}
}
#[tokio::test]
async fn xa_分支忘了收尾也不会污染连接池() {
for t in targets(&[(41, 1000)]).await {
{
let mut br = t.xa.begin(&t.pool, "xadrop", "01").await.unwrap();
let sql = t.be.q("UPDATE xa_acct SET bal = ? WHERE id = 41");
sqlx::query(&sql)
.bind(12345i64)
.execute(br.conn())
.await
.unwrap();
}
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert!(
hanging(&t, "xadrop").await.is_empty(),
"{}: 没 prepare 的分支不该留下 prepared 事务",
t.name
);
assert_eq!(
bal(&t, &[41]).await,
vec![1000],
"{}: 未收尾的改动必须回滚",
t.name
);
let n: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM xa_acct")
.fetch_one(&t.pool)
.await
.unwrap();
assert!(n >= 1, "{}: 连接池没被污染", t.name);
}
}
#[tokio::test]
async fn xa_启动自检能确认两阶段可用() {
for t in targets(&[(51, 0)]).await {
let n = t.xa.ensure_enabled(&t.pool).await.unwrap_or_else(|e| {
panic!("{}: 自检失败 {e}", t.name);
});
assert!(n > 0, "{}: 自检应该通过,实际 {n}", t.name);
}
}
#[tokio::test]
async fn xa_没解决的prepared事务会阻塞无关写入() {
for t in targets(&[(61, 1000)]).await {
let mut br = t.xa.begin(&t.pool, "xablock", "01").await.unwrap();
let sql = t.be.q("UPDATE xa_acct SET bal = bal - ? WHERE id = 61");
sqlx::query(&sql)
.bind(1i64)
.execute(br.conn())
.await
.unwrap();
let x = br.prepare().await.unwrap();
let h = t.xa.list_prepared(&t.pool).await.unwrap();
assert!(
h.iter().any(|p| p.xid == x),
"{}: list_prepared 要能看到",
t.name
);
let blocked = sqlx::query(&t.be.q("UPDATE xa_acct SET bal = 999 WHERE id = 61"))
.execute(&t.pool)
.await;
assert!(
blocked.is_err(),
"{}: 没解决的 prepared 事务必须阻塞无关写入 —— 这是 XA 的固有代价",
t.name
);
t.xa.rollback_prepared(&t.pool, &x).await.unwrap();
sqlx::query(&t.be.q("UPDATE xa_acct SET bal = 999 WHERE id = 61"))
.execute(&t.pool)
.await
.unwrap_or_else(|e| panic!("{}: 解决之后应该能正常写: {e}", t.name));
assert_eq!(bal(&t, &[61]).await, vec![999], "{}", t.name);
}
}
fn require_real_db(缺的变量: &str) {
if std::env::var("DTMRS_TEST_REQUIRE_REAL_DB").is_ok() {
panic!(
"设了 DTMRS_TEST_REQUIRE_REAL_DB,却没有 {缺的变量} —— \
这是 CI 配置坏了(容器没起来?变量名打错?),不是可以跳过的情况"
);
}
}