Skip to main content

xmtp_db/encrypted_store/
notifications.rs

1//! Durable notification configuration and confirmed subscription uploads.
2
3use super::{
4    ConnectionExt, DbConnection,
5    schema::{groups, push_uploaded_topic, user_preferences},
6};
7use crate::{StorageError, impl_fetch, impl_store};
8use diesel::prelude::*;
9use xmtp_proto::types::GroupId;
10
11/// Notification fields on the singleton preferences row.
12/// Keep credentials out of debug output.
13#[derive(Clone, Default, Queryable, Selectable, AsChangeset)]
14#[diesel(table_name = user_preferences, treat_none_as_null = true)]
15pub struct StoredNotification {
16    pub push_recipient_id: Option<Vec<u8>>,
17    pub push_recipient_secret: Option<Vec<u8>>,
18    pub push_state: i32,
19    pub push_failed_error: Option<Vec<u8>>,
20    pub push_config: Option<Vec<u8>>,
21    pub push_deadlines: Option<Vec<u8>>,
22    pub push_last_state: Option<Vec<u8>>,
23    pub push_repairing: bool,
24    pub push_generation: i64,
25    pub push_suppressed: Option<Vec<u8>>,
26}
27
28/// One confirmed topic upload. Stale rows belong to an active repair pass.
29#[derive(Clone, Debug, PartialEq, Eq, Queryable, Selectable, Insertable, Identifiable)]
30#[diesel(table_name = push_uploaded_topic, primary_key(topic))]
31pub struct UploadedTopic {
32    pub topic: Vec<u8>,
33    pub hmac_epoch_base: Option<i64>,
34    pub include_commits: bool,
35    pub root_key_fingerprint: Vec<u8>,
36    pub stale: bool,
37}
38
39impl_fetch!(UploadedTopic, push_uploaded_topic, Vec<u8>);
40impl_store!(UploadedTopic, push_uploaded_topic);
41
42pub trait QueryNotifications {
43    fn notification_record(&self) -> Result<StoredNotification, StorageError>;
44    fn save_notification_record(&self, record: &StoredNotification) -> Result<(), StorageError>;
45    /// Atomically disable notifications and return the cleared uploads.
46    /// Keep the recipient identity, conversation overrides, and task retry state.
47    fn disable_notifications(
48        &self,
49    ) -> Result<(StoredNotification, Vec<UploadedTopic>), StorageError>;
50    fn uploaded_topics(&self) -> Result<Vec<UploadedTopic>, StorageError>;
51    fn confirm_uploaded_topics(
52        &self,
53        adds: &[UploadedTopic],
54        removes: &[Vec<u8>],
55    ) -> Result<(), StorageError>;
56    fn clear_uploaded_topics(&self) -> Result<(), StorageError>;
57    fn mark_uploaded_topics_stale(&self) -> Result<(), StorageError>;
58    fn notification_groups(&self) -> Result<Vec<super::group::StoredGroup>, StorageError>;
59    fn set_notification_override(
60        &self,
61        group: &GroupId,
62        value: Option<i32>,
63    ) -> Result<(), StorageError>;
64}
65
66impl<T: QueryNotifications + ?Sized> QueryNotifications for &T {
67    fn notification_record(&self) -> Result<StoredNotification, StorageError> {
68        (**self).notification_record()
69    }
70    fn save_notification_record(&self, record: &StoredNotification) -> Result<(), StorageError> {
71        (**self).save_notification_record(record)
72    }
73    fn uploaded_topics(&self) -> Result<Vec<UploadedTopic>, StorageError> {
74        (**self).uploaded_topics()
75    }
76    fn disable_notifications(
77        &self,
78    ) -> Result<(StoredNotification, Vec<UploadedTopic>), StorageError> {
79        (**self).disable_notifications()
80    }
81    fn confirm_uploaded_topics(
82        &self,
83        adds: &[UploadedTopic],
84        removes: &[Vec<u8>],
85    ) -> Result<(), StorageError> {
86        (**self).confirm_uploaded_topics(adds, removes)
87    }
88    fn clear_uploaded_topics(&self) -> Result<(), StorageError> {
89        (**self).clear_uploaded_topics()
90    }
91    fn mark_uploaded_topics_stale(&self) -> Result<(), StorageError> {
92        (**self).mark_uploaded_topics_stale()
93    }
94    fn notification_groups(&self) -> Result<Vec<super::group::StoredGroup>, StorageError> {
95        (**self).notification_groups()
96    }
97    fn set_notification_override(
98        &self,
99        group: &GroupId,
100        value: Option<i32>,
101    ) -> Result<(), StorageError> {
102        (**self).set_notification_override(group, value)
103    }
104}
105
106impl<C: ConnectionExt> QueryNotifications for DbConnection<C> {
107    #[xmtp_common::db_span]
108    fn notification_record(&self) -> Result<StoredNotification, StorageError> {
109        Ok(self.raw_query(|conn| {
110            user_preferences::table
111                .select(StoredNotification::as_select())
112                .first(conn)
113        })?)
114    }
115
116    #[xmtp_common::db_span]
117    fn save_notification_record(&self, record: &StoredNotification) -> Result<(), StorageError> {
118        self.raw_query(|conn| {
119            diesel::update(user_preferences::table)
120                .set(record)
121                .execute(conn)
122        })?;
123        Ok(())
124    }
125
126    #[xmtp_common::db_span]
127    fn uploaded_topics(&self) -> Result<Vec<UploadedTopic>, StorageError> {
128        Ok(self.raw_query(|conn| {
129            push_uploaded_topic::table
130                .order(push_uploaded_topic::topic.asc())
131                .load(conn)
132        })?)
133    }
134
135    #[xmtp_common::db_span]
136    fn disable_notifications(
137        &self,
138    ) -> Result<(StoredNotification, Vec<UploadedTopic>), StorageError> {
139        self.raw_query(|conn| {
140            Ok(conn.transaction::<_, StorageError, _>(|conn| {
141                let mut record = user_preferences::table
142                    .select(StoredNotification::as_select())
143                    .first(conn)?;
144                record.push_generation = record
145                    .push_generation
146                    .checked_add(1)
147                    .ok_or(StorageError::DbSerialize)?;
148                record.push_state = 0;
149                record.push_config = None;
150                record.push_failed_error = None;
151                record.push_deadlines = None;
152                record.push_last_state = None;
153                record.push_repairing = false;
154                record.push_suppressed = None;
155                let cleared = push_uploaded_topic::table
156                    .order(push_uploaded_topic::topic.asc())
157                    .load(conn)?;
158                diesel::delete(push_uploaded_topic::table).execute(conn)?;
159                diesel::update(user_preferences::table)
160                    .set(&record)
161                    .execute(conn)?;
162                Ok((record, cleared))
163            }))
164        })?
165    }
166
167    #[xmtp_common::db_span]
168    fn confirm_uploaded_topics(
169        &self,
170        adds: &[UploadedTopic],
171        removes: &[Vec<u8>],
172    ) -> Result<(), StorageError> {
173        self.raw_query(|conn| {
174            diesel::delete(
175                push_uploaded_topic::table.filter(push_uploaded_topic::topic.eq_any(removes)),
176            )
177            .execute(conn)?;
178            for row in adds {
179                diesel::replace_into(push_uploaded_topic::table)
180                    .values(row)
181                    .execute(conn)?;
182            }
183            Ok(())
184        })?;
185        Ok(())
186    }
187
188    #[xmtp_common::db_span]
189    fn clear_uploaded_topics(&self) -> Result<(), StorageError> {
190        self.raw_query(|conn| diesel::delete(push_uploaded_topic::table).execute(conn))?;
191        Ok(())
192    }
193
194    #[xmtp_common::db_span]
195    fn mark_uploaded_topics_stale(&self) -> Result<(), StorageError> {
196        self.raw_query(|conn| {
197            diesel::update(push_uploaded_topic::table)
198                .set(push_uploaded_topic::stale.eq(true))
199                .execute(conn)
200        })?;
201        Ok(())
202    }
203
204    #[xmtp_common::db_span]
205    fn notification_groups(&self) -> Result<Vec<super::group::StoredGroup>, StorageError> {
206        Ok(self.raw_query(|conn| {
207            groups::table
208                .select(super::group::StoredGroup::as_select())
209                .order(groups::id.asc())
210                .load(conn)
211        })?)
212    }
213
214    #[xmtp_common::db_span]
215    fn set_notification_override(
216        &self,
217        group: &GroupId,
218        value: Option<i32>,
219    ) -> Result<(), StorageError> {
220        self.raw_query(|conn| {
221            diesel::update(groups::table.find(group))
222                .set(groups::push_override.eq(value))
223                .execute(conn)
224        })?;
225        Ok(())
226    }
227}
228
229#[cfg(test)]
230mod tests {
231    use super::*;
232    use crate::{Store, TestDb, XmtpTestDb, group::tests::generate_group};
233    use diesel::connection::SimpleConnection;
234
235    #[xmtp_common::test(unwrap_try = true)]
236    async fn notification_disable_clears_state_and_keeps_identity_and_overrides() {
237        let store = TestDb::create_persistent_store(None).await;
238        let db = store.db();
239        let group = generate_group(None);
240        group.store(&db)?;
241        db.set_notification_override(&group.id, Some(0))?;
242        let before = StoredNotification {
243            push_recipient_id: Some(vec![1; 32]),
244            push_recipient_secret: Some(vec![2; 32]),
245            push_state: 2,
246            push_failed_error: Some(vec![3]),
247            push_config: Some(vec![4]),
248            push_deadlines: Some(vec![5]),
249            push_last_state: Some(vec![6]),
250            push_repairing: true,
251            push_generation: 9,
252            push_suppressed: Some(vec![7]),
253        };
254        db.save_notification_record(&before)?;
255        let uploaded = UploadedTopic {
256            topic: vec![8],
257            hmac_epoch_base: Some(42),
258            include_commits: true,
259            root_key_fingerprint: vec![9],
260            stale: true,
261        };
262        db.confirm_uploaded_topics(std::slice::from_ref(&uploaded), &[])?;
263        let (disabled, cleared) = db.disable_notifications()?;
264        assert_eq!(cleared, vec![uploaded]);
265        assert!(db.uploaded_topics()?.is_empty());
266        assert_eq!(disabled.push_generation, 10);
267        let disabled = db.notification_record()?;
268        assert_eq!(disabled.push_generation, 10);
269        assert_eq!(disabled.push_state, 0);
270        assert_eq!(disabled.push_recipient_id, before.push_recipient_id);
271        assert_eq!(disabled.push_recipient_secret, before.push_recipient_secret);
272        assert!(disabled.push_config.is_none());
273        assert!(disabled.push_failed_error.is_none());
274        assert!(disabled.push_deadlines.is_none());
275        assert!(disabled.push_last_state.is_none());
276        assert!(disabled.push_suppressed.is_none());
277        assert!(!disabled.push_repairing);
278        assert_eq!(db.notification_groups()?[0].push_override, Some(0));
279    }
280
281    #[xmtp_common::test(unwrap_try = true)]
282    async fn notification_disable_rolls_back_cleared_uploads_if_state_write_fails() {
283        let store = TestDb::create_persistent_store(None).await;
284        let db = store.db();
285        let before = StoredNotification {
286            push_state: 1,
287            push_config: Some(vec![1]),
288            push_generation: 9,
289            ..Default::default()
290        };
291        db.save_notification_record(&before)?;
292        let uploaded = UploadedTopic {
293            topic: vec![8],
294            hmac_epoch_base: None,
295            include_commits: false,
296            root_key_fingerprint: vec![9],
297            stale: false,
298        };
299        db.confirm_uploaded_topics(std::slice::from_ref(&uploaded), &[])?;
300        db.raw_query(|conn| {
301            conn.batch_execute(
302                "CREATE TRIGGER reject_notification_disable BEFORE UPDATE ON user_preferences
303                 BEGIN SELECT RAISE(ABORT, 'test state write failure'); END;",
304            )
305        })?;
306        assert!(db.disable_notifications().is_err());
307        assert_eq!(db.uploaded_topics()?, vec![uploaded]);
308        let after = db.notification_record()?;
309        assert_eq!(after.push_state, before.push_state);
310        assert_eq!(after.push_config, before.push_config);
311        assert_eq!(after.push_generation, before.push_generation);
312    }
313}