Skip to main content

xmtp_db/encrypted_store/
key_package_history.rs

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    /// Highest confirmed publication receipt. Unknown publication never retires another key.
27    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    /// Retire keys by confirmed publication order, not local creation order.
50    /// The latest published key stays usable; duplicate receipts keep the first retirement deadline.
51    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    /// Soonest pending `delete_at_ns` across all key packages marked for deletion,
60    /// or `None` if none are marked. The KpDeletion task's reschedule source.
61    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) // convert ConnectionError into StorageError
210    }
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            // Aggregate MIN over an empty/unmarked table is NULL -> None.
336            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            // Now delete it
359            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}