xmtp_mls/groups/
group_membership.rs1use 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 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 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}