1use std::{borrow::Borrow, collections::HashMap, iter::once, sync::Arc};
2
3use futures::{FutureExt, StreamExt};
4use ruma::{
5 CanonicalJsonObject, CanonicalJsonValue, OwnedEventId, OwnedServerName, RoomId,
6 RoomOrAliasId, RoomVersionId, UserId,
7 api::federation::{self, membership::RawStrippedState},
8 canonical_json::to_canonical_value,
9 events::{
10 StateEventType,
11 room::member::{MembershipState, RoomMemberEventContent},
12 },
13};
14use tuwunel_core::{
15 Err, Event, PduCount, Result, async_noinline, at, debug, debug_info, debug_warn, err,
16 implement, info,
17 matrix::event::gen_event_id,
18 pdu::{PduBuilder, PduEvent},
19 trace, utils, warn,
20};
21
22use super::{
23 Service, StrippedCreateVerdict, enforce_stripped_create, into_client_stripped, v12_room_ids,
24};
25use crate::{
26 membership::join::get_servers_for_room,
27 rooms::{
28 state::RoomMutexGuard,
29 state_cache::MembershipUpdate,
30 state_compressor::{CompressedState, HashSetCompressStateEvent},
31 },
32};
33
34#[implement(Service)]
35#[async_noinline]
36#[tracing::instrument(
37 name = "knock",
38 level = "debug",
39 skip_all,
40 fields(%sender_user, %room_id)
41)]
42pub async fn knock<'a>(
43 &'a self,
44 sender_user: &'a UserId,
45 room_id: &'a RoomId,
46 orig_server_name: Option<&'a RoomOrAliasId>,
47 reason: Option<String>,
48 servers: &'a [OwnedServerName],
49 state_lock: &'a RoomMutexGuard,
50) -> Result {
51 let servers =
52 get_servers_for_room(&self.services, sender_user, room_id, orig_server_name, servers)
53 .await?;
54
55 if self
56 .services
57 .state_cache
58 .is_invited(sender_user, room_id)
59 .await
60 {
61 debug_warn!("{sender_user} is already invited in {room_id} but attempted to knock");
62 return Err!(Request(Forbidden(
63 "You cannot knock on a room you are already invited/accepted to."
64 )));
65 }
66
67 if self
68 .services
69 .state_cache
70 .is_joined(sender_user, room_id)
71 .await
72 {
73 debug_warn!("{sender_user} is already joined in {room_id} but attempted to knock");
74 return Err!(Request(Forbidden("You cannot knock on a room you are already joined in.")));
75 }
76
77 let server_in_room = self
78 .services
79 .state_cache
80 .server_in_room(self.services.globals.server_name(), room_id)
81 .await;
82
83 if server_in_room
85 && self
86 .services
87 .state_cache
88 .is_knocked(sender_user, room_id)
89 .await
90 {
91 debug_warn!("{sender_user} is already knocked in {room_id}");
92 return Ok(());
93 }
94
95 if let Ok(membership) = self
96 .services
97 .state_accessor
98 .get_member(room_id, sender_user)
99 .await && membership.membership == MembershipState::Ban
100 {
101 debug_warn!("{sender_user} is banned from {room_id} but attempted to knock");
102 return Err!(Request(Forbidden("You cannot knock on a room you are banned from.")));
103 }
104
105 let local_knock = server_in_room
106 || servers.is_empty()
107 || (servers.len() == 1 && self.services.globals.server_is_ours(&servers[0]));
108
109 if local_knock {
110 self.knock_room_helper_local(sender_user, room_id, reason, &servers, state_lock)
111 .boxed()
112 .await
113 } else {
114 self.knock_room_helper_remote(sender_user, room_id, reason, &servers, state_lock)
115 .boxed()
116 .await
117 }
118}
119
120#[implement(Service)]
121async fn knock_room_helper_local(
122 &self,
123 sender_user: &UserId,
124 room_id: &RoomId,
125 reason: Option<String>,
126 servers: &[OwnedServerName],
127 state_lock: &RoomMutexGuard,
128) -> Result {
129 debug_info!("We can knock locally");
130
131 let room_version_id = self
132 .services
133 .state
134 .get_room_version(room_id)
135 .await?;
136
137 ensure_room_version_supports_knock(&room_version_id)?;
138
139 let mut content = RoomMemberEventContent {
140 reason: reason.clone(),
141 ..RoomMemberEventContent::new(MembershipState::Knock)
142 };
143
144 self.services
145 .profile
146 .fill_profile_data(sender_user, &mut content)
147 .await;
148
149 let Err(error) = self
150 .services
151 .timeline
152 .build_and_append_pdu(
153 PduBuilder::state(sender_user.to_string(), &content),
154 sender_user,
155 room_id,
156 state_lock,
157 )
158 .await
159 else {
160 return Ok(());
161 };
162
163 if servers.is_empty()
164 || (servers.len() == 1 && self.services.globals.server_is_ours(&servers[0]))
165 {
166 return Err(error);
167 }
168
169 warn!("We couldn't do the knock locally, maybe federation can help to satisfy the knock");
170
171 self.knock_room_local_federation_fallback(sender_user, room_id, reason, servers, state_lock)
172 .boxed()
173 .await
174}
175
176fn ensure_room_version_supports_knock(room_version_id: &RoomVersionId) -> Result {
177 if matches!(
178 room_version_id,
179 RoomVersionId::V1
180 | RoomVersionId::V2
181 | RoomVersionId::V3
182 | RoomVersionId::V4
183 | RoomVersionId::V5
184 | RoomVersionId::V6
185 ) {
186 return Err!(Request(Forbidden("This room does not support knocking.")));
187 }
188
189 Ok(())
190}
191
192#[implement(Service)]
193async fn knock_room_local_federation_fallback(
194 &self,
195 sender_user: &UserId,
196 room_id: &RoomId,
197 reason: Option<String>,
198 servers: &[OwnedServerName],
199 state_lock: &RoomMutexGuard,
200) -> Result {
201 let (make_knock_response, remote_server) = self
202 .make_knock_request(sender_user, room_id, servers)
203 .await?;
204
205 info!("make_knock finished");
206
207 let room_version_id = make_knock_response.room_version.clone();
208
209 if !self
210 .services
211 .config
212 .supported_room_version(&room_version_id)
213 {
214 return Err!(BadServerResponse(
215 "Remote room version {room_version_id} is not supported by tuwunel"
216 ));
217 }
218
219 let (knock_event, event_id) = self
220 .build_knock_event(sender_user, room_id, reason, &make_knock_response, &room_version_id)
221 .await?;
222
223 let send_knock_response = self
224 .execute_send_knock(&remote_server, room_id, &event_id, &knock_event, &room_version_id)
225 .await?;
226
227 self.services
228 .short
229 .get_or_create_shortroomid(room_id)
230 .await;
231
232 self.finalize_knock_membership(
233 room_id,
234 sender_user,
235 &event_id,
236 knock_event,
237 send_knock_response,
238 state_lock,
239 )
240 .await
241}
242
243#[implement(Service)]
244async fn finalize_knock_membership(
245 &self,
246 room_id: &RoomId,
247 sender_user: &UserId,
248 event_id: &OwnedEventId,
249 knock_event: CanonicalJsonObject,
250 send_knock_response: federation::membership::create_knock_event::v1::Response,
251 state_lock: &RoomMutexGuard,
252) -> Result {
253 info!("Parsing knock event");
254 let parsed_knock_pdu = PduEvent::from_object_and_eventid(event_id, knock_event.clone())
255 .map_err(|e| err!(BadServerResponse("Invalid knock event PDU: {e:?}")))?;
256
257 info!("Updating membership locally to knock state with provided stripped state events");
258 let count = self.services.globals.next_count();
259 let membership_event = parsed_knock_pdu
260 .get_content::<RoomMemberEventContent>()
261 .expect("we just created this");
262
263 let last_state = send_knock_response
264 .knock_room_state
265 .into_iter()
266 .filter_map(|state| into_client_stripped(room_id, state))
267 .collect();
268
269 self.services
270 .state_cache
271 .update_membership(MembershipUpdate {
272 room_id,
273 user_id: sender_user,
274 membership_event,
275 sender: sender_user,
276 last_state: Some(last_state),
277 invite_via: None,
278 update_joined_count: false,
279 count: PduCount::Normal(*count),
280 })
281 .await?;
282
283 info!("Appending room knock event locally");
284 self.services
285 .timeline
286 .append_pdu(
287 &parsed_knock_pdu,
288 knock_event,
289 once(parsed_knock_pdu.event_id.borrow()),
290 state_lock,
291 )
292 .await?;
293
294 Ok(())
295}
296
297#[implement(Service)]
298async fn knock_room_helper_remote(
299 &self,
300 sender_user: &UserId,
301 room_id: &RoomId,
302 reason: Option<String>,
303 servers: &[OwnedServerName],
304 state_lock: &RoomMutexGuard,
305) -> Result {
306 info!("Knocking {room_id} over federation.");
307
308 let (make_knock_response, remote_server) = self
309 .make_knock_request(sender_user, room_id, servers)
310 .await?;
311
312 info!("make_knock finished");
313
314 let room_version_id = make_knock_response.room_version.clone();
315
316 if !self
317 .services
318 .config
319 .supported_room_version(&room_version_id)
320 {
321 return Err!(BadServerResponse(
322 "Remote room version {room_version_id} is not supported by tuwunel"
323 ));
324 }
325
326 let (knock_event, event_id) = self
327 .build_knock_event(sender_user, room_id, reason, &make_knock_response, &room_version_id)
328 .await?;
329
330 let send_knock_response = self
331 .execute_send_knock(&remote_server, room_id, &event_id, &knock_event, &room_version_id)
332 .await?;
333
334 self.services
335 .short
336 .get_or_create_shortroomid(room_id)
337 .await;
338
339 info!("Parsing knock event");
340 let parsed_knock_pdu = PduEvent::from_object_and_eventid(&event_id, knock_event.clone())
341 .map_err(|e| err!(BadServerResponse("Invalid knock event PDU: {e:?}")))?;
342
343 let state_map = self
344 .ingest_send_knock_state(room_id, &send_knock_response, &room_version_id)
345 .await?;
346
347 self.apply_send_knock_state(room_id, &state_map, state_lock)
348 .await?;
349
350 let statehash_after_knock = self
351 .services
352 .state
353 .append_to_state(&parsed_knock_pdu)
354 .await?;
355
356 info!("Updating membership locally to knock state with provided stripped state events");
357 let count = self.services.globals.next_count();
358 let membership_event = parsed_knock_pdu
359 .get_content::<RoomMemberEventContent>()
360 .expect("we just created this");
361
362 let last_state = send_knock_response
363 .knock_room_state
364 .into_iter()
365 .filter_map(|state| into_client_stripped(room_id, state))
366 .collect();
367
368 self.services
369 .state_cache
370 .update_membership(MembershipUpdate {
371 room_id,
372 user_id: sender_user,
373 membership_event,
374 sender: sender_user,
375 last_state: Some(last_state),
376 invite_via: None,
377 update_joined_count: false,
378 count: PduCount::Normal(*count),
379 })
380 .await?;
381
382 info!("Appending room knock event locally");
383 self.services
384 .timeline
385 .append_pdu(
386 &parsed_knock_pdu,
387 knock_event,
388 once(parsed_knock_pdu.event_id.borrow()),
389 state_lock,
390 )
391 .await?;
392
393 info!("Setting final room state for new room");
394 self.services
397 .state
398 .set_room_state(room_id, statehash_after_knock, state_lock);
399
400 Ok(())
401}
402
403#[implement(Service)]
404async fn build_knock_event(
405 &self,
406 sender_user: &UserId,
407 room_id: &RoomId,
408 reason: Option<String>,
409 make_knock_response: &federation::membership::prepare_knock_event::v1::Response,
410 room_version_id: &RoomVersionId,
411) -> Result<(CanonicalJsonObject, OwnedEventId)> {
412 let mut knock_event_stub: CanonicalJsonObject =
413 serde_json::from_str(make_knock_response.event.get()).map_err(|e| {
414 err!(BadServerResponse("Invalid make_knock event json received from server: {e:?}"))
415 })?;
416
417 let mut content = RoomMemberEventContent {
418 reason,
419 ..RoomMemberEventContent::new(MembershipState::Knock)
420 };
421
422 self.services
423 .profile
424 .fill_profile_data(sender_user, &mut content)
425 .await;
426
427 knock_event_stub.insert(
428 "origin".into(),
429 CanonicalJsonValue::String(
430 self.services
431 .globals
432 .server_name()
433 .as_str()
434 .to_owned(),
435 ),
436 );
437 knock_event_stub.insert(
438 "origin_server_ts".into(),
439 CanonicalJsonValue::Integer(
440 utils::millis_since_unix_epoch()
441 .try_into()
442 .expect("Timestamp is valid js_int value"),
443 ),
444 );
445 knock_event_stub.insert(
446 "content".into(),
447 to_canonical_value(content).expect("event is valid, we just created it"),
448 );
449
450 knock_event_stub
451 .insert("room_id".into(), CanonicalJsonValue::String(room_id.as_str().into()));
452
453 knock_event_stub
454 .insert("state_key".into(), CanonicalJsonValue::String(sender_user.as_str().into()));
455
456 knock_event_stub
457 .insert("sender".into(), CanonicalJsonValue::String(sender_user.as_str().into()));
458
459 knock_event_stub.insert("type".into(), CanonicalJsonValue::String("m.room.member".into()));
460
461 self.services
464 .server_keys
465 .hash_and_sign_event(&mut knock_event_stub, room_version_id)?;
466
467 let event_id = gen_event_id(&knock_event_stub, room_version_id)?;
468
469 knock_event_stub
470 .insert("event_id".into(), CanonicalJsonValue::String(event_id.clone().into()));
471
472 Ok((knock_event_stub, event_id))
473}
474
475#[implement(Service)]
476async fn execute_send_knock(
477 &self,
478 remote_server: &OwnedServerName,
479 room_id: &RoomId,
480 event_id: &OwnedEventId,
481 knock_event: &CanonicalJsonObject,
482 room_version_id: &RoomVersionId,
483) -> Result<federation::membership::create_knock_event::v1::Response> {
484 info!("Asking {remote_server} for send_knock in room {room_id}");
485 let send_knock_request = federation::membership::create_knock_event::v1::Request {
486 room_id: room_id.to_owned(),
487 event_id: event_id.clone(),
488 pdu: self
489 .services
490 .federation
491 .format_pdu_into(knock_event.clone(), Some(room_version_id))
492 .await,
493 };
494
495 let response = self
496 .services
497 .federation
498 .execute(remote_server, send_knock_request)
499 .await?;
500
501 info!("send_knock finished");
502 Ok(response)
503}
504
505#[implement(Service)]
506#[expect(
507 deprecated,
508 reason = "Matrix 1.16 still permits receiving the legacy stripped variant for backwards \
509 compatibility."
510)]
511async fn ingest_send_knock_state(
512 &self,
513 room_id: &RoomId,
514 send_knock_response: &federation::membership::create_knock_event::v1::Response,
515 room_version_id: &RoomVersionId,
516) -> Result<HashMap<u64, OwnedEventId>> {
517 info!("Going through send_knock response knock state events");
518
519 let verdict = self
520 .validate_stripped_create(&send_knock_response.knock_room_state, room_id, room_version_id)
521 .await?;
522
523 let enforce = self
524 .services
525 .config
526 .enforce_stripped_state_pdu_validation;
527
528 let drop_create = enforce_stripped_create(verdict, v12_room_ids(room_version_id), enforce);
529
530 if verdict != StrippedCreateVerdict::Valid {
531 debug_warn!(?verdict, %room_id, drop_create, "MSC4311 knock create-event validation failed");
532 }
533
534 let state = send_knock_response
535 .knock_room_state
536 .iter()
537 .filter_map(|event| match event {
538 | RawStrippedState::Pdu(raw) =>
539 serde_json::from_str::<CanonicalJsonObject>(raw.get()).ok(),
540 | RawStrippedState::Stripped(raw) =>
541 serde_json::from_str::<CanonicalJsonObject>(raw.json().get()).ok(),
542 });
543
544 let mut state_map: HashMap<u64, OwnedEventId> = HashMap::new();
545
546 for event in state {
547 let Some(state_key) = event.get("state_key") else {
548 debug_warn!("send_knock stripped state event missing state_key: {event:?}");
549 continue;
550 };
551 let Some(event_type) = event.get("type") else {
552 debug_warn!("send_knock stripped state event missing event type: {event:?}");
553 continue;
554 };
555
556 let Ok(state_key) = serde_json::from_value::<String>(state_key.clone().into()) else {
557 debug_warn!("send_knock stripped state event has invalid state_key: {event:?}");
558 continue;
559 };
560 let Ok(event_type) = serde_json::from_value::<StateEventType>(event_type.clone().into())
561 else {
562 debug_warn!("send_knock stripped state event has invalid event type: {event:?}");
563 continue;
564 };
565
566 if drop_create && event_type == StateEventType::RoomCreate && state_key.is_empty() {
568 debug_warn!(%room_id, "dropping unvalidated create event from knock state");
569 continue;
570 }
571
572 let event_id = gen_event_id(&event, room_version_id)?;
573 let shortstatekey = self
574 .services
575 .short
576 .get_or_create_shortstatekey(&event_type, &state_key)
577 .await;
578
579 self.services
580 .timeline
581 .add_pdu_outlier(&event_id, &event);
582
583 state_map.insert(shortstatekey, event_id.clone());
584 }
585
586 Ok(state_map)
587}
588
589#[implement(Service)]
590async fn apply_send_knock_state(
591 &self,
592 room_id: &RoomId,
593 state_map: &HashMap<u64, OwnedEventId>,
594 state_lock: &RoomMutexGuard,
595) -> Result {
596 info!("Compressing state from send_knock");
597 let compressed: CompressedState = self
598 .services
599 .state_compressor
600 .compress_state_events(
601 state_map
602 .iter()
603 .map(|(ssk, eid)| (ssk, eid.borrow())),
604 )
605 .collect()
606 .await;
607
608 debug!("Saving compressed state");
609 let HashSetCompressStateEvent {
610 shortstatehash: statehash_before_knock,
611 added,
612 removed,
613 } = self
614 .services
615 .state_compressor
616 .save_state(room_id, Arc::new(compressed))
617 .await?;
618
619 debug!("Forcing state for new room");
620 self.services
621 .state
622 .force_state(room_id, statehash_before_knock, added, removed, state_lock)
623 .await?;
624
625 Ok(())
626}
627
628#[implement(Service)]
629async fn make_knock_request(
630 &self,
631 sender_user: &UserId,
632 room_id: &RoomId,
633 servers: &[OwnedServerName],
634) -> Result<(federation::membership::prepare_knock_event::v1::Response, OwnedServerName)> {
635 let mut make_knock_response_and_server =
636 Err!(BadServerResponse("No server available to assist in knocking."));
637
638 let mut make_knock_counter: usize = 0;
639
640 for remote_server in servers {
641 if self
642 .services
643 .globals
644 .server_is_ours(remote_server)
645 {
646 continue;
647 }
648
649 info!("Asking {remote_server} for make_knock ({make_knock_counter})");
650
651 let make_knock_response = self
652 .services
653 .federation
654 .execute(remote_server, federation::membership::prepare_knock_event::v1::Request {
655 room_id: room_id.to_owned(),
656 user_id: sender_user.to_owned(),
657 ver: self
658 .services
659 .config
660 .supported_room_versions()
661 .map(at!(0))
662 .collect(),
663 })
664 .await;
665
666 trace!("make_knock response: {make_knock_response:?}");
667 make_knock_counter = make_knock_counter.saturating_add(1);
668
669 make_knock_response_and_server = make_knock_response.map(|r| (r, remote_server.clone()));
670
671 if make_knock_response_and_server.is_ok() {
672 break;
673 }
674
675 if make_knock_counter > 40 {
676 warn!(
677 "50 servers failed to provide valid make_knock response, assuming no server can \
678 assist in knocking."
679 );
680 make_knock_response_and_server =
681 Err!(BadServerResponse("No server available to assist in knocking."));
682
683 return make_knock_response_and_server;
684 }
685 }
686
687 make_knock_response_and_server
688}