Skip to main content

dtmrs_store/
lib.rs

1//! 存储层。TC 本身无状态,所有状态都在这里 —— 所以 TC 可以多实例、可以随时重启。
2//!
3//! # 一套 SQL 同时跑 sqlite / postgres / mysql
4//!
5//! 用 `sqlx::Any` + [`dtmrs_core::dialect`] 的模板渲染,而不是抽 `Store` trait
6//! 写三份实现。方言差异(占位符、冲突忽略、列类型、索引写法)全在 dialect 那层,
7//! 各家实测出来的坑也记在那个文件头,这里只遵守它的两条写法约定:
8//!
9//! 1. **模板里统一写 `?`**,由 [`Backend::q`] 渲染成各后端能吃的语句
10//!    (非 MySQL 转成 `$1..$n`,MySQL 原样保留)
11//! 2. **模板的字符串字面量里不能出现 `?`** —— 会被当成占位符
12//!
13//! 顺带一条只有 sqlite 有的老坑:它把 `$4` 当命名参数,所以同一个 `$N` 不能
14//! 复用。`q()` 逐个 `?` 顺序编号,天然不会复用。
15//!
16//! 时间统一用 **unix 秒(i64)** 存,不用数据库的 datetime 类型 ——
17//! 跨库的时间类型映射是反复踩坑的地方,整数没有这个问题。
18//! 列类型用 `BIGINT`:postgres 的 `INTEGER` 只有 4 字节,装不下时间戳。
19
20pub use dtmrs_core::Backend;
21
22use dtmrs_core::dialect::check_len;
23use dtmrs_core::{BranchOp, BranchStatus, GlobalStatus, TransType};
24use sqlx::any::{AnyPoolOptions, AnyRow};
25use sqlx::{AnyPool, Row};
26use std::sync::Once;
27
28pub type Result<T> = std::result::Result<T, sqlx::Error>;
29
30/// payload 列的字符上限(`trans_global.payload`)
31pub const BIG: usize = 8192;
32/// url / reason 一类中等长度列的字符上限
33pub const MID: usize = 1024;
34
35/// 把超长字段变成错误。
36///
37/// **不能省**:MySQL 的 `INSERT IGNORE` 遇到超长值会静默截断而不是报错,
38/// 详见 [`dtmrs_core::dialect::check_len`]。宁可提交时报错,也不能让一笔
39/// 内容被悄悄改过的事务落库。
40fn len_ok(col: &'static str, val: &str, max: usize) -> Result<()> {
41    check_len(col, val, max).map_err(|e| sqlx::Error::Encode(Box::new(e)))
42}
43
44pub fn now() -> i64 {
45    std::time::SystemTime::now()
46        .duration_since(std::time::UNIX_EPOCH)
47        .map(|d| d.as_secs() as i64)
48        .unwrap_or(0)
49}
50
51#[derive(Debug, Clone)]
52pub struct GlobalRow {
53    pub gid: String,
54    pub trans_type: TransType,
55    pub status: GlobalStatus,
56    pub payload: String,
57    pub next_cron_time: i64,
58    pub next_cron_interval: i64,
59    pub owner: String,
60    pub rollback_reason: String,
61    /// 二阶段消息的回查地址。进程在 prepare 和 submit 之间崩了,
62    /// TC 靠它问业务方"这单本地事务到底提交了没有"
63    pub query_prepared: String,
64    pub create_time: i64,
65    pub finish_time: Option<i64>,
66}
67
68#[derive(Debug, Clone)]
69pub struct BranchRow {
70    pub gid: String,
71    pub branch_id: String,
72    pub op: BranchOp,
73    pub url: String,
74    pub payload: String,
75    pub status: BranchStatus,
76}
77
78#[derive(Clone)]
79pub struct SqlStore {
80    pool: AnyPool,
81    be: Backend,
82}
83
84static DRIVERS: Once = Once::new();
85
86impl SqlStore {
87    /// `url` 可以是:
88    /// - `sqlite:dtmrs.db` / `sqlite::memory:`
89    /// - `postgres://user:pass@host:5432/db`
90    pub async fn open(url: &str) -> Result<Self> {
91        DRIVERS.call_once(sqlx::any::install_default_drivers);
92
93        // sqlite 默认只读打开,不会建文件。AnyConnectOptions 没法像
94        // SqliteConnectOptions 那样设 create_if_missing,只能走 URL 参数。
95        let mut url = url.to_string();
96        if url.starts_with("sqlite") && !url.contains("mode=") && !url.contains(":memory:") {
97            url.push_str(if url.contains('?') {
98                "&mode=rwc"
99            } else {
100                "?mode=rwc"
101            });
102        }
103        // 内存库必须单连接,否则每条连接看到的是各自独立的库
104        let max = if url.contains(":memory:") { 1 } else { 8 };
105        let be = Backend::from_url(&url);
106        let pool = AnyPoolOptions::new()
107            .max_connections(max)
108            .connect(&url)
109            .await?;
110        let s = Self { pool, be };
111        s.migrate_racy().await?;
112        Ok(s)
113    }
114
115    /// 建表,容忍并发。
116    ///
117    /// **Postgres 的 `CREATE TABLE IF NOT EXISTS` 不是并发安全的** ——
118    /// 两个 TC 实例同时启动会在系统目录上撞唯一键:
119    /// `duplicate key value violates unique constraint "pg_type_typname_nsp_index"`。
120    /// 这是实测撞出来的(sqlite 单写不会暴露)。
121    ///
122    /// 输了的那个重试一次就好:这时表已经被对方建出来了,
123    /// `IF NOT EXISTS` 会正常跳过。
124    async fn migrate_racy(&self) -> Result<()> {
125        let mut last = None;
126        for attempt in 0..3 {
127            match self.migrate().await {
128                Ok(()) => return Ok(()),
129                Err(e) => {
130                    last = Some(e);
131                    // 让对方把 DDL 事务提交完
132                    tokio::time::sleep(std::time::Duration::from_millis(100 * (attempt + 1))).await;
133                }
134            }
135        }
136        Err(last.expect("循环至少失败一次"))
137    }
138
139    pub async fn migrate(&self) -> Result<()> {
140        let idt = self.be.id_text();
141        let ids = self.be.id_short();
142        // payload 要装下所有步骤的 URL;MySQL 上是 VARCHAR,有长度上限。
143        // 上限同时是写库前的校验依据(BIG/MID),改这里就得改那里 —— 所以是常量
144        let big = self.be.text(BIG);
145        let mid = self.be.text(MID);
146        // 索引二选一:MySQL 只能建表时内联,其它后端用独立的
147        // CREATE INDEX IF NOT EXISTS(MySQL 那个语法直接 1064)
148        let inline = self
149            .be
150            .inline_index("idx_status_cron", "status, next_cron_time");
151
152        sqlx::query(&format!(
153            "CREATE TABLE IF NOT EXISTS trans_global (
154              gid                {idt} NOT NULL,
155              trans_type         {ids} NOT NULL,
156              status             {ids} NOT NULL,
157              payload            {big} NOT NULL,
158              next_cron_time     BIGINT NOT NULL DEFAULT 0,
159              next_cron_interval BIGINT NOT NULL DEFAULT 0,
160              owner              {idt} NOT NULL,
161              rollback_reason    {mid} NOT NULL,
162              query_prepared     {mid} NOT NULL,
163              create_time        BIGINT NOT NULL,
164              update_time        BIGINT NOT NULL,
165              finish_time        BIGINT,
166              PRIMARY KEY (gid){inline}
167            )"
168        ))
169        .execute(&self.pool)
170        .await?;
171        // cron 靠这个索引扫待办,没它到量之后会全表扫
172        if let Some(sql) =
173            self.be
174                .create_index("idx_status_cron", "trans_global", "status, next_cron_time")
175        {
176            sqlx::query(&sql).execute(&self.pool).await?;
177        }
178        sqlx::query(&format!(
179            "CREATE TABLE IF NOT EXISTS trans_branch_op (
180              gid         {idt} NOT NULL,
181              branch_id   {idt} NOT NULL,
182              op          {ids} NOT NULL,
183              url         {mid} NOT NULL,
184              payload     {mid} NOT NULL,
185              status      {ids} NOT NULL,
186              create_time BIGINT NOT NULL,
187              update_time BIGINT NOT NULL,
188              finish_time BIGINT,
189              PRIMARY KEY (gid, branch_id, op)
190            )"
191        ))
192        .execute(&self.pool)
193        .await?;
194        Ok(())
195    }
196
197    pub fn backend(&self) -> Backend {
198        self.be
199    }
200
201    pub fn pool(&self) -> &AnyPool {
202        &self.pool
203    }
204
205    /// 建全局事务 + 所有分支,一个事务里做完。
206    ///
207    /// 返回 `false` 表示 gid 已存在 —— 这是**幂等提交**,不是错误:
208    /// 客户端重试提交时必须拿到"已受理"而不是报错。
209    pub async fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> Result<bool> {
210        // 先校验再落库:超长的值在 MySQL 上会被 INSERT IGNORE 静默截断
211        len_ok("gid", &g.gid, Backend::ID_MAX)?;
212        len_ok("payload", &g.payload, BIG)?;
213        len_ok("query_prepared", &g.query_prepared, MID)?;
214        for b in branches {
215            len_ok("branch_id", &b.branch_id, Backend::ID_MAX)?;
216            len_ok("url", &b.url, MID)?;
217            len_ok("payload", &b.payload, MID)?;
218        }
219        let mut tx = self.pool.begin().await?;
220        let t = now();
221        let n = sqlx::query(&self.be.q("{INS} trans_global
222             (gid,trans_type,status,payload,next_cron_time,next_cron_interval,
223              owner,rollback_reason,query_prepared,create_time,update_time)
224             VALUES (?,?,?,?,?,?,'','',?,?,?)
225             {NOCONFLICT}"))
226        .bind(&g.gid)
227        .bind(g.trans_type.to_string())
228        .bind(g.status.as_str())
229        .bind(&g.payload)
230        .bind(g.next_cron_time)
231        .bind(g.next_cron_interval)
232        .bind(&g.query_prepared)
233        .bind(t)
234        .bind(t)
235        .execute(&mut *tx)
236        .await?
237        .rows_affected();
238        if n == 0 {
239            tx.rollback().await?;
240            return Ok(false);
241        }
242        for b in branches {
243            sqlx::query(&self.be.q("{INS} trans_branch_op
244                 (gid,branch_id,op,url,payload,status,create_time,update_time)
245                 VALUES (?,?,?,?,?,?,?,?)
246                 {NOCONFLICT}"))
247            .bind(&b.gid)
248            .bind(&b.branch_id)
249            .bind(b.op.as_str())
250            .bind(&b.url)
251            .bind(&b.payload)
252            .bind(b.status.as_str())
253            .bind(t)
254            .bind(t)
255            .execute(&mut *tx)
256            .await?;
257        }
258        tx.commit().await?;
259        Ok(true)
260    }
261
262    pub async fn get_global(&self, gid: &str) -> Result<Option<GlobalRow>> {
263        let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
264            .bind(gid)
265            .fetch_optional(&self.pool)
266            .await?;
267        Ok(row.map(global_from_row))
268    }
269
270    pub async fn list_branches(&self, gid: &str) -> Result<Vec<BranchRow>> {
271        let rows = sqlx::query(&self.be.q(
272            "SELECT gid,branch_id,op,url,payload,status FROM trans_branch_op
273             WHERE gid=? ORDER BY branch_id, op",
274        ))
275        .bind(gid)
276        .fetch_all(&self.pool)
277        .await?;
278        Ok(rows
279            .into_iter()
280            .map(|r| BranchRow {
281                gid: r.get("gid"),
282                branch_id: r.get("branch_id"),
283                op: BranchOp::parse(r.get::<String, _>("op").as_str()).unwrap_or(BranchOp::Action),
284                url: r.get("url"),
285                payload: r.get("payload"),
286                status: BranchStatus::parse(r.get::<String, _>("status").as_str())
287                    .unwrap_or(BranchStatus::Prepared),
288            })
289            .collect())
290    }
291
292    pub async fn set_global_status(
293        &self,
294        gid: &str,
295        status: GlobalStatus,
296        reason: &str,
297    ) -> Result<()> {
298        let t = now();
299        let fin = if status.is_final() { Some(t) } else { None };
300        // reason 是诊断信息,**截断而不是报错**:这条 UPDATE 是状态机的收尾,
301        // 让它因为一句话太长而失败,事务就永远推不到终态了(MySQL strict mode
302        // 下超长 UPDATE 直接报 1406,不像 INSERT IGNORE 那样只是截断)。
303        let reason: String = reason.chars().take(MID).collect();
304        let reason = reason.as_str();
305        // 注意 $4/$5 都绑 reason —— 不能复用同一个 $N,见文件头注释
306        sqlx::query(&self.be.q(
307            "UPDATE trans_global SET status=?, update_time=?, finish_time=?,
308             rollback_reason = CASE WHEN ? <> '' THEN ? ELSE rollback_reason END
309             WHERE gid=?",
310        ))
311        .bind(status.as_str())
312        .bind(t)
313        .bind(fin)
314        .bind(reason)
315        .bind(reason)
316        .bind(gid)
317        .execute(&self.pool)
318        .await?;
319        Ok(())
320    }
321
322    /// 落一个分支的状态**和结果数据**。
323    ///
324    /// workflow 模式的重放靠这个:函数崩溃后会从头再跑一遍,已完成的分支
325    /// 不重新执行,而是把上次存的 `payload` 原样还给它。所以这个值必须跟
326    /// 「分支已成功」在**同一条 UPDATE 里**落盘 —— 分两步写的话,中间崩了
327    /// 就会出现「标了成功但结果丢了」,重放时拿不到返回值。
328    pub async fn set_branch_result(
329        &self,
330        gid: &str,
331        branch_id: &str,
332        op: BranchOp,
333        status: BranchStatus,
334        payload: &str,
335    ) -> Result<()> {
336        len_ok("payload", payload, MID)?;
337        let t = now();
338        sqlx::query(&self.be.q(
339            "UPDATE trans_branch_op SET status=?, payload=?, update_time=?,
340             finish_time = CASE WHEN ? <> 'prepared' THEN ? ELSE finish_time END
341             WHERE gid=? AND branch_id=? AND op=?",
342        ))
343        .bind(status.as_str())
344        .bind(payload)
345        .bind(t)
346        .bind(status.as_str())
347        .bind(t)
348        .bind(gid)
349        .bind(branch_id)
350        .bind(op.as_str())
351        .execute(&self.pool)
352        .await?;
353        Ok(())
354    }
355
356    pub async fn set_branch_status(
357        &self,
358        gid: &str,
359        branch_id: &str,
360        op: BranchOp,
361        status: BranchStatus,
362    ) -> Result<()> {
363        let t = now();
364        sqlx::query(
365            &self
366                .be
367                .q("UPDATE trans_branch_op SET status=?, update_time=?,
368             finish_time = CASE WHEN ? <> 'prepared' THEN ? ELSE finish_time END
369             WHERE gid=? AND branch_id=? AND op=?"),
370        )
371        .bind(status.as_str())
372        .bind(t)
373        .bind(status.as_str())
374        .bind(t)
375        .bind(gid)
376        .bind(branch_id)
377        .bind(op.as_str())
378        .execute(&self.pool)
379        .await?;
380        Ok(())
381    }
382
383    /// 抢一个到期的待办事务,**抢占式更新,原子的**。
384    ///
385    /// 多个 TC 实例同时跑也不会重复推进同一个事务:谁的 UPDATE 生效谁持有租约。
386    /// 持租约的实例崩了,`next_cron_time` 到期后别的实例接手 —— 这就是崩溃恢复。
387    pub async fn lock_one_due(&self, owner: &str, lease: i64) -> Result<Option<GlobalRow>> {
388        let mut tx = self.pool.begin().await?;
389        let t = now();
390        let gid: Option<String> = sqlx::query_scalar(&self.be.q("SELECT gid FROM trans_global
391             WHERE (status IN ('submitted','aborting')
392                    OR (status = 'prepared' AND trans_type = 'msg'))
393               AND next_cron_time <= ?
394             ORDER BY next_cron_time LIMIT 1"))
395        .bind(t)
396        .fetch_optional(&mut *tx)
397        .await?;
398        let Some(gid) = gid else {
399            tx.rollback().await?;
400            return Ok(None);
401        };
402        // 立刻把 next_cron_time 推到租约之后,等于占坑
403        let n = sqlx::query(&self.be.q(
404            "UPDATE trans_global SET owner=?, next_cron_time=?, update_time=?
405             WHERE gid=? AND next_cron_time <= ?",
406        ))
407        .bind(owner)
408        .bind(t + lease)
409        .bind(t)
410        .bind(&gid)
411        .bind(t)
412        .execute(&mut *tx)
413        .await?
414        .rows_affected();
415        if n == 0 {
416            tx.rollback().await?;
417            return Ok(None); // 被别人抢走了
418        }
419        let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
420            .bind(&gid)
421            .fetch_one(&mut *tx)
422            .await?;
423        tx.commit().await?;
424        Ok(Some(global_from_row(row)))
425    }
426
427    /// 推进失败后设置下次重试时间(指数退避)
428    pub async fn schedule_retry(&self, gid: &str, interval: i64) -> Result<()> {
429        let t = now();
430        sqlx::query(&self.be.q(
431            "UPDATE trans_global SET next_cron_interval=?, next_cron_time=?, update_time=?
432             WHERE gid=?",
433        ))
434        .bind(interval)
435        .bind(t + interval)
436        .bind(t)
437        .bind(gid)
438        .execute(&self.pool)
439        .await?;
440        Ok(())
441    }
442
443    /// 让某个事务立刻可被调度(提交/中止之后叫一下,不用等 cron 周期)
444    pub async fn schedule_now(&self, gid: &str) -> Result<()> {
445        sqlx::query(
446            &self
447                .be
448                .q("UPDATE trans_global SET next_cron_time=?, next_cron_interval=0 WHERE gid=?"),
449        )
450        .bind(now())
451        .bind(gid)
452        .execute(&self.pool)
453        .await?;
454        Ok(())
455    }
456
457    /// TCC 的 try 阶段:客户端在调 try 之前先来登记这个分支的 confirm/cancel。
458    ///
459    /// **必须先登记再调 try**。反过来的话:try 成功了但登记失败,
460    /// TC 就不知道有这个分支,回滚时不会 cancel 它 —— 资源永久泄漏。
461    ///
462    /// 冲突时忽略,所以重复登记是幂等的(客户端重试很常见)。
463    pub async fn register_branch(
464        &self,
465        gid: &str,
466        branch_id: &str,
467        ops: &[(BranchOp, String)],
468    ) -> Result<()> {
469        len_ok("gid", gid, Backend::ID_MAX)?;
470        len_ok("branch_id", branch_id, Backend::ID_MAX)?;
471        for (_, url) in ops {
472            len_ok("url", url, MID)?;
473        }
474        let mut tx = self.pool.begin().await?;
475        let t = now();
476        for (op, url) in ops {
477            sqlx::query(&self.be.q("{INS} trans_branch_op
478                 (gid,branch_id,op,url,payload,status,create_time,update_time)
479                 VALUES (?,?,?,?,'',?,?,?)
480                 {NOCONFLICT}"))
481            .bind(gid)
482            .bind(branch_id)
483            .bind(op.as_str())
484            .bind(url)
485            .bind(BranchStatus::Prepared.as_str())
486            .bind(t)
487            .bind(t)
488            .execute(&mut *tx)
489            .await?;
490        }
491        tx.commit().await?;
492        Ok(())
493    }
494
495    pub async fn list_recent(&self, limit: i64) -> Result<Vec<GlobalRow>> {
496        let rows = sqlx::query(&self.be.q(&format!(
497            "{SELECT_GLOBAL} ORDER BY create_time DESC LIMIT ?"
498        )))
499        .bind(limit)
500        .fetch_all(&self.pool)
501        .await?;
502        Ok(rows.into_iter().map(global_from_row).collect())
503    }
504}
505
506/// 列清单只写一处 —— 三个地方读 trans_global,列顺序漂移过一次就够难查了
507const SELECT_GLOBAL: &str = "SELECT gid,trans_type,status,payload,next_cron_time,
508    next_cron_interval,owner,rollback_reason,query_prepared,create_time,finish_time
509    FROM trans_global";
510
511fn global_from_row(r: AnyRow) -> GlobalRow {
512    GlobalRow {
513        gid: r.get("gid"),
514        trans_type: TransType::parse(r.get::<String, _>("trans_type").as_str())
515            .unwrap_or(TransType::Saga),
516        status: GlobalStatus::parse(r.get::<String, _>("status").as_str())
517            .unwrap_or(GlobalStatus::Prepared),
518        payload: r.get("payload"),
519        next_cron_time: r.get("next_cron_time"),
520        next_cron_interval: r.get("next_cron_interval"),
521        owner: r.get("owner"),
522        rollback_reason: r.get("rollback_reason"),
523        query_prepared: r.get("query_prepared"),
524        create_time: r.get("create_time"),
525        finish_time: r.get("finish_time"),
526    }
527}
528
529#[cfg(test)]
530mod tests {
531    use super::*;
532
533    /// 每个测试都在**所有可用后端**上跑一遍:sqlite / postgres / mysql / redis。
534    ///
535    /// 真库靠环境变量开启(`DTMRS_TEST_PG` / `DTMRS_TEST_MYSQL` / `DTMRS_TEST_REDIS`)——
536    /// 没配就只跑 sqlite,这样没数据库的机器也能 `cargo test`。
537    /// 但**别把这当成"它们也过了"** —— 没配就是没测。
538    ///
539    /// Redis 跟另外三个不是同一类东西(不是 SQL),能共用这一套断言恰恰是
540    /// 我们要的证据:两种后端的**行为**必须一致,哪怕实现天差地别。
541    /// Postgres 测试必须串行 —— `lock_one_due` 和 `list_recent` 是**全局查询**,
542    /// 并行跑会互相看见对方的事务,断言就没意义了。
543    /// (光给各测试不同的 gid 不够:捞待办是不按 gid 过滤的。)
544    static PG_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
545
546    /// 返回 (串行锁, 各后端)。锁要持到测试结束,所以由调用方接着。
547    ///
548    /// 每次进来把 Postgres 的表清空(只 DELETE 不 DDL —— 并发 DDL 会撞上
549    /// Postgres 的 `pg_type` 竞态,见 `migrate_racy`)。
550    async fn backends() -> (
551        tokio::sync::MutexGuard<'static, ()>,
552        Vec<(&'static str, Store)>,
553    ) {
554        let guard = PG_LOCK.lock().await;
555        let mut v = vec![("sqlite", Store::open("sqlite::memory:").await.unwrap())];
556        // 每种真数据库都配一个环境变量。**没配就是没测**,不是"通过"。
557        for (name, env) in [("postgres", "DTMRS_TEST_PG"), ("mysql", "DTMRS_TEST_MYSQL")] {
558            if let Ok(url) = std::env::var(env) {
559                let s = Store::open(&url)
560                    .await
561                    .unwrap_or_else(|e| panic!("连不上 {env}: {e}"));
562                for t in ["trans_branch_op", "trans_global"] {
563                    sqlx::query(&format!("DELETE FROM {t}"))
564                        .execute(s.pool().expect("SQL 后端才有连接池"))
565                        .await
566                        .expect("清表");
567                }
568                v.push((name, s));
569            }
570        }
571        #[cfg(feature = "redis")]
572        if let Ok(url) = std::env::var("DTMRS_TEST_REDIS") {
573            let s = Store::open(&url)
574                .await
575                .unwrap_or_else(|e| panic!("连不上 DTMRS_TEST_REDIS: {e}"));
576            // 每次进来清干净。Redis 没有"表",按前缀删
577            s.as_redis()
578                .unwrap()
579                .flush_prefix()
580                .await
581                .expect("清 redis");
582            v.push(("redis", s));
583        }
584        (guard, v)
585    }
586
587    fn g(gid: &str) -> GlobalRow {
588        GlobalRow {
589            gid: gid.into(),
590            trans_type: TransType::Saga,
591            status: GlobalStatus::Submitted,
592            payload: "{}".into(),
593            next_cron_time: 0,
594            next_cron_interval: 0,
595            owner: String::new(),
596            rollback_reason: String::new(),
597            query_prepared: String::new(),
598            create_time: 0,
599            finish_time: None,
600        }
601    }
602
603    #[tokio::test]
604    async fn 重复提交同一个gid是幂等的() {
605        let (_g, bes) = backends().await;
606        for (name, s) in bes {
607            assert!(s.create_global(&g("t1"), &[]).await.unwrap(), "{name}");
608            // 第二次返回 false 而不是报错 —— 客户端重试不该失败
609            assert!(!s.create_global(&g("t1"), &[]).await.unwrap(), "{name}");
610            assert_eq!(s.list_recent(10).await.unwrap().len(), 1, "{name}");
611        }
612    }
613
614    #[tokio::test]
615    async fn 租约只能被抢到一次() {
616        let (_g, bes) = backends().await;
617        for (name, s) in bes {
618            s.create_global(&g("t2"), &[]).await.unwrap();
619            let a = s.lock_one_due("worker-a", 60).await.unwrap();
620            assert!(a.is_some(), "{name}: 第一个实例应该抢到");
621            // 同一个事务不能被第二个实例同时抢到,否则会重复推进
622            let b = s.lock_one_due("worker-b", 60).await.unwrap();
623            assert!(b.is_none(), "{name}: 租约期内不能被别人抢走");
624        }
625    }
626
627    #[tokio::test]
628    async fn 终态不再被调度() {
629        let (_g, bes) = backends().await;
630        for (name, s) in bes {
631            s.create_global(&g("t3"), &[]).await.unwrap();
632            s.set_global_status("t3", GlobalStatus::Succeed, "")
633                .await
634                .unwrap();
635            assert!(s.lock_one_due("w", 60).await.unwrap().is_none(), "{name}");
636            let got = s.get_global("t3").await.unwrap().unwrap();
637            assert_eq!(got.status, GlobalStatus::Succeed, "{name}");
638            assert!(got.finish_time.is_some(), "{name}: 终态要落 finish_time");
639        }
640    }
641
642    #[tokio::test]
643    async fn 分支状态可更新() {
644        let (_g, bes) = backends().await;
645        for (name, s) in bes {
646            let b = BranchRow {
647                gid: "t4".into(),
648                branch_id: "01".into(),
649                op: BranchOp::Action,
650                url: "http://x/a".into(),
651                payload: "{}".into(),
652                status: BranchStatus::Prepared,
653            };
654            s.create_global(&g("t4"), std::slice::from_ref(&b))
655                .await
656                .unwrap();
657            s.set_branch_status("t4", "01", BranchOp::Action, BranchStatus::Succeed)
658                .await
659                .unwrap();
660            let got = s.list_branches("t4").await.unwrap();
661            assert_eq!(got.len(), 1, "{name}");
662            assert_eq!(got[0].status, BranchStatus::Succeed, "{name}");
663        }
664    }
665
666    #[tokio::test]
667    async fn 回滚原因和回查地址能存取() {
668        // 这两列是后加的,跨库的字符串/空值处理最容易在这儿出问题
669        let (_g, bes) = backends().await;
670        for (name, s) in bes {
671            let mut row = g("t5");
672            row.query_prepared = "http://busi/query".into();
673            s.create_global(&row, &[]).await.unwrap();
674            s.set_global_status("t5", GlobalStatus::Aborting, "分支 02 返回 FAILURE")
675                .await
676                .unwrap();
677            let got = s.get_global("t5").await.unwrap().unwrap();
678            assert_eq!(got.query_prepared, "http://busi/query", "{name}");
679            assert_eq!(got.rollback_reason, "分支 02 返回 FAILURE", "{name}");
680            assert!(
681                got.finish_time.is_none(),
682                "{name}: 非终态不该有 finish_time"
683            );
684
685            // 空 reason 不能把已有的原因冲掉
686            s.set_global_status("t5", GlobalStatus::Failed, "")
687                .await
688                .unwrap();
689            let got = s.get_global("t5").await.unwrap().unwrap();
690            assert_eq!(
691                got.rollback_reason, "分支 02 返回 FAILURE",
692                "{name}: 空原因不能覆盖"
693            );
694        }
695    }
696
697    #[tokio::test]
698    async fn msg的prepared会被捞tcc的不会() {
699        let (_g, bes) = backends().await;
700        for (name, s) in bes {
701            let mut m = g("m1");
702            m.trans_type = TransType::Msg;
703            m.status = GlobalStatus::Prepared;
704            s.create_global(&m, &[]).await.unwrap();
705            let mut t = g("c1");
706            t.trans_type = TransType::Tcc;
707            t.status = GlobalStatus::Prepared;
708            s.create_global(&t, &[]).await.unwrap();
709
710            let got = s.lock_one_due("w", 60).await.unwrap();
711            assert_eq!(
712                got.map(|x| x.gid),
713                Some("m1".to_string()),
714                "{name}: 只该捞到 msg"
715            );
716            // 再捞一次应该没有了(msg 被租约占住,tcc 不该被碰)
717            assert!(s.lock_one_due("w2", 60).await.unwrap().is_none(), "{name}");
718        }
719    }
720
721    #[tokio::test]
722    async fn 分支登记是幂等的() {
723        let (_g, bes) = backends().await;
724        for (name, s) in bes {
725            let mut t = g("c2");
726            t.trans_type = TransType::Tcc;
727            s.create_global(&t, &[]).await.unwrap();
728            let ops = [
729                (BranchOp::Confirm, "http://x/c".to_string()),
730                (BranchOp::Cancel, "http://x/n".to_string()),
731            ];
732            s.register_branch("c2", "01", &ops).await.unwrap();
733            s.register_branch("c2", "01", &ops).await.unwrap(); // 客户端重试
734            assert_eq!(
735                s.list_branches("c2").await.unwrap().len(),
736                2,
737                "{name}: 不该重复插入"
738            );
739        }
740    }
741}
742
743// ==================== 后端分发 ====================
744
745#[cfg(feature = "redis")]
746pub mod redis_store;
747#[cfg(feature = "redis")]
748pub use redis_store::RedisStore;
749
750/// 存储后端。
751///
752/// # 为什么现在才抽这一层
753///
754/// 这个项目原本**刻意没有抽 `Store` trait**,理由写在 DESIGN.md 里:
755/// sqlite / postgres / mysql 的差异小到一层 SQL 模板就能吸收,抽象是过早的。
756/// 那个判断在当时是对的。
757///
758/// **Redis 让前提不成立了** —— 它根本不是 SQL,没有表、没有事务、没有 WHERE,
759/// 模板吸收不了。所以这里加了一层分发。
760///
761/// 用 enum 而不是 trait:调用方拿到的还是同一个 `Store` 具体类型,
762/// 四十多个调用点一行都不用改,也不用到处写泛型或 `dyn`。
763#[derive(Clone)]
764enum Inner {
765    Sql(SqlStore),
766    #[cfg(feature = "redis")]
767    Redis(RedisStore),
768}
769
770/// 存储层的统一入口。按 URL 前缀自动选后端:
771///
772/// ```text
773/// sqlite:...     / postgres://...  / mysql://...   → SQL 后端
774/// redis://...    / rediss://...                    → Redis 后端(要开 redis feature)
775/// ```
776///
777/// ⚠ Redis 后端跟 SQL 后端有**实打实的语义差异**(持久性更弱、终态会过期),
778/// 用之前务必读 [`redis_store`] 的模块说明。
779#[derive(Clone)]
780pub struct Store {
781    inner: Inner,
782}
783
784/// 存储层的错误。
785///
786/// 两种后端的原生错误类型不同,统一收口到这里;`sqlx::Error` 仍然直接透出,
787/// 免得改动现有调用方对错误的处理。
788pub type StoreError = sqlx::Error;
789
790#[cfg(feature = "redis")]
791fn redis_err(e: redis::RedisError) -> sqlx::Error {
792    sqlx::Error::Configuration(Box::new(e))
793}
794
795/// 这个 URL 是不是要走 Redis
796pub fn is_redis_url(url: &str) -> bool {
797    let u = url.trim().to_ascii_lowercase();
798    u.starts_with("redis://") || u.starts_with("rediss://") || u.starts_with("redis+unix:")
799}
800
801impl Store {
802    /// 按 URL 选后端并连上。
803    pub async fn open(url: &str) -> Result<Self> {
804        if is_redis_url(url) {
805            #[cfg(feature = "redis")]
806            {
807                let r = RedisStore::open(url).await.map_err(redis_err)?;
808                return Ok(Self {
809                    inner: Inner::Redis(r),
810                });
811            }
812            #[cfg(not(feature = "redis"))]
813            {
814                // 明确报错,而不是把 redis:// 当成 sqlite 文件名去建库 ——
815                // 那会静默跑起来然后数据全落在一个叫 "redis:" 的文件里
816                return Err(sqlx::Error::Configuration(
817                    "这个 URL 要 Redis 后端,但构建时没开 dtmrs-store 的 `redis` feature".into(),
818                ));
819            }
820        }
821        Ok(Self {
822            inner: Inner::Sql(SqlStore::open(url).await?),
823        })
824    }
825
826    /// 底层是不是 Redis
827    pub fn is_redis(&self) -> bool {
828        match &self.inner {
829            Inner::Sql(_) => false,
830            #[cfg(feature = "redis")]
831            Inner::Redis(_) => true,
832        }
833    }
834
835    /// SQL 后端的连接池。Redis 后端返回 `None` ——
836    /// 调用方(主要是测试和屏障)要自己处理这种情况
837    pub fn pool(&self) -> Option<&AnyPool> {
838        match &self.inner {
839            Inner::Sql(s) => Some(s.pool()),
840            #[cfg(feature = "redis")]
841            Inner::Redis(_) => None,
842        }
843    }
844
845    /// SQL 方言。Redis 后端没有方言可言,返回 `None`
846    pub fn backend(&self) -> Option<Backend> {
847        match &self.inner {
848            Inner::Sql(s) => Some(s.backend()),
849            #[cfg(feature = "redis")]
850            Inner::Redis(_) => None,
851        }
852    }
853
854    /// 拿底层的 Redis store(比如为了调 `with_ttl`)
855    #[cfg(feature = "redis")]
856    pub fn as_redis(&self) -> Option<&RedisStore> {
857        match &self.inner {
858            Inner::Redis(r) => Some(r),
859            _ => None,
860        }
861    }
862}
863
864/// 把 13 个方法逐个手写分发太啰嗦,而且漏一个编译器不会提醒 ——
865/// 用宏保证两边签名严格一致
866macro_rules! dispatch {
867    ($( $(#[$m:meta])* fn $name:ident (&self $(, $arg:ident : $ty:ty)* ) -> $ret:ty; )*) => {
868        impl Store {
869            $(
870                $(#[$m])*
871                pub async fn $name(&self $(, $arg: $ty)*) -> Result<$ret> {
872                    match &self.inner {
873                        Inner::Sql(s) => s.$name($($arg),*).await,
874                        #[cfg(feature = "redis")]
875                        Inner::Redis(r) => r.$name($($arg),*).await.map_err(redis_err),
876                    }
877                }
878            )*
879        }
880    };
881}
882
883dispatch! {
884    /// 建表(Redis 后端是空操作)
885    fn migrate(&self) -> ();
886    /// 建全局事务 + 分支。已存在返回 `false`,**不覆盖**
887    fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> bool;
888    fn get_global(&self, gid: &str) -> Option<GlobalRow>;
889    fn list_branches(&self, gid: &str) -> Vec<BranchRow>;
890    /// 抢一个到期事务。多实例不重复推进就靠它的原子性
891    fn lock_one_due(&self, owner: &str, lease: i64) -> Option<GlobalRow>;
892    fn set_global_status(&self, gid: &str, status: GlobalStatus, reason: &str) -> ();
893    fn set_branch_result(&self, gid: &str, branch_id: &str, op: BranchOp, status: BranchStatus, payload: &str) -> ();
894    fn set_branch_status(&self, gid: &str, branch_id: &str, op: BranchOp, status: BranchStatus) -> ();
895    fn schedule_retry(&self, gid: &str, interval: i64) -> ();
896    fn schedule_now(&self, gid: &str) -> ();
897    fn register_branch(&self, gid: &str, branch_id: &str, ops: &[(BranchOp, String)]) -> ();
898    fn list_recent(&self, limit: i64) -> Vec<GlobalRow>;
899}