tuwunel_service/rooms/threads/
mod.rs1use std::{collections::BTreeMap, pin::pin, sync::Arc};
2
3use futures::{Stream, StreamExt, TryFutureExt, future::join3};
4use ruma::{
5 CanonicalJsonObject, CanonicalJsonValue, EventId, OwnedEventId, OwnedUserId, RoomId, UInt,
6 UserId,
7 api::{Direction, client::threads::get_threads::v1::IncludeThreads},
8 events::{AnySyncMessageLikeEvent, TimelineEventType, relation::RelationType},
9 serde::Raw,
10 uint,
11};
12use serde::Deserialize;
13use serde_json::json;
14use tuwunel_core::{
15 Event, Result, err,
16 matrix::pdu::{PduCount, PduEvent, PduId, RawPduId},
17 utils::{
18 ReadyExt,
19 stream::{TryIgnore, WidebandExt, automatic_width},
20 },
21};
22use tuwunel_database::{Deserialized, Map, Txn};
23
24#[cfg(test)]
25mod tests;
26
27const MAX_THREAD_HOPS: usize = 3;
30
31#[derive(Deserialize)]
32struct ExtractThreadRelation {
33 #[serde(rename = "m.relates_to")]
34 relates_to: ThreadRelation,
35}
36
37#[derive(Deserialize)]
38struct ThreadRelation {
39 rel_type: RelationType,
40 event_id: OwnedEventId,
41}
42
43fn canonical_object_field<'a>(
44 object: &'a mut CanonicalJsonObject,
45 field: &str,
46) -> &'a mut CanonicalJsonObject {
47 if !matches!(object.get(field), Some(CanonicalJsonValue::Object(_))) {
48 object.insert(field.into(), CanonicalJsonValue::Object(BTreeMap::new()));
49 }
50
51 let Some(CanonicalJsonValue::Object(value)) = object.get_mut(field) else {
52 unreachable!("canonical object field was initialized as an object");
53 };
54
55 value
56}
57
58fn update_thread_bundle<E>(unsigned: &mut CanonicalJsonObject, event: &E)
61where
62 E: Event,
63{
64 let latest_event = event.to_sync_message_like_without_unsigned();
65 update_thread_bundle_raw(unsigned, &latest_event);
66}
67
68fn update_thread_bundle_raw(
69 unsigned: &mut CanonicalJsonObject,
70 latest_event: &Raw<AnySyncMessageLikeEvent>,
71) {
72 let relations = canonical_object_field(unsigned, "m.relations");
73 let thread = canonical_object_field(relations, "m.thread");
74 let count = thread
75 .get("count")
76 .cloned()
77 .and_then(|count| serde_json::from_value::<UInt>(count.into()).ok());
78
79 let count = count.map_or_else(|| uint!(1), |count| count.saturating_add(uint!(1)));
80 let latest_event = serde_json::from_str(latest_event.json().get())
81 .expect("thread latest event should be canonical JSON");
82
83 thread.insert("latest_event".into(), latest_event);
84 thread.insert(
85 "count".into(),
86 json!(count)
87 .try_into()
88 .expect("thread count is canonical JSON"),
89 );
90
91 if !matches!(thread.get("current_user_participated"), Some(CanonicalJsonValue::Bool(_))) {
92 thread.insert("current_user_participated".into(), CanonicalJsonValue::Bool(true));
93 }
94}
95
96pub struct Service {
97 db: Data,
98 services: Arc<crate::services::OnceServices>,
99}
100
101pub(super) struct Data {
102 threadid_userids: Arc<Map>,
103 threadactivityid_rootid: Arc<Map>,
104 threadrootid_latestcount: Arc<Map>,
105}
106
107impl crate::Service for Service {
108 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
109 Ok(Arc::new(Self {
110 db: Data {
111 threadid_userids: args.db["threadid_userids"].clone(),
112 threadactivityid_rootid: args.db["threadactivityid_rootid"].clone(),
113 threadrootid_latestcount: args.db["threadrootid_latestcount"].clone(),
114 },
115 services: args.services.clone(),
116 }))
117 }
118
119 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
120}
121
122impl Service {
123 pub async fn get_thread_id<E>(&self, event: &E) -> Option<OwnedEventId>
129 where
130 E: Event,
131 {
132 let initial = match event.get_content::<ExtractThreadRelation>() {
133 | Ok(t) => Some(t.relates_to),
134 | Err(_) => self.relates_to_via_redaction_target(event).await,
135 };
136
137 let mut relates_to = initial?;
138
139 for _ in 0..MAX_THREAD_HOPS {
140 if relates_to.rel_type == RelationType::Thread {
141 return Some(relates_to.event_id);
142 }
143
144 relates_to = self
145 .services
146 .timeline
147 .get_pdu(&relates_to.event_id)
148 .await
149 .ok()?
150 .get_content::<ExtractThreadRelation>()
151 .ok()?
152 .relates_to;
153 }
154
155 None
156 }
157
158 async fn relates_to_via_redaction_target<E>(&self, event: &E) -> Option<ThreadRelation>
162 where
163 E: Event,
164 {
165 if *event.kind() != TimelineEventType::RoomRedaction {
166 return None;
167 }
168
169 let room_rules = self
170 .services
171 .state
172 .get_room_version_rules(event.room_id())
173 .await
174 .ok()?;
175
176 let target_id = event.redacts_id(&room_rules)?;
177
178 self.services
179 .timeline
180 .get_pdu(&target_id)
181 .await
182 .ok()?
183 .get_content::<ExtractThreadRelation>()
184 .ok()
185 .map(|t| t.relates_to)
186 }
187
188 pub async fn get_thread_id_for_event(&self, event_id: &EventId) -> Option<OwnedEventId> {
191 let pdu = self
192 .services
193 .timeline
194 .get_pdu(event_id)
195 .await
196 .ok()?;
197
198 self.get_thread_id(&pdu).await
199 }
200
201 pub async fn add_to_thread<E>(
202 &self,
203 root_event_id: &EventId,
204 pdu_id: RawPduId,
205 event: &E,
206 ) -> Result
207 where
208 E: Event,
209 {
210 let root_id = self
211 .services
212 .timeline
213 .get_pdu_id(root_event_id)
214 .await
215 .map_err(|e| {
216 err!(Request(InvalidParam("Invalid event_id in thread message: {e:?}")))
217 })?;
218
219 let root_pdu = self
220 .services
221 .timeline
222 .get_pdu_from_id(&root_id)
223 .await
224 .map_err(|e| err!(Request(InvalidParam("Thread root not found: {e:?}"))))?;
225
226 if root_pdu.room_id() != event.room_id() {
227 return Ok(());
228 }
229
230 let mut root_pdu_json = self
231 .services
232 .timeline
233 .get_pdu_json_from_id(&root_id)
234 .await
235 .map_err(|e| err!(Request(InvalidParam("Thread root pdu not found: {e:?}"))))?;
236
237 let mut users = self
238 .get_participants(&root_id)
239 .await
240 .unwrap_or_else(|_| vec![root_pdu.sender().to_owned()]);
241
242 users.push(event.sender().to_owned());
243
244 let mut txn = self.services.db.txn();
245
246 self.update_participants(&mut txn, &root_id, &users);
247
248 let count = pdu_id.pdu_count();
249
250 if matches!(count, PduCount::Normal(_)) {
251 txn.insert_raw(&self.db.threadactivityid_rootid, pdu_id, root_id);
252 txn.insert_raw(&self.db.threadrootid_latestcount, root_id, count.to_be_bytes());
253 }
254
255 if let CanonicalJsonValue::Object(unsigned) = root_pdu_json
256 .entry("unsigned".into())
257 .or_insert_with(|| CanonicalJsonValue::Object(BTreeMap::default()))
258 {
259 update_thread_bundle(unsigned, event);
260
261 self.services
262 .timeline
263 .stage_replace_pdu(&mut txn, &root_id, &root_pdu_json);
264 }
265
266 txn.execute();
267 Ok(())
268 }
269
270 pub fn threads_until<'a>(
271 &'a self,
272 user_id: &'a UserId,
273 room_id: &'a RoomId,
274 count: PduCount,
275 include: &'a IncludeThreads,
276 ) -> impl Stream<Item = Result<(PduCount, PduEvent)>> + Send {
277 let participated = matches!(include, IncludeThreads::Participated);
278
279 self.services
280 .short
281 .get_shortroomid(room_id)
282 .map_ok(move |shortroomid| PduId {
283 shortroomid,
284 count: count.saturating_sub(1),
285 })
286 .map_ok(Into::into)
287 .map_ok(move |current: RawPduId| {
288 self.db
289 .threadactivityid_rootid
290 .rev_raw_stream_from(¤t)
291 .ignore_err()
292 .map(|(key, root_id)| (RawPduId::from(key), RawPduId::from(root_id)))
293 .ready_take_while(move |(activity_id, _)| {
294 activity_id.shortroomid() == current.shortroomid()
295 })
296 .map(move |(activity_id, root_id)| {
297 (activity_id, root_id, user_id, participated)
298 })
299 .wide_filter_map(async |(activity_id, root_id, user_id, participated)| {
300 self.live_thread(user_id, participated, activity_id, root_id)
301 .await
302 })
303 .map(Ok)
304 })
305 .try_flatten_stream()
306 }
307
308 async fn live_thread(
311 &self,
312 user_id: &UserId,
313 participated: bool,
314 activity_id: RawPduId,
315 root_id: RawPduId,
316 ) -> Option<(PduCount, PduEvent)> {
317 let count = activity_id.pdu_count();
318
319 let pointer = self
320 .db
321 .threadrootid_latestcount
322 .get(&root_id)
323 .await
324 .deserialized()
325 .map(PduCount::from_unsigned)
326 .ok()?;
327
328 if count != pointer {
329 if count < pointer {
332 self.db
333 .threadactivityid_rootid
334 .remove(&activity_id);
335 }
336
337 return None;
338 }
339
340 if participated && !self.is_participant(&root_id, user_id).await {
341 return None;
342 }
343
344 let mut pdu = self
345 .services
346 .timeline
347 .get_pdu_from_id(&root_id)
348 .await
349 .ok()?;
350
351 pdu.remove_transaction_id_unless_sender(Some(user_id))
352 .ok()?;
353
354 Some((count, pdu))
355 }
356
357 async fn is_participant(&self, root_id: &RawPduId, user_id: &UserId) -> bool {
358 self.db
359 .threadid_userids
360 .get(root_id)
361 .await
362 .is_ok_and(|participants| {
363 participants
364 .split(|&byte| byte == 0xFF)
365 .any(|user| user == user_id.as_bytes())
366 })
367 }
368
369 pub(super) fn update_participants(
370 &self,
371 txn: &mut Txn,
372 root_id: &RawPduId,
373 participants: &[OwnedUserId],
374 ) {
375 let users = participants
376 .iter()
377 .map(|user| user.as_bytes())
378 .collect::<Vec<_>>()
379 .join(&[0xFF][..]);
380
381 txn.insert_raw(&self.db.threadid_userids, root_id, &users);
382 }
383
384 pub(super) async fn get_participants(&self, root_id: &RawPduId) -> Result<Vec<OwnedUserId>> {
385 self.db
386 .threadid_userids
387 .get(root_id)
388 .await
389 .deserialized()
390 }
391
392 pub async fn user_participated(&self, root_event_id: &EventId, user_id: &UserId) -> bool {
395 let Ok(root_id) = self
396 .services
397 .timeline
398 .get_pdu_id(root_event_id)
399 .await
400 else {
401 return false;
402 };
403
404 self.is_participant(&root_id, user_id).await
405 }
406
407 #[tracing::instrument(skip(self), level = "debug")]
408 pub(super) async fn delete_all_rooms_threads(&self, room_id: &RoomId) -> Result {
409 let Ok(shortroomid) = self.services.short.get_shortroomid(room_id).await else {
410 return Ok(());
411 };
412
413 join3(
414 self.db.threadid_userids.del_prefix(&shortroomid),
415 self.db
416 .threadactivityid_rootid
417 .del_prefix(&shortroomid),
418 self.db
419 .threadrootid_latestcount
420 .del_prefix(&shortroomid),
421 )
422 .await;
423
424 Ok(())
425 }
426
427 pub async fn rebuild_thread_activity(&self) -> Result {
431 self.db.threadactivityid_rootid.clear().await;
432 self.db.threadrootid_latestcount.clear().await;
433
434 self.db
435 .threadid_userids
436 .raw_keys()
437 .ignore_err()
438 .map(RawPduId::from)
439 .for_each_concurrent(automatic_width(), async |root_id| {
440 self.index_thread_activity(root_id).await;
441 })
442 .await;
443
444 Ok(())
445 }
446
447 async fn index_thread_activity(&self, root_id: RawPduId) {
448 let root: PduId = root_id.into();
449
450 let replies = self
451 .services
452 .pdu_metadata
453 .get_relations(root.shortroomid, root.count, None, Direction::Backward, None)
454 .ready_filter_map(|(count, pdu)| {
455 pdu.get_content()
456 .is_ok_and(|content: ExtractThreadRelation| {
457 content.relates_to.rel_type == RelationType::Thread
458 })
459 .then_some(count)
460 });
461
462 let mut replies = pin!(replies);
463
464 let latest = replies.next().await.unwrap_or(root.count);
465
466 let activity_id: RawPduId = PduId {
467 shortroomid: root.shortroomid,
468 count: latest,
469 }
470 .into();
471
472 let mut txn = self.services.db.txn();
473
474 txn.insert_raw(&self.db.threadactivityid_rootid, activity_id, root_id);
475 txn.insert_raw(&self.db.threadrootid_latestcount, root_id, latest.to_be_bytes());
476 txn.execute();
477 }
478}