Skip to main content

tuwunel_api/client/sync/v5/
extensions.rs

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