Skip to main content

xmtp_mls/messages/
enrichment.rs

1use crate::messages::decoded_message::{DecodedMessage, DeletedBy, MessageBody};
2use hex::ToHexExt;
3use std::collections::HashMap;
4use thiserror::Error;
5use xmtp_common::{ErrorCode, RetryableError};
6use xmtp_db::DbQuery;
7use xmtp_db::group_message::{
8    ContentType as DbContentType, Deletable, RelationCounts, RelationQuery, StoredGroupMessage,
9};
10use xmtp_db::message_deletion::StoredMessageDeletion;
11use xmtp_proto::xmtp::mls::message_contents::ContentTypeId;
12
13use xmtp_proto::types::GroupId;
14/// Content type ID for deleted message placeholders shown in enriched message lists
15pub fn deleted_message_content_type() -> ContentTypeId {
16    ContentTypeId {
17        authority_id: "xmtp.org".to_string(),
18        type_id: "deletedMessage".to_string(),
19        version_major: 1,
20        version_minor: 0,
21    }
22}
23
24#[derive(Debug, Error, ErrorCode)]
25pub enum EnrichMessageError {
26    #[error("DB error: {0}")]
27    #[error_code(inherit)]
28    DbConnection(#[from] xmtp_db::ConnectionError),
29    /// Codec decode error.
30    ///
31    /// Content type codec failed. Not retryable.
32    #[error("Decode error: {0}")]
33    CodecError(#[from] xmtp_content_types::CodecError),
34    /// Decode error.
35    ///
36    /// Protobuf decoding failed. Not retryable.
37    #[error("Decode error: {0}")]
38    DecodeError(#[from] prost::DecodeError),
39}
40
41impl RetryableError for EnrichMessageError {
42    fn is_retryable(&self) -> bool {
43        match self {
44            Self::DbConnection(e) => e.is_retryable(),
45            Self::CodecError(_) => false,
46            Self::DecodeError(_) => false,
47        }
48    }
49}
50
51// Mapping of reactions, keyed by the ID of the message being reacted to.
52type ReactionMap = HashMap<Vec<u8>, Vec<DecodedMessage>>;
53// Mapping of referenced messages, keyed by ID (stores both stored and decoded)
54type ReferencedMessageMap = HashMap<Vec<u8>, (StoredGroupMessage, DecodedMessage)>;
55// Mapping of deletions, keyed by the ID of the deleted message
56type DeletionMap = HashMap<Vec<u8>, StoredMessageDeletion>;
57
58/// Validates if a deletion should be applied. Checks group membership and authorization.
59pub(crate) fn is_deletion_valid(
60    deletion: &StoredMessageDeletion,
61    message: &StoredGroupMessage,
62    group_id: &GroupId,
63) -> bool {
64    if deletion.deleted_message_id != message.id {
65        return false;
66    }
67
68    if deletion.group_id != *group_id || message.group_id != *group_id {
69        return false;
70    }
71
72    if !message.kind.is_deletable() || !message.content_type.is_deletable() {
73        return false;
74    }
75
76    let is_sender = deletion.deleted_by_inbox_id == message.sender_inbox_id;
77    is_sender || deletion.is_super_admin_deletion
78}
79
80#[xmtp_common::mls_span]
81pub fn enrich_messages(
82    conn: impl DbQuery,
83    group_id: &GroupId,
84    messages: Vec<StoredGroupMessage>,
85) -> Result<Vec<DecodedMessage>, EnrichMessageError> {
86    let initial_message_ids: Vec<&[u8]> = messages.iter().map(|m| m.id.as_ref()).collect();
87
88    let reference_ids: Vec<&[u8]> = messages
89        .iter()
90        .filter_map(|m| m.reference_id.as_deref())
91        .collect();
92
93    let mut relations = get_relations(conn, group_id, &initial_message_ids, &reference_ids)?;
94
95    let messages: Vec<DecodedMessage> = messages
96        .into_iter()
97        .filter_map(|stored_message| {
98            let mut decoded = DecodedMessage::try_from(stored_message.clone())
99                .inspect_err(|err| tracing::warn!("Failed to decode message {:?}", err))
100                .ok()?;
101
102            let valid_deletion = relations
103                .deletions
104                .get(&decoded.metadata.id)
105                .filter(|deletion| is_deletion_valid(deletion, &stored_message, group_id));
106
107            if let Some(deletion) = valid_deletion {
108                let is_sender = deletion.deleted_by_inbox_id == stored_message.sender_inbox_id;
109                decoded.content = MessageBody::DeletedMessage {
110                    deleted_by: if is_sender {
111                        DeletedBy::Sender
112                    } else {
113                        DeletedBy::Admin(deletion.deleted_by_inbox_id.clone())
114                    },
115                };
116                decoded.metadata.content_type = deleted_message_content_type();
117                decoded.reactions = Vec::new();
118                decoded.num_replies = 0;
119            } else {
120                decoded.reactions = relations
121                    .reactions
122                    .remove(&decoded.metadata.id)
123                    .unwrap_or_default();
124
125                decoded.num_replies = relations
126                    .reply_counts
127                    .get(&decoded.metadata.id)
128                    .cloned()
129                    .unwrap_or(0);
130
131                // Handle Reply messages - populate in_reply_to field
132                if let MessageBody::Reply(mut reply_body) = decoded.content {
133                    let _ = hex::decode(&reply_body.reference_id)
134                        .inspect_err(|err| {
135                            // The reference is sender-controlled; truncate so a
136                            // malformed value can't flood the log line.
137                            let reference_id: String =
138                                reply_body.reference_id.chars().take(64).collect();
139                            tracing::warn!(
140                                group_id = %group_id,
141                                message_id = %hex::encode(&stored_message.id),
142                                sender_inbox_id = %stored_message.sender_inbox_id,
143                                reference_id = %reference_id,
144                                "could not parse reference ID as hex: {:?}",
145                                err
146                            )
147                        })
148                        .inspect(|id| {
149                            let mut in_reply_to = relations
150                                .referenced_messages
151                                .get(id)
152                                .map(|(_, decoded)| decoded.clone());
153
154                            if let Some(msg) = in_reply_to.as_mut()
155                                && let Some(deletion) = relations.deletions.get(id)
156                                && let Some((stored_msg, _)) = relations.referenced_messages.get(id)
157                                && is_deletion_valid(deletion, stored_msg, group_id)
158                            {
159                                let is_sender =
160                                    deletion.deleted_by_inbox_id == stored_msg.sender_inbox_id;
161                                msg.content = MessageBody::DeletedMessage {
162                                    deleted_by: if is_sender {
163                                        DeletedBy::Sender
164                                    } else {
165                                        DeletedBy::Admin(deletion.deleted_by_inbox_id.clone())
166                                    },
167                                };
168                                msg.reactions = Vec::new();
169                                msg.num_replies = 0;
170                            }
171                            reply_body.in_reply_to = in_reply_to.map(Box::new);
172                        });
173                    decoded.content = MessageBody::Reply(reply_body);
174                }
175            }
176
177            Some(decoded)
178        })
179        .collect();
180
181    Ok(messages)
182}
183
184fn get_relations(
185    conn: impl DbQuery,
186    group_id: &GroupId,
187    message_ids: &[&[u8]],
188    reference_ids: &[&[u8]],
189) -> Result<GetRelationsResults, EnrichMessageError> {
190    if message_ids.is_empty() {
191        return Ok(GetRelationsResults {
192            reactions: HashMap::new(),
193            referenced_messages: HashMap::new(),
194            reply_counts: HashMap::new(),
195            deletions: HashMap::new(),
196        });
197    }
198
199    let reactions_relations_query = RelationQuery::builder()
200        .content_types(Some(vec![DbContentType::Reaction]))
201        .build()
202        .unwrap_or_default();
203
204    let replies_count_query = RelationQuery::builder()
205        .content_types(Some(vec![DbContentType::Reply]))
206        .build()
207        .unwrap_or_default();
208
209    let reactions = conn.get_inbound_relations(group_id, message_ids, reactions_relations_query)?;
210    let referenced_messages = conn.get_outbound_relations(group_id, reference_ids)?;
211    let reply_counts =
212        conn.get_inbound_relation_counts(group_id, message_ids, replies_count_query)?;
213
214    // Get deletions for all messages AND referenced messages in a single batch query.
215    // This ensures that if a reply references a deleted message, we can properly show
216    // the deletion state in the reply chain.
217    let mut all_ids: Vec<Vec<u8>> = message_ids.iter().map(|id| id.to_vec()).collect();
218    all_ids.extend(reference_ids.iter().map(|id| id.to_vec()));
219    let deletions = conn.get_deletions_for_messages(all_ids)?;
220
221    Ok(GetRelationsResults {
222        reactions: get_reactions(reactions),
223        referenced_messages: get_referenced_messages(referenced_messages),
224        reply_counts,
225        deletions: get_deletions(deletions),
226    })
227}
228
229struct GetRelationsResults {
230    reactions: ReactionMap,
231    referenced_messages: ReferencedMessageMap,
232    reply_counts: RelationCounts,
233    deletions: DeletionMap,
234}
235
236fn get_referenced_messages(messages: HashMap<Vec<u8>, StoredGroupMessage>) -> ReferencedMessageMap {
237    messages
238        .into_iter()
239        .filter_map(|(id, stored_message)| {
240            let message_id = id.clone();
241            DecodedMessage::try_from(stored_message.clone())
242                .inspect_err(|err| {
243                    tracing::warn!(
244                        "Failed to decode reply root message with ID {} {:?}",
245                        message_id.encode_hex(),
246                        err
247                    );
248                })
249                .map(|decoded| (id, (stored_message, decoded)))
250                .ok()
251        })
252        .collect()
253}
254
255fn get_reactions(messages: HashMap<Vec<u8>, Vec<StoredGroupMessage>>) -> ReactionMap {
256    messages
257        .into_iter()
258        .map(|(id, reaction_messages)| {
259            let mapped_reactions: Vec<DecodedMessage> = reaction_messages
260                .into_iter()
261                .filter_map(|stored_msg| {
262                    DecodedMessage::try_from(stored_msg)
263                        .inspect_err(|err| {
264                            tracing::warn!(
265                                "Failed to decode message categorized as Reaction: {:?}",
266                                err
267                            );
268                        })
269                        .ok()
270                })
271                .collect();
272            (id, mapped_reactions)
273        })
274        .collect()
275}
276
277fn get_deletions(deletions: Vec<StoredMessageDeletion>) -> DeletionMap {
278    deletions
279        .into_iter()
280        .map(|deletion| (deletion.deleted_message_id.clone(), deletion))
281        .collect()
282}