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, PartialEq, Eq)]
83pub enum RegisterOutcome {
84 Registered,
86 Conflict {
88 op: BranchOp,
89 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 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#[derive(Debug, Clone)]
124pub struct TokenRow {
125 pub token_hash: String,
127 pub name: String,
129 pub create_time: i64,
130 pub last_used: i64,
132 pub use_count: i64,
133 pub last_ip: String,
134 pub revoked: i64,
136 pub secret: String,
139}
140
141pub 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 pub async fn open(url: &str) -> Result<Self> {
166 DRIVERS.call_once(sqlx::any::install_default_drivers);
167
168 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 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 for pragma in [
219 "PRAGMA journal_mode=WAL",
220 "PRAGMA synchronous=NORMAL",
221 "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 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 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 let big = self.be.text(BIG);
267 let mid = self.be.text(MID);
268 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 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 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 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 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 pub async fn create_global(&self, g: &GlobalRow, branches: &[BranchRow]) -> Result<bool> {
382 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 .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 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 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 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 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 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 let reason: String = reason.chars().take(MID).collect();
553 let reason = reason.as_str();
554 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 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 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 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 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); }
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 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 pub async fn submit_prepared(
724 &self,
725 gid: &str,
726 owner: &str,
727 next_cron_time: i64,
728 ) -> Result<SubmitOutcome> {
729 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 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 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 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 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 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 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
853const 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 static PG_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
905
906 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 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 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 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 let b = s.lock_one_due("worker-b", 60).await.unwrap();
1008 assert!(b.is_none(), "{name}: 租约期内不能被别人抢走");
1009 }
1010 }
1011
1012 #[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 let uniq: std::collections::HashSet<_> = got.iter().collect();
1046 assert_eq!(uniq.len(), got.len(), "{name}: 同一笔被抢到了两次");
1047
1048 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 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 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 #[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 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(); assert_eq!(
1229 s.list_branches("c2").await.unwrap().len(),
1230 2,
1231 "{name}: 不该重复插入"
1232 );
1233 }
1234 }
1235}
1236
1237#[cfg(feature = "redis")]
1240pub mod redis_store;
1241#[cfg(feature = "redis")]
1242pub use redis_store::RedisStore;
1243
1244#[derive(Clone)]
1258enum Inner {
1259 Sql(SqlStore),
1260 #[cfg(feature = "redis")]
1261 Redis(RedisStore),
1262}
1263
1264#[derive(Clone)]
1274pub struct Store {
1275 inner: Inner,
1276}
1277
1278pub 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
1289pub 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 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 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 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 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 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 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 #[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
1404macro_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 fn migrate(&self) -> ();
1426 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 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 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 fn register_branch(&self, gid: &str, branch_id: &str, ops: &[(BranchOp, String)]) -> RegisterOutcome;
1442 fn list_recent(&self, limit: i64) -> Vec<GlobalRow>;
1443}