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