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#[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 #[serde(default)]
71 pub own_profile_since: u64,
72
73 #[serde(default)]
80 pub profiles_fields_widened: bool,
81}
82
83#[derive(Clone, Debug, Default, Deserialize, Serialize)]
87pub struct Room {
88 pub roomsince: u64,
89 #[serde(default)]
90 pub config_hash: u64,
91 #[serde(default)]
93 pub required_state: RequiredState,
94}
95
96pub type RequiredState = SmallVec<[u64; 18]>;
100
101pub 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#[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#[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#[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#[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#[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#[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 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
544fn 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}