Skip to main content

tuwunel_service/rooms/pdu_metadata/
relations.rs

1use futures::{Stream, StreamExt};
2use ruma::{
3	EventId, OwnedUserId, UserId,
4	api::Direction,
5	events::{reaction::ReactionEventContent, relation::RelationType},
6};
7use tuwunel_core::{
8	PduId,
9	arrayvec::ArrayVec,
10	implement, is_equal_to,
11	matrix::{Event, Pdu, PduCount, RawPduId, event::RelationTypeEqual},
12	utils::{
13		stream::{ReadyExt, TryIgnore, WidebandExt},
14		u64_from_u8,
15	},
16};
17
18use super::Service;
19use crate::rooms::short::ShortRoomId;
20
21type StartKey = ArrayVec<u8, 16>;
22
23#[implement(Service)]
24#[tracing::instrument(skip(self, from, to), level = "debug")]
25pub fn add_relation(&self, from: PduCount, to: PduCount) {
26	const BUFSIZE: usize = size_of::<u64>() * 2;
27
28	match (from, to) {
29		| (PduCount::Normal(from), PduCount::Normal(to)) => {
30			let key: &[u64] = &[to, from];
31
32			self.db
33				.tofrom_relation
34				.aput_raw::<BUFSIZE, _, _>(key, []);
35		},
36		| _ => {}, // TODO: Relations with backfilled pdus
37	}
38}
39
40/// Query relations of an event to determine if matching any of the trailing
41/// arguments. When all criteria are None the mere presence of a relation causes
42/// this function to return true.
43#[implement(Service)]
44pub async fn event_has_relation(
45	&self,
46	event_id: &EventId,
47	user_id: Option<&UserId>,
48	rel_type: Option<&RelationType>,
49	key: Option<&str>,
50) -> bool {
51	let Ok(pdu_id) = self.services.timeline.get_pdu_id(event_id).await else {
52		return false;
53	};
54
55	self.has_relation(pdu_id.into(), user_id, rel_type, key)
56		.await
57}
58
59/// Query relations of an event by PduId to determine if matching any of the
60/// trailing arguments. When all criteria are None the mere presence of a
61/// relation causes this function to return true.
62#[implement(Service)]
63pub async fn has_relation(
64	&self,
65	target: PduId,
66	user_id: Option<&UserId>,
67	rel_type: Option<&RelationType>,
68	key: Option<&str>,
69) -> bool {
70	self.get_relations(target.shortroomid, target.count, None, Direction::Forward, None)
71		.ready_filter(|(_, pdu)| user_id.is_none_or(is_equal_to!(pdu.sender())))
72		.ready_filter(|(_, pdu)| {
73			debug_assert!(
74				key.is_none() || rel_type.is_none_or(is_equal_to!(&RelationType::Annotation)),
75				"key argument only applies to Annotation type relations."
76			);
77
78			// When key is supplied we don't need to double-parse the content here and
79			// below.
80			key.is_some() || rel_type.is_none_or(|rel_type| rel_type.relation_type_equal(&pdu))
81		})
82		.ready_filter(|(_, pdu)| {
83			key.is_none_or(|key| {
84				pdu.get_content()
85					.map(|content: ReactionEventContent| content.relates_to.key == key)
86					.unwrap_or(false)
87			})
88		})
89		.ready_any(|_| true)
90		.await
91}
92
93/// MSC3440 `related_by_*`: whether any event relates to `target` with a
94/// `rel_type` in `rel_types` and a `sender` in `senders`. An empty list is
95/// unconstrained on that axis; a single relating event must satisfy both.
96#[implement(Service)]
97pub async fn has_incoming_relation(
98	&self,
99	target: PduId,
100	senders: &[OwnedUserId],
101	rel_types: &[RelationType],
102) -> bool {
103	self.get_relations(target.shortroomid, target.count, None, Direction::Forward, None)
104		.ready_any(|(_, pdu)| {
105			let sender_matches =
106				senders.is_empty() || senders.iter().any(is_equal_to!(pdu.sender()));
107
108			let rel_type_matches = rel_types.is_empty()
109				|| rel_types
110					.iter()
111					.any(|rel_type| rel_type.relation_type_equal(&pdu));
112
113			sender_matches && rel_type_matches
114		})
115		.await
116}
117
118#[implement(Service)]
119pub fn get_relations<'a>(
120	&'a self,
121	shortroomid: ShortRoomId,
122	target: PduCount,
123	from: Option<PduCount>,
124	dir: Direction,
125	user_id: Option<&'a UserId>,
126) -> impl Stream<Item = (PduCount, Pdu)> + Send + '_ {
127	let target = target.to_be_bytes();
128	let from = from
129		.map(|from| from.saturating_inc(dir))
130		.unwrap_or_else(|| match dir {
131			| Direction::Backward => PduCount::max(),
132			| Direction::Forward => PduCount::default(),
133		})
134		.to_be_bytes();
135
136	let mut buf = StartKey::new();
137	let start = {
138		buf.extend(target);
139		buf.extend(from);
140		buf.as_slice()
141	};
142
143	match dir {
144		| Direction::Backward => self
145			.db
146			.tofrom_relation
147			.rev_raw_keys_from(start)
148			.left_stream(),
149
150		| Direction::Forward => self
151			.db
152			.tofrom_relation
153			.raw_keys_from(start)
154			.right_stream(),
155	}
156	.ignore_err()
157	.ready_take_while(move |key| key.starts_with(&target))
158	.map(|to_from| u64_from_u8(&to_from[8..16]))
159	.map(PduCount::from_unsigned)
160	.map(move |count| (user_id, shortroomid, count))
161	.wide_filter_map(async |(user_id, shortroomid, count)| {
162		let pdu_id: RawPduId = PduId { shortroomid, count }.into();
163		let mut pdu = self
164			.services
165			.timeline
166			.get_pdu_from_id(&pdu_id)
167			.await
168			.ok()?;
169
170		pdu.remove_transaction_id_unless_sender(user_id)
171			.ok()?;
172
173		Some((count, pdu))
174	})
175}