cratestack_sqlx/
isolated_run.rs1use std::future::Future;
6use std::sync::Arc;
7use std::time::Duration;
8
9use cratestack_core::{CratestackError, DbErrorInfo, TransactionAbort, TransactionIsolation};
10
11use crate::audit::dispatch_audit_sink;
12use crate::bound::BoundTx;
13use crate::descriptor::SqlxRuntime;
14use crate::error::cratestack_error_from_sqlx;
15use crate::retriable::retriable_sqlstate;
16use crate::sqlx;
17use crate::transaction::Tx;
18
19const ISOLATION_CONFLICT_DETAIL: &str =
21 "transaction could not be completed because of concurrent updates; retry the request";
22
23enum Attempt<T> {
24 Committed(T),
25 Retry(&'static str),
26 Failed(CratestackError),
27}
28
29impl SqlxRuntime {
30 pub(crate) fn bound(&self) -> Option<&BoundTx> {
33 self.bound.as_deref()
34 }
35
36 pub fn with_isolation_max_retries(mut self, max_retries: u32) -> Self {
41 self.isolation_max_retries = max_retries;
42 self
43 }
44
45 #[doc(hidden)]
73 pub async fn run_isolated<F, Fut, T>(
74 &self,
75 isolation: TransactionIsolation,
76 mut body: F,
77 ) -> Result<T, CratestackError>
78 where
79 F: FnMut(SqlxRuntime) -> Fut,
80 Fut: Future<Output = Result<T, CratestackError>>,
81 {
82 if let Some(bound) = self.bound() {
83 return crate::bound::join_bound(self, bound, isolation, body).await;
84 }
85 let mut attempt = 0u32;
86 loop {
87 attempt += 1;
88 let begin = format!("BEGIN ISOLATION LEVEL {}", isolation.as_sql());
89 let tx = self
92 .pool()
93 .begin_with(sqlx::AssertSqlSafe(begin))
94 .await
95 .map_err(cratestack_error_from_sqlx)?;
96 let bound = Arc::new(BoundTx::new(Tx::new(tx), isolation));
97 let mut runtime = self.clone();
98 runtime.bound = Some(bound.clone());
99 let result = body(runtime)
100 .await
101 .map_err(CratestackError::propagate_transaction_abort);
102 match finish_attempt(&bound, result).await {
103 Attempt::Committed(value) => {
104 let (events, drain) = bound.take_deferred();
105 if drain {
106 let _ = self.drain_event_outbox().await;
107 }
108 dispatch_audit_sink(self, &events).await;
109 return Ok(value);
110 }
111 Attempt::Retry(sqlstate) if attempt <= self.isolation_max_retries => {
112 tracing::debug!(
113 target: "cratestack",
114 cratestack_isolation = isolation.as_sql(),
115 cratestack_sqlstate = sqlstate,
116 cratestack_attempt = attempt,
117 "retrying @isolation transaction",
118 );
119 backoff(attempt).await;
120 }
121 Attempt::Retry(sqlstate) => {
122 return Err(CratestackError::TransactionAborted(
123 TransactionAbort::exhausted(DbErrorInfo {
124 detail: ISOLATION_CONFLICT_DETAIL.to_owned(),
125 sqlstate: Some(sqlstate.to_owned()),
126 constraint: None,
127 }),
128 ));
129 }
130 Attempt::Failed(error) => return Err(error),
131 }
132 }
133 }
134}
135
136async fn finish_attempt<T>(bound: &BoundTx, result: Result<T, CratestackError>) -> Attempt<T> {
137 let Some(tx) = bound.take().await else {
138 return Attempt::Failed(CratestackError::Internal(
139 "@isolation transaction missing at commit".to_owned(),
140 ));
141 };
142 let tx = tx.into_inner();
143 let taint = bound.tainted();
144 let result = match (result, bound.take_poison()) {
148 (Ok(_), Some(poison)) => Err(poison),
149 (result, _) => result,
150 };
151 match result {
152 Ok(value) if taint.is_none() => match tx.commit().await {
153 Ok(()) => Attempt::Committed(value),
154 Err(error) => {
155 let error = cratestack_error_from_sqlx(error);
156 match retriable_sqlstate(&error) {
157 Some(sqlstate) => Attempt::Retry(sqlstate),
158 None => Attempt::Failed(error),
159 }
160 }
161 },
162 Ok(_) => {
163 let _ = tx.rollback().await;
164 Attempt::Retry(taint.unwrap_or("40001"))
165 }
166 Err(error) => {
167 let _ = tx.rollback().await;
168 match retriable_sqlstate(&error).or(taint) {
169 Some(sqlstate) => Attempt::Retry(sqlstate),
170 None => Attempt::Failed(error),
171 }
172 }
173 }
174}
175
176async fn backoff(retry: u32) {
179 let base_ms = (2u64 << retry.saturating_sub(1).min(5)).min(64);
180 let nanos = std::time::SystemTime::now()
181 .duration_since(std::time::UNIX_EPOCH)
182 .map(|elapsed| u64::from(elapsed.subsec_nanos()))
183 .unwrap_or(0);
184 let jitter = Duration::from_nanos(nanos % (base_ms * 1_000_000 + 1));
185 tokio::time::sleep(Duration::from_millis(base_ms) + jitter).await;
186}