Skip to main content

xmtp_db/
xmtp_openmls_provider.rs

1use crate::ConnectionExt;
2use crate::MlsProviderExt;
3use crate::TransactionalKeyStore;
4use crate::sql_key_store::SqlKeyStoreError;
5use openmls_rust_crypto::RustCrypto;
6use openmls_traits::OpenMlsProvider;
7use openmls_traits::storage::CURRENT_VERSION;
8use openmls_traits::storage::{Entity, StorageProvider};
9use xmtp_common::{MaybeSend, MaybeSync};
10
11/// Outcome of a [`XmtpMlsStorageProvider::transaction`] or
12/// [`XmtpMlsStorageProvider::savepoint`] closure.
13///
14/// Returning `Ok(TransactionOutcome::Continue(value))` persists the transaction
15/// and returns `Ok(value)` to the caller.
16///
17/// Returning `Ok(TransactionOutcome::Rollback)` rolls back the transaction
18/// *without* recording a span error — the rollback was intentional.
19///
20/// Returning `Err(e)` rolls back the transaction *and* records `status=error`
21/// on the enclosing `#[db_span]` / `#[rpc_span]` span — the error was real.
22#[derive(Debug, Clone, PartialEq, Eq)]
23pub enum TransactionOutcome<T> {
24    /// Persist the transaction and return the enclosed value.
25    Continue(T),
26    /// Roll back the transaction without treating it as an error.
27    Rollback,
28}
29
30impl<T> TransactionOutcome<T> {
31    /// Unwrap the persisted value for call sites that never roll back.
32    ///
33    /// Panics if this is a `Rollback` (a bug at that call site).
34    pub fn into_continued(self) -> T {
35        match self {
36            TransactionOutcome::Continue(v) => v,
37            TransactionOutcome::Rollback => {
38                unreachable!("transaction caller never returns TransactionOutcome::Rollback")
39            }
40        }
41    }
42}
43
44/// Convenience super trait to constrain the storage provider to a
45/// specific error type and version
46/// This storage provider is likewise implemented on both &T and T references,
47/// to allow creating a referenced or owned provider.
48// constraining the error type here will avoid leaking
49// the associated type parameter, so we don't need to define it on every function.
50pub trait XmtpMlsStorageProvider:
51    MaybeSend + MaybeSync + StorageProvider<CURRENT_VERSION, Error = SqlKeyStoreError>
52{
53    /// An Opaque Database connection type. Can be anything.
54    type Connection: ConnectionExt;
55
56    type TxQuery: TransactionalKeyStore;
57
58    type DbQuery<'a>: crate::DbQuery
59    where
60        Self::Connection: 'a;
61
62    fn db<'a>(&'a self) -> Self::DbQuery<'a>;
63
64    /// Start a new transaction.
65    ///
66    /// The closure returns `Ok(TransactionOutcome::Continue(v))` to persist or
67    /// `Ok(TransactionOutcome::Rollback)` to roll back without an error.
68    /// Returning `Err(e)` also rolls back and propagates `e` as a real error.
69    fn transaction<T, E, F>(&self, f: F) -> Result<TransactionOutcome<T>, E>
70    where
71        F: FnOnce(&mut Self::TxQuery) -> Result<TransactionOutcome<T>, E>,
72        E: From<diesel::result::Error> + From<crate::ConnectionError> + std::error::Error;
73
74    /// Start a savepoint within a transaction.
75    ///
76    /// Must only be used when already in a transaction.
77    // TODO: enforce that this is only used within transactions
78    // otherwise we run into sqlite race conditions b/c this does not
79    // use BEGIN IMMEDIATE.
80    // we can ensure this by checking sqlite transaction depth.
81    fn savepoint<T, E, F>(&self, f: F) -> Result<TransactionOutcome<T>, E>
82    where
83        F: FnOnce(&mut Self::TxQuery) -> Result<TransactionOutcome<T>, E>,
84        E: From<diesel::result::Error> + From<crate::ConnectionError> + std::error::Error;
85
86    fn _disable_lint_for_self<'a>(_: Self::DbQuery<'a>) {}
87
88    fn read<V: Entity<CURRENT_VERSION>>(
89        &self,
90        label: &[u8],
91        key: &[u8],
92    ) -> Result<Option<V>, SqlKeyStoreError>;
93
94    fn read_list<V: Entity<CURRENT_VERSION>>(
95        &self,
96        label: &[u8],
97        key: &[u8],
98    ) -> Result<Vec<V>, <Self as StorageProvider<CURRENT_VERSION>>::Error>;
99
100    fn delete(
101        &self,
102        label: &[u8],
103        key: &[u8],
104    ) -> Result<(), <Self as StorageProvider<CURRENT_VERSION>>::Error>;
105
106    fn write(
107        &self,
108        label: &[u8],
109        key: &[u8],
110        value: &[u8],
111    ) -> Result<(), <Self as StorageProvider<CURRENT_VERSION>>::Error>;
112
113    #[cfg(feature = "test-utils")]
114    fn hash_all(&self) -> Result<Vec<u8>, SqlKeyStoreError>;
115}
116
117pub struct XmtpOpenMlsProvider<S> {
118    crypto: RustCrypto,
119    mls_storage: S,
120}
121
122impl<S> XmtpOpenMlsProvider<S> {
123    pub fn new(mls_storage: S) -> Self {
124        Self {
125            crypto: RustCrypto::default(),
126            mls_storage,
127        }
128    }
129}
130
131impl<S> XmtpOpenMlsProvider<S> {
132    pub fn new_crypto() -> RustCrypto {
133        RustCrypto::default()
134    }
135}
136
137impl<S> MlsProviderExt for XmtpOpenMlsProvider<S>
138where
139    S: XmtpMlsStorageProvider,
140{
141    type XmtpStorage = S;
142
143    fn key_store(&self) -> &Self::XmtpStorage {
144        &self.mls_storage
145    }
146}
147
148impl<S> OpenMlsProvider for XmtpOpenMlsProvider<S>
149where
150    S: XmtpMlsStorageProvider,
151{
152    type CryptoProvider = RustCrypto;
153    type RandProvider = RustCrypto;
154    type StorageProvider = S;
155    fn crypto(&self) -> &Self::CryptoProvider {
156        &self.crypto
157    }
158
159    fn rand(&self) -> &Self::RandProvider {
160        &self.crypto
161    }
162
163    fn storage(&self) -> &Self::StorageProvider {
164        &self.mls_storage
165    }
166}
167
168pub struct XmtpOpenMlsProviderRef<'a, S> {
169    crypto: RustCrypto,
170    mls_storage: &'a S,
171}
172
173impl<'a, S> MlsProviderExt for XmtpOpenMlsProviderRef<'a, S>
174where
175    S: XmtpMlsStorageProvider,
176{
177    type XmtpStorage = S;
178
179    fn key_store(&self) -> &Self::XmtpStorage {
180        self.mls_storage
181    }
182}
183
184impl<'a, S> OpenMlsProvider for XmtpOpenMlsProviderRef<'a, S>
185where
186    S: XmtpMlsStorageProvider,
187{
188    type CryptoProvider = RustCrypto;
189    type RandProvider = RustCrypto;
190    type StorageProvider = S;
191    fn crypto(&self) -> &Self::CryptoProvider {
192        &self.crypto
193    }
194
195    fn rand(&self) -> &Self::RandProvider {
196        &self.crypto
197    }
198
199    fn storage(&self) -> &Self::StorageProvider {
200        self.mls_storage
201    }
202}
203
204impl<'a, S> XmtpOpenMlsProviderRef<'a, S> {
205    pub fn new(mls_storage: &'a S) -> Self {
206        Self {
207            crypto: RustCrypto::default(),
208            mls_storage,
209        }
210    }
211}
212
213pub struct XmtpOpenMlsProviderRefMut<'a, S> {
214    crypto: RustCrypto,
215    mls_storage: &'a mut S,
216}
217
218impl<'a, S> XmtpOpenMlsProviderRefMut<'a, S> {
219    pub fn new(mls_storage: &'a mut S) -> Self {
220        Self {
221            crypto: RustCrypto::default(),
222            mls_storage,
223        }
224    }
225}
226
227impl<'a, S> MlsProviderExt for XmtpOpenMlsProviderRefMut<'a, S>
228where
229    S: XmtpMlsStorageProvider,
230{
231    type XmtpStorage = S;
232
233    fn key_store(&self) -> &Self::XmtpStorage {
234        self.mls_storage
235    }
236}
237
238impl<'a, S> OpenMlsProvider for XmtpOpenMlsProviderRefMut<'a, S>
239where
240    S: XmtpMlsStorageProvider,
241{
242    type CryptoProvider = RustCrypto;
243    type RandProvider = RustCrypto;
244    type StorageProvider = S;
245    fn crypto(&self) -> &Self::CryptoProvider {
246        &self.crypto
247    }
248
249    fn rand(&self) -> &Self::RandProvider {
250        &self.crypto
251    }
252
253    fn storage(&self) -> &Self::StorageProvider {
254        self.mls_storage
255    }
256}