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/// [`Store::submit_prepared`] 的结果。
52///
53/// 之所以要区分三种情况:`submit` 既要处理「tcc/msg/xa 把 prepared 推成
54/// submitted」,也要处理「saga 第一次提交,事务还不存在」,还要保证
55/// **重复提交返回成功而不是报错**(见 `api::submit` 的注释)。
56#[derive(Debug, Clone)]
57pub enum SubmitOutcome {
58    /// gid 不存在 —— 调用方该按新事务建
59    Missing,
60    /// 本来停在 prepared,已经推成 submitted 并排进调度队列。
61    ///
62    /// **带上事务体**:调用方要是顺便占了租约(见 `submit_prepared` 的
63    /// `owner` 参数),可以拿它直接开推,不用再读一次。为这个多带的返回值,
64    /// 两种后端都没有多付往返 —— Redis 是脚本尾巴上加一个 `HGETALL`,
65    /// SQL 是把本来就要发的那条 SELECT 从「只取 status」改成取全行
66    Advanced(Box<GlobalRow>),
67    /// 已经提交过了。**必须当成功返回**,否则客户端会以为没受理
68    Already,
69}
70
71/// [`Store::register_branch`] 的结果。
72///
73/// 之所以不能只返回 `()`:登记走的是「冲突忽略」,而**冲突有两种,结论完全相反**——
74///
75/// * 客户端重试,URL 跟上次一模一样 → 必须当成功(`register_branch` 是幂等的)
76/// * 客户端把两个不同的分支都编成了同一个号 → 第二个分支的 URL **根本没存进去**
77///
78/// 后者原先也返回 SUCCESS,是个很难查的坑:客户端以为登记成功了,
79/// 接着去调那个分支的 try 把资源冻结上,而 TC 压根不知道有这个分支 ——
80/// confirm 和 cancel 都不会调,**那份资源永久泄漏**。
81/// 实测过:两次登记 branch_id="01",库里只留下第一个的 URL,两次都回 SUCCESS。
82#[derive(Debug, Clone, PartialEq, Eq)]
83pub enum RegisterOutcome {
84    /// 登记成功,或客户端重试且 URL 完全一致
85    Registered,
86    /// 这个 branch_id 已经被**另一组 URL** 占了。调用方必须拒绝,不能当成功
87    Conflict {
88        op: BranchOp,
89        /// 库里已经存着的 URL
90        existing: String,
91    },
92}
93
94#[derive(Debug, Clone)]
95pub struct GlobalRow {
96    pub gid: String,
97    pub trans_type: TransType,
98    pub status: GlobalStatus,
99    pub payload: String,
100    pub next_cron_time: i64,
101    pub next_cron_interval: i64,
102    pub owner: String,
103    pub rollback_reason: String,
104    /// 二阶段消息的回查地址。进程在 prepare 和 submit 之间崩了,
105    /// TC 靠它问业务方"这单本地事务到底提交了没有"
106    pub query_prepared: String,
107    pub create_time: i64,
108    pub finish_time: Option<i64>,
109}
110
111#[derive(Debug, Clone)]
112pub struct BranchRow {
113    pub gid: String,
114    pub branch_id: String,
115    pub op: BranchOp,
116    pub url: String,
117    pub payload: String,
118    pub status: BranchStatus,
119}
120
121
122/// 访问令牌的展示信息。**不含明文** —— 明文只在生成那一刻返回一次
123#[derive(Debug, Clone)]
124pub struct TokenRow {
125    /// SHA-256 十六进制,同时是主键
126    pub token_hash: String,
127    /// 人给的名字,用来认出「这个 token 是给谁的」
128    pub name: String,
129    pub create_time: i64,
130    /// 0 表示从没被用过
131    pub last_used: i64,
132    pub use_count: i64,
133    pub last_ip: String,
134    /// 0 表示有效,否则是作废时刻
135    pub revoked: i64,
136    /// 令牌明文的**密文**(`nonce||ciphertext` 的十六进制)。
137    /// 空串表示没保存 —— 没配 `DTMRS_TOKEN_KEY` 时就是这样,退化成「只显示一次」
138    pub secret: String,
139}
140
141/// 令牌的哈希。**存哈希不存明文**:库被看到也拿不到能用的凭据。
142///
143/// 这里刻意用朴素的 SHA-256 而不是 bcrypt/argon2 —— 那些是给**低熵的人类密码**
144/// 用的,慢是特性。这里的令牌是 24 字节随机数(192 位熵),爆破不可行,
145/// 而认证在热路径上,每请求跑一次 argon2 会直接毁掉吞吐。
146pub fn hash_token(raw: &str) -> String {
147    use sha2::{Digest, Sha256};
148    let mut h = Sha256::new();
149    h.update(raw.as_bytes());
150    h.finalize().iter().map(|b| format!("{b:02x}")).collect()
151}
152
153#[derive(Clone)]
154pub struct SqlStore {
155    pool: AnyPool,
156    be: Backend,
157}
158
159static DRIVERS: Once = Once::new();
160
161impl SqlStore {
162    /// `url` 可以是:
163    /// - `sqlite:dtmrs.db` / `sqlite::memory:`
164    /// - `postgres://user:pass@host:5432/db`
165    pub async fn open(url: &str) -> Result<Self> {
166        DRIVERS.call_once(sqlx::any::install_default_drivers);
167
168        // sqlite 默认只读打开,不会建文件。AnyConnectOptions 没法像
169        // SqliteConnectOptions 那样设 create_if_missing,只能走 URL 参数。
170        let mut url = url.to_string();
171        if url.starts_with("sqlite") && !url.contains("mode=") && !url.contains(":memory:") {
172            url.push_str(if url.contains('?') {
173                "&mode=rwc"
174            } else {
175                "?mode=rwc"
176            });
177        }
178        // 内存库必须单连接,否则每条连接看到的是各自独立的库
179        //
180        // 非内存库默认 32:这个池子是**推进器和 HTTP/gRPC 接口共用**的,
181        // 而推进一笔事务要好几次往返。原来写死 8,比推进 worker 数还少,
182        // 于是提交请求和推进器互相抢连接。
183        //
184        // 32 是「够用就行」:默认 16 个 worker + HTTP/gRPC 接口,实测再往上
185        // 加池子已经不涨了(Postgres 64 worker 配 32 的池子 3184 笔/秒,
186        // 池子开到 64 也是 3227)。
187        //
188        // 后端连接数吃紧(比如和业务共用一个 Postgres,默认才 100 条)
189        // 就用 `DTMRS_DB_POOL` 调小
190        let max = if url.contains(":memory:") {
191            1
192        } else {
193            std::env::var("DTMRS_DB_POOL")
194                .ok()
195                .and_then(|v| v.parse::<u32>().ok())
196                .filter(|v| *v > 0)
197                .unwrap_or(32)
198        };
199        let be = Backend::from_url(&url);
200        let is_file_sqlite = be == Backend::Sqlite && !url.contains(":memory:");
201        let pool = AnyPoolOptions::new()
202            .max_connections(max)
203            .after_connect(move |conn, _| {
204                Box::pin(async move {
205                    if is_file_sqlite {
206                        // ⚠ **别去掉这两条。**
207                        //
208                        // sqlite 默认是 rollback journal + synchronous=FULL,
209                        // 每笔事务一次 fsync,而且写事务会锁住整个库 ——
210                        // 实测提交吞吐只有约 13 笔/秒,并发 20 就大量
211                        // `database is locked` 并把请求拖到超时。
212                        //
213                        // WAL 让读写不互斥、synchronous=NORMAL 把每事务 fsync
214                        // 降成 checkpoint 时才 fsync。代价是**断电可能丢最后
215                        // 几笔已提交事务**(进程崩溃不丢,WAL 还在)——
216                        // sqlite 后端本来就只建议单机/开发用,这个取舍是划算的。
217                        // 要严格持久性就用 Postgres。
218                        for pragma in [
219                            "PRAGMA journal_mode=WAL",
220                            "PRAGMA synchronous=NORMAL",
221                            // 拿不到锁时先自旋 5 秒再报错,别让偶发争用直接失败
222                            "PRAGMA busy_timeout=5000",
223                        ] {
224                            sqlx::query(pragma).execute(&mut *conn).await?;
225                        }
226                    }
227                    Ok(())
228                })
229            })
230            .connect(&url)
231            .await?;
232        let s = Self { pool, be };
233        s.migrate_racy().await?;
234        Ok(s)
235    }
236
237    /// 建表,容忍并发。
238    ///
239    /// **Postgres 的 `CREATE TABLE IF NOT EXISTS` 不是并发安全的** ——
240    /// 两个 TC 实例同时启动会在系统目录上撞唯一键:
241    /// `duplicate key value violates unique constraint "pg_type_typname_nsp_index"`。
242    /// 这是实测撞出来的(sqlite 单写不会暴露)。
243    ///
244    /// 输了的那个重试一次就好:这时表已经被对方建出来了,
245    /// `IF NOT EXISTS` 会正常跳过。
246    async fn migrate_racy(&self) -> Result<()> {
247        let mut last = None;
248        for attempt in 0..3 {
249            match self.migrate().await {
250                Ok(()) => return Ok(()),
251                Err(e) => {
252                    last = Some(e);
253                    // 让对方把 DDL 事务提交完
254                    tokio::time::sleep(std::time::Duration::from_millis(100 * (attempt + 1))).await;
255                }
256            }
257        }
258        Err(last.expect("循环至少失败一次"))
259    }
260
261    pub async fn migrate(&self) -> Result<()> {
262        let idt = self.be.id_text();
263        let ids = self.be.id_short();
264        // payload 要装下所有步骤的 URL;MySQL 上是 VARCHAR,有长度上限。
265        // 上限同时是写库前的校验依据(BIG/MID),改这里就得改那里 —— 所以是常量
266        let big = self.be.text(BIG);
267        let mid = self.be.text(MID);
268        // 索引二选一:MySQL 只能建表时内联,其它后端用独立的
269        // CREATE INDEX IF NOT EXISTS(MySQL 那个语法直接 1064)
270        let inline = self
271            .be
272            .inline_index("idx_status_cron", "status, next_cron_time");
273
274        sqlx::query(&format!(
275            "CREATE TABLE IF NOT EXISTS trans_global (
276              gid                {idt} NOT NULL,
277              trans_type         {ids} NOT NULL,
278              status             {ids} NOT NULL,
279              payload            {big} NOT NULL,
280              next_cron_time     BIGINT NOT NULL DEFAULT 0,
281              next_cron_interval BIGINT NOT NULL DEFAULT 0,
282              owner              {idt} NOT NULL,
283              rollback_reason    {mid} NOT NULL,
284              query_prepared     {mid} NOT NULL,
285              create_time        BIGINT NOT NULL,
286              update_time        BIGINT NOT NULL,
287              finish_time        BIGINT,
288              PRIMARY KEY (gid){inline}
289            )"
290        ))
291        .execute(&self.pool)
292        .await?;
293        // cron 靠这个索引扫待办,没它到量之后会全表扫
294        if let Some(sql) =
295            self.be
296                .create_index("idx_status_cron", "trans_global", "status, next_cron_time")
297        {
298            sqlx::query(&sql).execute(&self.pool).await?;
299        }
300        sqlx::query(&format!(
301            "CREATE TABLE IF NOT EXISTS trans_branch_op (
302              gid         {idt} NOT NULL,
303              branch_id   {idt} NOT NULL,
304              op          {ids} NOT NULL,
305              url         {mid} NOT NULL,
306              payload     {mid} NOT NULL,
307              status      {ids} NOT NULL,
308              create_time BIGINT NOT NULL,
309              update_time BIGINT NOT NULL,
310              finish_time BIGINT,
311              PRIMARY KEY (gid, branch_id, op)
312            )"
313        ))
314        .execute(&self.pool)
315        .await?;
316        // 业务端的访问令牌。**存 SHA-256 不存明文** —— 库被看到也拿不到能用的凭据,
317        // 明文只在生成那一刻返回给用户一次
318        sqlx::query(&format!(
319            "CREATE TABLE IF NOT EXISTS auth_token (
320              token_hash  {idt} NOT NULL,
321              name        {ids} NOT NULL,
322              create_time BIGINT NOT NULL,
323              last_used   BIGINT NOT NULL DEFAULT 0,
324              use_count   BIGINT NOT NULL DEFAULT 0,
325              last_ip     {ids} NOT NULL DEFAULT '',
326              revoked     BIGINT NOT NULL DEFAULT 0,
327              secret      {mid} NOT NULL DEFAULT '',
328              PRIMARY KEY (token_hash)
329            )"
330        ))
331        .execute(&self.pool)
332        .await?;
333        self.add_missing_columns().await?;
334        Ok(())
335    }
336
337    /// 给**已存在**的表补新增的列。
338    ///
339    /// ⚠ 这个方法存在的理由:`CREATE TABLE IF NOT EXISTS` 对已有的表**什么都不做**,
340    /// 所以给表加一列时老库升级上来会直接报 `no such column`。
341    /// 在 auth_token 加 `secret` 列时踩到过 —— 服务起得来,但一调令牌接口就 500。
342    ///
343    /// 三种方言对「列已存在」的处理都不一样(sqlite/MySQL 的 ADD COLUMN 没有
344    /// IF NOT EXISTS),所以统一的做法是:**照发不误,把「列已存在」这个错吞掉**。
345    /// 别的错误照常抛出去。
346    ///
347    /// 加新列时在这个数组里加一行就行。
348    async fn add_missing_columns(&self) -> Result<()> {
349        let mid = self.be.text(MID);
350        let adds: [(&str, String); 1] = [(
351            "auth_token",
352            format!("secret {mid} NOT NULL DEFAULT ''"),
353        )];
354        for (table, coldef) in adds {
355            let sql = format!("ALTER TABLE {table} ADD COLUMN {coldef}");
356            if let Err(e) = sqlx::query(&sql).execute(&self.pool).await {
357                let m = e.to_string().to_lowercase();
358                // sqlite: "duplicate column name" / postgres: "already exists"
359                // / mysql: "duplicate column name"
360                let already = m.contains("duplicate column") || m.contains("already exists");
361                if !already {
362                    return Err(e);
363                }
364            }
365        }
366        Ok(())
367    }
368
369    pub fn backend(&self) -> Backend {
370        self.be
371    }
372
373    pub fn pool(&self) -> &AnyPool {
374        &self.pool
375    }
376
377    /// 建全局事务 + 所有分支,一个事务里做完。
378    ///
379    /// 返回 `false` 表示 gid 已存在 —— 这是**幂等提交**,不是错误:
380    /// 客户端重试提交时必须拿到"已受理"而不是报错。
381    pub async fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> Result<bool> {
382        // 先校验再落库:超长的值在 MySQL 上会被 INSERT IGNORE 静默截断
383        len_ok("gid", &g.gid, Backend::ID_MAX)?;
384        len_ok("payload", &g.payload, BIG)?;
385        len_ok("query_prepared", &g.query_prepared, MID)?;
386        for b in branches {
387            len_ok("branch_id", &b.branch_id, Backend::ID_MAX)?;
388            len_ok("url", &b.url, MID)?;
389            len_ok("payload", &b.payload, MID)?;
390        }
391        let mut tx = self.pool.begin().await?;
392        let t = now();
393        let n = sqlx::query(&self.be.q("{INS} trans_global
394             (gid,trans_type,status,payload,next_cron_time,next_cron_interval,
395              owner,rollback_reason,query_prepared,create_time,update_time)
396             VALUES (?,?,?,?,?,?,?,'',?,?,?)
397             {NOCONFLICT}"))
398        .bind(&g.gid)
399        .bind(g.trans_type.to_string())
400        .bind(g.status.as_str())
401        .bind(&g.payload)
402        .bind(g.next_cron_time)
403        .bind(g.next_cron_interval)
404        // ⚠ owner 要真的写进去,不能像原来那样写死空串。
405        // 提交方可以在建事务时就把租约占在自己手上(owner=自己、
406        // next_cron_time=现在+租约),这样它能直接开推,
407        // **省掉一次抢占往返** —— 见 `Api::submit`
408        .bind(&g.owner)
409        .bind(&g.query_prepared)
410        .bind(t)
411        .bind(t)
412        .execute(&mut *tx)
413        .await?
414        .rows_affected();
415        if n == 0 {
416            tx.rollback().await?;
417            return Ok(false);
418        }
419        for b in branches {
420            sqlx::query(&self.be.q("{INS} trans_branch_op
421                 (gid,branch_id,op,url,payload,status,create_time,update_time)
422                 VALUES (?,?,?,?,?,?,?,?)
423                 {NOCONFLICT}"))
424            .bind(&b.gid)
425            .bind(&b.branch_id)
426            .bind(b.op.as_str())
427            .bind(&b.url)
428            .bind(&b.payload)
429            .bind(b.status.as_str())
430            .bind(t)
431            .bind(t)
432            .execute(&mut *tx)
433            .await?;
434        }
435        tx.commit().await?;
436        Ok(true)
437    }
438
439    // ---------------- 访问令牌 ----------------
440
441    pub async fn create_token(&self, hash: &str, name: &str, secret: &str) -> Result<()> {
442        len_ok("name", name, MID)?;
443        len_ok("secret", secret, MID)?;
444        sqlx::query(&self.be.q(
445            "INSERT INTO auth_token(token_hash,name,create_time,last_used,use_count,last_ip,revoked,secret)
446             VALUES(?,?,?,0,0,'',0,?)",
447        ))
448        .bind(hash)
449        .bind(name)
450        .bind(now())
451        .bind(secret)
452        .execute(&self.pool)
453        .await?;
454        Ok(())
455    }
456
457    pub async fn list_tokens(&self) -> Result<Vec<TokenRow>> {
458        let rows = sqlx::query(&self.be.q(
459            "SELECT token_hash,name,create_time,last_used,use_count,last_ip,revoked,secret
460             FROM auth_token ORDER BY create_time DESC",
461        ))
462        .fetch_all(&self.pool)
463        .await?;
464        Ok(rows.iter().map(token_from_row).collect())
465    }
466
467    /// 作废。**不删行** —— 留着才能在管理台看到「这个 token 什么时候被谁作废的」
468    pub async fn revoke_token(&self, hash: &str) -> Result<bool> {
469        let r = sqlx::query(
470            &self
471                .be
472                .q("UPDATE auth_token SET revoked=? WHERE token_hash=? AND revoked=0"),
473        )
474        .bind(now())
475        .bind(hash)
476        .execute(&self.pool)
477        .await?;
478        Ok(r.rows_affected() > 0)
479    }
480
481    /// 当前有效的令牌哈希。认证的热路径**不查这个** —— 上层按 TTL 缓存,
482    /// 见 `dtmrs_server::auth`
483    pub async fn active_token_hashes(&self) -> Result<Vec<String>> {
484        let rows = sqlx::query(&self.be.q(
485            "SELECT token_hash FROM auth_token WHERE revoked=0",
486        ))
487        .fetch_all(&self.pool)
488        .await?;
489        Ok(rows.iter().map(|r| r.get::<String, _>("token_hash")).collect())
490    }
491
492    /// 记一次使用。**尽力而为**:失败只吞掉不影响请求 ——
493    /// 统计信息不值得让一次正常的业务调用失败
494    pub async fn touch_token(&self, hash: &str, ip: &str) -> Result<()> {
495        sqlx::query(&self.be.q(
496            "UPDATE auth_token SET last_used=?, use_count=use_count+1, last_ip=? WHERE token_hash=?",
497        ))
498        .bind(now())
499        .bind(ip)
500        .bind(hash)
501        .execute(&self.pool)
502        .await?;
503        Ok(())
504    }
505
506    pub async fn get_global(&self, gid: &str) -> Result<Option<GlobalRow>> {
507        let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
508            .bind(gid)
509            .fetch_optional(&self.pool)
510            .await?;
511        Ok(row.map(global_from_row))
512    }
513
514    pub async fn list_branches(&self, gid: &str) -> Result<Vec<BranchRow>> {
515        let rows = sqlx::query(&self.be.q(
516            "SELECT gid,branch_id,op,url,payload,status FROM trans_branch_op
517             WHERE gid=? ORDER BY branch_id, op",
518        ))
519        .bind(gid)
520        .fetch_all(&self.pool)
521        .await?;
522        Ok(rows
523            .into_iter()
524            .map(|r| BranchRow {
525                gid: r.get("gid"),
526                branch_id: r.get("branch_id"),
527                op: BranchOp::parse(r.get::<String, _>("op").as_str()).unwrap_or(BranchOp::Action),
528                url: r.get("url"),
529                payload: r.get("payload"),
530                status: BranchStatus::parse(r.get::<String, _>("status").as_str())
531                    .unwrap_or(BranchStatus::Prepared),
532            })
533            .collect())
534    }
535
536    /// 落全局状态。
537    ///
538    /// `trans_type` 这一层用不上(UPDATE 不需要它),但 Redis 后端靠它把
539    /// 「落终态」这条热路径从 Lua 脚本降级成一次 MULTI —— 两边签名要一致
540    pub async fn set_global_status(
541        &self,
542        gid: &str,
543        status: GlobalStatus,
544        _trans_type: TransType,
545        reason: &str,
546    ) -> Result<()> {
547        let t = now();
548        let fin = if status.is_final() { Some(t) } else { None };
549        // reason 是诊断信息,**截断而不是报错**:这条 UPDATE 是状态机的收尾,
550        // 让它因为一句话太长而失败,事务就永远推不到终态了(MySQL strict mode
551        // 下超长 UPDATE 直接报 1406,不像 INSERT IGNORE 那样只是截断)。
552        let reason: String = reason.chars().take(MID).collect();
553        let reason = reason.as_str();
554        // 注意 $4/$5 都绑 reason —— 不能复用同一个 $N,见文件头注释
555        sqlx::query(&self.be.q(
556            "UPDATE trans_global SET status=?, update_time=?, finish_time=?,
557             rollback_reason = CASE WHEN ? <> '' THEN ? ELSE rollback_reason END
558             WHERE gid=?",
559        ))
560        .bind(status.as_str())
561        .bind(t)
562        .bind(fin)
563        .bind(reason)
564        .bind(reason)
565        .bind(gid)
566        .execute(&self.pool)
567        .await?;
568        Ok(())
569    }
570
571    /// 落一个分支的状态**和结果数据**。
572    ///
573    /// workflow 模式的重放靠这个:函数崩溃后会从头再跑一遍,已完成的分支
574    /// 不重新执行,而是把上次存的 `payload` 原样还给它。所以这个值必须跟
575    /// 「分支已成功」在**同一条 UPDATE 里**落盘 —— 分两步写的话,中间崩了
576    /// 就会出现「标了成功但结果丢了」,重放时拿不到返回值。
577    pub async fn set_branch_result(
578        &self,
579        gid: &str,
580        branch_id: &str,
581        op: BranchOp,
582        status: BranchStatus,
583        payload: &str,
584    ) -> Result<()> {
585        len_ok("payload", payload, MID)?;
586        let t = now();
587        sqlx::query(&self.be.q(
588            "UPDATE trans_branch_op SET status=?, payload=?, update_time=?,
589             finish_time = CASE WHEN ? <> 'prepared' THEN ? ELSE finish_time END
590             WHERE gid=? AND branch_id=? AND op=?",
591        ))
592        .bind(status.as_str())
593        .bind(payload)
594        .bind(t)
595        .bind(status.as_str())
596        .bind(t)
597        .bind(gid)
598        .bind(branch_id)
599        .bind(op.as_str())
600        .execute(&self.pool)
601        .await?;
602        Ok(())
603    }
604
605    pub async fn set_branch_status(
606        &self,
607        gid: &str,
608        branch_id: &str,
609        op: BranchOp,
610        status: BranchStatus,
611    ) -> Result<()> {
612        let t = now();
613        sqlx::query(
614            &self
615                .be
616                .q("UPDATE trans_branch_op SET status=?, update_time=?,
617             finish_time = CASE WHEN ? <> 'prepared' THEN ? ELSE finish_time END
618             WHERE gid=? AND branch_id=? AND op=?"),
619        )
620        .bind(status.as_str())
621        .bind(t)
622        .bind(status.as_str())
623        .bind(t)
624        .bind(gid)
625        .bind(branch_id)
626        .bind(op.as_str())
627        .execute(&self.pool)
628        .await?;
629        Ok(())
630    }
631
632    /// 抢一个到期的待办事务,**抢占式更新,原子的**。
633    ///
634    /// 多个 TC 实例同时跑也不会重复推进同一个事务:谁的 UPDATE 生效谁持有租约。
635    /// 持租约的实例崩了,`next_cron_time` 到期后别的实例接手 —— 这就是崩溃恢复。
636    pub async fn lock_one_due(&self, owner: &str, lease: i64) -> Result<Option<GlobalRow>> {
637        let mut tx = self.pool.begin().await?;
638        let t = now();
639        // ⚠ 两个细节都不能改,改了并行推进就退化成串行:
640        //
641        // 1. 结尾的 `FOR UPDATE SKIP LOCKED`(sqlite 上是空串)。没有它的话
642        //    每个 worker 都选中同一行,然后挤在下面那条 UPDATE 上排队,
643        //    只有一个能成。见 `Backend::skip_locked`
644        //
645        // 2. **不能加 `ORDER BY next_cron_time`。** 索引是
646        //    (status, next_cron_time),而 WHERE 里 status 是个 IN 范围,
647        //    所以按 next_cron_time 排序用不上索引 —— MySQL 的执行计划里会
648        //    出现 `Using filesort`,意味着它要**把所有命中的行都读出来并加锁**
649        //    才能排序,然后才 LIMIT 1。于是第一个 worker 锁光全部待办,
650        //    其余 worker 全部 SKIP 掉、一笔都抢不到(实测 6 并发只成 1 笔)。
651        //
652        //    去掉 ORDER BY 后走索引范围扫描,天然就是按 (status, next_cron_time)
653        //    顺序取第一条:同一状态内仍然是**最早到期的先跑**,只是不再跨状态
654        //    全局排序。不会饿死 —— 抢到的行会把 next_cron_time 推到租约之后,
655        //    自动排到队尾。
656        //    (Redis 那边是 ZRANGEBYSCORE,严格按到期时间。可调度的**集合**
657        //    两边完全一致,只是取用顺序不同,这个差异是可以接受的。)
658        let gid: Option<String> = sqlx::query_scalar(&self.be.q(&format!(
659            "SELECT gid FROM trans_global
660             WHERE (status IN ('submitted','aborting')
661                    OR (status = 'prepared' AND trans_type = 'msg'))
662               AND next_cron_time <= ?
663             LIMIT 1{}",
664            self.be.skip_locked()
665        )))
666        .bind(t)
667        .fetch_optional(&mut *tx)
668        .await?;
669        let Some(gid) = gid else {
670            tx.rollback().await?;
671            return Ok(None);
672        };
673        // 立刻把 next_cron_time 推到租约之后,等于占坑
674        let n = sqlx::query(&self.be.q(
675            "UPDATE trans_global SET owner=?, next_cron_time=?, update_time=?
676             WHERE gid=? AND next_cron_time <= ?",
677        ))
678        .bind(owner)
679        .bind(t + lease)
680        .bind(t)
681        .bind(&gid)
682        .bind(t)
683        .execute(&mut *tx)
684        .await?
685        .rows_affected();
686        if n == 0 {
687            tx.rollback().await?;
688            return Ok(None); // 被别人抢走了
689        }
690        let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
691            .bind(&gid)
692            .fetch_one(&mut *tx)
693            .await?;
694        tx.commit().await?;
695        Ok(Some(global_from_row(row)))
696    }
697
698    /// 推进失败后设置下次重试时间(指数退避)
699    pub async fn schedule_retry(&self, gid: &str, interval: i64) -> Result<()> {
700        let t = now();
701        sqlx::query(&self.be.q(
702            "UPDATE trans_global SET next_cron_interval=?, next_cron_time=?, update_time=?
703             WHERE gid=?",
704        ))
705        .bind(interval)
706        .bind(t + interval)
707        .bind(t)
708        .bind(gid)
709        .execute(&self.pool)
710        .await?;
711        Ok(())
712    }
713
714    /// 把停在 prepared 的事务推成 submitted,并立刻排进调度队列。
715    ///
716    /// **一次调用做完原来三次的活**(`get_global` + `set_global_status` +
717    /// `schedule_now`)。Redis 后端上这是一个 Lua 脚本,11 条命令降到 3 条 ——
718    /// 那边是单线程 CPU 瓶颈,命令数直接决定吞吐。
719    ///
720    /// `owner` / `next_cron_time` 让提交方**顺便把租约占下来**:传自己的
721    /// owner 和「现在 + 租约」,这一条 UPDATE 之后事务就归调用方推了,
722    /// 不用再走一次抢占。不想占就传空 owner 和 `now()`。
723    pub async fn submit_prepared(
724        &self,
725        gid: &str,
726        owner: &str,
727        next_cron_time: i64,
728    ) -> Result<SubmitOutcome> {
729        // 先查一次。saga 第一次提交时事务还不存在,这是最常见的路径,
730        // 一次 SELECT 就该返回,不值得为它先空跑一条 UPDATE。
731        // 取全行而不只是 status —— 反正这条 SELECT 免不了,顺手把事务体带回去,
732        // 调用方就能直接开推(见 `SubmitOutcome::Advanced`)
733        let row = sqlx::query(&self.be.q(&format!("{SELECT_GLOBAL} WHERE gid=?")))
734            .bind(gid)
735            .fetch_optional(&self.pool)
736            .await?;
737        let Some(row) = row else {
738            return Ok(SubmitOutcome::Missing);
739        };
740        let mut g = global_from_row(row);
741        if g.status != GlobalStatus::Prepared {
742            return Ok(SubmitOutcome::Already);
743        }
744        let t = now();
745        // 状态、退避、排队一条 UPDATE 落完。
746        // ⚠ `AND status=?`(prepared)不能省:并发重复提交时,
747        // 别把已经在推进的事务硬拽回队首
748        sqlx::query(&self.be.q("UPDATE trans_global SET status=?, update_time=?,
749             next_cron_time=?, next_cron_interval=0, owner=?
750             WHERE gid=? AND status=?"))
751        .bind(GlobalStatus::Submitted.as_str())
752        .bind(t)
753        .bind(next_cron_time)
754        .bind(owner)
755        .bind(gid)
756        .bind(GlobalStatus::Prepared.as_str())
757        .execute(&self.pool)
758        .await?;
759        // 把刚写下去的三个字段补到返回的事务体上,省掉一次回读
760        g.status = GlobalStatus::Submitted;
761        g.next_cron_time = next_cron_time;
762        g.owner = owner.to_string();
763        Ok(SubmitOutcome::Advanced(Box::new(g)))
764    }
765
766    /// 让某个事务立刻可被调度(提交/中止之后叫一下,不用等 cron 周期)
767    pub async fn schedule_now(&self, gid: &str) -> Result<()> {
768        sqlx::query(
769            &self
770                .be
771                .q("UPDATE trans_global SET next_cron_time=?, next_cron_interval=0 WHERE gid=?"),
772        )
773        .bind(now())
774        .bind(gid)
775        .execute(&self.pool)
776        .await?;
777        Ok(())
778    }
779
780    /// TCC 的 try 阶段:客户端在调 try 之前先来登记这个分支的 confirm/cancel。
781    ///
782    /// **必须先登记再调 try**。反过来的话:try 成功了但登记失败,
783    /// TC 就不知道有这个分支,回滚时不会 cancel 它 —— 资源永久泄漏。
784    ///
785    /// 冲突时忽略,所以重复登记是幂等的(客户端重试很常见)。
786    pub async fn register_branch(
787        &self,
788        gid: &str,
789        branch_id: &str,
790        ops: &[(BranchOp, String)],
791    ) -> Result<RegisterOutcome> {
792        len_ok("gid", gid, Backend::ID_MAX)?;
793        len_ok("branch_id", branch_id, Backend::ID_MAX)?;
794        for (_, url) in ops {
795            len_ok("url", url, MID)?;
796        }
797        let mut tx = self.pool.begin().await?;
798        let t = now();
799        for (op, url) in ops {
800            sqlx::query(&self.be.q("{INS} trans_branch_op
801                 (gid,branch_id,op,url,payload,status,create_time,update_time)
802                 VALUES (?,?,?,?,'',?,?,?)
803                 {NOCONFLICT}"))
804            .bind(gid)
805            .bind(branch_id)
806            .bind(op.as_str())
807            .bind(url)
808            .bind(BranchStatus::Prepared.as_str())
809            .bind(t)
810            .bind(t)
811            .execute(&mut *tx)
812            .await?;
813
814            // ⚠ 插完必须回读一次,**不能看 rows_affected**。
815            //
816            // MySQL 上 `INSERT IGNORE` 遇到冲突返回 0、成功返回 1,看似能用;
817            // 但这里要区分的不是「插没插进去」而是「里面躺的是不是同一个 URL」,
818            // 那个数字回答不了。回读还顺带解决了并发:两个请求同时登记同一个号时,
819            // 输的那个在自己事务里读到的是赢家已提交的行,照样能发现冲突。
820            let stored: Option<String> = sqlx::query_scalar(&self.be.q(
821                "SELECT url FROM trans_branch_op WHERE gid=? AND branch_id=? AND op=?",
822            ))
823            .bind(gid)
824            .bind(branch_id)
825            .bind(op.as_str())
826            .fetch_optional(&mut *tx)
827            .await?;
828            if let Some(existing) = stored {
829                if existing != *url {
830                    // 不 commit,前面几个 op 一并回滚 —— 半登记的分支比没登记更难查
831                    return Ok(RegisterOutcome::Conflict {
832                        op: *op,
833                        existing,
834                    });
835                }
836            }
837        }
838        tx.commit().await?;
839        Ok(RegisterOutcome::Registered)
840    }
841
842    pub async fn list_recent(&self, limit: i64) -> Result<Vec<GlobalRow>> {
843        let rows = sqlx::query(&self.be.q(&format!(
844            "{SELECT_GLOBAL} ORDER BY create_time DESC LIMIT ?"
845        )))
846        .bind(limit)
847        .fetch_all(&self.pool)
848        .await?;
849        Ok(rows.into_iter().map(global_from_row).collect())
850    }
851}
852
853/// 列清单只写一处 —— 三个地方读 trans_global,列顺序漂移过一次就够难查了
854const SELECT_GLOBAL: &str = "SELECT gid,trans_type,status,payload,next_cron_time,
855    next_cron_interval,owner,rollback_reason,query_prepared,create_time,finish_time
856    FROM trans_global";
857
858fn token_from_row(r: &AnyRow) -> TokenRow {
859    TokenRow {
860        token_hash: r.get("token_hash"),
861        name: r.get("name"),
862        create_time: r.get("create_time"),
863        last_used: r.get("last_used"),
864        use_count: r.get("use_count"),
865        last_ip: r.get("last_ip"),
866        revoked: r.get("revoked"),
867        secret: r.get("secret"),
868    }
869}
870
871fn global_from_row(r: AnyRow) -> GlobalRow {
872    GlobalRow {
873        gid: r.get("gid"),
874        trans_type: TransType::parse(r.get::<String, _>("trans_type").as_str())
875            .unwrap_or(TransType::Saga),
876        status: GlobalStatus::parse(r.get::<String, _>("status").as_str())
877            .unwrap_or(GlobalStatus::Prepared),
878        payload: r.get("payload"),
879        next_cron_time: r.get("next_cron_time"),
880        next_cron_interval: r.get("next_cron_interval"),
881        owner: r.get("owner"),
882        rollback_reason: r.get("rollback_reason"),
883        query_prepared: r.get("query_prepared"),
884        create_time: r.get("create_time"),
885        finish_time: r.get("finish_time"),
886    }
887}
888
889#[cfg(test)]
890mod tests {
891    use super::*;
892
893    /// 每个测试都在**所有可用后端**上跑一遍:sqlite / postgres / mysql / redis。
894    ///
895    /// 真库靠环境变量开启(`DTMRS_TEST_PG` / `DTMRS_TEST_MYSQL` / `DTMRS_TEST_REDIS`)——
896    /// 没配就只跑 sqlite,这样没数据库的机器也能 `cargo test`。
897    /// 但**别把这当成"它们也过了"** —— 没配就是没测。
898    ///
899    /// Redis 跟另外三个不是同一类东西(不是 SQL),能共用这一套断言恰恰是
900    /// 我们要的证据:两种后端的**行为**必须一致,哪怕实现天差地别。
901    /// Postgres 测试必须串行 —— `lock_one_due` 和 `list_recent` 是**全局查询**,
902    /// 并行跑会互相看见对方的事务,断言就没意义了。
903    /// (光给各测试不同的 gid 不够:捞待办是不按 gid 过滤的。)
904    static PG_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
905
906    /// 返回 (串行锁, 各后端)。锁要持到测试结束,所以由调用方接着。
907    ///
908    /// 每次进来把 Postgres 的表清空(只 DELETE 不 DDL —— 并发 DDL 会撞上
909    /// Postgres 的 `pg_type` 竞态,见 `migrate_racy`)。
910    /// 「跳过 ≠ 通过」的闸门。
911    ///
912    /// 这些测试没配环境变量时直接返回,**仍然显示为 passed**。所以只要 CI 里
913    /// 某个数据库容器没起来、或者环境变量名打错一个字母,那个 job 会
914    /// **安安静静地全绿** —— 而真库那部分其实一行没跑。
915    ///
916    /// CI 的真库 job 里设 `DTMRS_TEST_REQUIRE_REAL_DB=1`,把「悄悄没测」
917    /// 变成「响亮地失败」。本地开发不设这个变量,跳过行为不变。
918    fn require_real_db(缺的变量: &str) {
919        if std::env::var("DTMRS_TEST_REQUIRE_REAL_DB").is_ok() {
920            panic!(
921                "设了 DTMRS_TEST_REQUIRE_REAL_DB,却没有 {缺的变量} —— \
922                 这是 CI 配置坏了(容器没起来?变量名打错?),不是可以跳过的情况"
923            );
924        }
925    }
926
927    async fn backends() -> (
928        tokio::sync::MutexGuard<'static, ()>,
929        Vec<(&'static str, Store)>,
930    ) {
931        let guard = PG_LOCK.lock().await;
932        let mut v = vec![("sqlite", Store::open("sqlite::memory:").await.unwrap())];
933        // 每种真数据库都配一个环境变量。**没配就是没测**,不是"通过"。
934        for (name, env) in [("postgres", "DTMRS_TEST_PG"), ("mysql", "DTMRS_TEST_MYSQL")] {
935            if std::env::var(env).is_err() {
936                require_real_db(env);
937                continue;
938            }
939            if let Ok(url) = std::env::var(env) {
940                let s = Store::open(&url)
941                    .await
942                    .unwrap_or_else(|e| panic!("连不上 {env}: {e}"));
943                for t in ["trans_branch_op", "trans_global"] {
944                    sqlx::query(&format!("DELETE FROM {t}"))
945                        .execute(s.pool().expect("SQL 后端才有连接池"))
946                        .await
947                        .expect("清表");
948                }
949                v.push((name, s));
950            }
951        }
952        #[cfg(feature = "redis")]
953        if std::env::var("DTMRS_TEST_REDIS").is_err() {
954            require_real_db("DTMRS_TEST_REDIS");
955        }
956        #[cfg(feature = "redis")]
957        if let Ok(url) = std::env::var("DTMRS_TEST_REDIS") {
958            let s = Store::open(&url)
959                .await
960                .unwrap_or_else(|e| panic!("连不上 DTMRS_TEST_REDIS: {e}"));
961            // 每次进来清干净。Redis 没有"表",按前缀删
962            s.as_redis()
963                .unwrap()
964                .flush_prefix()
965                .await
966                .expect("清 redis");
967            v.push(("redis", s));
968        }
969        (guard, v)
970    }
971
972    fn g(gid: &str) -> GlobalRow {
973        GlobalRow {
974            gid: gid.into(),
975            trans_type: TransType::Saga,
976            status: GlobalStatus::Submitted,
977            payload: "{}".into(),
978            next_cron_time: 0,
979            next_cron_interval: 0,
980            owner: String::new(),
981            rollback_reason: String::new(),
982            query_prepared: String::new(),
983            create_time: 0,
984            finish_time: None,
985        }
986    }
987
988    #[tokio::test]
989    async fn 重复提交同一个gid是幂等的() {
990        let (_g, bes) = backends().await;
991        for (name, s) in bes {
992            assert!(s.create_global(&g("t1"), &[]).await.unwrap(), "{name}");
993            // 第二次返回 false 而不是报错 —— 客户端重试不该失败
994            assert!(!s.create_global(&g("t1"), &[]).await.unwrap(), "{name}");
995            assert_eq!(s.list_recent(10).await.unwrap().len(), 1, "{name}");
996        }
997    }
998
999    #[tokio::test]
1000    async fn 租约只能被抢到一次() {
1001        let (_g, bes) = backends().await;
1002        for (name, s) in bes {
1003            s.create_global(&g("t2"), &[]).await.unwrap();
1004            let a = s.lock_one_due("worker-a", 60).await.unwrap();
1005            assert!(a.is_some(), "{name}: 第一个实例应该抢到");
1006            // 同一个事务不能被第二个实例同时抢到,否则会重复推进
1007            let b = s.lock_one_due("worker-b", 60).await.unwrap();
1008            assert!(b.is_none(), "{name}: 租约期内不能被别人抢走");
1009        }
1010    }
1011
1012    /// 并发抢占要抢到**不同的**事务,而不是全挤在同一笔上。
1013    ///
1014    /// 这条钉的是 `FOR UPDATE SKIP LOCKED`(见 `Backend::skip_locked`)。
1015    /// 少了它,N 个 worker 的 SELECT 会同时选中队首那一行,然后在 UPDATE
1016    /// 上排队,最后只有一个成功 —— 不会算错,但并行推进等于白做:
1017    /// 实测 Postgres 上 8 个 worker 只跑出 1 个 worker 的 1.8 倍。
1018    ///
1019    /// sqlite 例外:它没有行锁,写本来就是全库串行的。所以那边只要求
1020    /// 「不重复」(安全性),不要求「都能抢到」(并行度)。
1021    #[tokio::test]
1022    async fn 并发抢占要各拿各的不能全挤在同一笔上() {
1023        const K: usize = 6;
1024        let (_g, bes) = backends().await;
1025        for (name, s) in bes {
1026            for i in 0..K {
1027                s.create_global(&g(&format!("par-{i}")), &[]).await.unwrap();
1028            }
1029
1030            let mut hs = Vec::new();
1031            for i in 0..K {
1032                let s = s.clone();
1033                hs.push(tokio::spawn(async move {
1034                    s.lock_one_due(&format!("w-{i}"), 60).await.unwrap()
1035                }));
1036            }
1037            let mut got: Vec<String> = Vec::new();
1038            for h in hs {
1039                if let Some(row) = h.await.unwrap() {
1040                    got.push(row.gid);
1041                }
1042            }
1043
1044            // 安全性:所有后端都不能把同一笔交给两个 owner
1045            let uniq: std::collections::HashSet<_> = got.iter().collect();
1046            assert_eq!(uniq.len(), got.len(), "{name}: 同一笔被抢到了两次");
1047
1048            // 并行度:有行锁的后端应该 K 个各拿各的
1049            if name != "sqlite" {
1050                assert_eq!(
1051                    got.len(),
1052                    K,
1053                    "{name}: 并发抢占退化成串行了(SKIP LOCKED 没生效?)"
1054                );
1055            }
1056        }
1057    }
1058
1059    #[tokio::test]
1060    async fn 终态不再被调度() {
1061        let (_g, bes) = backends().await;
1062        for (name, s) in bes {
1063            s.create_global(&g("t3"), &[]).await.unwrap();
1064            s.set_global_status("t3", GlobalStatus::Succeed, TransType::Saga, "")
1065                .await
1066                .unwrap();
1067            assert!(s.lock_one_due("w", 60).await.unwrap().is_none(), "{name}");
1068            let got = s.get_global("t3").await.unwrap().unwrap();
1069            assert_eq!(got.status, GlobalStatus::Succeed, "{name}");
1070            assert!(got.finish_time.is_some(), "{name}: 终态要落 finish_time");
1071        }
1072    }
1073
1074    #[tokio::test]
1075    async fn 分支状态可更新() {
1076        let (_g, bes) = backends().await;
1077        for (name, s) in bes {
1078            let b = BranchRow {
1079                gid: "t4".into(),
1080                branch_id: "01".into(),
1081                op: BranchOp::Action,
1082                url: "http://x/a".into(),
1083                payload: "{}".into(),
1084                status: BranchStatus::Prepared,
1085            };
1086            s.create_global(&g("t4"), std::slice::from_ref(&b))
1087                .await
1088                .unwrap();
1089            s.set_branch_status("t4", "01", BranchOp::Action, BranchStatus::Succeed)
1090                .await
1091                .unwrap();
1092            let got = s.list_branches("t4").await.unwrap();
1093            assert_eq!(got.len(), 1, "{name}");
1094            assert_eq!(got[0].status, BranchStatus::Succeed, "{name}");
1095        }
1096    }
1097
1098    #[tokio::test]
1099    async fn 回滚原因和回查地址能存取() {
1100        // 这两列是后加的,跨库的字符串/空值处理最容易在这儿出问题
1101        let (_g, bes) = backends().await;
1102        for (name, s) in bes {
1103            let mut row = g("t5");
1104            row.query_prepared = "http://busi/query".into();
1105            s.create_global(&row, &[]).await.unwrap();
1106            s.set_global_status(
1107                "t5",
1108                GlobalStatus::Aborting,
1109                TransType::Saga,
1110                "分支 02 返回 FAILURE",
1111            )
1112            .await
1113            .unwrap();
1114            let got = s.get_global("t5").await.unwrap().unwrap();
1115            assert_eq!(got.query_prepared, "http://busi/query", "{name}");
1116            assert_eq!(got.rollback_reason, "分支 02 返回 FAILURE", "{name}");
1117            assert!(
1118                got.finish_time.is_none(),
1119                "{name}: 非终态不该有 finish_time"
1120            );
1121
1122            // 空 reason 不能把已有的原因冲掉
1123            s.set_global_status("t5", GlobalStatus::Failed, TransType::Saga, "")
1124                .await
1125                .unwrap();
1126            let got = s.get_global("t5").await.unwrap().unwrap();
1127            assert_eq!(
1128                got.rollback_reason, "分支 02 返回 FAILURE",
1129                "{name}: 空原因不能覆盖"
1130            );
1131        }
1132    }
1133
1134    /// 重号登记的判定必须在**每个后端**上一致。
1135    ///
1136    /// 三家的机制各不相同:SQL 走「冲突忽略 + 回读比对」,其中 MySQL 是
1137    /// `INSERT IGNORE`、其它是 `ON CONFLICT DO NOTHING`;Redis 走
1138    /// `hset_nx` 的返回值 + 回读。机制不同,结论必须逐条一样 ——
1139    /// 所以这条测试挂在 `backends()` 上而不是只测 sqlite。
1140    #[tokio::test]
1141    async fn 重号登记要报冲突而同号重试要幂等() {
1142        let (_g, backends) = backends().await;
1143        for (name, s) in backends {
1144            let 库存 = [
1145                (BranchOp::Confirm, "http://kucun/confirm".to_string()),
1146                (BranchOp::Cancel, "http://kucun/cancel".to_string()),
1147            ];
1148            let 订单 = [
1149                (BranchOp::Confirm, "http://dingdan/confirm".to_string()),
1150                (BranchOp::Cancel, "http://dingdan/cancel".to_string()),
1151            ];
1152
1153            assert_eq!(
1154                s.register_branch("dup1", "01", &库存).await.unwrap(),
1155                RegisterOutcome::Registered,
1156                "{name}: 首次登记"
1157            );
1158            assert_eq!(
1159                s.register_branch("dup1", "01", &库存).await.unwrap(),
1160                RegisterOutcome::Registered,
1161                "{name}: URL 一致的重复登记是客户端重试,必须幂等放行"
1162            );
1163            assert!(
1164                matches!(
1165                    s.register_branch("dup1", "01", &订单).await.unwrap(),
1166                    RegisterOutcome::Conflict { .. }
1167                ),
1168                "{name}: 重号必须报冲突 —— 放行的话订单的 URL 根本写不进去,\
1169                 客户端却以为登记成功并去冻结资源,那份资源永久泄漏"
1170            );
1171            assert_eq!(
1172                s.register_branch("dup1", "02", &订单).await.unwrap(),
1173                RegisterOutcome::Registered,
1174                "{name}: 各用各的号要互不影响"
1175            );
1176
1177            let rows = s.list_branches("dup1").await.unwrap();
1178            assert_eq!(rows.len(), 4, "{name}: 两个分支各两个 op");
1179            for r in &rows {
1180                let 期望 = if r.branch_id == "01" { "kucun" } else { "dingdan" };
1181                assert!(
1182                    r.url.contains(期望),
1183                    "{name}: 分支 {} 的地址串味了 —— {}",
1184                    r.branch_id,
1185                    r.url
1186                );
1187            }
1188        }
1189    }
1190
1191    #[tokio::test]
1192    async fn msg的prepared会被捞tcc的不会() {
1193        let (_g, bes) = backends().await;
1194        for (name, s) in bes {
1195            let mut m = g("m1");
1196            m.trans_type = TransType::Msg;
1197            m.status = GlobalStatus::Prepared;
1198            s.create_global(&m, &[]).await.unwrap();
1199            let mut t = g("c1");
1200            t.trans_type = TransType::Tcc;
1201            t.status = GlobalStatus::Prepared;
1202            s.create_global(&t, &[]).await.unwrap();
1203
1204            let got = s.lock_one_due("w", 60).await.unwrap();
1205            assert_eq!(
1206                got.map(|x| x.gid),
1207                Some("m1".to_string()),
1208                "{name}: 只该捞到 msg"
1209            );
1210            // 再捞一次应该没有了(msg 被租约占住,tcc 不该被碰)
1211            assert!(s.lock_one_due("w2", 60).await.unwrap().is_none(), "{name}");
1212        }
1213    }
1214
1215    #[tokio::test]
1216    async fn 分支登记是幂等的() {
1217        let (_g, bes) = backends().await;
1218        for (name, s) in bes {
1219            let mut t = g("c2");
1220            t.trans_type = TransType::Tcc;
1221            s.create_global(&t, &[]).await.unwrap();
1222            let ops = [
1223                (BranchOp::Confirm, "http://x/c".to_string()),
1224                (BranchOp::Cancel, "http://x/n".to_string()),
1225            ];
1226            s.register_branch("c2", "01", &ops).await.unwrap();
1227            s.register_branch("c2", "01", &ops).await.unwrap(); // 客户端重试
1228            assert_eq!(
1229                s.list_branches("c2").await.unwrap().len(),
1230                2,
1231                "{name}: 不该重复插入"
1232            );
1233        }
1234    }
1235}
1236
1237// ==================== 后端分发 ====================
1238
1239#[cfg(feature = "redis")]
1240pub mod redis_store;
1241#[cfg(feature = "redis")]
1242pub use redis_store::RedisStore;
1243
1244/// 存储后端。
1245///
1246/// # 为什么现在才抽这一层
1247///
1248/// 这个项目原本**刻意没有抽 `Store` trait**,理由写在 DESIGN.md 里:
1249/// sqlite / postgres / mysql 的差异小到一层 SQL 模板就能吸收,抽象是过早的。
1250/// 那个判断在当时是对的。
1251///
1252/// **Redis 让前提不成立了** —— 它根本不是 SQL,没有表、没有事务、没有 WHERE,
1253/// 模板吸收不了。所以这里加了一层分发。
1254///
1255/// 用 enum 而不是 trait:调用方拿到的还是同一个 `Store` 具体类型,
1256/// 四十多个调用点一行都不用改,也不用到处写泛型或 `dyn`。
1257#[derive(Clone)]
1258enum Inner {
1259    Sql(SqlStore),
1260    #[cfg(feature = "redis")]
1261    Redis(RedisStore),
1262}
1263
1264/// 存储层的统一入口。按 URL 前缀自动选后端:
1265///
1266/// ```text
1267/// sqlite:...     / postgres://...  / mysql://...   → SQL 后端
1268/// redis://...    / rediss://...                    → Redis 后端(要开 redis feature)
1269/// ```
1270///
1271/// ⚠ Redis 后端跟 SQL 后端有**实打实的语义差异**(持久性更弱、终态会过期),
1272/// 用之前务必读 [`redis_store`] 的模块说明。
1273#[derive(Clone)]
1274pub struct Store {
1275    inner: Inner,
1276}
1277
1278/// 存储层的错误。
1279///
1280/// 两种后端的原生错误类型不同,统一收口到这里;`sqlx::Error` 仍然直接透出,
1281/// 免得改动现有调用方对错误的处理。
1282pub type StoreError = sqlx::Error;
1283
1284#[cfg(feature = "redis")]
1285fn redis_err(e: redis::RedisError) -> sqlx::Error {
1286    sqlx::Error::Configuration(Box::new(e))
1287}
1288
1289/// 这个 URL 是不是要走 Redis
1290pub fn is_redis_url(url: &str) -> bool {
1291    let u = url.trim().to_ascii_lowercase();
1292    u.starts_with("redis://") || u.starts_with("rediss://") || u.starts_with("redis+unix:")
1293}
1294
1295impl Store {
1296    /// 按 URL 选后端并连上。
1297    pub async fn open(url: &str) -> Result<Self> {
1298        if is_redis_url(url) {
1299            #[cfg(feature = "redis")]
1300            {
1301                let r = RedisStore::open(url).await.map_err(redis_err)?;
1302                return Ok(Self {
1303                    inner: Inner::Redis(r),
1304                });
1305            }
1306            #[cfg(not(feature = "redis"))]
1307            {
1308                // 明确报错,而不是把 redis:// 当成 sqlite 文件名去建库 ——
1309                // 那会静默跑起来然后数据全落在一个叫 "redis:" 的文件里
1310                return Err(sqlx::Error::Configuration(
1311                    "这个 URL 要 Redis 后端,但构建时没开 dtmrs-store 的 `redis` feature".into(),
1312                ));
1313            }
1314        }
1315        Ok(Self {
1316            inner: Inner::Sql(SqlStore::open(url).await?),
1317        })
1318    }
1319
1320    // ---------------- 访问令牌 ----------------
1321    //
1322    // 两个后端的语义必须逐条一致:作废是打标记不删、列举按创建时间倒序、
1323    // 重复作废返回 false。Redis 侧的令牌 key **不设 TTL** ——
1324    // 事务是流水可以过期,令牌是配置,过期消失等于凭据莫名失效。
1325
1326    pub async fn create_token(&self, hash: &str, name: &str, secret: &str) -> Result<()> {
1327        match &self.inner {
1328            Inner::Sql(s) => s.create_token(hash, name, secret).await,
1329            #[cfg(feature = "redis")]
1330            Inner::Redis(r) => r.create_token(hash, name, secret).await.map_err(redis_err),
1331        }
1332    }
1333
1334    pub async fn list_tokens(&self) -> Result<Vec<TokenRow>> {
1335        match &self.inner {
1336            Inner::Sql(s) => s.list_tokens().await,
1337            #[cfg(feature = "redis")]
1338            Inner::Redis(r) => r.list_tokens().await.map_err(redis_err),
1339        }
1340    }
1341
1342    pub async fn revoke_token(&self, hash: &str) -> Result<bool> {
1343        match &self.inner {
1344            Inner::Sql(s) => s.revoke_token(hash).await,
1345            #[cfg(feature = "redis")]
1346            Inner::Redis(r) => r.revoke_token(hash).await.map_err(redis_err),
1347        }
1348    }
1349
1350    pub async fn active_token_hashes(&self) -> Result<Vec<String>> {
1351        match &self.inner {
1352            Inner::Sql(s) => s.active_token_hashes().await,
1353            #[cfg(feature = "redis")]
1354            Inner::Redis(r) => r.active_token_hashes().await.map_err(redis_err),
1355        }
1356    }
1357
1358    pub async fn touch_token(&self, hash: &str, ip: &str) -> Result<()> {
1359        match &self.inner {
1360            Inner::Sql(s) => s.touch_token(hash, ip).await,
1361            #[cfg(feature = "redis")]
1362            Inner::Redis(r) => r.touch_token(hash, ip).await.map_err(redis_err),
1363        }
1364    }
1365
1366    /// 底层是不是 Redis
1367    pub fn is_redis(&self) -> bool {
1368        match &self.inner {
1369            Inner::Sql(_) => false,
1370            #[cfg(feature = "redis")]
1371            Inner::Redis(_) => true,
1372        }
1373    }
1374
1375    /// SQL 后端的连接池。Redis 后端返回 `None` ——
1376    /// 调用方(主要是测试和屏障)要自己处理这种情况
1377    pub fn pool(&self) -> Option<&AnyPool> {
1378        match &self.inner {
1379            Inner::Sql(s) => Some(s.pool()),
1380            #[cfg(feature = "redis")]
1381            Inner::Redis(_) => None,
1382        }
1383    }
1384
1385    /// SQL 方言。Redis 后端没有方言可言,返回 `None`
1386    pub fn backend(&self) -> Option<Backend> {
1387        match &self.inner {
1388            Inner::Sql(s) => Some(s.backend()),
1389            #[cfg(feature = "redis")]
1390            Inner::Redis(_) => None,
1391        }
1392    }
1393
1394    /// 拿底层的 Redis store(比如为了调 `with_ttl`)
1395    #[cfg(feature = "redis")]
1396    pub fn as_redis(&self) -> Option<&RedisStore> {
1397        match &self.inner {
1398            Inner::Redis(r) => Some(r),
1399            _ => None,
1400        }
1401    }
1402}
1403
1404/// 把 13 个方法逐个手写分发太啰嗦,而且漏一个编译器不会提醒 ——
1405/// 用宏保证两边签名严格一致
1406macro_rules! dispatch {
1407    ($( $(#[$m:meta])* fn $name:ident (&self $(, $arg:ident : $ty:ty)* ) -> $ret:ty; )*) => {
1408        impl Store {
1409            $(
1410                $(#[$m])*
1411                pub async fn $name(&self $(, $arg: $ty)*) -> Result<$ret> {
1412                    match &self.inner {
1413                        Inner::Sql(s) => s.$name($($arg),*).await,
1414                        #[cfg(feature = "redis")]
1415                        Inner::Redis(r) => r.$name($($arg),*).await.map_err(redis_err),
1416                    }
1417                }
1418            )*
1419        }
1420    };
1421}
1422
1423dispatch! {
1424    /// 建表(Redis 后端是空操作)
1425    fn migrate(&self) -> ();
1426    /// 建全局事务 + 分支。已存在返回 `false`,**不覆盖**
1427    fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> bool;
1428    fn get_global(&self, gid: &str) -> Option<GlobalRow>;
1429    fn list_branches(&self, gid: &str) -> Vec<BranchRow>;
1430    /// 抢一个到期事务。多实例不重复推进就靠它的原子性
1431    fn lock_one_due(&self, owner: &str, lease: i64) -> Option<GlobalRow>;
1432    fn set_global_status(&self, gid: &str, status: GlobalStatus, trans_type: TransType, reason: &str) -> ();
1433    /// 把 prepared 推成 submitted 并排进调度队列,一次调用做完。见 [`SubmitOutcome`]
1434    fn submit_prepared(&self, gid: &str, owner: &str, next_cron_time: i64) -> SubmitOutcome;
1435    fn set_branch_result(&self, gid: &str, branch_id: &str, op: BranchOp, status: BranchStatus, payload: &str) -> ();
1436    fn set_branch_status(&self, gid: &str, branch_id: &str, op: BranchOp, status: BranchStatus) -> ();
1437    fn schedule_retry(&self, gid: &str, interval: i64) -> ();
1438    fn schedule_now(&self, gid: &str) -> ();
1439    /// 登记分支。**重号但 URL 不同时返回 [`RegisterOutcome::Conflict`]**,
1440    /// 调用方必须拒绝 —— 见那个类型的文档
1441    fn register_branch(&self, gid: &str, branch_id: &str, ops: &[(BranchOp, String)]) -> RegisterOutcome;
1442    fn list_recent(&self, limit: i64) -> Vec<GlobalRow>;
1443}