Skip to main content

systemprompt_database/services/
transaction.rs

1//! A transaction wrapper over [`PgDbPool`] that retries the whole closure on a
2//! serialization failure, without going through the dyn-safe trait.
3//!
4//! Copyright (c) systemprompt.io — Business Source License 1.1.
5//! See <https://systemprompt.io> for licensing details.
6
7use crate::error::RepositoryError;
8use crate::repository::PgDbPool;
9use crate::resilience::classify::Outcome;
10use crate::resilience::config::RetryConfig;
11use crate::resilience::retry::retry_async;
12use sqlx::{Postgres, Transaction};
13use std::future::Future;
14use std::pin::Pin;
15use std::time::Duration;
16
17pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
18
19pub async fn with_transaction_retry<F, T>(
20    pool: &PgDbPool,
21    max_retries: u32,
22    f: F,
23) -> Result<T, RepositoryError>
24where
25    T: Send,
26    F: for<'c> Fn(&'c mut Transaction<'_, Postgres>) -> BoxFuture<'c, Result<T, RepositoryError>>
27        + Send
28        + Sync,
29{
30    let cfg = RetryConfig {
31        max_attempts: max_retries.saturating_add(1),
32        base_delay: Duration::from_millis(20),
33        max_delay: Duration::from_millis(640),
34        jitter: false,
35    };
36    let classify = |err: &RepositoryError| {
37        if err.is_serialization_failure() {
38            Outcome::Transient { retry_after: None }
39        } else {
40            Outcome::Permanent
41        }
42    };
43    let attempt = || async {
44        let mut tx = pool.begin().await?;
45        match f(&mut tx).await {
46            Ok(result) => {
47                tx.commit().await?;
48                Ok(result)
49            },
50            Err(e) => {
51                if let Err(rollback_err) = tx.rollback().await {
52                    tracing::error!(error = %rollback_err, "Transaction rollback failed");
53                }
54                Err(e)
55            },
56        }
57    };
58    retry_async(&cfg, "transaction", classify, attempt).await
59}