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
37pub type PrivateReadEvents = SmallVec<[Raw<AnySyncEphemeralRoomEvent>; 1]>;
41
42type ThreadKind = SmallString<[u8; 48]>;
47
48#[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 #[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 #[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 #[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 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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
469fn 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
482pub 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}