Skip to main content

tuwunel_api/client/sync/v5/
extensions.rs

1mod account_data;
2mod e2ee;
3mod profiles;
4mod receipts;
5mod to_device;
6mod typing;
7
8use std::{collections::BTreeMap, fmt::Debug};
9
10use futures::{FutureExt, future::join4};
11use ruma::{
12	OwnedRoomId, RoomId,
13	api::client::sync::sync_events::v5::{
14		ListId,
15		request::ExtensionRoomConfig,
16		response::{Extensions, Room as ResponseRoom},
17	},
18};
19use tuwunel_core::{Result, apply, at, extract_variant, utils::BoolExt};
20use tuwunel_service::sync::Connection;
21
22use self::{
23	account_data::{
24		collect as collect_account_data, collect_ranges as collect_account_data_ranges,
25	},
26	profiles::collect as collect_profiles,
27	receipts::collect_ranges as collect_receipt_ranges,
28};
29use super::{SyncInfo, Window, WindowRoom, range::Results, share_encrypted_room};
30
31pub(super) struct Collected {
32	response: Extensions,
33	typing: typing::Collected,
34}
35
36impl Collected {
37	pub(super) fn into_response(
38		mut self,
39		payloads: &BTreeMap<OwnedRoomId, ResponseRoom>,
40	) -> Extensions {
41		self.response.typing = self.typing.into_response(payloads);
42		self.response
43	}
44}
45
46#[tracing::instrument(
47	name = "extensions",
48	level = "debug",
49	skip_all,
50	fields(
51		next_batch = conn.next_batch,
52		window = window.len(),
53		rooms = conn.rooms.len(),
54		subs = conn.subscriptions.len(),
55	)
56)]
57pub(super) async fn handle(
58	sync_info: SyncInfo<'_>,
59	conn: &Connection,
60	window: &Window,
61) -> Result<Collected> {
62	let account_data = conn
63		.extensions
64		.account_data
65		.enabled
66		.unwrap_or(false)
67		.then_async(|| collect_account_data(sync_info, conn));
68
69	let typing = conn
70		.extensions
71		.typing
72		.enabled
73		.unwrap_or(false)
74		.then_async(|| typing::collect(sync_info, conn, window));
75
76	let to_device = conn
77		.extensions
78		.to_device
79		.enabled
80		.unwrap_or(false)
81		.then_async(|| to_device::collect(sync_info, conn));
82
83	let e2ee = conn
84		.extensions
85		.e2ee
86		.enabled
87		.unwrap_or(false)
88		.then_async(|| e2ee::collect(sync_info, conn));
89
90	let (account_data, typing, to_device, e2ee) = join4(account_data, typing, to_device, e2ee)
91		.map(apply!(4, |t: Option<_>| t.unwrap_or(Ok(Default::default()))))
92		.await;
93
94	// Receipt and room account-data payloads only exist as bounded room-range
95	// outputs, applied by `apply_ranges` after the ranges resolve.
96	let response = Extensions {
97		account_data: account_data?,
98		receipts: Default::default(),
99		typing: Default::default(),
100		to_device: to_device?,
101		e2ee: e2ee?,
102		profiles: Default::default(),
103	};
104
105	Ok(Collected { response, typing: typing? })
106}
107
108#[tracing::instrument(level = "trace", skip_all)]
109pub(super) async fn apply_profiles(
110	sync_info: SyncInfo<'_>,
111	conn: &Connection,
112	window: &Window,
113	ranges: &Results,
114	extensions: &mut Collected,
115) -> Result {
116	if conn.extensions.profiles.enabled.unwrap_or(false) {
117		extensions.response.profiles = collect_profiles(sync_info, conn, window, ranges).await?;
118	}
119
120	Ok(())
121}
122
123pub(super) fn apply_ranges(
124	conn: &Connection,
125	window: &Window,
126	ranges: &mut Results,
127	extensions: &mut Collected,
128) {
129	if conn.extensions.receipts.enabled.unwrap_or(false) {
130		extensions.response.receipts = collect_receipt_ranges(conn, window, ranges);
131	}
132
133	if conn
134		.extensions
135		.account_data
136		.enabled
137		.unwrap_or(false)
138	{
139		extensions.response.account_data.rooms =
140			collect_account_data_ranges(conn, window, ranges);
141	}
142}
143
144#[tracing::instrument(
145	name = "selector",
146	level = "trace",
147	skip_all,
148	fields(?implicit, ?explicit),
149)]
150fn selector<'a, ListIter, SubsIter>(
151	conn: &'a Connection,
152	window: &'a Window,
153	implicit: Option<ListIter>,
154	explicit: Option<SubsIter>,
155) -> impl Iterator<Item = &'a RoomId> + Send + Sync + 'a
156where
157	ListIter: Iterator<Item = &'a ListId> + Clone + Debug + Send + Sync + 'a,
158	SubsIter: Iterator<Item = &'a ExtensionRoomConfig> + Clone + Debug + Send + Sync + 'a,
159{
160	let has_all_subscribed = explicit
161		.clone()
162		.into_iter()
163		.flatten()
164		.any(|erc| matches!(erc, ExtensionRoomConfig::AllSubscribed));
165	let implicit_subscribed = implicit.clone();
166	let implicit_explicit = implicit.clone();
167
168	let all_subscribed = has_all_subscribed
169		.then(|| {
170			window
171				.keys()
172				.filter(|room_id| conn.subscriptions.contains_key(*room_id))
173				.filter(move |room_id| {
174					window
175						.get(*room_id)
176						.is_some_and(|room| !implicit_match(room, implicit_subscribed.as_ref()))
177				})
178		})
179		.into_iter()
180		.flatten()
181		.map(AsRef::as_ref);
182
183	let rooms_explicit = has_all_subscribed
184		.is_false()
185		.then(move || {
186			explicit
187				.into_iter()
188				.flatten()
189				.filter_map(|erc| extract_variant!(erc, ExtensionRoomConfig::Room))
190				.filter(move |room_id| {
191					window
192						.get::<RoomId>(room_id.as_ref())
193						.is_some_and(|room| !implicit_match(room, implicit_explicit.as_ref()))
194				})
195				.map(AsRef::as_ref)
196		})
197		.into_iter()
198		.flatten();
199
200	let rooms_selected = window
201		.iter()
202		.filter(move |(_, room)| implicit_match(room, implicit.as_ref()))
203		.map(at!(0))
204		.map(AsRef::as_ref);
205
206	all_subscribed
207		.chain(rooms_explicit)
208		.chain(rooms_selected)
209}
210
211fn implicit_match<'a, ListIter>(room: &WindowRoom, implicit: Option<&ListIter>) -> bool
212where
213	ListIter: Iterator<Item = &'a ListId> + Clone,
214{
215	implicit.is_none_or(|lists| {
216		lists
217			.clone()
218			.any(|list| room.lists.contains(list))
219	})
220}
221
222#[cfg(test)]
223mod tests {
224	use std::slice::Iter;
225
226	use ruma::room_id;
227
228	use super::*;
229	use crate::client::sync::v5::ListIds;
230
231	#[test]
232	fn explicit_room_outside_the_window_is_rejected() {
233		let selected = room_id!("!selected:example.com");
234		let foreign = room_id!("!foreign:example.com");
235		let window = window(selected, 0);
236		let conn = Connection::default();
237		let list = ListId::from("unmatched");
238		let lists = [list];
239		let rooms = [ExtensionRoomConfig::Room(foreign.to_owned())];
240
241		assert!(
242			selector(&conn, &window, Some(lists.iter()), Some(rooms.iter()))
243				.next()
244				.is_none()
245		);
246	}
247
248	#[test]
249	fn all_subscribed_room_outside_the_window_is_rejected() {
250		let selected = room_id!("!selected:example.com");
251		let foreign = room_id!("!foreign:example.com");
252		let window = window(selected, 0);
253		let mut conn = Connection::default();
254		conn.subscriptions
255			.insert(foreign.to_owned(), Default::default());
256		let list = ListId::from("unmatched");
257		let lists = [list];
258		let rooms = [ExtensionRoomConfig::AllSubscribed];
259
260		assert!(
261			selector(&conn, &window, Some(lists.iter()), Some(rooms.iter()))
262				.next()
263				.is_none()
264		);
265	}
266
267	#[test]
268	fn stale_payload_room_in_the_window_remains_extension_eligible() {
269		let room_id = room_id!("!stale:example.com");
270		let window = window(room_id, 1);
271		let mut conn = Connection::default();
272		conn.rooms
273			.entry(room_id.to_owned())
274			.or_default()
275			.roomsince = 9;
276
277		let lists: Option<Iter<'_, ListId>> = None;
278		let rooms: Option<Iter<'_, ExtensionRoomConfig>> = None;
279
280		let selected = selector(&conn, &window, lists, rooms).collect::<Vec<_>>();
281
282		assert_eq!(selected, [room_id]);
283	}
284
285	#[test]
286	fn explicit_room_already_selected_by_list_is_not_duplicated() {
287		let room_id = room_id!("!explicit-overlap:example.com");
288		let list = ListId::from("main");
289		let mut window = window(room_id, 0);
290
291		window
292			.get_mut(room_id)
293			.expect("test room should be present")
294			.lists
295			.push(list.clone());
296
297		let conn = Connection::default();
298		let lists = [list];
299		let rooms = [ExtensionRoomConfig::Room(room_id.to_owned())];
300		let selected =
301			selector(&conn, &window, Some(lists.iter()), Some(rooms.iter())).collect::<Vec<_>>();
302
303		assert_eq!(selected, [room_id]);
304	}
305
306	#[test]
307	fn subscribed_room_already_selected_by_list_is_not_duplicated() {
308		let room_id = room_id!("!subscribed-overlap:example.com");
309		let list = ListId::from("main");
310		let mut window = window(room_id, 0);
311
312		window
313			.get_mut(room_id)
314			.expect("test room should be present")
315			.lists
316			.push(list.clone());
317
318		let mut conn = Connection::default();
319		conn.subscriptions
320			.insert(room_id.to_owned(), Default::default());
321
322		let lists = [list];
323		let rooms = [ExtensionRoomConfig::AllSubscribed];
324		let selected =
325			selector(&conn, &window, Some(lists.iter()), Some(rooms.iter())).collect::<Vec<_>>();
326
327		assert_eq!(selected, [room_id]);
328	}
329
330	fn window(room_id: &RoomId, payload_count: u64) -> Window {
331		let room = WindowRoom {
332			room_id: room_id.to_owned(),
333			membership: None,
334			lists: ListIds::new(),
335			event_count: 0,
336			payload_count,
337		};
338
339		[(room_id.to_owned(), room)].into()
340	}
341}