1use std::collections::HashSet;
9
10use futures::{
11 FutureExt, StreamExt,
12 future::{join, join3},
13};
14use ruma::{
15 OwnedServerName, RoomId, UserId,
16 events::{
17 AnyStrippedStateEvent, AnySyncStateEvent, GlobalAccountDataEventType,
18 RoomAccountDataEventType, StateEventType,
19 direct::DirectEvent,
20 room::{
21 create::RoomCreateEventContent,
22 member::{MembershipState, RoomMemberEventContent},
23 },
24 },
25 serde::Raw,
26};
27use tuwunel_core::{
28 Result, at, implement, is_not_empty,
29 matrix::PduCount,
30 utils::{
31 BoolExt, FutureBoolExt, ReadyExt,
32 future::ReadyBoolExt,
33 result::{LogErr, NotFound},
34 },
35 warn,
36};
37use tuwunel_database::{Json, Txn, keyval::ValBuf, serialize_key, serialize_val};
38
39pub type StrippedRoomState = Option<Vec<Raw<AnyStrippedStateEvent>>>;
45
46pub struct MembershipUpdate<'a> {
51 pub room_id: &'a RoomId,
55
56 pub user_id: &'a UserId,
60
61 pub membership_event: RoomMemberEventContent,
65
66 pub sender: &'a UserId,
70
71 pub last_state: StrippedRoomState,
76
77 pub invite_via: Option<Vec<OwnedServerName>>,
82
83 pub update_joined_count: bool,
88
89 pub count: PduCount,
93}
94
95#[implement(super::Service)]
103#[tracing::instrument(
104 level = "debug",
105 skip_all,
106 fields(
107 %room_id,
108 %user_id,
109 %sender,
110 %count,
111 ?membership_event,
112 ),
113 )]
114pub async fn update_membership(
115 &self,
116 MembershipUpdate {
117 room_id,
118 user_id,
119 membership_event,
120 sender,
121 last_state,
122 invite_via,
123 update_joined_count,
124 count,
125 }: MembershipUpdate<'_>,
126) -> Result {
127 let membership = membership_event.membership;
128
129 self.ensure_remote_user(user_id).await?;
130
131 match membership {
132 | MembershipState::Join => {
133 self.handle_join(room_id, user_id, count).await?;
134 },
135 | MembershipState::Invite => {
136 self.mark_as_invited(user_id, room_id, count, last_state, invite_via)
139 .await?;
140 },
141 | MembershipState::Leave | MembershipState::Ban => {
142 self.handle_leave(room_id, user_id, count).await;
143
144 if self.services.globals.user_is_local(user_id) {
146 self.services
147 .sending
148 .refresh_push_badge(user_id)
149 .await
150 .log_err()
151 .ok();
152 }
153 },
154 | MembershipState::Knock => {
155 self.mark_as_knocked(user_id, room_id, count, last_state);
156 },
157 | _ => {},
158 }
159
160 if update_joined_count {
161 self.update_joined_count(room_id).await;
162 }
163
164 Ok(())
165}
166
167#[implement(super::Service)]
174#[tracing::instrument(level = "debug", skip(self))]
175pub async fn update_joined_count(&self, room_id: &RoomId) {
176 let joined = self.joined_count(room_id);
177 let invited = self.room_members_invited(room_id).count();
178 let knocked = self.room_members_knocked(room_id).count();
179 let ((joinedcount, joined_servers), invitedcount, knockedcount) =
181 join3(joined, invited, knocked).await;
182
183 let invitedcount = u64::try_from(invitedcount).unwrap_or(0);
184 let knockedcount = u64::try_from(knockedcount).unwrap_or(0);
185 let txn = Txn::insert_each_slice(&[
186 (&self.db.roomid_joinedcount, room_id, joinedcount.to_be_bytes()),
187 (&self.db.roomid_invitedcount, room_id, invitedcount.to_be_bytes()),
188 (&self.db.roomid_knockedcount, room_id, knockedcount.to_be_bytes()),
189 ]);
190
191 let (txn, joined_servers) = self
192 .room_servers(room_id)
193 .ready_fold((txn, joined_servers), |(mut txn, mut servers), old_server| {
194 if !servers.remove(old_server) {
195 txn.del(&self.db.roomserverids, (room_id, old_server));
196 txn.del(&self.db.serverroomids, (old_server, room_id));
197 }
198
199 (txn, servers)
200 })
201 .await;
202
203 joined_servers
204 .iter()
205 .fold(txn, |mut txn, server| {
206 let roomserver_id =
207 serialize_key((room_id, server)).expect("failed to serialize roomserver_id");
208
209 let serverroom_id =
210 serialize_key((server, room_id)).expect("failed to serialize serverroom_id");
211
212 txn.insert_raw(&self.db.roomserverids, roomserver_id, []);
213 txn.insert_raw(&self.db.serverroomids, serverroom_id, []);
214 txn
215 })
216 .execute();
217
218 self.appservice_in_room_cache
219 .write()
220 .expect("locked")
221 .remove(room_id);
222}
223
224#[implement(super::Service)]
225fn joined_count<'a>(
226 &'a self,
227 room_id: &'a RoomId,
228) -> impl Future<Output = (u64, HashSet<OwnedServerName>)> + Send + 'a {
229 self.room_members(room_id).ready_fold(
230 (0_u64, HashSet::new()),
231 |(count, mut servers), joined| {
232 servers.insert(joined.server_name().to_owned());
233 (count.saturating_add(1), servers)
234 },
235 )
236}
237
238#[implement(super::Service)]
245#[tracing::instrument(skip(self), level = "debug")]
246pub(crate) fn mark_as_joined(&self, user_id: &UserId, room_id: &RoomId, count: PduCount) {
247 let userroom_id = (user_id, room_id);
248 let userroom_id = serialize_key(userroom_id).expect("failed to serialize userroom_id");
249
250 let roomuser_id = (room_id, user_id);
251 let roomuser_id = serialize_key(roomuser_id).expect("failed to serialize roomuser_id");
252
253 let count = count.into_unsigned().to_be_bytes();
254 let mut txn = self.services.db.txn();
255
256 txn.insert_raw(&self.db.userroomid_joinedcount, &userroom_id, count);
257 txn.insert_raw(&self.db.roomuserid_joinedcount, &roomuser_id, count);
258 txn.del_raw(&self.db.userroomid_invitestate, &userroom_id);
259 txn.del_raw(&self.db.roomuserid_invitecount, &roomuser_id);
260 txn.del_raw(&self.db.userroomid_leftstate, &userroom_id);
261 txn.del_raw(&self.db.roomuserid_leftcount, &roomuser_id);
262 txn.del_raw(&self.db.userroomid_knockedstate, &userroom_id);
263 txn.del_raw(&self.db.roomuserid_knockedcount, &roomuser_id);
264 txn.execute();
265}
266
267#[implement(super::Service)]
274#[tracing::instrument(skip(self), level = "debug")]
275pub(crate) fn mark_as_left(&self, user_id: &UserId, room_id: &RoomId, count: PduCount) {
276 let userroom_id = (user_id, room_id);
277 let userroom_id = serialize_key(userroom_id).expect("failed to serialize userroom_id");
278
279 let roomuser_id = (room_id, user_id);
280 let roomuser_id = serialize_key(roomuser_id).expect("failed to serialize roomuser_id");
281
282 let leftstate = serialize_val(Json(Vec::<Raw<AnySyncStateEvent>>::new()))
283 .expect("failed to serialize left state");
284
285 let count = count.into_unsigned().to_be_bytes();
286 let mut txn = self.services.db.txn();
287
288 txn.insert_raw(&self.db.userroomid_leftstate, &userroom_id, leftstate);
289 txn.insert_raw(&self.db.roomuserid_leftcount, &roomuser_id, count);
290 txn.del_raw(&self.db.userroomid_joinedcount, &userroom_id);
291 txn.del_raw(&self.db.roomuserid_joinedcount, &roomuser_id);
292 txn.del_raw(&self.db.userroomid_invitestate, &userroom_id);
293 txn.del_raw(&self.db.roomuserid_invitecount, &roomuser_id);
294 txn.del_raw(&self.db.userroomid_knockedstate, &userroom_id);
295 txn.del_raw(&self.db.roomuserid_knockedcount, &roomuser_id);
296 txn.execute();
297}
298
299#[implement(super::Service)]
305#[tracing::instrument(skip(self), level = "debug")]
306pub(crate) fn mark_as_knocked(
307 &self,
308 user_id: &UserId,
309 room_id: &RoomId,
310 count: PduCount,
311 knocked_state: StrippedRoomState,
312) {
313 let userroom_id = (user_id, room_id);
314 let userroom_id = serialize_key(userroom_id).expect("failed to serialize userroom_id");
315
316 let roomuser_id = (room_id, user_id);
317 let roomuser_id = serialize_key(roomuser_id).expect("failed to serialize roomuser_id");
318
319 let knocked_state = serialize_val(Json(knocked_state.unwrap_or_default()))
320 .expect("failed to serialize knocked state");
321
322 let count = count.into_unsigned().to_be_bytes();
323 let mut txn = self.services.db.txn();
324
325 txn.insert_raw(&self.db.userroomid_knockedstate, &userroom_id, knocked_state);
326 txn.insert_raw(&self.db.roomuserid_knockedcount, &roomuser_id, count);
327 txn.del_raw(&self.db.userroomid_joinedcount, &userroom_id);
328 txn.del_raw(&self.db.roomuserid_joinedcount, &roomuser_id);
329 txn.del_raw(&self.db.userroomid_invitestate, &userroom_id);
330 txn.del_raw(&self.db.roomuserid_invitecount, &roomuser_id);
331 txn.del_raw(&self.db.userroomid_leftstate, &userroom_id);
332 txn.del_raw(&self.db.roomuserid_leftcount, &roomuser_id);
333 txn.execute();
334}
335
336#[implement(super::Service)]
342#[tracing::instrument(skip(self), level = "debug")]
343pub fn forget(&self, room_id: &RoomId, user_id: &UserId) {
344 let userroom_id = (user_id, room_id);
345 let roomuser_id = (room_id, user_id);
346 let mut txn = self.services.db.txn();
347
348 txn.del(&self.db.userroomid_leftstate, userroom_id);
349 txn.del(&self.db.roomuserid_leftcount, roomuser_id);
350 txn.execute();
351}
352
353#[implement(super::Service)]
354#[tracing::instrument(level = "debug", skip(self))]
355fn mark_as_once_joined(&self, user_id: &UserId, room_id: &RoomId) {
356 let key = (user_id, room_id);
357 let key = serialize_key(key).expect("failed to serialize roomuseroncejoinedid");
358 let mut txn = self.services.db.txn();
359
360 txn.insert_raw(&self.db.roomuseroncejoinedids, key, []);
361 txn.execute();
362}
363
364pub(super) const EMPTY_INVITE_STATE: &[u8] = b"[]";
370
371#[implement(super::Service)]
378#[tracing::instrument(level = "debug", skip(self, last_state, invite_via))]
379pub(crate) async fn mark_as_invited(
380 &self,
381 user_id: &UserId,
382 room_id: &RoomId,
383 count: PduCount,
384 last_state: StrippedRoomState,
385 invite_via: Option<Vec<OwnedServerName>>,
386) -> Result {
387 let userroom_id = (user_id, room_id);
388 let userroom_id = serialize_key(userroom_id).expect("failed to serialize userroom_id");
389
390 let roomuser_id = (room_id, user_id);
391 let roomuser_id = serialize_key(roomuser_id).expect("failed to serialize roomuser_id");
392
393 let last_state = last_state.filter(is_not_empty!());
396 let stored = last_state
397 .is_none()
398 .then_async(|| self.db.userroomid_invitestate.get(&userroom_id))
399 .await
400 .transpose()
401 .optional()?
402 .flatten();
403
404 let invite_state = stored.as_deref().map_or_else(
405 || {
406 serialize_val(Json(last_state.unwrap_or_default()))
407 .expect("failed to serialize invite state")
408 },
409 ValBuf::from_slice,
410 );
411
412 let count = count.into_unsigned().to_be_bytes();
413 let mut txn = self.services.db.txn();
414
415 txn.insert_raw(&self.db.userroomid_invitestate, &userroom_id, invite_state);
416 txn.insert_raw(&self.db.roomuserid_invitecount, &roomuser_id, count);
417 txn.del_raw(&self.db.userroomid_joinedcount, &userroom_id);
418 txn.del_raw(&self.db.roomuserid_joinedcount, &roomuser_id);
419 txn.del_raw(&self.db.userroomid_leftstate, &userroom_id);
420 txn.del_raw(&self.db.roomuserid_leftcount, &roomuser_id);
421 txn.del_raw(&self.db.userroomid_knockedstate, &userroom_id);
422 txn.del_raw(&self.db.roomuserid_knockedcount, &roomuser_id);
423
424 if let Some(servers) = invite_via.filter(is_not_empty!()) {
425 self.add_servers_invite_via(&mut txn, room_id, servers)
426 .await;
427 }
428
429 txn.execute();
430
431 Ok(())
432}
433
434#[implement(super::Service)]
435#[tracing::instrument(skip(self), level = "debug")]
436async fn ensure_remote_user(&self, user_id: &UserId) -> Result {
437 if self.services.globals.user_is_local(user_id) || self.services.users.exists(user_id).await {
438 return Ok(());
439 }
440
441 self.services
442 .users
443 .create(user_id, None, None)
444 .await
445}
446
447#[implement(super::Service)]
448async fn handle_join(&self, room_id: &RoomId, user_id: &UserId, count: PduCount) -> Result {
449 if !self.once_joined(user_id, room_id).await {
450 self.mark_as_once_joined(user_id, room_id);
451 self.copy_predecessor_data(room_id, user_id)
452 .await?;
453 }
454
455 self.mark_as_joined(user_id, room_id, count);
456
457 Ok(())
458}
459
460#[implement(super::Service)]
461async fn copy_predecessor_data(&self, room_id: &RoomId, user_id: &UserId) -> Result {
462 let Ok(Some(predecessor)) = self
463 .services
464 .state_accessor
465 .room_state_get_content(room_id, &StateEventType::RoomCreate, "")
466 .await
467 .map(|content: RoomCreateEventContent| content.predecessor)
468 else {
469 return Ok(());
470 };
471
472 join(
473 self.copy_predecessor_tags(room_id, user_id, &predecessor.room_id),
474 self.copy_predecessor_direct(room_id, user_id, &predecessor.room_id),
475 )
476 .map(at!(1))
477 .await
478}
479
480#[implement(super::Service)]
481#[tracing::instrument(skip(self), level = "debug")]
482async fn copy_predecessor_tags(&self, room_id: &RoomId, user_id: &UserId, predecessor: &RoomId) {
483 let Ok(tag_event) = self
484 .services
485 .account_data
486 .get_room(predecessor, user_id, RoomAccountDataEventType::Tag)
487 .await
488 else {
489 return;
490 };
491
492 self.services
493 .account_data
494 .update(Some(room_id), user_id, RoomAccountDataEventType::Tag, &tag_event)
495 .await
496 .ok();
497}
498
499#[implement(super::Service)]
500#[tracing::instrument(skip(self), level = "debug")]
501async fn copy_predecessor_direct(
502 &self,
503 room_id: &RoomId,
504 user_id: &UserId,
505 predecessor: &RoomId,
506) -> Result {
507 let Ok(mut direct_event) = self
508 .services
509 .account_data
510 .get_global::<DirectEvent>(user_id, GlobalAccountDataEventType::Direct)
511 .await
512 else {
513 return Ok(());
514 };
515
516 let room_ids_updated =
517 direct_event
518 .content
519 .0
520 .values_mut()
521 .fold(false, |updated, room_ids| {
522 if !room_ids
523 .iter()
524 .any(|direct_room_id| direct_room_id == predecessor)
525 {
526 return updated;
527 }
528
529 room_ids.push(room_id.to_owned());
530
531 true
532 });
533
534 if !room_ids_updated {
535 return Ok(());
536 }
537
538 let event_type = GlobalAccountDataEventType::Direct
539 .to_string()
540 .into();
541
542 let direct_event =
543 serde_json::to_value(&direct_event).expect("failed to serialize DirectEvent");
544
545 self.services
546 .account_data
547 .update(None, user_id, event_type, &direct_event)
548 .await
549}
550
551#[implement(super::Service)]
552#[tracing::instrument(skip(self), level = "debug")]
553async fn handle_leave(&self, room_id: &RoomId, user_id: &UserId, count: PduCount) {
554 self.mark_as_left(user_id, room_id, count);
555
556 if !self.services.globals.user_is_local(user_id) {
557 return;
558 }
559
560 if !self.services.config.forget_forced_upon_leave {
561 let not_disabled = self
562 .services
563 .metadata
564 .is_disabled(room_id)
565 .is_false();
566
567 let not_banned = self
568 .services
569 .metadata
570 .is_banned(room_id)
571 .is_false();
572
573 if not_disabled.and(not_banned).await {
574 return;
575 }
576 }
577
578 self.forget(room_id, user_id);
579}