1pub 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
30pub const BIG: usize = 8192;
32pub const MID: usize = 1024;
34
35fn len_ok(col: &'static str, val: &str, max: usize) -> Result<()> {
41 check_len(col, val, max).map_err(|e| sqlx::Error::Encode(Box::new(e)))
42}
43
44pub fn now() -> i64 {
45 std::time::SystemTime::now()
46 .duration_since(std::time::UNIX_EPOCH)
47 .map(|d| d.as_secs() as i64)
48 .unwrap_or(0)
49}
50
51#[derive(Debug, Clone)]
57pub enum SubmitOutcome {
58 Missing,
60 Advanced(Box<GlobalRow>),
67 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 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#[derive(Debug, Clone)]
101pub struct TokenRow {
102 pub token_hash: String,
104 pub name: String,
106 pub create_time: i64,
107 pub last_used: i64,
109 pub use_count: i64,
110 pub last_ip: String,
111 pub revoked: i64,
113 pub secret: String,
116}
117
118pub 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 pub async fn open(url: &str) -> Result<Self> {
143 DRIVERS.call_once(sqlx::any::install_default_drivers);
144
145 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 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 for pragma in [
196 "PRAGMA journal_mode=WAL",
197 "PRAGMA synchronous=NORMAL",
198 "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 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 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 let big = self.be.text(BIG);
244 let mid = self.be.text(MID);
245 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 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 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 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 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 pub async fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> Result<bool> {
359 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 .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 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 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 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 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 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 let reason: String = reason.chars().take(MID).collect();
530 let reason = reason.as_str();
531 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 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 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 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 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); }
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 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 pub async fn submit_prepared(
701 &self,
702 gid: &str,
703 owner: &str,
704 next_cron_time: i64,
705 ) -> Result<SubmitOutcome> {
706 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 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 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 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 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
806const 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 static PG_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
858
859 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 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 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 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 let b = s.lock_one_due("worker-b", 60).await.unwrap();
961 assert!(b.is_none(), "{name}: 租约期内不能被别人抢走");
962 }
963 }
964
965 #[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 let uniq: std::collections::HashSet<_> = got.iter().collect();
999 assert_eq!(uniq.len(), got.len(), "{name}: 同一笔被抢到了两次");
1000
1001 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 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 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 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(); assert_eq!(
1125 s.list_branches("c2").await.unwrap().len(),
1126 2,
1127 "{name}: 不该重复插入"
1128 );
1129 }
1130 }
1131}
1132
1133#[cfg(feature = "redis")]
1136pub mod redis_store;
1137#[cfg(feature = "redis")]
1138pub use redis_store::RedisStore;
1139
1140#[derive(Clone)]
1154enum Inner {
1155 Sql(SqlStore),
1156 #[cfg(feature = "redis")]
1157 Redis(RedisStore),
1158}
1159
1160#[derive(Clone)]
1170pub struct Store {
1171 inner: Inner,
1172}
1173
1174pub 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
1185pub 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 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 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 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 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 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 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 #[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
1300macro_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 fn migrate(&self) -> ();
1322 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 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 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}