Skip to main content

xmtp_db/encrypted_store/group/
dms.rs

1use crate::ConnectionExt;
2
3use super::*;
4use crate::ConnectionError;
5
6use xmtp_proto::types::GroupId;
7pub trait QueryDms {
8    /// Same behavior as fetched, but will stitch DM groups
9    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    /// Load the other DMs that are stitched into this group
16    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    /// Same behavior as fetched, but will stitch DM groups
41    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        // Is this group a DM?
50        let Some(StoredGroup {
51            dm_id: Some(dm_id), ..
52        }) = group
53        else {
54            // If not, return the group
55            return Ok(group);
56        };
57
58        // Otherwise, return the stitched DM
59        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    /// Load the other DMs that are stitched into this group
81    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        // Grab the dm_id of the group
87        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    /// Generate a test dm group
115    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; // 1 second ago
164
165            // Create DM groups with same dm_id but different timestamps
166            let dm_id = "dm:alice:bob";
167
168            // Oldest DM (should be filtered out)
169            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            // Middle DM (should be filtered out)
181            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            // Latest DM (should be kept)
193            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            // Create another DM with different dm_id (should always be kept)
205            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            // Create a regular group (non-DM, should always be kept)
217            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) // No dm_id = regular group
224                .build()
225                .unwrap();
226            regular_group.store(conn).unwrap();
227
228            // Test with include_duplicate_dms = false (default deduplication)
229            let deduplicated_groups = conn
230                .find_groups(GroupQueryArgs {
231                    include_duplicate_dms: false,
232                    ..Default::default()
233                })
234                .unwrap();
235
236            // Should have 3 groups: latest DM, different DM, and regular group
237            assert_eq!(deduplicated_groups.len(), 3);
238
239            // Verify the latest DM is kept (highest last_message_ns for dm_id)
240            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            // Verify different DM is kept
248            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            // Verify regular group is kept
255            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            // Test with include_duplicate_dms = true (no deduplication)
262            let all_groups = conn
263                .find_groups(GroupQueryArgs {
264                    include_duplicate_dms: true,
265                    ..Default::default()
266                })
267                .unwrap();
268
269            // Should have all 5 groups
270            assert_eq!(all_groups.len(), 5);
271        })
272    }
273}