Skip to main content

tuwunel_service/rooms/read_receipt/
mod.rs

1mod data;
2#[cfg(test)]
3mod tests;
4
5use std::{collections::BTreeMap, sync::Arc};
6
7use futures::{Stream, StreamExt, TryStreamExt};
8use ruma::{
9	MilliSecondsSinceUnixEpoch, OwnedEventId, OwnedUserId, RoomId, UInt, UserId,
10	api::appservice::event::push_events::v1::EphemeralData,
11	events::{
12		AnySyncEphemeralRoomEvent, SyncEphemeralRoomEvent,
13		receipt::{
14			Receipt, ReceiptEvent, ReceiptEventContent, ReceiptThread, ReceiptType, Receipts,
15		},
16	},
17	serde::Raw,
18};
19use serde_json::value::to_raw_value;
20use tuwunel_core::{
21	Result, debug,
22	debug::INFO_SPAN_LEVEL,
23	err,
24	matrix::{
25		Event,
26		pdu::{PduCount, PduId, RawPduId},
27	},
28	result::NotFound,
29	smallstr::SmallString,
30	smallvec::SmallVec,
31	utils::{BoolExt, IterStream},
32	warn,
33};
34
35use self::data::{Data, ReceiptItem};
36
37/// Private read receipts surfaced by `private_read_get`. One legacy
38/// unthreaded row plus zero or more per-thread rows; inline-1 catches the
39/// dominant case (a single unthreaded marker) without a heap alloc.
40pub type PrivateReadEvents = SmallVec<[Raw<AnySyncEphemeralRoomEvent>; 1]>;
41
42/// Stored thread-kind tag: `""` for `Unthreaded`, `"main"` for `Main`, or
43/// the event-id string for `Thread(...)`. v3+ event ids are 44 bytes
44/// including the leading `$`; 48 bytes inline matches the project's
45/// `StateKey` budget and stays inline for every realistic thread root.
46type ThreadKind = SmallString<[u8; 48]>;
47
48/// A private read marker write for one `(room, user, thread)` context.
49///
50/// `count` is the timeline position the marker addresses and `ts` the receipt
51/// timestamp. `announce` opens the sync gate, carrying the marker to the
52/// user's other devices; a marker the server writes on the user's behalf
53/// leaves it closed.
54#[derive(Clone, Copy, Debug)]
55pub struct PrivateRead<'a> {
56	pub room_id: &'a RoomId,
57	pub user_id: &'a UserId,
58	pub count: u64,
59	pub ts: MilliSecondsSinceUnixEpoch,
60	pub thread: &'a ReceiptThread,
61	pub announce: bool,
62}
63
64pub struct Service {
65	services: Arc<crate::services::OnceServices>,
66	db: Data,
67}
68
69impl crate::Service for Service {
70	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
71		Ok(Arc::new(Self {
72			services: args.services.clone(),
73			db: Data::new(args),
74		}))
75	}
76
77	fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
78}
79
80impl Service {
81	/// Replaces the previous read receipt when the incoming one advances.
82	///
83	/// Returns whether the receipt was stored. A re-posted marker allocates no
84	/// stream position, so appservice and federation delivery are both skipped.
85	#[tracing::instrument(
86		name = "receipt"
87		level = INFO_SPAN_LEVEL,
88		skip_all,
89		fields(
90			%room_id,
91			%user_id,
92			?event.content
93		)
94	)]
95	pub async fn readreceipt_update(
96		&self,
97		user_id: &UserId,
98		room_id: &RoomId,
99		event: &ReceiptEvent,
100	) -> bool {
101		if self
102			.db
103			.readreceipt_update(user_id, room_id, event)
104			.await
105			.is_false()
106		{
107			return false;
108		}
109
110		self.services
111			.sending
112			.send_edu_room_appservices(room_id, |buf| {
113				let edu = EphemeralData::Receipt(ReceiptEvent {
114					content: event.content.clone(),
115					room_id: room_id.to_owned(),
116				});
117
118				Ok(serde_json::to_writer(buf, &edu)?)
119			})
120			.await
121			.expect("edu serialization or flush failed");
122
123		if self.services.globals.user_is_local(user_id) {
124			self.services
125				.sending
126				.flush_room(room_id)
127				.await
128				.expect("room flush failed");
129		}
130
131		true
132	}
133
134	/// Gets every stored private read receipt for `(room, user)`. Returns
135	/// one ephemeral event per stored row (legacy unthreaded plus per-thread
136	/// rows). An empty result means no marker is set.
137	#[tracing::instrument(skip(self), level = "debug", name = "get_private")]
138	pub async fn private_read_get(
139		&self,
140		room_id: &RoomId,
141		user_id: &UserId,
142	) -> Result<PrivateReadEvents> {
143		let shortroomid = self
144			.services
145			.short
146			.get_shortroomid(room_id)
147			.await
148			.map_err(|e| {
149				err!(Database(warn!(
150					"Short room ID does not exist in database for {room_id}: {e}"
151				)))
152			})?;
153
154		let legacy = self
155			.private_read_get_count(room_id, user_id)
156			.await
157			.ok()
158			.map(|(count, ts)| (ThreadKind::new(), count, ts));
159
160		let events = legacy
161			.into_iter()
162			.stream()
163			.chain(
164				self.db
165					.private_read_threaded_stream(room_id, user_id),
166			)
167			.filter_map(async |(kind, count, ts)| {
168				self.build_private_read_event(shortroomid, count, ts, user_id, &kind)
169					.await
170			})
171			.collect()
172			.await;
173
174		Ok(events)
175	}
176
177	/// Gets the complete announced private read snapshot for `update` without
178	/// suppressing malformed rows.
179	///
180	/// A snapshot older than the gate predates this durable mirror, so those
181	/// rows fall back to the tolerant active-state read they always published
182	/// rather than withholding the whole room after an upgrade. A newer
183	/// snapshot means a concurrent announce; that fails the bounded room range
184	/// so its cursor remains pinned. A marker naming an event that no longer
185	/// resolves is skipped rather than failing the range, which would
186	/// otherwise repeat on every request and withhold the room indefinitely.
187	#[tracing::instrument(skip(self), level = "debug", name = "get_private_fallible")]
188	pub async fn private_read_get_fallible(
189		&self,
190		room_id: &RoomId,
191		user_id: &UserId,
192		update: u64,
193	) -> Result<PrivateReadEvents> {
194		let snapshot = self
195			.db
196			.private_read_sync_update_fallible(user_id, room_id)
197			.await?;
198
199		if snapshot < update {
200			debug!(%room_id, %user_id, "Serving pre-mirror private read from the active store.");
201			return self.private_read_get(room_id, user_id).await;
202		}
203
204		if snapshot > update {
205			return Err(err!(Database(
206				"Private read snapshot advanced while assembling a bounded sync range."
207			)));
208		}
209
210		let shortroomid = async {
211			self.services
212				.short
213				.get_shortroomid(room_id)
214				.await
215				.map_err(|e| {
216					err!(Database(warn!(
217						"Short room ID does not exist in database for {room_id}: {e}"
218					)))
219				})
220		};
221
222		let shortroomid = shortroomid.await?;
223		let events = self
224			.db
225			.private_read_sync_stream_fallible(room_id, user_id)
226			.try_filter_map(async |(kind, count, ts)| {
227				self.build_private_read_event_skippable(shortroomid, count, ts, user_id, &kind)
228					.await
229			})
230			.try_collect()
231			.await?;
232
233		let confirmed = self
234			.db
235			.private_read_sync_update_fallible(user_id, room_id)
236			.await?;
237
238		if confirmed != update {
239			return Err(err!(Database(
240				"Private read snapshot changed while assembling a bounded sync range."
241			)));
242		}
243
244		Ok(events)
245	}
246
247	/// Builds one announced private read row, skipping a marker whose event
248	/// no longer resolves.
249	///
250	/// The absent case returns `Ok(None)` so the bounded range assembles
251	/// without it. Decode failures and invalid timestamps still fail the
252	/// range as real inconsistencies.
253	async fn build_private_read_event_skippable(
254		&self,
255		shortroomid: u64,
256		count: u64,
257		ts: Option<u64>,
258		user_id: &UserId,
259		thread_kind: &str,
260	) -> Result<Option<Raw<AnySyncEphemeralRoomEvent>>> {
261		let skip = || {
262			debug!(
263				count,
264				thread_kind,
265				%user_id,
266				"Skipping a private read marker naming a missing event."
267			);
268
269			None
270		};
271
272		self.build_private_read_event_fallible(shortroomid, count, ts, user_id, thread_kind)
273			.await
274			.optional()
275			.map(|event| event.or_else(skip))
276	}
277
278	async fn build_private_read_event(
279		&self,
280		shortroomid: u64,
281		count: u64,
282		ts: Option<u64>,
283		user_id: &UserId,
284		thread_kind: &str,
285	) -> Option<Raw<AnySyncEphemeralRoomEvent>> {
286		let thread = thread_kind_to_receipt(thread_kind).unwrap_or(ReceiptThread::Unthreaded);
287		let ts = ts
288			.and_then(UInt::new)
289			.map(MilliSecondsSinceUnixEpoch);
290
291		self.build_private_read_event_from(shortroomid, count, ts, user_id, thread)
292			.await
293			.ok()
294	}
295
296	async fn build_private_read_event_fallible(
297		&self,
298		shortroomid: u64,
299		count: u64,
300		ts: Option<u64>,
301		user_id: &UserId,
302		thread_kind: &str,
303	) -> Result<Raw<AnySyncEphemeralRoomEvent>> {
304		let thread = thread_kind_to_receipt(thread_kind)?;
305		let ts = ts
306			.map(|ts| {
307				UInt::new(ts)
308					.map(MilliSecondsSinceUnixEpoch)
309					.ok_or_else(|| err!(Database("Invalid private receipt timestamp {ts}.")))
310			})
311			.transpose()?;
312
313		self.build_private_read_event_from(shortroomid, count, ts, user_id, thread)
314			.await
315	}
316
317	async fn build_private_read_event_from(
318		&self,
319		shortroomid: u64,
320		count: u64,
321		ts: Option<MilliSecondsSinceUnixEpoch>,
322		user_id: &UserId,
323		thread: ReceiptThread,
324	) -> Result<Raw<AnySyncEphemeralRoomEvent>> {
325		let pdu_id: RawPduId = PduId {
326			shortroomid,
327			count: PduCount::Normal(count),
328		}
329		.into();
330		let pdu = self
331			.services
332			.timeline
333			.get_pdu_from_id(&pdu_id)
334			.await?;
335
336		let event_id: OwnedEventId = pdu.event_id().to_owned();
337		let user_id: OwnedUserId = user_id.to_owned();
338		let content: BTreeMap<OwnedEventId, Receipts> = BTreeMap::from_iter([(
339			event_id,
340			BTreeMap::from_iter([(
341				ReceiptType::ReadPrivate,
342				BTreeMap::from_iter([(user_id, Receipt { ts, thread })]),
343			)]),
344		)]);
345
346		let receipt_event_content = ReceiptEventContent(content);
347		let receipt_sync_event = SyncEphemeralRoomEvent { content: receipt_event_content };
348		let event = to_raw_value(&receipt_sync_event)?;
349
350		Ok(Raw::from_json(event))
351	}
352
353	/// Returns an iterator over the most recent read_receipts in a room that
354	/// happened after the event with id `since`.
355	#[tracing::instrument(skip(self), level = "debug")]
356	pub fn readreceipts_since<'a>(
357		&'a self,
358		room_id: &'a RoomId,
359		since: u64,
360		to: Option<u64>,
361	) -> impl Stream<Item = ReceiptItem<'_>> + Send + 'a {
362		self.db.readreceipts_since(room_id, since, to)
363	}
364
365	/// Returns read receipts in a bounded room range without suppressing
366	/// failures.
367	///
368	/// The lower bound is exclusive and the optional upper bound is inclusive.
369	/// Cursor, decode, and serialization failures remain in the stream for an
370	/// atomic caller to handle.
371	#[tracing::instrument(skip(self), level = "debug")]
372	pub fn readreceipts_since_fallible<'a>(
373		&'a self,
374		room_id: &'a RoomId,
375		since: u64,
376		to: Option<u64>,
377	) -> impl Stream<Item = Result<ReceiptItem<'_>>> + Send + 'a {
378		self.db
379			.readreceipts_since_fallible(room_id, since, to)
380	}
381
382	/// Sets a private read marker at PDU `count` for the given thread.
383	///
384	/// Unthreaded writes supersede prior per-thread rows so the room-wide
385	/// receipt subsumes thread state. Returns whether the marker advanced; a
386	/// position at or behind the stored one writes nothing.
387	#[tracing::instrument(skip(self), level = "debug", name = "set_private")]
388	pub async fn private_read_set(&self, private_read: PrivateRead<'_>) -> bool {
389		self.db.private_read_set(private_read).await
390	}
391
392	/// Returns the private read marker PDU count.
393	#[tracing::instrument(
394		name = "get_private_count",
395		level = "debug",
396		skip(self),
397		ret(level = "trace")
398	)]
399	pub async fn private_read_get_count(
400		&self,
401		room_id: &RoomId,
402		user_id: &UserId,
403	) -> Result<(u64, Option<u64>)> {
404		self.db
405			.private_read_get_count(room_id, user_id)
406			.await
407	}
408
409	/// Returns the announced unthreaded private read marker PDU count.
410	#[tracing::instrument(
411		name = "get_private_sync_count",
412		level = "debug",
413		skip(self),
414		ret(level = "trace")
415	)]
416	pub async fn private_read_sync_get_count(
417		&self,
418		room_id: &RoomId,
419		user_id: &UserId,
420	) -> Result<(u64, Option<u64>)> {
421		self.db
422			.private_read_sync_get_count(room_id, user_id)
423			.await
424	}
425
426	/// Returns the PDU count of the last private read update in this room.
427	///
428	/// Missing or unreadable update rows return zero for legacy callers.
429	/// Bounded sync callers use the fallible variant below so failures retain
430	/// the room cursor.
431	#[tracing::instrument(
432		name = "get_private_last",
433		level = "debug",
434		skip(self),
435		ret(level = "trace")
436	)]
437	pub async fn last_privateread_update(&self, user_id: &UserId, room_id: &RoomId) -> u64 {
438		self.db
439			.last_privateread_update(user_id, room_id)
440			.await
441	}
442
443	/// Returns the bounded-sync token for the last private read update.
444	///
445	/// A missing token is returned as zero. Database and decode failures are
446	/// preserved so a caller can retain its room cursor and retry the complete
447	/// range.
448	#[tracing::instrument(
449		name = "get_private_last_fallible",
450		level = "debug",
451		skip(self),
452		ret(level = "trace")
453	)]
454	pub async fn last_privateread_update_fallible(
455		&self,
456		user_id: &UserId,
457		room_id: &RoomId,
458	) -> Result<u64> {
459		self.db
460			.last_privateread_update_fallible(user_id, room_id)
461			.await
462	}
463
464	pub async fn delete_all_read_receipts(&self, room_id: &RoomId) -> Result {
465		self.db.delete_all_read_receipts(room_id).await
466	}
467}
468
469/// Reverse of `ReceiptThread::as_str`: parse a stored thread tag into the
470/// enum. Empty string maps to `Unthreaded`; `"main"` to `Main`; all other
471/// values must be valid event IDs.
472fn thread_kind_to_receipt(thread_kind: &str) -> Result<ReceiptThread> {
473	match thread_kind {
474		| "" => Ok(ReceiptThread::Unthreaded),
475		| "main" => Ok(ReceiptThread::Main),
476		| _ => OwnedEventId::try_from(thread_kind)
477			.map(ReceiptThread::Thread)
478			.map_err(|error| err!(Database("Invalid private receipt thread: {error}"))),
479	}
480}
481
482/// Packs read receipts into one sync event without suppressing malformed input.
483///
484/// Every input must deserialize as a receipt event. The caller receives any
485/// parse or serialization failure and can retain the bounded room cursor
486/// instead of publishing a partial event.
487pub fn pack_receipts_fallible<I>(
488	mut receipts: I,
489) -> Result<Raw<SyncEphemeralRoomEvent<ReceiptEventContent>>>
490where
491	I: Iterator<Item = Raw<AnySyncEphemeralRoomEvent>>,
492{
493	let json = receipts.try_fold(BTreeMap::new(), |mut json, value| -> Result<_> {
494		let value = serde_json::from_str::<SyncEphemeralRoomEvent<ReceiptEventContent>>(
495			value.json().get(),
496		)?;
497
498		for (event, receipt) in value.content {
499			json.insert(event, receipt);
500		}
501
502		Ok(json)
503	})?;
504
505	let content = ReceiptEventContent(json);
506	let event = to_raw_value(&SyncEphemeralRoomEvent { content })?;
507
508	Ok(Raw::from_json(event))
509}