1use 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#[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#[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 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}