tuwunel_service/rooms/pdu_metadata/
bundling.rs1use std::collections::BTreeSet;
2
3use futures::{Stream, StreamExt, TryFutureExt, pin_mut};
4use ruma::{OwnedUserId, UserId, api::Direction, events::room::encrypted::Relation};
5use tuwunel_core::{
6 PduId,
7 arrayvec::ArrayVec,
8 implement,
9 matrix::{Event, Pdu, PduCount, RawPduId},
10 result::LogErr,
11 utils::{
12 BoolExt,
13 stream::{ReadyExt, TryIgnore},
14 u64_from_u8,
15 },
16};
17
18use super::{
19 ExtractRelatesTo, IgnoredThreadView,
20 IgnoredThreadView::{Adjusted, Unchanged, WithoutSummary},
21 Service,
22 typed_relations::{CHILD_COUNT_OFFSET, KEY_LEN, Tag, prefix},
23};
24
25type Seek = ArrayVec<u8, KEY_LEN>;
26
27#[implement(Service)]
38#[tracing::instrument(skip_all, level = "trace")]
39pub async fn bundle_aggregations(&self, sender_user: &UserId, mut pdu: Pdu) -> Pdu {
40 if let Some(pruned) = self
43 .services
44 .state_accessor
45 .erased_view(sender_user, &pdu)
46 .await
47 {
48 return pruned;
49 }
50
51 let has_thread = pdu.has_thread_bundle().log_err().unwrap_or(true);
52
53 if has_thread {
54 if pdu
55 .remove_thread_latest_transaction_id_unless_sender(sender_user)
56 .log_err()
57 .is_err()
58 {
59 drop_thread_bundle(&mut pdu);
60 } else {
61 let participated = self
62 .services
63 .threads
64 .user_participated(pdu.event_id(), sender_user)
65 .await;
66
67 let participation = pdu
68 .set_thread_participated(participated)
69 .log_err();
70
71 if thread_result_or_drop(&mut pdu, participation).is_some() {
72 self.erase_thread_latest(sender_user, &mut pdu)
73 .await;
74
75 if self.services.server.config.bundle_edit_relations {
76 self.bundle_thread_latest_edit(sender_user, &mut pdu)
77 .await;
78 }
79 }
80 }
81 }
82
83 let replacement = self
84 .services
85 .server
86 .config
87 .bundle_edit_relations
88 .then_async(|| self.newest_replacement(&pdu))
89 .await
90 .flatten();
91
92 if let Some(mut replacement) = replacement
93 && !self
94 .services
95 .state_accessor
96 .erased_for(sender_user, &replacement)
97 .await
98 && replacement
99 .remove_transaction_id_unless_sender(Some(sender_user))
100 .log_err()
101 .is_ok()
102 {
103 pdu.set_replacement_bundle(&replacement.into_format())
104 .log_err()
105 .ok();
106 }
107
108 let references = self
109 .services
110 .server
111 .config
112 .bundle_reference_relations
113 .then_async(|| self.references(&pdu))
114 .await
115 .unwrap_or_default();
116
117 if !references.is_empty() {
118 pdu.set_reference_bundle(&references)
119 .log_err()
120 .ok();
121 }
122
123 pdu
124}
125
126#[implement(Service)]
130#[tracing::instrument(skip_all, level = "trace")]
131async fn erase_thread_latest(&self, sender_user: &UserId, pdu: &mut Pdu) {
132 let identity = pdu.thread_latest_event().log_err();
133
134 let Some((event_id, sender)) = thread_result_or_drop(pdu, identity).flatten() else {
135 return;
136 };
137
138 if !self.services.users.is_erased(&sender).await {
139 return;
140 }
141
142 let latest = self.services.timeline.get_pdu(&event_id).await;
143 let Some(latest) = thread_result_or_drop(pdu, latest) else {
144 return;
145 };
146
147 let Some(pruned) = self
148 .services
149 .state_accessor
150 .erased_view(sender_user, &latest)
151 .await
152 else {
153 return;
154 };
155
156 if pdu
157 .set_thread_latest_event(&pruned.into_format())
158 .log_err()
159 .is_err()
160 {
161 drop_thread_bundle(pdu);
162 }
163}
164
165fn thread_result_or_drop<T, E>(pdu: &mut Pdu, result: Result<T, E>) -> Option<T> {
166 result
167 .inspect_err(|_| drop_thread_bundle(pdu))
168 .ok()
169}
170
171fn drop_thread_bundle(pdu: &mut Pdu) { pdu.remove_thread_bundle().log_err().ok(); }
172
173#[implement(Service)]
178#[tracing::instrument(skip_all, level = "trace")]
179async fn bundle_thread_latest_edit(&self, sender_user: &UserId, pdu: &mut Pdu) {
180 let identity = pdu.thread_latest_event().log_err();
181
182 let Some((event_id, _)) = thread_result_or_drop(pdu, identity).flatten() else {
183 return;
184 };
185
186 let latest_event = self
187 .services
188 .timeline
189 .get_pdu(&event_id)
190 .await
191 .log_err();
192
193 let Some(mut latest_event) = thread_result_or_drop(pdu, latest_event) else {
194 return;
195 };
196
197 if self
198 .services
199 .state_accessor
200 .erased_for(sender_user, &latest_event)
201 .await
202 {
203 return;
204 }
205
206 let Some(mut replacement_event) = self.newest_replacement(&latest_event).await else {
207 return;
208 };
209
210 if self
211 .services
212 .state_accessor
213 .erased_for(sender_user, &replacement_event)
214 .await
215 {
216 return;
217 }
218
219 let sanitized = replacement_event
220 .remove_transaction_id_unless_sender(Some(sender_user))
221 .and_then(|()| latest_event.remove_transaction_id_unless_sender(Some(sender_user)))
222 .log_err();
223
224 if thread_result_or_drop(pdu, sanitized).is_none() {
225 return;
226 }
227
228 let replacement = latest_event
229 .set_replacement_bundle(&replacement_event.into_format())
230 .log_err();
231
232 if thread_result_or_drop(pdu, replacement).is_none() {
233 return;
234 }
235
236 let latest = pdu
237 .set_thread_latest_event(&latest_event.into_format())
238 .log_err();
239
240 thread_result_or_drop(pdu, latest);
241}
242
243#[implement(Service)]
248#[tracing::instrument(skip_all, level = "trace")]
249async fn newest_replacement(&self, parent: &Pdu) -> Option<Pdu> {
250 if parent.is_redacted() {
251 return None;
252 }
253
254 let parent_id: PduId = self
255 .services
256 .timeline
257 .get_pdu_id(parent.event_id())
258 .map_ok(Into::into)
259 .await
260 .ok()?;
261
262 let replacements = self.replacement_children(parent, parent_id);
263
264 pin_mut!(replacements);
265 replacements.next().await
266}
267
268#[implement(Service)]
272fn replacement_children<'a>(
273 &'a self,
274 parent: &'a Pdu,
275 parent_id: PduId,
276) -> impl Stream<Item = Pdu> + Send + 'a {
277 let shortroomid = parent_id.shortroomid;
278 let prefix = prefix(shortroomid, parent_id.count, Tag::Replace);
279
280 let mut seek = Seek::new();
281
282 seek.extend(prefix.iter().copied());
283 seek.extend([u8::MAX; size_of::<u64>() * 2]);
284
285 self.db
286 .relatesto_typed
287 .rev_raw_keys_from(seek.as_slice())
288 .ignore_err()
289 .ready_take_while(move |key| key.starts_with(&prefix))
290 .map(|key| u64_from_u8(&key[CHILD_COUNT_OFFSET..KEY_LEN]))
291 .map(PduCount::from_unsigned)
292 .map(move |count| (shortroomid, count))
293 .filter_map(async |(shortroomid, count)| {
294 let child_id: RawPduId = PduId { shortroomid, count }.into();
295 self.services
296 .timeline
297 .get_pdu_from_id(&child_id)
298 .await
299 .ok()
300 .filter(|child| !child.is_redacted())
301 .filter(|child| child.sender() == parent.sender())
302 .filter(|child| child.kind() == parent.kind())
303 })
304}
305
306#[implement(Service)]
313#[tracing::instrument(skip_all, level = "trace")]
314pub async fn ignored_thread_view(
315 &self,
316 sender_user: &UserId,
317 ignored: &BTreeSet<OwnedUserId>,
318 root: &Pdu,
319) -> IgnoredThreadView {
320 let Ok(root_id) = self
321 .services
322 .timeline
323 .get_pdu_id(root.event_id())
324 .await
325 else {
326 return Unchanged;
327 };
328
329 let participants = self
330 .services
331 .threads
332 .get_participants(&root_id)
333 .await
334 .unwrap_or_default();
335
336 if !participants
337 .iter()
338 .any(|user| ignored.contains(user))
339 {
340 return Unchanged;
341 }
342
343 let root_pid: PduId = root_id.into();
344 let replies = self
345 .get_relations(
346 root_pid.shortroomid,
347 root_pid.count,
348 None,
349 Direction::Backward,
350 Some(sender_user),
351 )
352 .ready_filter_map(|(_, pdu)| {
353 pdu.get_content()
354 .is_ok_and(|content: ExtractRelatesTo| {
355 matches!(content.relates_to, Relation::Thread(_))
356 })
357 .then_some(pdu)
358 });
359
360 let fold = |(total, unignored, latest): (usize, usize, Option<Pdu>), pdu: Pdu| match ignored
361 .contains(pdu.sender())
362 {
363 | true => (total.saturating_add(1), unignored, latest),
364 | false => (total.saturating_add(1), unignored.saturating_add(1), latest.or(Some(pdu))),
365 };
366
367 let (total, unignored, latest) = replies.ready_fold((0, 0, None), fold).await;
368
369 if total == 0 {
370 return match self.redacted_root(ignored, root).await {
371 | None => Unchanged,
372 | root => Adjusted { root, count: None, latest: None },
373 };
374 }
375
376 if unignored == 0 {
377 return WithoutSummary {
378 root: self.redacted_root(ignored, root).await,
379 };
380 }
381
382 let swap = root
383 .thread_latest_event()
384 .ok()
385 .flatten()
386 .is_some_and(|(_, sender)| ignored.contains(&sender));
387
388 let latest = match swap.then_some(latest).flatten() {
389 | None => None,
390 | Some(reply) => {
391 let reply = self
394 .services
395 .state_accessor
396 .erased_view(sender_user, &reply)
397 .await
398 .unwrap_or(reply);
399
400 Some(reply.into_format())
401 },
402 };
403
404 let count = unignored.ne(&total).then_some(unignored);
405
406 let root = self.redacted_root(ignored, root).await;
407
408 if root.is_none() && count.is_none() && latest.is_none() {
409 return Unchanged;
410 }
411
412 Adjusted { root, count, latest }
413}
414
415#[implement(Service)]
419#[tracing::instrument(skip_all, level = "trace")]
420async fn redacted_root(&self, ignored: &BTreeSet<OwnedUserId>, root: &Pdu) -> Option<Box<Pdu>> {
421 ignored
422 .contains(root.sender())
423 .then_async(async || {
424 self.services
425 .state
426 .get_room_version_rules(root.room_id())
427 .await
428 .log_err()
429 .ok()
430 .and_then(|rules| root.redacted(&rules.redaction).log_err().ok())
431 .map(Box::new)
432 })
433 .await
434 .flatten()
435}
436
437#[cfg(test)]
438mod tests {
439 use serde_json::json;
440
441 use super::*;
442
443 #[test]
444 fn thread_latest_load_error_drops_bundle() {
445 let mut pdu: Pdu = serde_json::from_value(json!({
446 "type": "m.room.member",
447 "content": { "membership": "join" },
448 "event_id": "$member:example.com",
449 "room_id": "!room:example.com",
450 "sender": "@alice:example.com",
451 "state_key": "@alice:example.com",
452 "prev_events": ["$prev:example.com"],
453 "auth_events": ["$auth:example.com"],
454 "origin_server_ts": 1_838_188_000,
455 "depth": 12,
456 "hashes": { "sha256": "thishashcoversallfieldsincasethisisredacted" },
457 "unsigned": {
458 "age": 4612,
459 "m.relations": {
460 "m.replace": { "event_id": "$edit:example.com" },
461 "m.thread": { "count": 3 },
462 },
463 },
464 }))
465 .expect("test fixture should deserialize as a valid PDU");
466
467 let latest: Option<()> = thread_result_or_drop(&mut pdu, Err("missing latest event"));
468
469 assert!(latest.is_none(), "load failure returned a latest event");
470
471 let unsigned: serde_json::Value = serde_json::from_str(
472 pdu.unsigned
473 .as_ref()
474 .expect("sibling unsigned data should remain")
475 .json()
476 .get(),
477 )
478 .expect("remaining unsigned data should be valid JSON");
479
480 assert!(
481 unsigned["m.relations"].get("m.thread").is_none(),
482 "failed load retained the thread bundle",
483 );
484
485 assert_eq!(
486 unsigned["m.relations"]["m.replace"],
487 json!({ "event_id": "$edit:example.com" }),
488 "failed load removed a sibling relation",
489 );
490 assert_eq!(unsigned["age"], 4612, "failed load removed outer unsigned data");
491 }
492}