1use std::borrow::Borrow;
2
3use futures::future::{join, join3};
4use ruma::{
5 AnyKeyName, EventId, SigningKeyId, UserId,
6 events::room::member::MembershipState,
7 room_version_rules::AuthorizationRules,
8 serde::{Base64, base64::Standard},
9 signatures::verify_canonical_json_bytes,
10};
11use tuwunel_core::{Err, Result, err, matrix::Event};
12
13#[cfg(test)]
14mod tests;
15
16#[cfg(test)]
17use super::test_utils;
18use super::{
19 FetchState, auth_input_error,
20 events::{
21 JoinRule, RoomCreateEvent, RoomMemberEvent, RoomPowerLevelsIntField,
22 member::ThirdPartyInvite, power_levels::RoomPowerLevelsEventOptionExt,
23 },
24};
25
26#[tracing::instrument(level = "trace", skip_all)]
33pub(super) async fn check_room_member<Fetch, Pdu>(
34 room_member_event: &RoomMemberEvent<Pdu>,
35 rules: &AuthorizationRules,
36 room_create_event: &RoomCreateEvent<Fetch::Pdu>,
37 fetch_state: Fetch,
38) -> Result
39where
40 Fetch: FetchState,
41 Pdu: Event,
42{
43 let Some(state_key) = room_member_event.state_key() else {
46 return Err!("missing `state_key` field in `m.room.member` event");
47 };
48
49 let target_user = <&UserId>::try_from(state_key)
50 .map_err(|e| err!("invalid `state_key` field in `m.room.member` event: {e}"))?;
51
52 if !room_create_event
56 .federate()
57 .map_err(auth_input_error)?
58 && target_user.server_name() != room_create_event.sender().server_name()
59 {
60 return Err!(
61 "MSC4361: room is not federated and target user domain does not match \
62 `m.room.create` event's sender domain"
63 );
64 }
65
66 let target_membership = room_member_event.membership()?;
67
68 match target_membership {
76 | MembershipState::Join =>
78 check_room_member_join(
79 room_member_event,
80 target_user,
81 rules,
82 room_create_event,
83 fetch_state,
84 )
85 .await,
86
87 | MembershipState::Invite =>
89 check_room_member_invite(
90 room_member_event,
91 target_user,
92 rules,
93 room_create_event,
94 fetch_state,
95 )
96 .await,
97
98 | MembershipState::Leave =>
100 check_room_member_leave(
101 room_member_event,
102 target_user,
103 rules,
104 room_create_event,
105 fetch_state,
106 )
107 .await,
108
109 | MembershipState::Ban =>
111 check_room_member_ban(
112 room_member_event,
113 target_user,
114 rules,
115 room_create_event,
116 fetch_state,
117 )
118 .await,
119
120 | MembershipState::Knock if rules.knocking =>
122 check_room_member_knock(room_member_event, target_user, rules, fetch_state).await,
123
124 | _ => Err!("unknown membership"),
126 }
127}
128
129#[tracing::instrument(level = "trace", skip_all)]
132async fn check_room_member_join<Fetch, Pdu>(
133 room_member_event: &RoomMemberEvent<Pdu>,
134 target_user: &UserId,
135 rules: &AuthorizationRules,
136 room_create_event: &RoomCreateEvent<Fetch::Pdu>,
137 fetch_state: Fetch,
138) -> Result
139where
140 Fetch: FetchState,
141 Pdu: Event,
142{
143 let creator = room_create_event
144 .creator(rules)
145 .map_err(auth_input_error)?;
146
147 let creators = room_create_event
148 .creators(rules)
149 .map_err(auth_input_error)?;
150
151 let mut prev_events = room_member_event.prev_events();
152
153 let prev_event_is_room_create_event = prev_events.next().is_some_and(|event_id| {
154 <EventId as Borrow<str>>::borrow(event_id)
155 == <EventId as Borrow<str>>::borrow(room_create_event.event_id())
156 });
157
158 let prev_event_is_only_room_create_event =
159 prev_event_is_room_create_event && prev_events.next().is_none();
160
161 if prev_event_is_only_room_create_event && *target_user == *creator {
167 return Ok(());
168 }
169
170 if room_member_event.sender() != target_user {
172 return Err!("sender of join event must match target user");
173 }
174
175 let (current_membership, join_rule) =
176 join(fetch_state.user_membership(target_user), fetch_state.join_rule()).await;
177
178 let current_membership = current_membership?;
180 if current_membership == MembershipState::Ban {
181 return Err!("banned user cannot join room");
182 }
183
184 let join_rule = join_rule?;
189 if (join_rule == JoinRule::Invite || rules.knocking && join_rule == JoinRule::Knock)
190 && matches!(current_membership, MembershipState::Invite | MembershipState::Join)
191 {
192 return Ok(());
193 }
194
195 if rules.restricted_join_rule && matches!(join_rule, JoinRule::Restricted)
198 || rules.knock_restricted_join_rule && matches!(join_rule, JoinRule::KnockRestricted)
199 {
200 if matches!(current_membership, MembershipState::Join | MembershipState::Invite) {
202 return Ok(());
203 }
204
205 let Some(authorized_via_user) = room_member_event.join_authorised_via_users_server()?
210 else {
211 return Err!(
213 "cannot join restricted room without `join_authorised_via_users_server` field \
214 if not invited"
215 );
216 };
217
218 let authorized_via_user_membership = fetch_state
220 .user_membership(&authorized_via_user)
221 .await?;
222
223 if authorized_via_user_membership != MembershipState::Join {
224 return Err!("`join_authorised_via_users_server` is not joined");
225 }
226
227 let room_power_levels_event = fetch_state.room_power_levels_event().await?;
228
229 let authorized_via_user_power_level = room_power_levels_event
230 .user_power_level(&authorized_via_user, creators, rules)
231 .map_err(auth_input_error)?;
232
233 let invite_power_level = room_power_levels_event
234 .get_as_int_or_default(RoomPowerLevelsIntField::Invite, rules)
235 .map_err(auth_input_error)?;
236
237 if authorized_via_user_power_level < invite_power_level {
238 return Err!("`join_authorised_via_users_server` does not have enough power");
239 }
240
241 return Ok(());
242 }
243
244 if join_rule != JoinRule::Public {
246 return Err!("cannot join a room that is not `public`");
247 }
248
249 Ok(())
250}
251
252#[tracing::instrument(level = "trace", skip_all)]
255async fn check_room_member_invite<Fetch, Pdu>(
256 room_member_event: &RoomMemberEvent<Pdu>,
257 target_user: &UserId,
258 rules: &AuthorizationRules,
259 room_create_event: &RoomCreateEvent<Fetch::Pdu>,
260 fetch_state: Fetch,
261) -> Result
262where
263 Fetch: FetchState,
264 Pdu: Event,
265{
266 let third_party_invite = room_member_event.third_party_invite()?;
267
268 if let Some(third_party_invite) = third_party_invite {
270 return check_third_party_invite(
271 room_member_event,
272 &third_party_invite,
273 target_user,
274 fetch_state,
275 )
276 .await;
277 }
278
279 let sender_user = room_member_event.sender();
280 let (sender_membership, current_target_user_membership, room_power_levels_event) = join3(
281 fetch_state.user_membership(sender_user),
282 fetch_state.user_membership(target_user),
283 fetch_state.room_power_levels_event(),
284 )
285 .await;
286
287 let sender_membership = sender_membership?;
289 if sender_membership != MembershipState::Join {
290 return Err!("cannot invite user if sender is not joined");
291 }
292
293 let current_target_user_membership = current_target_user_membership?;
295 if matches!(current_target_user_membership, MembershipState::Join | MembershipState::Ban) {
296 return Err!("cannot invite user that is joined or banned");
297 }
298
299 let room_power_levels_event = room_power_levels_event?;
300
301 let creators = room_create_event
302 .creators(rules)
303 .map_err(auth_input_error)?;
304
305 let sender_power_level = room_power_levels_event
306 .user_power_level(room_member_event.sender(), creators, rules)
307 .map_err(auth_input_error)?;
308
309 let invite_power_level = room_power_levels_event
310 .get_as_int_or_default(RoomPowerLevelsIntField::Invite, rules)
311 .map_err(auth_input_error)?;
312
313 if sender_power_level < invite_power_level {
316 return Err!("sender does not have enough power to invite");
317 }
318
319 Ok(())
320}
321
322#[tracing::instrument(level = "trace", skip_all)]
325async fn check_third_party_invite<Fetch, Pdu>(
326 room_member_event: &RoomMemberEvent<Pdu>,
327 third_party_invite: &ThirdPartyInvite,
328 target_user: &UserId,
329 fetch_state: Fetch,
330) -> Result
331where
332 Fetch: FetchState,
333 Pdu: Event,
334{
335 let current_target_user_membership = fetch_state.user_membership(target_user).await?;
336
337 if current_target_user_membership == MembershipState::Ban {
339 return Err!("cannot invite user that is banned");
340 }
341
342 let third_party_invite_token = third_party_invite.token()?;
345 let third_party_invite_mxid = third_party_invite.mxid()?;
346
347 if target_user != third_party_invite_mxid {
349 return Err!("third-party invite mxid does not match target user");
350 }
351
352 let Some(room_third_party_invite_event) = fetch_state
355 .room_third_party_invite_event(third_party_invite_token)
356 .await?
357 else {
358 return Err!("no `m.room.third_party_invite` in room state matches the token");
359 };
360
361 if room_member_event.sender() != room_third_party_invite_event.sender() {
364 return Err!(
365 "sender of `m.room.third_party_invite` does not match sender of `m.room.member`"
366 );
367 }
368
369 let signatures = third_party_invite.signatures()?;
370 let public_keys = room_third_party_invite_event
371 .public_keys()
372 .map_err(auth_input_error)?;
373
374 let signed_canonical_json = third_party_invite.signed_canonical_json()?;
375
376 for entity_signatures_value in signatures.values() {
379 let Some(entity_signatures) = entity_signatures_value.as_object() else {
380 return Err!(Request(InvalidParam(
381 "unexpected format of `signatures` field in `third_party_invite.signed` of \
382 `m.room.member` event: expected a map of string to object, got \
383 {entity_signatures_value:?}"
384 )));
385 };
386
387 for (key_id, signature_value) in entity_signatures {
391 let Ok(parsed_key_id) = <&SigningKeyId<AnyKeyName>>::try_from(key_id.as_str()) else {
392 continue;
393 };
394
395 let Some(signature_str) = signature_value.as_str() else {
396 continue;
397 };
398
399 let Ok(signature) = Base64::<Standard>::parse(signature_str) else {
400 continue;
401 };
402
403 let algorithm = parsed_key_id.algorithm();
404 for encoded_public_key in &public_keys {
405 let Ok(public_key) = encoded_public_key.decode() else {
406 continue;
407 };
408
409 if verify_canonical_json_bytes(
410 &algorithm,
411 &public_key,
412 signature.as_bytes(),
413 signed_canonical_json.as_bytes(),
414 )
415 .is_ok()
416 {
417 return Ok(());
418 }
419 }
420 }
421 }
422
423 Err!(
425 "no signature on third-party invite matches a public key in `m.room.third_party_invite` \
426 event"
427 )
428}
429
430#[tracing::instrument(level = "trace", skip_all)]
433async fn check_room_member_leave<Fetch, Pdu>(
434 room_member_event: &RoomMemberEvent<Pdu>,
435 target_user: &UserId,
436 rules: &AuthorizationRules,
437 room_create_event: &RoomCreateEvent<Fetch::Pdu>,
438 fetch_state: Fetch,
439) -> Result
440where
441 Fetch: FetchState,
442 Pdu: Event,
443{
444 let (sender_membership, room_power_levels_event, current_target_user_membership) = join3(
445 fetch_state.user_membership(room_member_event.sender()),
446 fetch_state.room_power_levels_event(),
447 fetch_state.user_membership(target_user),
448 )
449 .await;
450
451 let sender_membership = sender_membership?;
452
453 if room_member_event.sender() == target_user {
458 let membership_is_invite_or_join =
459 matches!(sender_membership, MembershipState::Join | MembershipState::Invite);
460 let membership_is_knock = rules.knocking && sender_membership == MembershipState::Knock;
461
462 return if membership_is_invite_or_join || membership_is_knock {
463 Ok(())
464 } else {
465 Err!("cannot leave if not joined, invited or knocked")
466 };
467 }
468
469 if sender_membership != MembershipState::Join {
471 return Err!("cannot kick if sender is not joined");
472 }
473
474 let creators = room_create_event
475 .creators(rules)
476 .map_err(auth_input_error)?;
477
478 let current_target_user_membership = current_target_user_membership?;
479 let room_power_levels_event = room_power_levels_event?;
480
481 let sender_power_level = room_power_levels_event
482 .user_power_level(room_member_event.sender(), creators.clone(), rules)
483 .map_err(auth_input_error)?;
484
485 let ban_power_level = room_power_levels_event
486 .get_as_int_or_default(RoomPowerLevelsIntField::Ban, rules)
487 .map_err(auth_input_error)?;
488
489 if current_target_user_membership == MembershipState::Ban
492 && sender_power_level < ban_power_level
493 {
494 return Err!("sender does not have enough power to unban");
495 }
496
497 let kick_power_level = room_power_levels_event
498 .get_as_int_or_default(RoomPowerLevelsIntField::Kick, rules)
499 .map_err(auth_input_error)?;
500
501 let target_user_power_level = room_power_levels_event
502 .user_power_level(target_user, creators, rules)
503 .map_err(auth_input_error)?;
504
505 if sender_power_level >= kick_power_level && target_user_power_level < sender_power_level {
511 Ok(())
512 } else {
513 Err!("sender does not have enough power to kick target user")
514 }
515}
516
517#[tracing::instrument(level = "trace", skip_all)]
520async fn check_room_member_ban<Fetch, Pdu>(
521 room_member_event: &RoomMemberEvent<Pdu>,
522 target_user: &UserId,
523 rules: &AuthorizationRules,
524 room_create_event: &RoomCreateEvent<Fetch::Pdu>,
525 fetch_state: Fetch,
526) -> Result
527where
528 Fetch: FetchState,
529 Pdu: Event,
530{
531 let (sender_membership, room_power_levels_event) = join(
532 fetch_state.user_membership(room_member_event.sender()),
533 fetch_state.room_power_levels_event(),
534 )
535 .await;
536
537 let sender_membership = sender_membership?;
539 if sender_membership != MembershipState::Join {
540 return Err!("cannot ban if sender is not joined");
541 }
542
543 let room_power_levels_event = room_power_levels_event?;
544
545 let creators = room_create_event
546 .creators(rules)
547 .map_err(auth_input_error)?;
548
549 let sender_power_level = room_power_levels_event
550 .user_power_level(room_member_event.sender(), creators.clone(), rules)
551 .map_err(auth_input_error)?;
552
553 let ban_power_level = room_power_levels_event
554 .get_as_int_or_default(RoomPowerLevelsIntField::Ban, rules)
555 .map_err(auth_input_error)?;
556
557 let target_user_power_level = room_power_levels_event
558 .user_power_level(target_user, creators, rules)
559 .map_err(auth_input_error)?;
560
561 if sender_power_level >= ban_power_level && target_user_power_level < sender_power_level {
566 Ok(())
567 } else {
568 Err!("sender does not have enough power to ban target user")
569 }
570}
571
572#[tracing::instrument(level = "trace", skip_all)]
575async fn check_room_member_knock<Fetch, Pdu>(
576 room_member_event: &RoomMemberEvent<Pdu>,
577 target_user: &UserId,
578 rules: &AuthorizationRules,
579 fetch_state: Fetch,
580) -> Result
581where
582 Fetch: FetchState,
583 Pdu: Event,
584{
585 let sender = room_member_event.sender();
586 let (join_rule, sender_membership) =
587 join(fetch_state.join_rule(), fetch_state.user_membership(sender)).await;
588
589 let join_rule = join_rule?;
593 let supports_knock = matches!(join_rule, JoinRule::Knock)
594 || (rules.knock_restricted_join_rule && matches!(join_rule, JoinRule::KnockRestricted));
595
596 if !supports_knock {
597 return Err!(
598 "join rule is not set to knock or knock_restricted, knocking is not allowed"
599 );
600 }
601
602 if room_member_event.sender() != target_user {
604 return Err!("cannot make another user knock, sender does not match target user");
605 }
606
607 let sender_membership = sender_membership?;
610 if !matches!(
611 sender_membership,
612 MembershipState::Ban | MembershipState::Invite | MembershipState::Join
613 ) {
614 Ok(())
615 } else {
616 Err!("cannot knock if user is banned, invited or joined")
617 }
618}