Skip to main content

tuwunel_service/membership/
knock.rs

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	// Trust a local knock; re-drive a remote one in case we missed a kick.
93	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	// We set the room state after inserting the pdu, so that we never have a moment
404	// in time where events in the current room state do not exist
405	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	// In order to create a compatible ref hash (EventID) the `hashes` field needs
470	// to be present
471	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	// Settled here so all three readers of this array agree.
512	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		// MSC4311: drop a create event that failed validation when policy enforces.
579		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}