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