Skip to main content

xmtp_mls/groups/
group_membership.rs

1use crate::groups::intents::Installation;
2use openmls::key_packages::KeyPackage;
3use prost::{DecodeError, Message};
4use std::collections::{HashMap, HashSet};
5use xmtp_proto::xmtp::mls::message_contents::GroupMembership as GroupMembershipProto;
6
7#[derive(Debug, Clone, PartialEq)]
8pub struct GroupMembership {
9    pub(crate) members: HashMap<String, u64>,
10    pub(crate) failed_installations: Vec<Vec<u8>>,
11}
12
13impl GroupMembership {
14    pub fn new() -> Self {
15        GroupMembership {
16            members: HashMap::new(),
17            failed_installations: Vec::new(),
18        }
19    }
20
21    pub fn add(&mut self, inbox_id: String, last_sequence_id: u64) {
22        self.members.insert(inbox_id, last_sequence_id);
23    }
24
25    pub fn remove<InboxId: AsRef<str>>(&mut self, inbox_id: InboxId) {
26        self.members.remove(inbox_id.as_ref());
27    }
28
29    pub fn get<InboxId: AsRef<str>>(&self, inbox_id: InboxId) -> Option<&u64> {
30        self.members.get(inbox_id.as_ref())
31    }
32
33    pub fn inbox_ids(&self) -> Vec<&str> {
34        self.members.keys().map(AsRef::as_ref).collect()
35    }
36
37    // Convert the mapping to a vector of `inbox_id`/`sequence_id` tuples
38    pub fn to_filters(&self) -> Vec<(&str, i64)> {
39        self.members
40            .iter()
41            .map(|(inbox_id, sequence_id)| (inbox_id.as_str(), *sequence_id as i64))
42            .collect()
43    }
44
45    pub fn diff<'inbox_id>(
46        &'inbox_id self,
47        new_group_membership: &'inbox_id Self,
48    ) -> MembershipDiff<'inbox_id> {
49        let mut removed_inboxes: Vec<&String> = vec![];
50        let mut updated_inboxes: Vec<&String> = vec![];
51
52        for (inbox_id, last_sequence_id) in self.members.iter() {
53            match new_group_membership.get(inbox_id) {
54                Some(new_last_sequence_id) => {
55                    if new_last_sequence_id.ne(last_sequence_id) {
56                        updated_inboxes.push(inbox_id);
57                    }
58                }
59                None => {
60                    removed_inboxes.push(inbox_id);
61                }
62            }
63        }
64
65        let added_inboxes = new_group_membership
66            .members
67            .keys()
68            .filter(|&inbox_id| !self.members.contains_key(inbox_id))
69            .collect::<Vec<&String>>();
70
71        MembershipDiff {
72            added_inboxes,
73            removed_inboxes,
74            updated_inboxes,
75        }
76    }
77}
78
79impl Default for GroupMembership {
80    fn default() -> Self {
81        GroupMembership::new()
82    }
83}
84
85impl TryFrom<Vec<u8>> for GroupMembership {
86    type Error = DecodeError;
87
88    fn try_from(value: Vec<u8>) -> Result<Self, Self::Error> {
89        let membership_proto = GroupMembershipProto::decode(value.as_slice())?;
90
91        Ok(GroupMembership {
92            members: membership_proto.members,
93            failed_installations: membership_proto.failed_installations,
94        })
95    }
96}
97
98impl From<&GroupMembership> for Vec<u8> {
99    fn from(value: &GroupMembership) -> Self {
100        let membership_proto = GroupMembershipProto {
101            members: value.members.clone(),
102            failed_installations: value.failed_installations.clone(),
103        };
104
105        membership_proto.encode_to_vec()
106    }
107}
108
109#[derive(Debug, Clone)]
110pub struct MembershipDiff<'inbox_id> {
111    pub added_inboxes: Vec<&'inbox_id String>,
112    pub removed_inboxes: Vec<&'inbox_id String>,
113    pub updated_inboxes: Vec<&'inbox_id String>,
114}
115
116#[derive(Debug)]
117pub struct MembershipDiffWithKeyPackages {
118    pub new_installations: Vec<Installation>,
119    pub new_key_packages: Vec<KeyPackage>,
120    pub removed_installations: HashSet<Vec<u8>>,
121    pub failed_installations: Vec<Vec<u8>>,
122}
123
124impl MembershipDiffWithKeyPackages {
125    pub fn new(
126        new_installations: Vec<Installation>,
127        new_key_packages: Vec<KeyPackage>,
128        removed_installations: HashSet<Vec<u8>>,
129        failed_installations: Vec<Vec<u8>>,
130    ) -> MembershipDiffWithKeyPackages {
131        MembershipDiffWithKeyPackages {
132            new_installations,
133            new_key_packages,
134            removed_installations,
135            failed_installations,
136        }
137    }
138}
139
140#[cfg(test)]
141pub(crate) mod tests {
142    use super::GroupMembership;
143
144    #[xmtp_common::test]
145    fn test_equality_works() {
146        let inbox_id_1 = "inbox_1".to_string();
147        let sequence_id_1: u64 = 1;
148        let mut member_map_1 = GroupMembership::new();
149        let mut member_map_2 = GroupMembership::new();
150
151        member_map_1.add(inbox_id_1.clone(), sequence_id_1);
152
153        assert!(member_map_1.ne(&member_map_2));
154
155        member_map_2.add(inbox_id_1.clone(), sequence_id_1);
156        assert!(member_map_1.eq(&member_map_2));
157
158        // Now change the sequence ID and make sure it is not equal again
159        member_map_2.add(inbox_id_1.clone(), 2);
160        assert!(member_map_1.ne(&member_map_2));
161    }
162
163    #[xmtp_common::test]
164    fn test_diff() {
165        let mut initial_members = GroupMembership::new();
166        initial_members.add("inbox_1".into(), 1);
167        initial_members.add("inbox_2".into(), 1);
168
169        let mut updated_list = initial_members.clone();
170        updated_list.remove("inbox_1");
171        updated_list.add("inbox_2".into(), 2);
172        updated_list.add("inbox_3".into(), 1);
173
174        let diff = initial_members.diff(&updated_list);
175        assert_eq!(diff.added_inboxes, vec!["inbox_3"]);
176        assert_eq!(diff.updated_inboxes, vec!["inbox_2"]);
177        assert_eq!(diff.removed_inboxes, vec!["inbox_1"]);
178    }
179}