xmtp_db/
xmtp_openmls_provider.rs1use 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#[derive(Debug, Clone, PartialEq, Eq)]
23pub enum TransactionOutcome<T> {
24 Continue(T),
26 Rollback,
28}
29
30impl<T> TransactionOutcome<T> {
31 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
44pub trait XmtpMlsStorageProvider:
51 MaybeSend + MaybeSync + StorageProvider<CURRENT_VERSION, Error = SqlKeyStoreError>
52{
53 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 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 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}