xmtp_db/encrypted_store/group/
dms.rs1use crate::ConnectionExt;
2
3use super::*;
4use crate::ConnectionError;
5
6use xmtp_proto::types::GroupId;
7pub trait QueryDms {
8 fn fetch_stitched(&self, key: &GroupId) -> Result<Option<StoredGroup>, ConnectionError>;
10
11 fn find_active_dm_group<M>(&self, members: M) -> Result<Option<StoredGroup>, ConnectionError>
12 where
13 M: std::fmt::Display;
14
15 fn other_dms(&self, group_id: &GroupId) -> Result<Vec<StoredGroup>, ConnectionError>;
17}
18
19impl<T> QueryDms for &T
20where
21 T: QueryDms,
22{
23 fn fetch_stitched(&self, key: &GroupId) -> Result<Option<StoredGroup>, ConnectionError> {
24 (**self).fetch_stitched(key)
25 }
26
27 fn find_active_dm_group<M>(&self, members: M) -> Result<Option<StoredGroup>, ConnectionError>
28 where
29 M: std::fmt::Display,
30 {
31 (**self).find_active_dm_group(members)
32 }
33
34 fn other_dms(&self, group_id: &GroupId) -> Result<Vec<StoredGroup>, ConnectionError> {
35 (**self).other_dms(group_id)
36 }
37}
38
39impl<C: ConnectionExt> QueryDms for DbConnection<C> {
40 fn fetch_stitched(&self, key: &GroupId) -> Result<Option<StoredGroup>, ConnectionError> {
42 let group = self.raw_query(|conn| {
43 groups::table
44 .filter(groups::id.eq(key))
45 .first::<StoredGroup>(conn)
46 .optional()
47 })?;
48
49 let Some(StoredGroup {
51 dm_id: Some(dm_id), ..
52 }) = group
53 else {
54 return Ok(group);
56 };
57
58 self.raw_query(|conn| {
60 groups::table
61 .filter(groups::dm_id.eq(dm_id))
62 .order_by(groups::last_message_ns.desc())
63 .first::<StoredGroup>(conn)
64 .optional()
65 })
66 }
67
68 fn find_active_dm_group<M>(&self, members: M) -> Result<Option<StoredGroup>, ConnectionError>
69 where
70 M: std::fmt::Display,
71 {
72 let query = dsl::groups
73 .filter(dsl::dm_id.eq(Some(members.to_string())))
74 .filter(dsl::membership_state.ne(GroupMembershipState::Restored))
75 .order_by(dsl::last_message_ns.desc());
76
77 self.raw_query(|conn| query.first(conn).optional())
78 }
79
80 fn other_dms(&self, group_id: &GroupId) -> Result<Vec<StoredGroup>, ConnectionError> {
82 let query = dsl::groups.filter(dsl::id.eq(group_id));
83
84 let groups: Vec<StoredGroup> = self.raw_query(|conn| query.load(conn))?;
85
86 let Some(StoredGroup {
88 id,
89 dm_id: Some(dm_id),
90 ..
91 }) = groups.into_iter().next()
92 else {
93 return Ok(vec![]);
94 };
95
96 let query = dsl::groups
97 .filter(dsl::dm_id.eq(dm_id))
98 .filter(dsl::id.ne(id));
99
100 let other_dms: Vec<StoredGroup> = self.raw_query(|conn| query.load(conn))?;
101 Ok(other_dms)
102 }
103}
104
105#[cfg(test)]
106pub(super) mod tests {
107 use super::*;
108 use crate::{Store, test_utils::with_connection};
109 use std::sync::atomic::{AtomicU16, Ordering};
110 use xmtp_common::{Generate, time::now_ns};
111
112 static TARGET_INBOX_ID: AtomicU16 = AtomicU16::new(2);
113
114 pub fn generate_dm(state: Option<GroupMembershipState>) -> StoredGroup {
116 let target = TARGET_INBOX_ID.fetch_add(1, Ordering::SeqCst).to_string();
117 StoredGroup::builder()
118 .id(GroupId::generate())
119 .created_at_ns(now_ns())
120 .membership_state(state.unwrap_or(GroupMembershipState::Allowed))
121 .added_by_inbox_id("placeholder_address")
122 .dm_id(format!(
123 "dm:placeholder_inbox_id_1:placeholder_inbox_id_{target}",
124 ))
125 .build()
126 .unwrap()
127 }
128
129 #[xmtp_common::test]
130 fn test_dm_stitching() {
131 with_connection(|conn| {
132 StoredGroup::builder()
133 .id(GroupId::generate())
134 .created_at_ns(now_ns())
135 .membership_state(GroupMembershipState::Allowed)
136 .added_by_inbox_id("placeholder_address")
137 .dm_id(Some("dm:some_wise_guy:thats_me".to_string()))
138 .build()
139 .unwrap()
140 .store(conn)
141 .unwrap();
142
143 StoredGroup::builder()
144 .id(GroupId::generate())
145 .created_at_ns(now_ns())
146 .membership_state(GroupMembershipState::Allowed)
147 .added_by_inbox_id("placeholder_address")
148 .dm_id(Some("dm:some_wise_guy:thats_me".to_string()))
149 .build()
150 .unwrap()
151 .store(conn)
152 .unwrap();
153 let all_groups = conn.find_groups(GroupQueryArgs::default()).unwrap();
154
155 assert_eq!(all_groups.len(), 1);
156 })
157 }
158
159 #[xmtp_common::test]
160 fn test_dm_deduplication() {
161 with_connection(|conn| {
162 let now = now_ns();
163 let base_time = now - 1_000_000_000; let dm_id = "dm:alice:bob";
167
168 let oldest_dm = StoredGroup::builder()
170 .id(GroupId::generate())
171 .created_at_ns(base_time)
172 .last_message_ns(base_time)
173 .membership_state(GroupMembershipState::Allowed)
174 .added_by_inbox_id("alice")
175 .dm_id(Some(dm_id.to_string()))
176 .build()
177 .unwrap();
178 oldest_dm.store(conn).unwrap();
179
180 let middle_dm = StoredGroup::builder()
182 .id(GroupId::generate())
183 .created_at_ns(base_time + 1_000_000)
184 .last_message_ns(base_time + 1_000_000)
185 .membership_state(GroupMembershipState::Allowed)
186 .added_by_inbox_id("bob")
187 .dm_id(Some(dm_id.to_string()))
188 .build()
189 .unwrap();
190 middle_dm.store(conn).unwrap();
191
192 let latest_dm = StoredGroup::builder()
194 .id(GroupId::generate())
195 .created_at_ns(base_time + 2_000_000)
196 .last_message_ns(base_time + 2_000_000)
197 .membership_state(GroupMembershipState::Allowed)
198 .added_by_inbox_id("alice")
199 .dm_id(Some(dm_id.to_string()))
200 .build()
201 .unwrap();
202 latest_dm.store(conn).unwrap();
203
204 let different_dm = StoredGroup::builder()
206 .id(GroupId::generate())
207 .created_at_ns(base_time + 500_000)
208 .last_message_ns(base_time + 500_000)
209 .membership_state(GroupMembershipState::Allowed)
210 .added_by_inbox_id("charlie")
211 .dm_id(Some("dm:charlie:dave".to_string()))
212 .build()
213 .unwrap();
214 different_dm.store(conn).unwrap();
215
216 let regular_group = StoredGroup::builder()
218 .id(GroupId::generate())
219 .created_at_ns(base_time + 1_500_000)
220 .last_message_ns(base_time + 1_500_000)
221 .membership_state(GroupMembershipState::Allowed)
222 .added_by_inbox_id("alice")
223 .dm_id(None) .build()
225 .unwrap();
226 regular_group.store(conn).unwrap();
227
228 let deduplicated_groups = conn
230 .find_groups(GroupQueryArgs {
231 include_duplicate_dms: false,
232 ..Default::default()
233 })
234 .unwrap();
235
236 assert_eq!(deduplicated_groups.len(), 3);
238
239 let kept_dm = deduplicated_groups
241 .iter()
242 .find(|g| g.dm_id.as_deref() == Some(dm_id))
243 .expect("Should find the DM group");
244 assert_eq!(kept_dm.id, latest_dm.id);
245 assert_eq!(kept_dm.last_message_ns, Some(base_time + 2_000_000));
246
247 let kept_different_dm = deduplicated_groups
249 .iter()
250 .find(|g| g.dm_id.as_deref() == Some("dm:charlie:dave"))
251 .expect("Should find the different DM group");
252 assert_eq!(kept_different_dm.id, different_dm.id);
253
254 let kept_regular = deduplicated_groups
256 .iter()
257 .find(|g| g.dm_id.is_none())
258 .expect("Should find the regular group");
259 assert_eq!(kept_regular.id, regular_group.id);
260
261 let all_groups = conn
263 .find_groups(GroupQueryArgs {
264 include_duplicate_dms: true,
265 ..Default::default()
266 })
267 .unwrap();
268
269 assert_eq!(all_groups.len(), 5);
271 })
272 }
273}