1use super::{
2 ConnectionExt, StorageError, db_connection::DbConnection, schema::key_package_history,
3};
4use crate::{StoreOrIgnore, impl_store_or_ignore};
5use diesel::prelude::*;
6use xmtp_common::time::now_ns;
7use xmtp_configuration::KEYS_EXPIRATION_INTERVAL_NS;
8use xmtp_proto::types::Cursor;
9
10#[derive(Insertable, Debug, Clone)]
11#[diesel(table_name = key_package_history)]
12pub struct NewKeyPackageHistoryEntry {
13 pub key_package_hash_ref: Vec<u8>,
14 pub post_quantum_public_key: Option<Vec<u8>>,
15 pub created_at_ns: i64,
16}
17
18#[derive(Queryable, Selectable, Debug, Clone)]
19#[diesel(table_name = key_package_history)]
20pub struct StoredKeyPackageHistoryEntry {
21 pub id: i32,
22 pub key_package_hash_ref: Vec<u8>,
23 pub created_at_ns: i64,
24 pub delete_at_ns: Option<i64>,
25 pub post_quantum_public_key: Option<Vec<u8>>,
26 pub published_sequence_id: Option<i64>,
28}
29
30impl_store_or_ignore!(NewKeyPackageHistoryEntry, key_package_history);
31
32pub trait QueryKeyPackageHistory {
33 fn store_key_package_history_entry(
34 &self,
35 key_package_hash_ref: Vec<u8>,
36 post_quantum_public_key: Option<Vec<u8>>,
37 ) -> Result<StoredKeyPackageHistoryEntry, StorageError>;
38
39 fn find_key_package_history_entry_by_hash_ref(
40 &self,
41 hash_ref: Vec<u8>,
42 ) -> Result<StoredKeyPackageHistoryEntry, StorageError>;
43
44 fn find_key_package_history_entries_before_id(
45 &self,
46 id: i32,
47 ) -> Result<Vec<StoredKeyPackageHistoryEntry>, StorageError>;
48
49 fn record_key_package_publication(
52 &self,
53 history_id: i32,
54 sequence: Cursor,
55 ) -> Result<(), StorageError>;
56
57 fn get_expired_key_packages(&self) -> Result<Vec<StoredKeyPackageHistoryEntry>, StorageError>;
58
59 fn min_key_package_delete_at_ns(&self) -> Result<Option<i64>, StorageError>;
62
63 fn delete_key_package_entry_with_id(&self, id: i32) -> Result<(), StorageError>;
64}
65
66impl<T> QueryKeyPackageHistory for &T
67where
68 T: QueryKeyPackageHistory,
69{
70 fn store_key_package_history_entry(
71 &self,
72 key_package_hash_ref: Vec<u8>,
73 post_quantum_public_key: Option<Vec<u8>>,
74 ) -> Result<StoredKeyPackageHistoryEntry, StorageError> {
75 (**self).store_key_package_history_entry(key_package_hash_ref, post_quantum_public_key)
76 }
77
78 fn find_key_package_history_entry_by_hash_ref(
79 &self,
80 hash_ref: Vec<u8>,
81 ) -> Result<StoredKeyPackageHistoryEntry, StorageError> {
82 (**self).find_key_package_history_entry_by_hash_ref(hash_ref)
83 }
84
85 fn find_key_package_history_entries_before_id(
86 &self,
87 id: i32,
88 ) -> Result<Vec<StoredKeyPackageHistoryEntry>, StorageError> {
89 (**self).find_key_package_history_entries_before_id(id)
90 }
91
92 fn record_key_package_publication(
93 &self,
94 history_id: i32,
95 sequence: Cursor,
96 ) -> Result<(), StorageError> {
97 (**self).record_key_package_publication(history_id, sequence)
98 }
99
100 fn get_expired_key_packages(&self) -> Result<Vec<StoredKeyPackageHistoryEntry>, StorageError> {
101 (**self).get_expired_key_packages()
102 }
103
104 fn min_key_package_delete_at_ns(&self) -> Result<Option<i64>, StorageError> {
105 (**self).min_key_package_delete_at_ns()
106 }
107
108 fn delete_key_package_entry_with_id(&self, id: i32) -> Result<(), StorageError> {
109 (**self).delete_key_package_entry_with_id(id)
110 }
111}
112
113impl<C: ConnectionExt> QueryKeyPackageHistory for DbConnection<C> {
114 fn record_key_package_publication(
115 &self,
116 history_id: i32,
117 sequence: Cursor,
118 ) -> Result<(), StorageError> {
119 use crate::schema::key_package_history::dsl;
120 let sequence = i64::try_from(sequence.0)
121 .ok()
122 .filter(|sequence| *sequence > 0)
123 .ok_or(crate::stream_storage::StreamStorageError::InvalidBatch)?;
124 super::stream_storage::stream_transaction(self, |conn| {
125 let previous = dsl::key_package_history
126 .find(history_id)
127 .select(dsl::published_sequence_id)
128 .first::<Option<i64>>(conn)
129 .optional()?
130 .ok_or(crate::NotFound::KeyPackageHistory(history_id))?;
131 diesel::update(dsl::key_package_history.find(history_id))
132 .set(dsl::published_sequence_id.eq(previous.unwrap_or(0).max(sequence)))
133 .execute(conn)?;
134 let latest = dsl::key_package_history
135 .select(diesel::dsl::max(dsl::published_sequence_id))
136 .first::<Option<i64>>(conn)?
137 .ok_or(StorageError::DbDeserialize)?;
138 let delete_at = now_ns()
139 .checked_add(KEYS_EXPIRATION_INTERVAL_NS)
140 .ok_or(StorageError::DbSerialize)?;
141 diesel::update(
142 dsl::key_package_history
143 .filter(dsl::published_sequence_id.lt(latest))
144 .filter(dsl::delete_at_ns.is_null()),
145 )
146 .set(dsl::delete_at_ns.eq(delete_at))
147 .execute(conn)?;
148 diesel::update(
149 dsl::key_package_history.filter(
150 dsl::published_sequence_id
151 .eq(latest)
152 .or(dsl::published_sequence_id.is_null()),
153 ),
154 )
155 .set(dsl::delete_at_ns.eq(None::<i64>))
156 .execute(conn)?;
157 Ok(())
158 })
159 }
160
161 fn store_key_package_history_entry(
162 &self,
163 key_package_hash_ref: Vec<u8>,
164 post_quantum_public_key: Option<Vec<u8>>,
165 ) -> Result<StoredKeyPackageHistoryEntry, StorageError> {
166 let entry = NewKeyPackageHistoryEntry {
167 key_package_hash_ref: key_package_hash_ref.clone(),
168 post_quantum_public_key: post_quantum_public_key.clone(),
169 created_at_ns: now_ns(),
170 };
171 entry.store_or_ignore(self)?;
172
173 self.find_key_package_history_entry_by_hash_ref(key_package_hash_ref)
174 }
175
176 fn find_key_package_history_entry_by_hash_ref(
177 &self,
178 hash_ref: Vec<u8>,
179 ) -> Result<StoredKeyPackageHistoryEntry, StorageError> {
180 let result = self.raw_query(|conn| {
181 key_package_history::dsl::key_package_history
182 .filter(key_package_history::dsl::key_package_hash_ref.eq(hash_ref))
183 .first::<StoredKeyPackageHistoryEntry>(conn)
184 })?;
185
186 Ok(result)
187 }
188
189 fn find_key_package_history_entries_before_id(
190 &self,
191 id: i32,
192 ) -> Result<Vec<StoredKeyPackageHistoryEntry>, StorageError> {
193 let result = self.raw_query(|conn| {
194 key_package_history::dsl::key_package_history
195 .filter(key_package_history::dsl::id.lt(id))
196 .load::<StoredKeyPackageHistoryEntry>(conn)
197 })?;
198
199 Ok(result)
200 }
201
202 fn get_expired_key_packages(&self) -> Result<Vec<StoredKeyPackageHistoryEntry>, StorageError> {
203 use crate::schema::key_package_history::dsl;
204 self.raw_query(|conn| {
205 dsl::key_package_history
206 .filter(dsl::delete_at_ns.le(now_ns()))
207 .load::<StoredKeyPackageHistoryEntry>(conn)
208 })
209 .map_err(StorageError::from) }
211
212 fn min_key_package_delete_at_ns(&self) -> Result<Option<i64>, StorageError> {
213 use crate::schema::key_package_history::dsl;
214 use diesel::dsl::min;
215 let v: Option<i64> = self.raw_query(|conn| {
216 dsl::key_package_history
217 .filter(dsl::delete_at_ns.is_not_null())
218 .select(min(dsl::delete_at_ns))
219 .first::<Option<i64>>(conn)
220 })?;
221 Ok(v)
222 }
223
224 fn delete_key_package_entry_with_id(&self, id: i32) -> Result<(), StorageError> {
225 self.raw_query(|conn| {
226 diesel::delete(
227 key_package_history::dsl::key_package_history
228 .filter(key_package_history::dsl::id.eq(id)),
229 )
230 .execute(conn)
231 })?;
232
233 Ok(())
234 }
235}
236
237#[cfg(test)]
238mod tests {
239 #[xmtp_common::test(unwrap_try = true)]
240 async fn duplicate_publication_preserves_the_first_retirement_deadline() {
241 use crate::{TestDb, XmtpTestDb};
242 use xmtp_proto::types::Cursor;
243 let store = TestDb::create_persistent_store(None).await;
244 let db = store.db();
245 let old = db.store_key_package_history_entry(vec![1], None)?;
246 let latest = db.store_key_package_history_entry(vec![2], None)?;
247 db.record_key_package_publication(old.id, Cursor(10))?;
248 db.record_key_package_publication(latest.id, Cursor(20))?;
249 let retired = db.find_key_package_history_entry_by_hash_ref(vec![1])?;
250 assert!(retired.delete_at_ns.is_some());
251 db.record_key_package_publication(old.id, Cursor(5))?;
252 db.record_key_package_publication(latest.id, Cursor(20))?;
253 let repeated = db.find_key_package_history_entry_by_hash_ref(vec![1])?;
254 assert_eq!(repeated.published_sequence_id, Some(10));
255 assert_eq!(repeated.delete_at_ns, retired.delete_at_ns);
256 assert_eq!(db.min_key_package_delete_at_ns()?, retired.delete_at_ns);
257 assert!(db.get_expired_key_packages()?.is_empty());
258 assert!(
259 db.find_key_package_history_entry_by_hash_ref(vec![2])?
260 .delete_at_ns
261 .is_none()
262 );
263 }
264
265 #[xmtp_common::test(unwrap_try = true)]
266 async fn failed_publication_receipt_rolls_back_retirement() {
267 use crate::{ConnectionExt, TestDb, XmtpTestDb};
268 use diesel::connection::SimpleConnection;
269 use xmtp_proto::types::Cursor;
270 let store = TestDb::create_persistent_store(None).await;
271 let db = store.db();
272 let old = db.store_key_package_history_entry(vec![1], None)?;
273 let next = db.store_key_package_history_entry(vec![2], None)?;
274 db.record_key_package_publication(old.id, Cursor(10))?;
275 db.raw_query(|conn| conn.batch_execute("CREATE TEMP TRIGGER fail_key_retirement BEFORE UPDATE OF delete_at_ns ON key_package_history WHEN NEW.delete_at_ns IS NOT NULL BEGIN SELECT RAISE(ABORT, 'injected retirement failure'); END;"))?;
276 assert!(
277 db.record_key_package_publication(next.id, Cursor(20))
278 .is_err()
279 );
280 let old = db.find_key_package_history_entry_by_hash_ref(vec![1])?;
281 let next = db.find_key_package_history_entry_by_hash_ref(vec![2])?;
282 assert_eq!(old.published_sequence_id, Some(10));
283 assert!(old.delete_at_ns.is_none());
284 assert!(next.published_sequence_id.is_none());
285 assert!(next.delete_at_ns.is_none());
286 }
287
288 #[xmtp_common::test(unwrap_try = true)]
289 async fn publication_order_preserves_the_last_advertised_key() {
290 use crate::{TestDb, XmtpTestDb};
291 use xmtp_proto::types::Cursor;
292 let store = TestDb::create_persistent_store(None).await;
293 let db = store.db();
294 let first = db.store_key_package_history_entry(vec![1], None)?;
295 let second = db.store_key_package_history_entry(vec![2], None)?;
296 let unknown = db.store_key_package_history_entry(vec![3], None)?;
297 db.record_key_package_publication(second.id, Cursor(10))?;
298 db.record_key_package_publication(first.id, Cursor(20))?;
299 db.record_key_package_publication(second.id, Cursor(10))?;
300 assert!(
301 db.find_key_package_history_entry_by_hash_ref(first.key_package_hash_ref.clone())?
302 .delete_at_ns
303 .is_none()
304 );
305 assert!(
306 db.find_key_package_history_entry_by_hash_ref(second.key_package_hash_ref.clone())?
307 .delete_at_ns
308 .is_some()
309 );
310 assert!(
311 db.find_key_package_history_entry_by_hash_ref(unknown.key_package_hash_ref)?
312 .delete_at_ns
313 .is_none()
314 );
315 db.record_key_package_publication(second.id, Cursor(30))?;
316 assert!(
317 db.find_key_package_history_entry_by_hash_ref(first.key_package_hash_ref)?
318 .delete_at_ns
319 .is_some()
320 );
321 assert!(
322 db.find_key_package_history_entry_by_hash_ref(second.key_package_hash_ref)?
323 .delete_at_ns
324 .is_none()
325 );
326 }
327
328 use crate::prelude::*;
329 use crate::test_utils::with_connection;
330 use xmtp_common::rand_vec;
331
332 #[xmtp_common::test]
333 fn min_key_package_delete_at_ns_none_when_empty() {
334 with_connection(|conn| {
335 assert_eq!(conn.min_key_package_delete_at_ns().unwrap(), None);
337 })
338 }
339
340 #[xmtp_common::test]
341 fn test_store_key_package_history_entry() {
342 with_connection(|conn| {
343 let hash_ref = rand_vec::<24>();
344 let post_quantum_public_key = rand_vec::<32>();
345 let new_entry = conn
346 .store_key_package_history_entry(
347 hash_ref.clone(),
348 Some(post_quantum_public_key.clone()),
349 )
350 .unwrap();
351 assert_eq!(new_entry.key_package_hash_ref, hash_ref);
352 assert_eq!(
353 new_entry.post_quantum_public_key,
354 Some(post_quantum_public_key)
355 );
356 assert_eq!(new_entry.id, 1);
357
358 conn.delete_key_package_entry_with_id(1).unwrap();
360 let all_entries = conn
361 .find_key_package_history_entries_before_id(100)
362 .unwrap();
363 assert!(all_entries.is_empty());
364 })
365 }
366
367 #[xmtp_common::test]
368 fn test_store_multiple() {
369 with_connection(|conn| {
370 let post_quantum_public_key = rand_vec::<32>();
371 let hash_ref1 = rand_vec::<24>();
372 let hash_ref2 = rand_vec::<24>();
373 let hash_ref3 = rand_vec::<24>();
374
375 conn.store_key_package_history_entry(
376 hash_ref1.clone(),
377 Some(post_quantum_public_key.clone()),
378 )
379 .unwrap();
380 conn.store_key_package_history_entry(
381 hash_ref2.clone(),
382 Some(post_quantum_public_key.clone()),
383 )
384 .unwrap();
385 let entry_3 = conn
386 .store_key_package_history_entry(hash_ref3.clone(), None)
387 .unwrap();
388
389 let all_entries = conn
390 .find_key_package_history_entries_before_id(100)
391 .unwrap();
392
393 assert_eq!(all_entries.len(), 3);
394
395 let earlier_entries = conn
396 .find_key_package_history_entries_before_id(entry_3.id)
397 .unwrap();
398 assert_eq!(earlier_entries.len(), 2);
399 })
400 }
401}