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
25pub type StrippedRoomState = Option<Vec<Raw<AnyStrippedStateEvent>>>;
27
28pub struct MembershipUpdate<'a> {
33 pub room_id: &'a RoomId,
37
38 pub user_id: &'a UserId,
42
43 pub membership_event: RoomMemberEventContent,
47
48 pub sender: &'a UserId,
52
53 pub last_state: StrippedRoomState,
57
58 pub invite_via: Option<Vec<OwnedServerName>>,
63
64 pub update_joined_count: bool,
69
70 pub count: PduCount,
74}
75
76#[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 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 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 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#[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#[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#[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#[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}