Skip to main content

xmtp_db/encrypted_store/group/
version.rs

1use crate::ConnectionExt;
2
3use super::*;
4
5use xmtp_proto::types::GroupId;
6pub trait QueryGroupVersion {
7    fn set_group_paused(&self, group_id: &GroupId, min_version: &str) -> Result<(), StorageError>;
8
9    fn unpause_group(&self, group_id: &GroupId) -> Result<(), StorageError>;
10
11    fn get_group_paused_version(&self, group_id: &GroupId) -> Result<Option<String>, StorageError>;
12
13    /// Return every group currently flagged as paused, with the
14    /// `paused_for_version` floor it's pinned to. Used by the
15    /// startup/sweep recovery path to re-evaluate paused groups
16    /// against the now-current `pkg_version` without having to sync
17    /// each group individually.
18    fn get_paused_groups_with_versions(&self) -> Result<Vec<(GroupId, String)>, StorageError>;
19}
20
21impl<T> QueryGroupVersion for &T
22where
23    T: QueryGroupVersion,
24{
25    fn set_group_paused(&self, group_id: &GroupId, min_version: &str) -> Result<(), StorageError> {
26        (**self).set_group_paused(group_id, min_version)
27    }
28
29    fn unpause_group(&self, group_id: &GroupId) -> Result<(), StorageError> {
30        (**self).unpause_group(group_id)
31    }
32
33    fn get_group_paused_version(&self, group_id: &GroupId) -> Result<Option<String>, StorageError> {
34        (**self).get_group_paused_version(group_id)
35    }
36
37    fn get_paused_groups_with_versions(&self) -> Result<Vec<(GroupId, String)>, StorageError> {
38        (**self).get_paused_groups_with_versions()
39    }
40}
41
42impl<C: ConnectionExt> QueryGroupVersion for DbConnection<C> {
43    fn set_group_paused(&self, group_id: &GroupId, min_version: &str) -> Result<(), StorageError> {
44        use crate::schema::groups::dsl;
45
46        self.raw_query(|conn| {
47            diesel::update(dsl::groups.filter(dsl::id.eq(group_id)))
48                .set(dsl::paused_for_version.eq(Some(min_version.to_string())))
49                .execute(conn)
50        })?;
51
52        Ok(())
53    }
54
55    fn unpause_group(&self, group_id: &GroupId) -> Result<(), StorageError> {
56        use crate::schema::groups::dsl;
57
58        self.raw_query(|conn| {
59            diesel::update(dsl::groups.filter(dsl::id.eq(group_id)))
60                .set(dsl::paused_for_version.eq::<Option<String>>(None))
61                .execute(conn)
62        })?;
63
64        Ok(())
65    }
66
67    fn get_group_paused_version(&self, group_id: &GroupId) -> Result<Option<String>, StorageError> {
68        use crate::schema::groups::dsl;
69
70        let paused_version = self.raw_query(|conn| {
71            dsl::groups
72                .select(dsl::paused_for_version)
73                .filter(dsl::id.eq(group_id))
74                .first::<Option<String>>(conn)
75        })?;
76
77        Ok(paused_version)
78    }
79
80    fn get_paused_groups_with_versions(&self) -> Result<Vec<(GroupId, String)>, StorageError> {
81        use crate::schema::groups::dsl;
82
83        let rows: Vec<(Vec<u8>, Option<String>)> = self.raw_query(|conn| {
84            dsl::groups
85                .select((dsl::id, dsl::paused_for_version))
86                .filter(dsl::paused_for_version.is_not_null())
87                .load::<(Vec<u8>, Option<String>)>(conn)
88        })?;
89
90        Ok(rows
91            .into_iter()
92            .filter_map(|(id, version)| {
93                let v = version?;
94                match GroupId::try_from(id.as_slice()) {
95                    Ok(group_id) => Some((group_id, v)),
96                    Err(err) => {
97                        tracing::warn!(
98                            error = %err,
99                            id_hex = %hex::encode(&id),
100                            id_len = id.len(),
101                            "get_paused_groups_with_versions: skipping row with \
102                             unparseable group id (not 16 bytes)"
103                        );
104                        None
105                    }
106                }
107            })
108            .collect())
109    }
110}