Skip to main content

tuwunel_api/server/
utils.rs

1use std::pin::pin;
2
3use futures::{FutureExt, StreamExt, future::join3};
4use ruma::{EventId, OwnedRoomId, RoomId, ServerName};
5use serde::Deserialize;
6use tuwunel_core::{
7	Err, Result, err, implement, is_false,
8	utils::{FutureBoolExt, future::ReadyBoolExt, option::OptionExt},
9};
10use tuwunel_service::Services;
11
12pub(super) struct AccessCheck<'a> {
13	pub(super) services: &'a Services,
14	pub(super) origin: &'a ServerName,
15	pub(super) room_id: &'a RoomId,
16	pub(super) event_id: Option<&'a EventId>,
17}
18
19#[implement(AccessCheck, params = "<'_>")]
20pub(super) async fn check(&self) -> Result {
21	let acl_check = self
22		.services
23		.event_handler
24		.acl_check(self.origin, self.room_id)
25		.map(|result| result.is_ok());
26
27	let server_in_room = self
28		.services
29		.state_cache
30		.server_in_room(self.origin, self.room_id);
31
32	let world_readable = self
33		.services
34		.state_accessor
35		.is_world_readable(self.room_id);
36
37	// if any user on our homeserver is trying to knock this room, we'll need to
38	// acknowledge bans or leaves
39	let user_is_knocking = async {
40		let knocked = self
41			.services
42			.state_cache
43			.room_members_knocked(self.room_id);
44		let mut knocked = pin!(knocked);
45
46		knocked.next().await.is_some()
47	};
48
49	let server_can_see = self.event_id.map_async(|event_id| {
50		self.services
51			.state_accessor
52			.server_can_see_event(self.origin, self.room_id, event_id)
53	});
54
55	// The cheap membership probe leads; a hit there elides the other reads.
56	let room_unreachable = server_in_room
57		.is_false()
58		.and2(world_readable.is_false(), user_is_knocking.is_false());
59
60	let (acl_check, room_unreachable, server_can_see) =
61		join3(acl_check, room_unreachable, server_can_see).await;
62
63	if !acl_check {
64		return Err!(Request(Forbidden("Server access denied.")));
65	}
66
67	if room_unreachable {
68		return Err!(Request(Forbidden("Server is not in room.")));
69	}
70
71	if server_can_see.is_some_and(is_false!()) {
72		return Err!(Request(Forbidden("Server is not allowed to see event.")));
73	}
74
75	Ok(())
76}
77
78pub(super) async fn require_known_room(
79	services: &Services,
80	room_id: &RoomId,
81	origin: &ServerName,
82) -> Result {
83	if !services.metadata.exists(room_id).await {
84		return Err!(Request(NotFound("Room is unknown to this server.")));
85	}
86
87	services
88		.event_handler
89		.acl_check(origin, room_id)
90		.await
91}
92
93pub(super) async fn require_event_in_room(
94	services: &Services,
95	event_id: &EventId,
96	room_id: &RoomId,
97) -> Result {
98	#[derive(Deserialize)]
99	struct PduRoomId {
100		room_id: OwnedRoomId,
101	}
102
103	services
104		.timeline
105		.get::<PduRoomId>(event_id)
106		.await
107		.is_ok_and(|pdu| pdu.room_id == room_id)
108		.then_some(())
109		.ok_or_else(|| err!(Request(NotFound("Event not found."))))
110}