Skip to main content

tuwunel_service/rooms/threads/
mod.rs

1use 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
27/// Maximum relation hops walked when resolving thread membership, per
28/// the Matrix v1.4 spec recommendation (also MSC3771/MSC3773).
29const 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
58/// Persist a latest event whose embedded sender and event ID come from the
59/// same validated event used to index the thread activity.
60fn 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	/// Resolves the thread root for `event` by walking up `m.relates_to`
124	/// links, bounded at `MAX_THREAD_HOPS`. Returns `None` for events
125	/// that belong to the main timeline. Redaction events carry no
126	/// `m.relates_to` of their own; their thread is resolved from the
127	/// redacted target event per MSC3771/MSC3773.
128	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	/// Resolve a redaction event's thread by looking through to the
159	/// redacted target. Returns `None` for non-redaction events and for
160	/// redactions whose target is unknown or carries no thread relation.
161	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	/// `get_thread_id` for an event referenced by id; events missing
189	/// locally resolve to `None` (the main timeline).
190	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(&current)
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	/// Resolve one activity row to its thread root, skipping and reaping rows
309	/// the validity pointer has left behind.
310	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			// A row ahead of the pointer is a write in flight; only rows behind
330			// the pointer are dead and safe to reap.
331			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	/// MSC3816: whether `user_id` has participated in the thread rooted at
393	/// `root_event_id`, having sent the root event or a threaded reply to it.
394	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	/// Rebuild the thread activity index from every thread root. Run once at
428	/// startup behind a `global` marker, and on demand from the admin command.
429	/// Clears first so a partial or stale index is replaced wholesale.
430	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}