Skip to main content

tuwunel_service/rooms/state_cache/
update.rs

1use std::collections::HashSet;
2
3use futures::StreamExt;
4use ruma::{
5	OwnedServerName, RoomId, UserId,
6	events::{
7		AnyStrippedStateEvent, AnySyncStateEvent, GlobalAccountDataEventType,
8		RoomAccountDataEventType, StateEventType,
9		direct::DirectEvent,
10		room::{
11			create::RoomCreateEventContent,
12			member::{MembershipState, RoomMemberEventContent},
13		},
14	},
15	serde::Raw,
16};
17use tuwunel_core::{
18	Result, implement, is_not_empty,
19	matrix::PduCount,
20	utils::{ReadyExt, result::LogErr},
21	warn,
22};
23use tuwunel_database::{Json, serialize_key, serialize_val};
24
25/// Optional stripped room state attached to invite and knock transitions.
26pub type StrippedRoomState = Option<Vec<Raw<AnyStrippedStateEvent>>>;
27
28/// Parameters for one membership cache transition.
29///
30/// Borrowed identifiers remain valid only for the duration of the update. Owned
31/// event data is consumed by the selected transition.
32pub struct MembershipUpdate<'a> {
33	/// Room whose membership changed.
34	///
35	/// Membership indexes and aggregate counts are updated for this room.
36	pub room_id: &'a RoomId,
37
38	/// User whose membership changed.
39	///
40	/// Both local and remote users are represented in the membership indexes.
41	pub user_id: &'a UserId,
42
43	/// Membership event content driving the transition.
44	///
45	/// The membership state selects which indexes are written and cleared.
46	pub membership_event: RoomMemberEventContent,
47
48	/// User who sent the membership event.
49	///
50	/// Invite handling uses the sender when applying the ignored-user policy.
51	pub sender: &'a UserId,
52
53	/// Stripped room state associated with an invite or knock.
54	///
55	/// Other membership transitions leave this value unused.
56	pub last_state: StrippedRoomState,
57
58	/// Servers supplied as routing hints for an invite.
59	///
60	/// Invite handling stores only non-empty lists. The routing hints commit
61	/// with the membership indexes.
62	pub invite_via: Option<Vec<OwnedServerName>>,
63
64	/// Whether to rebuild the room's aggregate membership counts.
65	///
66	/// Bulk state updates can defer this rebuild until all transitions are
67	/// applied.
68	pub update_joined_count: bool,
69
70	/// Stream position associated with the membership event.
71	///
72	/// Count-indexed membership rows store its unsigned representation.
73	pub count: PduCount,
74}
75
76/// Update current membership data.
77#[implement(super::Service)]
78#[tracing::instrument(
79		level = "debug",
80		skip_all,
81		fields(
82			%room_id,
83			%user_id,
84			%sender,
85			%count,
86			?membership_event,
87		),
88	)]
89pub async fn update_membership(
90	&self,
91	MembershipUpdate {
92		room_id,
93		user_id,
94		membership_event,
95		sender,
96		last_state,
97		invite_via,
98		update_joined_count,
99		count,
100	}: MembershipUpdate<'_>,
101) -> Result {
102	let membership = membership_event.membership;
103
104	self.ensure_remote_user(user_id).await?;
105
106	match membership {
107		| MembershipState::Join => {
108			self.handle_join(room_id, user_id, count).await?;
109		},
110		| MembershipState::Invite => {
111			if self
112				.services
113				.users
114				.user_is_ignored(sender, user_id)
115				.await
116			{
117				return Ok(());
118			}
119
120			self.mark_as_invited(user_id, room_id, count, last_state, invite_via)
121				.await;
122		},
123		| MembershipState::Leave | MembershipState::Ban => {
124			self.handle_leave(room_id, user_id, count).await;
125
126			// A departure drops the room from the account-wide badge total.
127			if self.services.globals.user_is_local(user_id) {
128				self.services
129					.sending
130					.refresh_push_badge(user_id)
131					.await
132					.log_err()
133					.ok();
134			}
135		},
136		| MembershipState::Knock => {
137			self.mark_as_knocked(user_id, room_id, count, last_state);
138		},
139		| _ => {},
140	}
141
142	if update_joined_count {
143		self.update_joined_count(room_id).await;
144	}
145
146	Ok(())
147}
148
149#[implement(super::Service)]
150#[tracing::instrument(level = "debug", skip(self))]
151pub async fn update_joined_count(&self, room_id: &RoomId) {
152	let mut joinedcount = 0_u64;
153	let mut invitedcount = 0_u64;
154	let mut knockedcount = 0_u64;
155	let mut joined_servers = HashSet::new();
156
157	self.room_members(room_id)
158		.ready_for_each(|joined| {
159			joined_servers.insert(joined.server_name().to_owned());
160			joinedcount = joinedcount.saturating_add(1);
161		})
162		.await;
163
164	invitedcount = invitedcount.saturating_add(
165		self.room_members_invited(room_id)
166			.count()
167			.await
168			.try_into()
169			.unwrap_or(0),
170	);
171
172	knockedcount = knockedcount.saturating_add(
173		self.room_members_knocked(room_id)
174			.count()
175			.await
176			.try_into()
177			.unwrap_or(0),
178	);
179
180	let joinedcount = joinedcount.to_be_bytes();
181	let invitedcount = invitedcount.to_be_bytes();
182	let knockedcount = knockedcount.to_be_bytes();
183	let mut txn = self.services.db.txn();
184
185	txn.insert_raw(&self.db.roomid_joinedcount, room_id, joinedcount);
186	txn.insert_raw(&self.db.roomid_invitedcount, room_id, invitedcount);
187	txn.insert_raw(&self.db.roomid_knockedcount, room_id, knockedcount);
188
189	self.room_servers(room_id)
190		.ready_for_each(|old_joined_server| {
191			if joined_servers.remove(old_joined_server) {
192				return;
193			}
194
195			// Server not in room anymore
196			let roomserver_id = (room_id, old_joined_server);
197			let serverroom_id = (old_joined_server, room_id);
198
199			txn.del(&self.db.roomserverids, roomserver_id);
200			txn.del(&self.db.serverroomids, serverroom_id);
201		})
202		.await;
203
204	// Now only new servers are in joined_servers anymore
205	for server in &joined_servers {
206		let roomserver_id = (room_id, server);
207		let serverroom_id = (server, room_id);
208		let roomserver_id =
209			serialize_key(roomserver_id).expect("failed to serialize roomserver_id");
210
211		let serverroom_id =
212			serialize_key(serverroom_id).expect("failed to serialize serverroom_id");
213
214		txn.insert_raw(&self.db.roomserverids, roomserver_id, []);
215		txn.insert_raw(&self.db.serverroomids, serverroom_id, []);
216	}
217
218	txn.execute();
219
220	self.appservice_in_room_cache
221		.write()
222		.expect("locked")
223		.remove(room_id);
224}
225
226/// Direct DB function to directly mark a user as joined. It is not
227/// recommended to use this directly. You most likely should use
228/// `update_membership` instead
229#[implement(super::Service)]
230#[tracing::instrument(skip(self), level = "debug")]
231pub(crate) fn mark_as_joined(&self, user_id: &UserId, room_id: &RoomId, count: PduCount) {
232	let userroom_id = (user_id, room_id);
233	let userroom_id = serialize_key(userroom_id).expect("failed to serialize userroom_id");
234
235	let roomuser_id = (room_id, user_id);
236	let roomuser_id = serialize_key(roomuser_id).expect("failed to serialize roomuser_id");
237
238	let count = count.into_unsigned().to_be_bytes();
239	let mut txn = self.services.db.txn();
240
241	txn.insert_raw(&self.db.userroomid_joinedcount, &userroom_id, count);
242	txn.insert_raw(&self.db.roomuserid_joinedcount, &roomuser_id, count);
243	txn.del_raw(&self.db.userroomid_invitestate, &userroom_id);
244	txn.del_raw(&self.db.roomuserid_invitecount, &roomuser_id);
245	txn.del_raw(&self.db.userroomid_leftstate, &userroom_id);
246	txn.del_raw(&self.db.roomuserid_leftcount, &roomuser_id);
247	txn.del_raw(&self.db.userroomid_knockedstate, &userroom_id);
248	txn.del_raw(&self.db.roomuserid_knockedcount, &roomuser_id);
249	txn.execute();
250}
251
252/// Direct DB function to directly mark a user as left. It is not
253/// recommended to use this directly. You most likely should use
254/// `update_membership` instead
255#[implement(super::Service)]
256#[tracing::instrument(skip(self), level = "debug")]
257pub(crate) fn mark_as_left(&self, user_id: &UserId, room_id: &RoomId, count: PduCount) {
258	let userroom_id = (user_id, room_id);
259	let userroom_id = serialize_key(userroom_id).expect("failed to serialize userroom_id");
260
261	let roomuser_id = (room_id, user_id);
262	let roomuser_id = serialize_key(roomuser_id).expect("failed to serialize roomuser_id");
263
264	let leftstate = serialize_val(Json(Vec::<Raw<AnySyncStateEvent>>::new()))
265		.expect("failed to serialize left state");
266
267	let count = count.into_unsigned().to_be_bytes();
268	let mut txn = self.services.db.txn();
269
270	txn.insert_raw(&self.db.userroomid_leftstate, &userroom_id, leftstate);
271	txn.insert_raw(&self.db.roomuserid_leftcount, &roomuser_id, count);
272	txn.del_raw(&self.db.userroomid_joinedcount, &userroom_id);
273	txn.del_raw(&self.db.roomuserid_joinedcount, &roomuser_id);
274	txn.del_raw(&self.db.userroomid_invitestate, &userroom_id);
275	txn.del_raw(&self.db.roomuserid_invitecount, &roomuser_id);
276	txn.del_raw(&self.db.userroomid_knockedstate, &userroom_id);
277	txn.del_raw(&self.db.roomuserid_knockedcount, &roomuser_id);
278	txn.execute();
279}
280
281/// Direct DB function to directly mark a user as knocked. It is not
282/// recommended to use this directly. You most likely should use
283/// `update_membership` instead
284#[implement(super::Service)]
285#[tracing::instrument(skip(self), level = "debug")]
286pub(crate) fn mark_as_knocked(
287	&self,
288	user_id: &UserId,
289	room_id: &RoomId,
290	count: PduCount,
291	knocked_state: StrippedRoomState,
292) {
293	let userroom_id = (user_id, room_id);
294	let userroom_id = serialize_key(userroom_id).expect("failed to serialize userroom_id");
295
296	let roomuser_id = (room_id, user_id);
297	let roomuser_id = serialize_key(roomuser_id).expect("failed to serialize roomuser_id");
298
299	let knocked_state = serialize_val(Json(knocked_state.unwrap_or_default()))
300		.expect("failed to serialize knocked state");
301
302	let count = count.into_unsigned().to_be_bytes();
303	let mut txn = self.services.db.txn();
304
305	txn.insert_raw(&self.db.userroomid_knockedstate, &userroom_id, knocked_state);
306	txn.insert_raw(&self.db.roomuserid_knockedcount, &roomuser_id, count);
307	txn.del_raw(&self.db.userroomid_joinedcount, &userroom_id);
308	txn.del_raw(&self.db.roomuserid_joinedcount, &roomuser_id);
309	txn.del_raw(&self.db.userroomid_invitestate, &userroom_id);
310	txn.del_raw(&self.db.roomuserid_invitecount, &roomuser_id);
311	txn.del_raw(&self.db.userroomid_leftstate, &userroom_id);
312	txn.del_raw(&self.db.roomuserid_leftcount, &roomuser_id);
313	txn.execute();
314}
315
316/// Makes a user forget a room.
317#[implement(super::Service)]
318#[tracing::instrument(skip(self), level = "debug")]
319pub fn forget(&self, room_id: &RoomId, user_id: &UserId) {
320	let userroom_id = (user_id, room_id);
321	let roomuser_id = (room_id, user_id);
322	let mut txn = self.services.db.txn();
323
324	txn.del(&self.db.userroomid_leftstate, userroom_id);
325	txn.del(&self.db.roomuserid_leftcount, roomuser_id);
326	txn.execute();
327}
328
329#[implement(super::Service)]
330#[tracing::instrument(level = "debug", skip(self))]
331fn mark_as_once_joined(&self, user_id: &UserId, room_id: &RoomId) {
332	let key = (user_id, room_id);
333	let key = serialize_key(key).expect("failed to serialize roomuseroncejoinedid");
334	let mut txn = self.services.db.txn();
335
336	txn.insert_raw(&self.db.roomuseroncejoinedids, key, []);
337	txn.execute();
338}
339
340#[implement(super::Service)]
341#[tracing::instrument(level = "debug", skip(self, last_state, invite_via))]
342pub(crate) async fn mark_as_invited(
343	&self,
344	user_id: &UserId,
345	room_id: &RoomId,
346	count: PduCount,
347	last_state: StrippedRoomState,
348	invite_via: Option<Vec<OwnedServerName>>,
349) {
350	let userroom_id = (user_id, room_id);
351	let userroom_id = serialize_key(userroom_id).expect("failed to serialize userroom_id");
352
353	let roomuser_id = (room_id, user_id);
354	let roomuser_id = serialize_key(roomuser_id).expect("failed to serialize roomuser_id");
355
356	let invite_state = serialize_val(Json(last_state.unwrap_or_default()))
357		.expect("failed to serialize invite state");
358
359	let count = count.into_unsigned().to_be_bytes();
360	let mut txn = self.services.db.txn();
361
362	txn.insert_raw(&self.db.userroomid_invitestate, &userroom_id, invite_state);
363	txn.insert_raw(&self.db.roomuserid_invitecount, &roomuser_id, count);
364	txn.del_raw(&self.db.userroomid_joinedcount, &userroom_id);
365	txn.del_raw(&self.db.roomuserid_joinedcount, &roomuser_id);
366	txn.del_raw(&self.db.userroomid_leftstate, &userroom_id);
367	txn.del_raw(&self.db.roomuserid_leftcount, &roomuser_id);
368	txn.del_raw(&self.db.userroomid_knockedstate, &userroom_id);
369	txn.del_raw(&self.db.roomuserid_knockedcount, &roomuser_id);
370
371	if let Some(servers) = invite_via.filter(is_not_empty!()) {
372		self.add_servers_invite_via(&mut txn, room_id, servers)
373			.await;
374	}
375
376	txn.execute();
377}
378
379#[implement(super::Service)]
380#[tracing::instrument(skip(self), level = "debug")]
381async fn ensure_remote_user(&self, user_id: &UserId) -> Result {
382	if self.services.globals.user_is_local(user_id) || self.services.users.exists(user_id).await {
383		return Ok(());
384	}
385
386	self.services
387		.users
388		.create(user_id, None, None)
389		.await
390}
391
392#[implement(super::Service)]
393async fn handle_join(&self, room_id: &RoomId, user_id: &UserId, count: PduCount) -> Result {
394	if !self.once_joined(user_id, room_id).await {
395		self.mark_as_once_joined(user_id, room_id);
396		self.copy_predecessor_data(room_id, user_id)
397			.await?;
398	}
399
400	self.mark_as_joined(user_id, room_id, count);
401
402	Ok(())
403}
404
405#[implement(super::Service)]
406async fn copy_predecessor_data(&self, room_id: &RoomId, user_id: &UserId) -> Result {
407	let predecessor = self
408		.services
409		.state_accessor
410		.room_state_get_content(room_id, &StateEventType::RoomCreate, "")
411		.await
412		.map(|content: RoomCreateEventContent| content.predecessor);
413
414	let Ok(Some(predecessor)) = predecessor else {
415		return Ok(());
416	};
417
418	self.copy_predecessor_tags(room_id, user_id, &predecessor.room_id)
419		.await;
420
421	self.copy_predecessor_direct(room_id, user_id, &predecessor.room_id)
422		.await
423}
424
425#[implement(super::Service)]
426#[tracing::instrument(skip(self), level = "debug")]
427async fn copy_predecessor_tags(&self, room_id: &RoomId, user_id: &UserId, predecessor: &RoomId) {
428	let Ok(tag_event) = self
429		.services
430		.account_data
431		.get_room(predecessor, user_id, RoomAccountDataEventType::Tag)
432		.await
433	else {
434		return;
435	};
436
437	self.services
438		.account_data
439		.update(Some(room_id), user_id, RoomAccountDataEventType::Tag, &tag_event)
440		.await
441		.ok();
442}
443
444#[implement(super::Service)]
445#[tracing::instrument(skip(self), level = "debug")]
446async fn copy_predecessor_direct(
447	&self,
448	room_id: &RoomId,
449	user_id: &UserId,
450	predecessor: &RoomId,
451) -> Result {
452	let Ok(mut direct_event) = self
453		.services
454		.account_data
455		.get_global::<DirectEvent>(user_id, GlobalAccountDataEventType::Direct)
456		.await
457	else {
458		return Ok(());
459	};
460
461	let room_ids_updated =
462		direct_event
463			.content
464			.0
465			.values_mut()
466			.fold(false, |updated, room_ids| {
467				if !room_ids
468					.iter()
469					.any(|direct_room_id| direct_room_id == predecessor)
470				{
471					return updated;
472				}
473
474				room_ids.push(room_id.to_owned());
475
476				true
477			});
478
479	if !room_ids_updated {
480		return Ok(());
481	}
482
483	let event_type = GlobalAccountDataEventType::Direct
484		.to_string()
485		.into();
486
487	let direct_event = serde_json::to_value(&direct_event).expect("to json always works");
488
489	self.services
490		.account_data
491		.update(None, user_id, event_type, &direct_event)
492		.await
493}
494
495#[implement(super::Service)]
496#[tracing::instrument(skip(self), level = "debug")]
497async fn handle_leave(&self, room_id: &RoomId, user_id: &UserId, count: PduCount) {
498	self.mark_as_left(user_id, room_id, count);
499
500	if self.services.globals.user_is_local(user_id)
501		&& (self.services.config.forget_forced_upon_leave
502			|| self.services.metadata.is_banned(room_id).await
503			|| self.services.metadata.is_disabled(room_id).await)
504	{
505		self.forget(room_id, user_id);
506	}
507}