Skip to main content

tuwunel_service/sync/
mod.rs

1mod watch;
2
3#[cfg(test)]
4mod tests;
5
6use std::{
7	collections::{BTreeMap, btree_map::Entry},
8	sync::Arc,
9};
10
11use futures::{FutureExt, Stream};
12use ruma::{
13	DeviceId, OwnedDeviceId, OwnedRoomId, OwnedUserId, RoomId, UserId,
14	api::client::sync::sync_events::v5::{
15		ConnId as ConnectionId, ListId, Request, request,
16		request::{AccountData, E2EE, Profiles, Receipts, ToDevice, Typing},
17	},
18	profile::ProfileFieldName,
19};
20use serde::{Deserialize, Serialize};
21use tokio::sync::Mutex as TokioMutex;
22use tuwunel_core::{
23	Result, at, debug, err, implement, is_equal_to, smallvec::SmallVec, utils::stream::TryIgnore,
24};
25use tuwunel_database::{Cbor, Deserialized, Map};
26
27pub struct Service {
28	services: Arc<crate::services::OnceServices>,
29	connections: Connections,
30	db: Data,
31}
32
33struct Data {
34	userdeviceconnid_conn: Arc<Map>,
35	todeviceid_events: Arc<Map>,
36	userroomid_joined: Arc<Map>,
37	userroomid_invitestate: Arc<Map>,
38	userroomid_leftstate: Arc<Map>,
39	userroomid_knockedstate: Arc<Map>,
40	userroomid_notificationcount: Arc<Map>,
41	userroomid_highlightcount: Arc<Map>,
42	pduid_pdu: Arc<Map>,
43	keychangeid_userid: Arc<Map>,
44	profilechangeid_userid: Arc<Map>,
45	roomuserdataid_accountdata: Arc<Map>,
46	roomusertype_roomuserdataid: Arc<Map>,
47	readreceiptid_readreceipt: Arc<Map>,
48	userid_lastonetimekeyupdate: Arc<Map>,
49	roomuserid_lastnotificationread: Arc<Map>,
50}
51
52/// What a sliding sync connection carries from one request to the next.
53///
54/// It is stored after every response, so a client resuming from a position
55/// finds the lists, extension settings and per-room progress it left there.
56#[derive(Debug, Default, Deserialize, Serialize)]
57pub struct Connection {
58	pub globalsince: u64,
59	pub next_batch: u64,
60	pub lists: Lists,
61	pub extensions: request::Extensions,
62	pub subscriptions: Subscriptions,
63	pub rooms: Rooms,
64
65	/// Position of the response that last carried the syncing user's whole
66	/// profile in the MSC4262 profiles extension.
67	///
68	/// Zero until a response has carried it, and again whenever the extension is
69	/// switched off, so switching it back on sends the whole profile anew.
70	#[serde(default)]
71	pub own_profile_since: u64,
72
73	/// Whether the connection asked for a profile field it did not have before.
74	///
75	/// It rides the whole profile above, so the rooms replay their slice of the
76	/// change log until the client advances past the response carrying the
77	/// widened field set, and it is forgotten with everything else the extension
78	/// owes when the extension goes off.
79	#[serde(default)]
80	pub profiles_fields_widened: bool,
81}
82
83/// Delivery progress for one room on a Sliding Sync connection.
84///
85/// The cursor advances with complete ranges; configuration tracks room payloads.
86#[derive(Clone, Debug, Default, Deserialize, Serialize)]
87pub struct Room {
88	pub roomsince: u64,
89	#[serde(default)]
90	pub config_hash: u64,
91	/// Fingerprints of the required-state selectors last delivered for this room.
92	#[serde(default)]
93	pub required_state: RequiredState,
94}
95
96/// Fingerprints of delivered required-state selectors.
97///
98/// The inline capacity covers common room-list and open-room selections.
99pub type RequiredState = SmallVec<[u64; 18]>;
100
101/// A delivered room configuration and its required-state coverage.
102///
103/// Both values advance together only after a room payload is assembled.
104pub type RoomConfig = (u64, RequiredState);
105
106type Connections = TokioMutex<BTreeMap<ConnectionKey, ConnectionVal>>;
107pub type ConnectionVal = Arc<TokioMutex<Connection>>;
108pub type ConnectionKey = (OwnedUserId, Option<OwnedDeviceId>, Option<ConnectionId>);
109
110pub type Subscriptions = BTreeMap<OwnedRoomId, request::ListConfig>;
111pub type Lists = BTreeMap<ListId, request::List>;
112pub type Rooms = BTreeMap<OwnedRoomId, Room>;
113type RoomUpdate<'a> = (&'a RoomId, Option<RoomConfig>);
114
115impl crate::Service for Service {
116	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
117		Ok(Arc::new(Self {
118			db: Data {
119				userdeviceconnid_conn: args.db["userdeviceconnid_conn"].clone(),
120				todeviceid_events: args.db["todeviceid_events"].clone(),
121				userroomid_joined: args.db["userroomid_joined"].clone(),
122				userroomid_invitestate: args.db["userroomid_invitestate"].clone(),
123				userroomid_leftstate: args.db["userroomid_leftstate"].clone(),
124				userroomid_knockedstate: args.db["userroomid_knockedstate"].clone(),
125				userroomid_notificationcount: args.db["userroomid_notificationcount"].clone(),
126				userroomid_highlightcount: args.db["userroomid_highlightcount"].clone(),
127				pduid_pdu: args.db["pduid_pdu"].clone(),
128				keychangeid_userid: args.db["keychangeid_userid"].clone(),
129				profilechangeid_userid: args.db["profilechangeid_userid"].clone(),
130				roomuserdataid_accountdata: args.db["roomuserdataid_accountdata"].clone(),
131				roomusertype_roomuserdataid: args.db["roomusertype_roomuserdataid"].clone(),
132				readreceiptid_readreceipt: args.db["readreceiptid_readreceipt"].clone(),
133				userid_lastonetimekeyupdate: args.db["userid_lastonetimekeyupdate"].clone(),
134				roomuserid_lastnotificationread: args.db["roomuserid_lastnotificationread"]
135					.clone(),
136			},
137			services: args.services.clone(),
138			connections: Default::default(),
139		}))
140	}
141
142	fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
143}
144
145#[implement(Service)]
146#[tracing::instrument(level = "debug", skip(self))]
147pub async fn clear_connections(
148	&self,
149	user_id: Option<&UserId>,
150	device_id: Option<&DeviceId>,
151	conn_id: Option<&ConnectionId>,
152) {
153	self.connections
154		.lock()
155		.await
156		.retain(|(conn_user_id, conn_device_id, conn_conn_id), _| {
157			let retain = user_id.is_none_or(is_equal_to!(conn_user_id))
158				&& (device_id.is_none() || device_id == conn_device_id.as_deref())
159				&& (conn_id.is_none() || conn_id == conn_conn_id.as_ref());
160
161			if !retain {
162				self.db
163					.userdeviceconnid_conn
164					.del((conn_user_id, conn_device_id, conn_conn_id));
165			}
166
167			retain
168		});
169}
170
171#[implement(Service)]
172#[tracing::instrument(level = "debug", skip(self))]
173pub async fn drop_connection(&self, key: &ConnectionKey) {
174	let mut cache = self.connections.lock().await;
175
176	self.db.userdeviceconnid_conn.del(key);
177	cache.remove(key);
178}
179
180#[implement(Service)]
181#[tracing::instrument(level = "debug", skip(self))]
182pub async fn load_or_init_connection(&self, key: &ConnectionKey) -> ConnectionVal {
183	let mut cache = self.connections.lock().await;
184
185	match cache.entry(key.clone()) {
186		| Entry::Occupied(val) => val.get().clone(),
187		| Entry::Vacant(val) => {
188			let conn = self
189				.db
190				.userdeviceconnid_conn
191				.qry(key)
192				.boxed()
193				.await
194				.deserialized::<Cbor<_>>()
195				.map(at!(0))
196				.map(TokioMutex::new)
197				.map(Arc::new)
198				.unwrap_or_default();
199
200			val.insert(conn).clone()
201		},
202	}
203}
204
205#[implement(Service)]
206#[tracing::instrument(level = "debug", skip(self))]
207pub async fn load_connection(&self, key: &ConnectionKey) -> Result<ConnectionVal> {
208	let mut cache = self.connections.lock().await;
209
210	match cache.entry(key.clone()) {
211		| Entry::Occupied(val) => Ok(val.get().clone()),
212		| Entry::Vacant(val) => self
213			.db
214			.userdeviceconnid_conn
215			.qry(key)
216			.await
217			.deserialized::<Cbor<_>>()
218			.map(at!(0))
219			.map(TokioMutex::new)
220			.map(Arc::new)
221			.map(|conn| val.insert(conn).clone()),
222	}
223}
224
225#[implement(Service)]
226#[tracing::instrument(level = "debug", skip(self))]
227pub async fn get_loaded_connection(&self, key: &ConnectionKey) -> Result<ConnectionVal> {
228	self.connections
229		.lock()
230		.await
231		.get(key)
232		.cloned()
233		.ok_or_else(|| err!(Request(NotFound("Connection not found."))))
234}
235
236#[implement(Service)]
237#[tracing::instrument(level = "trace", skip(self))]
238pub async fn list_loaded_connections(&self) -> Vec<ConnectionKey> {
239	self.connections
240		.lock()
241		.await
242		.keys()
243		.cloned()
244		.collect()
245}
246
247#[implement(Service)]
248#[tracing::instrument(level = "trace", skip(self))]
249pub fn list_stored_connections(&self) -> impl Stream<Item = ConnectionKey> {
250	self.db.userdeviceconnid_conn.keys().ignore_err()
251}
252
253#[implement(Service)]
254#[tracing::instrument(level = "trace", skip(self))]
255pub async fn is_connection_loaded(&self, key: &ConnectionKey) -> bool {
256	self.connections.lock().await.contains_key(key)
257}
258
259#[implement(Service)]
260#[tracing::instrument(level = "trace", skip(self))]
261pub async fn is_connection_stored(&self, key: &ConnectionKey) -> bool {
262	self.db.userdeviceconnid_conn.contains(key).await
263}
264
265#[inline]
266pub fn into_connection_key<U, D, C>(
267	user_id: U,
268	device_id: Option<D>,
269	conn_id: Option<C>,
270) -> ConnectionKey
271where
272	U: Into<OwnedUserId>,
273	D: Into<OwnedDeviceId>,
274	C: Into<ConnectionId>,
275{
276	(user_id.into(), device_id.map(Into::into), conn_id.map(Into::into))
277}
278
279#[implement(Connection)]
280#[tracing::instrument(level = "debug", skip(self, service))]
281pub fn store(&self, service: &Service, key: &ConnectionKey) {
282	service
283		.db
284		.userdeviceconnid_conn
285		.put(key, Cbor(self));
286
287	debug!(
288		since = %self.globalsince,
289		next_batch = %self.next_batch,
290		"Persisted connection state"
291	);
292}
293
294#[implement(Connection)]
295#[tracing::instrument(level = "debug", skip(self))]
296pub fn update_rooms_prologue(&mut self, retard_since: Option<u64>) {
297	self.rooms.values_mut().for_each(|room| {
298		if let Some(retard_since) = retard_since
299			&& room.roomsince > retard_since
300		{
301			room.roomsince = retard_since;
302			room.config_hash = 0;
303			room.required_state.clear();
304		}
305	});
306}
307
308/// Advance the per-room cursor for each complete bounded room range.
309///
310/// `roomsince` is the lower bound of every content query for its room. Only
311/// rooms whose complete range was safely assembled may advance. A failed room
312/// keeps its cursor and retries the same range after a later wake.
313#[implement(Connection)]
314#[tracing::instrument(level = "debug", skip_all)]
315pub fn update_rooms_epilogue<'a, Complete>(&mut self, complete: Complete)
316where
317	Complete: Iterator<Item = RoomUpdate<'a>> + Send + 'a,
318{
319	let next_batch = self.next_batch;
320	complete.for_each(|(room_id, config)| {
321		let room = self.rooms.entry(room_id.into()).or_default();
322
323		room.roomsince = next_batch;
324		if let Some((config_hash, required_state)) = config {
325			room.config_hash = config_hash;
326			room.required_state = required_state;
327		}
328	});
329}
330
331/// Records that this pass carried the syncing user's whole profile.
332///
333/// Called once the pass has assembled its extensions, so a pass that failed
334/// leaves the profile owed to the next one.
335#[implement(Connection)]
336#[tracing::instrument(level = "debug", skip_all)]
337pub fn update_profiles_epilogue(&mut self) {
338	if self.own_profile_owed() {
339		self.own_profile_since = self.next_batch;
340	}
341}
342
343/// Whether a base is owed for a profile field the connection did not have
344/// before.
345///
346/// MSC4262 asks for the widened fields of every user in the room subset, which
347/// the room passes deliver by replaying their whole slice of the change log.
348#[implement(Connection)]
349#[inline]
350#[must_use]
351pub fn profiles_fields_owed(&self) -> bool {
352	self.profiles_fields_widened && self.own_profile_owed()
353}
354
355/// Whether the syncing user's whole profile is owed to the profiles extension.
356///
357/// It is owed while the extension is on and the client has not acknowledged a
358/// response carrying it: the connection is new, the extension was switched on
359/// after the connection began, or the client is replaying from before the
360/// response that carried it.
361#[implement(Connection)]
362#[inline]
363#[must_use]
364pub fn own_profile_owed(&self) -> bool {
365	self.extensions.profiles.enabled.unwrap_or(false)
366		&& (self.own_profile_since == 0 || self.own_profile_since > self.globalsince)
367}
368
369#[implement(Connection)]
370#[tracing::instrument(level = "debug", skip_all)]
371pub fn update_cache(&mut self, request: &Request) -> bool {
372	let lists_changed = Self::update_cache_lists(request, self);
373	let subscriptions_changed = Self::update_cache_subscriptions(request, self);
374
375	let fields_widened = Self::update_cache_extensions(request, self);
376
377	self.update_cache_profiles_owed(fields_widened);
378
379	lists_changed || subscriptions_changed
380}
381
382#[implement(Connection)]
383fn update_cache_lists(request: &Request, cached: &mut Self) -> bool {
384	request
385		.lists
386		.iter()
387		.fold(false, |changed, (list_id, request_list)| {
388			let list_changed = match cached.lists.get_mut(list_id) {
389				| Some(cached_list) => Self::update_cache_list(request_list, cached_list),
390				| None => {
391					cached
392						.lists
393						.insert(list_id.clone(), request_list.clone());
394
395					true
396				},
397			};
398
399			changed | list_changed
400		})
401}
402
403#[implement(Connection)]
404fn update_cache_list(request: &request::List, cached: &mut request::List) -> bool {
405	let ranges_changed = request.ranges != cached.ranges;
406	let timeline_limit_changed =
407		request.room_details.timeline_limit != cached.room_details.timeline_limit;
408
409	let required_state_changed = !request.room_details.required_state.is_empty()
410		&& request.room_details.required_state != cached.room_details.required_state;
411
412	let filters_changed = request.filters.as_ref().is_some_and(|request| {
413		cached
414			.filters
415			.as_ref()
416			.is_none_or(|cached| !list_filters_are_equal(request, cached))
417	});
418
419	let changed =
420		ranges_changed || timeline_limit_changed || required_state_changed || filters_changed;
421
422	if ranges_changed {
423		cached.ranges.clone_from(&request.ranges);
424	}
425
426	cached.room_details.timeline_limit = request.room_details.timeline_limit;
427
428	if required_state_changed {
429		cached
430			.room_details
431			.required_state
432			.clone_from(&request.room_details.required_state);
433	}
434
435	if filters_changed {
436		cached.filters.clone_from(&request.filters);
437	}
438
439	changed
440}
441
442#[implement(Connection)]
443fn update_cache_subscriptions(request: &Request, cached: &mut Self) -> bool {
444	let changed = !subscriptions_are_equal(&request.room_subscriptions, &cached.subscriptions);
445
446	if changed {
447		cached
448			.subscriptions
449			.clone_from(&request.room_subscriptions);
450	}
451
452	changed
453}
454
455fn subscriptions_are_equal(request: &Subscriptions, cached: &Subscriptions) -> bool {
456	request.len() == cached.len()
457		&& request
458			.iter()
459			.zip(cached)
460			.all(|(request, cached)| {
461				request.0 == cached.0 && list_config_is_equal(request.1, cached.1)
462			})
463}
464
465fn list_config_is_equal(request: &request::ListConfig, cached: &request::ListConfig) -> bool {
466	request.timeline_limit == cached.timeline_limit
467		&& request.required_state == cached.required_state
468}
469
470fn list_filters_are_equal(request: &request::ListFilters, cached: &request::ListFilters) -> bool {
471	request.is_dm == cached.is_dm
472		&& request.is_encrypted == cached.is_encrypted
473		&& request.is_invite == cached.is_invite
474		&& request.room_types == cached.room_types
475		&& request.not_room_types == cached.not_room_types
476		&& request.tags == cached.tags
477		&& request.not_tags == cached.not_tags
478		&& request.spaces == cached.spaces
479}
480
481/// Merges the request's extension settings into the connection.
482///
483/// Returns whether the MSC4262 field filter named a field the connection did
484/// not have, which the profiles extension owes a base for.
485#[implement(Connection)]
486fn update_cache_extensions(request: &Request, cached: &mut Self) -> bool {
487	let request = &request.extensions;
488	let cached = &mut cached.extensions;
489
490	Self::update_cache_account_data(&request.account_data, &mut cached.account_data);
491	Self::update_cache_receipts(&request.receipts, &mut cached.receipts);
492	Self::update_cache_typing(&request.typing, &mut cached.typing);
493	Self::update_cache_to_device(&request.to_device, &mut cached.to_device);
494	Self::update_cache_e2ee(&request.e2ee, &mut cached.e2ee);
495
496	Self::update_cache_profiles(&request.profiles, &mut cached.profiles)
497}
498
499#[implement(Connection)]
500fn update_cache_account_data(request: &AccountData, cached: &mut AccountData) {
501	some_or_sticky(request.enabled.as_ref(), &mut cached.enabled);
502	some_or_sticky(request.lists.as_ref(), &mut cached.lists);
503	some_or_sticky(request.rooms.as_ref(), &mut cached.rooms);
504}
505
506#[implement(Connection)]
507fn update_cache_receipts(request: &Receipts, cached: &mut Receipts) {
508	some_or_sticky(request.enabled.as_ref(), &mut cached.enabled);
509	some_or_sticky(request.rooms.as_ref(), &mut cached.rooms);
510	some_or_sticky(request.lists.as_ref(), &mut cached.lists);
511}
512
513#[implement(Connection)]
514fn update_cache_typing(request: &Typing, cached: &mut Typing) {
515	some_or_sticky(request.enabled.as_ref(), &mut cached.enabled);
516	some_or_sticky(request.rooms.as_ref(), &mut cached.rooms);
517	some_or_sticky(request.lists.as_ref(), &mut cached.lists);
518}
519
520#[implement(Connection)]
521fn update_cache_to_device(request: &ToDevice, cached: &mut ToDevice) {
522	some_or_sticky(request.enabled.as_ref(), &mut cached.enabled);
523	cached.since.clone_from(&request.since);
524}
525
526/// Merges the profiles extension settings into the connection.
527///
528/// Returns whether the request widened the field filter, which is read before
529/// the merge overwrites the filter it compares against.
530#[implement(Connection)]
531fn update_cache_profiles(request: &Profiles, cached: &mut Profiles) -> bool {
532	some_or_sticky(request.enabled.as_ref(), &mut cached.enabled);
533	some_or_sticky(request.rooms.as_ref(), &mut cached.rooms);
534	some_or_sticky(request.lists.as_ref(), &mut cached.lists);
535
536	// Compare against the cached filter before the merge below overwrites it.
537	let widened = fields_widened(request.fields.as_deref(), cached.fields.as_deref());
538
539	some_or_sticky(request.fields.as_ref(), &mut cached.fields);
540
541	widened
542}
543
544/// Whether the request names a profile field the connection did not ask for.
545///
546/// An absent cached filter already covers every field, and an absent request
547/// keeps the cached one, so neither widens anything.
548fn fields_widened(
549	request: Option<&[ProfileFieldName]>,
550	cached: Option<&[ProfileFieldName]>,
551) -> bool {
552	request
553		.zip(cached)
554		.is_some_and(|(request, cached)| request.iter().any(|name| !cached.contains(name)))
555}
556
557#[implement(Connection)]
558fn update_cache_e2ee(request: &E2EE, cached: &mut E2EE) {
559	some_or_sticky(request.enabled.as_ref(), &mut cached.enabled);
560}
561
562#[implement(Connection)]
563fn update_cache_profiles_owed(&mut self, fields_widened: bool) {
564	if fields_widened || !self.extensions.profiles.enabled.unwrap_or(false) {
565		self.own_profile_since = 0;
566	}
567
568	self.profiles_fields_widened =
569		fields_widened || (self.profiles_fields_widened && self.own_profile_owed());
570}
571
572fn some_or_sticky<T: Clone>(target: Option<&T>, cached: &mut Option<T>) {
573	if let Some(target) = target {
574		cached.replace(target.clone());
575	}
576}