Skip to main content

xmtp_db/encrypted_store/
user_preferences.rs

1use super::{
2    ConnectionExt,
3    schema::user_preferences::{self, dsl},
4};
5use crate::{StorageError, Store};
6use diesel::{insert_into, prelude::*};
7use xmtp_common::time::now_ns;
8
9#[derive(
10    Identifiable,
11    Insertable,
12    Queryable,
13    Selectable,
14    AsChangeset,
15    Debug,
16    Clone,
17    PartialEq,
18    Eq,
19    Default,
20)]
21#[diesel(table_name = user_preferences)]
22#[diesel(primary_key(id))]
23pub struct StoredUserPreferences {
24    pub id: i32,
25    /// HMAC key root
26    pub hmac_key: Option<Vec<u8>>,
27    pub hmac_key_cycled_at_ns: Option<i64>,
28    /// Whether DM group updates have been migrated.
29    pub dm_group_updates_migrated: bool,
30}
31
32impl<C> Store<C> for StoredUserPreferences
33where
34    C: ConnectionExt,
35{
36    type Output = ();
37    fn store(&self, conn: &C) -> Result<Self::Output, StorageError> {
38        conn.raw_query(|conn| {
39            diesel::update(dsl::user_preferences)
40                .set(self)
41                .execute(conn)
42        })?;
43
44        Ok(())
45    }
46}
47
48#[derive(Debug)]
49pub struct HmacKey {
50    // TODO: Use xmtp_cryptography::Secret for Zeroize support
51    pub key: [u8; 42],
52    // # of 30 day periods since unix epoch
53    pub epoch: i64,
54}
55
56impl HmacKey {
57    pub fn random_key() -> Vec<u8> {
58        xmtp_common::rand_vec::<42>()
59    }
60}
61
62impl StoredUserPreferences {
63    pub fn load(conn: impl ConnectionExt) -> Result<Self, StorageError> {
64        let pref = conn.raw_query(|conn| {
65            dsl::user_preferences
66                .select(Self::as_select())
67                .first(conn)
68                .optional()
69        })?;
70        Ok(pref.unwrap_or_default())
71    }
72
73    fn store(&self, conn: &impl crate::DbQuery) -> Result<(), StorageError> {
74        conn.raw_query(|conn| {
75            insert_into(dsl::user_preferences)
76                .values(self)
77                .on_conflict(user_preferences::id)
78                .do_update()
79                .set(self)
80                .execute(conn)
81        })?;
82
83        Ok(())
84    }
85
86    pub fn store_hmac_key(
87        conn: &impl crate::DbQuery,
88        key: &[u8],
89        cycled_at: Option<i64>,
90    ) -> Result<(), StorageError> {
91        if key.len() != 42 {
92            return Err(StorageError::InvalidHmacLength);
93        }
94
95        let mut preferences = Self::load(conn)?;
96
97        if let (Some(old), Some(new)) = (preferences.hmac_key_cycled_at_ns, cycled_at)
98            && old > new
99        {
100            return Ok(());
101        }
102
103        preferences.hmac_key = Some(key.to_vec());
104        preferences.hmac_key_cycled_at_ns = Some(cycled_at.unwrap_or_else(now_ns));
105        preferences.store(conn)?;
106
107        Ok(())
108    }
109}
110
111#[cfg(test)]
112mod tests {
113    use super::*;
114
115    #[xmtp_common::test]
116    fn test_insert_and_update_preferences() {
117        crate::test_utils::with_connection(|conn| {
118            let pref = StoredUserPreferences::load(conn).unwrap();
119            // by default, there is no key
120            assert!(pref.hmac_key.is_none());
121
122            // loads and stores a default
123            let pref = StoredUserPreferences::load(conn).unwrap();
124            // by default, there is no key
125            assert!(pref.hmac_key.is_none());
126
127            // set an hmac key
128            let hmac_key = HmacKey::random_key();
129            StoredUserPreferences::store_hmac_key(conn, &hmac_key, None).unwrap();
130            let pref = StoredUserPreferences::load(conn).unwrap();
131            // Make sure it saved
132            assert_eq!(hmac_key, pref.hmac_key.unwrap());
133
134            // check that there is only one preference stored
135            let query = dsl::user_preferences.order(dsl::id.desc());
136            let result = conn
137                .raw_query(|conn| {
138                    query
139                        .select(StoredUserPreferences::as_select())
140                        .load::<StoredUserPreferences>(conn)
141                })
142                .unwrap();
143            assert_eq!(result.len(), 1);
144        })
145    }
146}