1mod append;
8mod backfill;
9mod build;
10mod create;
11mod pdus;
12mod purge;
13mod redact;
14
15#[cfg(test)]
16mod tests;
17
18use std::{fmt::Write, sync::Arc};
19
20use async_trait::async_trait;
21use futures::{
22 FutureExt, StreamExt, TryFutureExt, TryStreamExt,
23 future::{Either, select, select_ok},
24 pin_mut,
25};
26use ruma::{
27 CanonicalJsonObject, EventId, MilliSecondsSinceUnixEpoch, OwnedEventId, OwnedRoomId, RoomId,
28 UserId, api::Direction, events::room::encrypted::Relation,
29};
30use serde::Deserialize;
31pub use tuwunel_core::matrix::pdu::{PduId, RawPduId};
35use tuwunel_core::{
36 Err, Error, Result, at, err, implement,
37 matrix::{
38 ShortEventId,
39 pdu::{PduCount, PduEvent},
40 },
41 utils::{
42 MutexMap, MutexMapGuard,
43 future::TryExtExt,
44 result::{LogErr, NotFound},
45 stream::{IterStream, TryReadyExt},
46 },
47 warn,
48};
49use tuwunel_database::{Database, Deserialized, Json, Map, Txn};
50
51pub use self::pdus::{PdusIterItem, bias_count};
55use crate::rooms::{
56 short::{ShortRoomId, ShortStateHash},
57 state_res::FetchEvent,
58};
59
60pub struct Service {
66 services: Arc<crate::services::OnceServices>,
67 db: Data,
68 pub mutex_insert: RoomMutexMap,
73}
74
75struct Data {
76 eventid_outlierpdu: Arc<Map>,
77 eventid_pduid: Arc<Map>,
78 pduid_pdu: Arc<Map>,
79 roomid_tscount_pducount: Arc<Map>,
80 db: Arc<Database>,
81}
82
83#[derive(Deserialize)]
85struct ExtractRelatesTo {
86 #[serde(rename = "m.relates_to")]
87 relates_to: Relation,
88}
89
90#[derive(Clone, Debug, Deserialize)]
91struct ExtractEventId {
92 event_id: OwnedEventId,
93}
94#[derive(Clone, Debug, Deserialize)]
95struct ExtractRelatesToEventId {
96 #[serde(rename = "m.relates_to")]
97 relates_to: ExtractEventId,
98}
99
100#[derive(Deserialize)]
101struct ExtractBody {
102 body: Option<String>,
103}
104
105type RoomMutexMap = MutexMap<OwnedRoomId, ()>;
106pub type RoomMutexGuard = MutexMapGuard<OwnedRoomId, ()>;
111
112#[async_trait]
113impl crate::Service for Service {
114 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
115 Ok(Arc::new(Self {
116 services: args.services.clone(),
117 db: Data {
118 eventid_outlierpdu: args.db["eventid_outlierpdu"].clone(),
119 eventid_pduid: args.db["eventid_pduid"].clone(),
120 pduid_pdu: args.db["pduid_pdu"].clone(),
121 roomid_tscount_pducount: args.db["roomid_tscount_pducount"].clone(),
122 db: args.db.clone(),
123 },
124 mutex_insert: RoomMutexMap::new(),
125 }))
126 }
127
128 async fn memory_usage(&self, out: &mut (dyn Write + Send)) -> Result {
129 let mutex_insert = self.mutex_insert.len();
130 writeln!(out, "- insert_mutex: {mutex_insert}")?;
131
132 Ok(())
133 }
134
135 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
136}
137
138#[implement(Service)]
143#[tracing::instrument(skip(self), level = "debug")]
144pub async fn replace_pdu(&self, pdu_id: &RawPduId, pdu_json: &CanonicalJsonObject) -> Result {
145 if self.db.pduid_pdu.get(pdu_id).await.is_not_found() {
146 return Err!(Request(NotFound("PDU does not exist.")));
147 }
148
149 self.db.pduid_pdu.raw_put(pdu_id, Json(pdu_json));
150
151 Ok(())
152}
153
154#[implement(Service)]
158pub(super) fn stage_replace_pdu(
159 &self,
160 txn: &mut Txn,
161 pdu_id: &RawPduId,
162 pdu_json: &CanonicalJsonObject,
163) {
164 txn.raw_put(&self.db.pduid_pdu, pdu_id, Json(pdu_json));
165}
166
167#[implement(Service)]
172#[tracing::instrument(skip(self, pdu), level = "debug")]
173pub fn add_pdu_outlier(&self, event_id: &EventId, pdu: &CanonicalJsonObject) {
174 self.db
175 .eventid_outlierpdu
176 .raw_put(event_id, Json(pdu));
177}
178
179#[implement(Service)]
183#[tracing::instrument(skip(self), level = "debug")]
184pub async fn first_pdu_in_room(&self, room_id: &RoomId) -> Result<PduEvent> {
185 self.first_item_in_room(room_id).await.map(at!(1))
186}
187
188#[implement(Service)]
193#[tracing::instrument(skip(self), level = "debug")]
194#[inline]
195pub async fn latest_pdu_in_room(&self, room_id: &RoomId) -> Result<PduEvent> {
196 self.latest_item_in_room(None, room_id).await
197}
198
199#[implement(Service)]
204#[tracing::instrument(skip(self), level = "debug")]
205pub async fn first_item_in_room(&self, room_id: &RoomId) -> Result<(PduCount, PduEvent)> {
206 let pdus = self.pdus(None, room_id, None);
207
208 pin_mut!(pdus);
209 pdus.try_next()
210 .await?
211 .ok_or_else(|| err!(Request(NotFound("No PDU found in room"))))
212}
213
214#[implement(Service)]
219#[tracing::instrument(skip(self), level = "debug")]
220pub async fn latest_item_in_room(
221 &self,
222 sender_user: Option<&UserId>,
223 room_id: &RoomId,
224) -> Result<PduEvent> {
225 let pdus_rev = self.pdus_rev(sender_user, room_id, None);
226
227 pin_mut!(pdus_rev);
228 pdus_rev
229 .try_next()
230 .await?
231 .map(at!(1))
232 .ok_or_else(|| err!(Request(NotFound("No PDU's found in room"))))
233}
234
235#[implement(Service)]
240#[tracing::instrument(skip(self), level = "debug")]
241pub async fn prev_shortstatehash(
242 &self,
243 room_id: &RoomId,
244 before: PduCount,
245) -> Result<ShortStateHash> {
246 let shortroomid: ShortRoomId = self
247 .services
248 .short
249 .get_shortroomid(room_id)
250 .await
251 .map_err(|e| err!(Request(NotFound("Room {room_id:?} not found: {e:?}"))))?;
252
253 let before = PduId { shortroomid, count: before };
254
255 let prev = PduId {
256 shortroomid,
257 count: self.prev_timeline_count(&before).await?,
258 };
259
260 let shorteventid = self.get_shorteventid_from_pdu_id(&prev).await?;
261
262 self.services
263 .state
264 .get_shortstatehash(shorteventid)
265 .await
266}
267
268#[implement(Service)]
273#[tracing::instrument(skip(self), level = "debug")]
274pub async fn next_shortstatehash(
275 &self,
276 room_id: &RoomId,
277 after: PduCount,
278) -> Result<ShortStateHash> {
279 let shortroomid: ShortRoomId = self
280 .services
281 .short
282 .get_shortroomid(room_id)
283 .await
284 .map_err(|e| err!(Request(NotFound("Room {room_id:?} not found: {e:?}"))))?;
285
286 let after = PduId { shortroomid, count: after };
287
288 let next = PduId {
289 shortroomid,
290 count: self.next_timeline_count(&after).await?,
291 };
292
293 let shorteventid = self.get_shorteventid_from_pdu_id(&next).await?;
294
295 self.services
296 .state
297 .get_shortstatehash(shorteventid)
298 .await
299}
300
301#[implement(Service)]
309#[tracing::instrument(skip(self), level = "debug")]
310pub async fn shortstatehash_after(
311 &self,
312 room_id: &RoomId,
313 count: PduCount,
314) -> Result<ShortStateHash> {
315 let shortroomid: ShortRoomId = self
316 .services
317 .short
318 .get_shortroomid(room_id)
319 .map_err(|e| err!(Request(NotFound("Room {room_id:?} not found: {e:?}"))))
320 .await?;
321
322 let after = PduId { shortroomid, count };
323 let count = match self.next_timeline_count(&after).await {
324 | Ok(count) => count,
325 | Err(e) if !e.is_not_found() => return Err(e),
326 | Err(_) => {
327 return self
328 .services
329 .state
330 .get_room_shortstatehash(room_id)
331 .await;
332 },
333 };
334
335 let next = PduId { shortroomid, count };
336 let shorteventid = self.get_shorteventid_from_pdu_id(&next).await?;
337
338 self.services
339 .state
340 .get_shortstatehash(shorteventid)
341 .await
342}
343
344#[implement(Service)]
349#[tracing::instrument(skip(self), level = "debug")]
350pub async fn get_shortstatehash(
351 &self,
352 room_id: &RoomId,
353 count: PduCount,
354) -> Result<ShortStateHash> {
355 let shortroomid: ShortRoomId = self
356 .services
357 .short
358 .get_shortroomid(room_id)
359 .await
360 .map_err(|e| err!(Request(NotFound("Room {room_id:?} not found: {e:?}"))))?;
361
362 let pdu_id = PduId { shortroomid, count };
363
364 let shorteventid = self.get_shorteventid_from_pdu_id(&pdu_id).await?;
365
366 self.services
367 .state
368 .get_shortstatehash(shorteventid)
369 .await
370}
371
372#[implement(Service)]
377#[tracing::instrument(skip(self), level = "debug")]
378pub async fn prev_timeline_count(&self, before: &PduId) -> Result<PduCount> {
379 let before = Self::pdu_count_to_id(before.shortroomid, before.count, Direction::Backward);
380
381 let pdu_ids = self
382 .db
383 .pduid_pdu
384 .rev_keys_raw_from(&before)
385 .ready_try_take_while(|pdu_id: &RawPduId| Ok(pdu_id.is_room_eq(before)))
386 .ready_and_then(|pdu_id: RawPduId| Ok(pdu_id.pdu_count()));
387
388 pin_mut!(pdu_ids);
389 pdu_ids
390 .try_next()
391 .await
392 .log_err()?
393 .ok_or_else(|| err!(Request(NotFound("No earlier PDU's found in room"))))
394}
395
396#[implement(Service)]
401#[tracing::instrument(skip(self), level = "debug")]
402pub async fn next_timeline_count(&self, after: &PduId) -> Result<PduCount> {
403 let after = Self::pdu_count_to_id(after.shortroomid, after.count, Direction::Forward);
404
405 let pdu_ids = self
406 .db
407 .pduid_pdu
408 .keys_raw_from(&after)
409 .ready_try_take_while(|pdu_id: &RawPduId| Ok(pdu_id.is_room_eq(after)))
410 .ready_and_then(|pdu_id: RawPduId| Ok(pdu_id.pdu_count()));
411
412 pin_mut!(pdu_ids);
413 pdu_ids
414 .try_next()
415 .await
416 .log_err()?
417 .ok_or(err!(Request(NotFound("No more PDU's found in room"))))
418}
419
420#[implement(Service)]
426#[tracing::instrument(skip(self), level = "debug")]
427pub async fn last_timeline_count(
428 &self,
429 sender_user: Option<&UserId>,
430 room_id: &RoomId,
431 upper_bound: Option<PduCount>,
432) -> Result<PduCount> {
433 let upper_bound = upper_bound.unwrap_or_else(PduCount::max);
434 let pdus_rev = self.pdus_rev(sender_user, room_id, None);
435
436 pin_mut!(pdus_rev);
437 let last_count = pdus_rev
438 .ready_try_skip_while(|&(pducount, _)| Ok(pducount > upper_bound))
439 .try_next()
440 .await?
441 .map(at!(0))
442 .filter(|&count| matches!(count, PduCount::Normal(_)))
443 .unwrap_or_else(PduCount::max);
444
445 Ok(last_count)
446}
447
448#[implement(Service)]
454pub async fn get_event_id_near_ts(
455 &self,
456 room_id: &RoomId,
457 ts: MilliSecondsSinceUnixEpoch,
458 dir: Direction,
459) -> Result<(MilliSecondsSinceUnixEpoch, OwnedEventId)> {
460 self.get_pdu_id_near_ts(room_id, ts, dir)
461 .and_then(async |(ts, pdu_id)| {
462 self.get_event_id_from_pdu_id(&pdu_id)
463 .map_ok(|event_id| (ts, event_id))
464 .await
465 })
466 .await
467}
468
469#[implement(Service)]
475pub async fn get_pdu_id_near_ts(
476 &self,
477 room_id: &RoomId,
478 ts: MilliSecondsSinceUnixEpoch,
479 dir: Direction,
480) -> Result<(MilliSecondsSinceUnixEpoch, PduId)> {
481 let pdu_ids = self.pdu_ids_near_ts(room_id, ts, dir);
482
483 pin_mut!(pdu_ids);
484 pdu_ids
485 .try_next()
486 .await?
487 .ok_or_else(|| err!(Request(NotFound("No event found near this timestamp."))))
488}
489
490#[implement(Service)]
496pub async fn get_pdu_near_ts(
497 &self,
498 _user_id: Option<&UserId>,
499 room_id: &RoomId,
500 ts: MilliSecondsSinceUnixEpoch,
501 dir: Direction,
502) -> Result<PdusIterItem> {
503 let pdus = self
504 .pdu_ids_near_ts(room_id, ts, dir)
505 .map_ok(|(ts, pdu_id)| (ts, pdu_id.into()))
506 .and_then(async |(_, pdu_id): (_, RawPduId)| {
507 self.get_pdu_from_id(&pdu_id)
508 .map_ok(|pdu| (pdu_id.pdu_count(), pdu))
509 .await
510 });
511
512 pin_mut!(pdus);
513 pdus.try_next()
514 .await?
515 .ok_or_else(|| err!(Request(NotFound("No event found near this timestamp."))))
516}
517
518#[implement(Service)]
519async fn count_to_id(
520 &self,
521 room_id: &RoomId,
522 count: PduCount,
523 dir: Direction,
524) -> Result<RawPduId> {
525 let shortroomid: ShortRoomId = self
526 .services
527 .short
528 .get_shortroomid(room_id)
529 .await
530 .map_err(|e| err!(Request(NotFound("Room {room_id:?} not found: {e:?}"))))?;
531
532 Ok(Self::pdu_count_to_id(shortroomid, count, dir))
533}
534
535#[implement(Service)]
536fn pdu_count_to_id(shortroomid: ShortRoomId, count: PduCount, dir: Direction) -> RawPduId {
537 let count = match (count, dir) {
540 | (PduCount::Backfilled(0), Direction::Forward) => count,
541 | _ => count.saturating_inc(dir),
542 };
543
544 let pdu_id = PduId { shortroomid, count };
545
546 pdu_id.into()
547}
548
549#[implement(Service)]
554pub async fn get_pdu_from_shorteventid(&self, shorteventid: ShortEventId) -> Result<PduEvent> {
555 let event_id: OwnedEventId = self
556 .services
557 .short
558 .get_eventid_from_short(shorteventid)
559 .await?;
560
561 self.get_pdu(&event_id).await
562}
563
564#[implement(Service)]
569pub async fn get_pdu(&self, event_id: &EventId) -> Result<PduEvent> { self.get(event_id).await }
570
571#[implement(Service)]
576pub async fn get_outlier_pdu(&self, event_id: &EventId) -> Result<PduEvent> {
577 self.get_outlier(event_id).await
578}
579
580#[implement(Service)]
585pub async fn get_non_outlier_pdu(&self, event_id: &EventId) -> Result<PduEvent> {
586 self.get_non_outlier(event_id).await
587}
588
589#[implement(Service)]
594pub async fn get_pdu_from_id(&self, pdu_id: &RawPduId) -> Result<PduEvent> {
595 self.get_from_id(pdu_id).await
596}
597
598#[implement(Service)]
603pub async fn get_pdu_json(&self, event_id: &EventId) -> Result<CanonicalJsonObject> {
604 self.get(event_id).await
605}
606
607#[implement(Service)]
612pub async fn get_outlier_pdu_json(&self, event_id: &EventId) -> Result<CanonicalJsonObject> {
613 self.get_outlier(event_id).await
614}
615
616#[implement(Service)]
621pub async fn get_non_outlier_pdu_json(&self, event_id: &EventId) -> Result<CanonicalJsonObject> {
622 self.get_non_outlier(event_id).await
623}
624
625#[implement(Service)]
630pub async fn get_pdu_json_from_id(&self, pdu_id: &RawPduId) -> Result<CanonicalJsonObject> {
631 self.get_from_id(pdu_id).await
632}
633
634#[implement(Service)]
639#[inline]
640pub async fn get<T>(&self, event_id: &EventId) -> Result<T>
641where
642 T: for<'de> Deserialize<'de>,
643{
644 let accepted = self.get_non_outlier(event_id);
645 let outlier = self.get_outlier(event_id);
646
647 pin_mut!(accepted, outlier);
648 select_ok([accepted.left_future(), outlier.right_future()])
649 .await
650 .map(at!(0))
651}
652
653impl FetchEvent for &Service {
654 async fn get<T>(self, event_id: &EventId) -> Result<T>
655 where
656 T: for<'de> Deserialize<'de> + Send,
657 {
658 Service::get(self, event_id).await
659 }
660
661 async fn exists(self, event_id: &EventId) -> Result<bool> {
662 let non_outlier = self.non_outlier_pdu_exists(event_id);
663 let outlier = self.outlier_pdu_exists(event_id);
664 let classify = |first: Error, second: Result| match second {
665 | Ok(()) => Ok(true),
666 | Err(second) if first.is_not_found() && second.is_not_found() => Ok(false),
667 | Err(second) if first.is_not_found() => Err(second),
668 | Err(_) => Err(first),
669 };
670
671 pin_mut!(non_outlier, outlier);
672 match select(non_outlier, outlier).await {
673 | Either::Left((Ok(()), _)) | Either::Right((Ok(()), _)) => Ok(true),
674 | Either::Left((Err(first), second)) => classify(first, second.await),
675 | Either::Right((Err(first), second)) => classify(first, second.await),
676 }
677 }
678}
679
680#[implement(Service)]
685#[inline]
686pub async fn get_outlier<T>(&self, event_id: &EventId) -> Result<T>
687where
688 T: for<'de> Deserialize<'de>,
689{
690 self.db
691 .eventid_outlierpdu
692 .get(event_id)
693 .await
694 .deserialized()
695}
696
697#[implement(Service)]
702#[inline]
703pub async fn get_non_outlier<T>(&self, event_id: &EventId) -> Result<T>
704where
705 T: for<'de> Deserialize<'de>,
706{
707 let pdu_id = self.get_pdu_id(event_id).await?;
708
709 self.get_from_id(&pdu_id).await
710}
711
712#[implement(Service)]
717#[inline]
718pub async fn get_from_id<T>(&self, pdu_id: &RawPduId) -> Result<T>
719where
720 T: for<'de> Deserialize<'de>,
721{
722 self.db.pduid_pdu.get(pdu_id).await.deserialized()
723}
724
725#[implement(Service)]
730pub async fn pdu_exists<'a>(&'a self, event_id: &'a EventId) -> bool {
731 let non_outlier = self.non_outlier_pdu_exists(event_id);
732 let outlier = self.outlier_pdu_exists(event_id);
733
734 pin_mut!(non_outlier, outlier);
735 select_ok([non_outlier.left_future(), outlier.right_future()])
736 .await
737 .map(at!(0))
738 .is_ok()
739}
740
741#[implement(Service)]
748pub async fn non_outlier_pdus_exist<'a, I>(&self, event_ids: I) -> bool
749where
750 I: Iterator<Item = &'a EventId> + Send,
751{
752 event_ids
753 .stream()
754 .all(|event_id| self.non_outlier_pdu_exists(event_id).is_ok())
755 .await
756}
757
758#[implement(Service)]
765pub fn watch_event<'a>(&'a self, event_id: &EventId) -> impl Future<Output = ()> + Send + 'a {
766 self.db
767 .eventid_pduid
768 .watch_raw_prefix_once(event_id)
769}
770
771#[implement(Service)]
776pub async fn non_outlier_pdu_exists(&self, event_id: &EventId) -> Result {
777 let pduid = self.get_pdu_id(event_id).await?;
778
779 self.db.pduid_pdu.exists(&pduid).await
780}
781
782#[implement(Service)]
787#[inline]
788pub async fn outlier_pdu_exists(&self, event_id: &EventId) -> Result {
789 self.db.eventid_outlierpdu.exists(event_id).await
790}
791
792#[implement(Service)]
796pub async fn get_pdu_count(&self, event_id: &EventId) -> Result<PduCount> {
797 self.get_pdu_id(event_id)
798 .await
799 .map(RawPduId::pdu_count)
800}
801
802#[implement(Service)]
807pub async fn get_shorteventid_from_pdu_id(&self, pdu_id: &PduId) -> Result<ShortEventId> {
808 let event_id = self.get_event_id_from_pdu_id(pdu_id).await?;
809
810 self.services
811 .short
812 .get_shorteventid(&event_id)
813 .await
814}
815
816#[implement(Service)]
820pub async fn get_event_id_from_pdu_id(&self, pdu_id: &PduId) -> Result<OwnedEventId> {
821 let pdu_id: RawPduId = (*pdu_id).into();
822
823 self.get_pdu_from_id(&pdu_id)
824 .map_ok(|pdu| pdu.event_id)
825 .await
826}
827
828#[implement(Service)]
833pub async fn get_pdu_id_from_shorteventid(&self, shorteventid: ShortEventId) -> Result<RawPduId> {
834 let event_id: OwnedEventId = self
835 .services
836 .short
837 .get_eventid_from_short(shorteventid)
838 .await?;
839
840 self.get_pdu_id(&event_id).await
841}
842
843#[implement(Service)]
848pub async fn get_pdu_id(&self, event_id: &EventId) -> Result<RawPduId> {
849 self.db
850 .eventid_pduid
851 .get(event_id)
852 .await
853 .map(|handle| RawPduId::from(&*handle))
854}