tuwunel_service/membership/
invite.rs1use futures::{FutureExt, future::join};
2use ruma::{
3 OwnedServerName, RoomId, UserId,
4 api::{
5 error::ErrorKind,
6 federation::membership::{RawStrippedState, create_invite},
7 },
8 events::{
9 invite_permission_config::InvitePermission,
10 room::{
11 join_rules::JoinRule,
12 member::{MembershipState, RoomMemberEventContent},
13 },
14 },
15};
16use tuwunel_core::{
17 Err, Result, at, err, implement,
18 matrix::event::gen_event_id_canonical_json,
19 pdu::PduBuilder,
20 utils::future::{ReadyBoolExt, and4},
21};
22
23use super::Service;
24
25#[implement(Service)]
26#[tracing::instrument(
27 level = "debug",
28 skip_all,
29 fields(%sender_user, %room_id, %user_id)
30)]
31pub async fn invite(
32 &self,
33 sender_user: &UserId,
34 user_id: &UserId,
35 room_id: &RoomId,
36 reason: Option<&String>,
37 is_direct: bool,
38) -> Result {
39 if self.services.globals.user_is_local(user_id) {
40 self.local_invite(sender_user, user_id, room_id, reason, is_direct)
41 .boxed()
42 .await?;
43 } else {
44 self.remote_invite(sender_user, user_id, room_id, reason, is_direct)
45 .boxed()
46 .await?;
47 }
48
49 Ok(())
50}
51
52#[implement(Service)]
60pub async fn join_needs_invite(&self, room_id: &RoomId, user_id: &UserId) -> bool {
61 let server_in_room = self
62 .services
63 .state_cache
64 .server_in_room(self.services.globals.server_name(), room_id);
65
66 let joined = self
67 .services
68 .state_cache
69 .is_joined(user_id, room_id);
70
71 let invited = self
72 .services
73 .state_cache
74 .is_invited(user_id, room_id);
75
76 let admitted = self
77 .services
78 .state_accessor
79 .get_join_rules(room_id)
80 .then(async |rule| self.admits_uninvited(&rule, user_id).await);
81
82 and4(server_in_room, joined.is_false(), invited.is_false(), admitted.is_false()).await
84}
85
86#[implement(Service)]
87async fn admits_uninvited(&self, rule: &JoinRule, user_id: &UserId) -> bool {
88 match rule {
89 | JoinRule::Public => true,
90 | JoinRule::Restricted(_) | JoinRule::KnockRestricted(_) =>
91 self.services
92 .state_cache
93 .is_joined_any(user_id, rule.allowed_room_ids())
94 .await,
95 | _ => false,
96 }
97}
98
99#[implement(Service)]
100#[tracing::instrument(name = "remote", level = "debug", skip_all)]
101async fn remote_invite(
102 &self,
103 sender_user: &UserId,
104 user_id: &UserId,
105 room_id: &RoomId,
106 reason: Option<&String>,
107 is_direct: bool,
108) -> Result {
109 let (pdu, pdu_json, invite_room_state, room_version_id) = {
110 let state_lock = self.services.state.mutex.lock(room_id).await;
111
112 let content = self
113 .services
114 .profile
115 .fill_content(user_id, RoomMemberEventContent {
116 is_direct,
117 reason: reason.cloned(),
118 ..RoomMemberEventContent::new(MembershipState::Invite)
119 })
120 .await;
121
122 let event = self.services.timeline.create_hash_and_sign_event(
123 PduBuilder::state(user_id.to_string(), &content),
124 sender_user,
125 room_id,
126 &state_lock,
127 );
128
129 let room_version_id = self.services.state.get_room_version(room_id);
130 let (event, room_version_id) = join(event, room_version_id).await;
131 let (pdu, pdu_json) = event?;
132 let room_version_id = room_version_id?;
133
134 let invite_room_state = self
135 .services
136 .state
137 .summary_pdus(&pdu, &pdu_json, &room_version_id)
138 .await;
139
140 drop(state_lock);
141
142 (pdu, pdu_json, invite_room_state, room_version_id)
143 };
144
145 let event = self
146 .services
147 .federation
148 .format_pdu_into(pdu_json.clone(), Some(&room_version_id));
149
150 let via = self
151 .services
152 .state_cache
153 .servers_route_via(room_id)
154 .map(Result::ok);
155
156 let (event, via) = join(event, via).await;
157
158 let response = self
159 .services
160 .federation
161 .execute(user_id.server_name(), create_invite::v2::Request {
162 room_id: room_id.to_owned(),
163 event_id: (*pdu.event_id).to_owned(),
164 room_version: room_version_id.clone(),
165 event,
166 invite_room_state: invite_room_state
167 .into_iter()
168 .map(RawStrippedState::Pdu)
169 .collect(),
170 via,
171 })
172 .await
173 .map_err(|e| match e.kind() {
174 | ErrorKind::IncompatibleRoomVersion { .. } | ErrorKind::UnsupportedRoomVersion =>
175 err!(Request(UnsupportedRoomVersion(
176 "Server {} does not support room version {room_version_id}.",
177 user_id.server_name(),
178 ))),
179 | ErrorKind::MissingParam => err!(BadServerResponse(
182 "Remote server could not validate the invite's create event."
183 )),
184 | _ => e,
185 })?;
186
187 let (event_id, value) = gen_event_id_canonical_json(&response.event, &room_version_id)
190 .map_err(|e| {
191 err!(Request(BadJson(warn!("Could not convert event to canonical JSON: {e}"))))
192 })?;
193
194 if pdu.event_id != event_id {
195 return Err!(Request(BadJson(warn!(
196 %pdu.event_id, %event_id,
197 "Server {} sent event with wrong event ID",
198 user_id.server_name()
199 ))));
200 }
201
202 let origin: OwnedServerName = serde_json::from_value(serde_json::to_value(
203 value
204 .get("origin")
205 .ok_or_else(|| err!(Request(BadJson("Event missing origin field."))))?,
206 )?)
207 .map_err(|e| {
208 err!(Request(BadJson(warn!("Origin field in event is not a valid server name: {e}"))))
209 })?;
210
211 let pdu_id = self
212 .services
213 .event_handler
214 .handle_incoming_pdu(&origin, room_id, &event_id, value, true)
215 .await?
216 .map(at!(0))
217 .ok_or_else(|| {
218 err!(Request(InvalidParam("Could not accept incoming PDU as timeline event.")))
219 })?;
220
221 self.services
222 .sending
223 .send_pdu_room(room_id, &pdu_id)
224 .await?;
225
226 Ok(())
227}
228
229#[implement(Service)]
230#[tracing::instrument(name = "local", level = "debug", skip_all)]
231async fn local_invite(
232 &self,
233 sender_user: &UserId,
234 user_id: &UserId,
235 room_id: &RoomId,
236 reason: Option<&String>,
237 is_direct: bool,
238) -> Result {
239 let blocked = self
240 .services
241 .users
242 .invite_permission(sender_user, user_id)
243 .map(|permission| permission.eq(&InvitePermission::Block));
244
245 let joined = self
246 .services
247 .state_cache
248 .is_joined(sender_user, room_id);
249
250 let (blocked, joined) = join(blocked, joined).await;
251
252 if blocked {
253 return Err!(Request(InviteBlocked("{user_id} has blocked this invite.")));
254 }
255
256 if !joined {
257 return Err!(Request(Forbidden(
258 "You must be joined in the room you are trying to invite from."
259 )));
260 }
261
262 let state_lock = self.services.state.mutex.lock(room_id).await;
263
264 let content = self
265 .services
266 .profile
267 .fill_content(user_id, RoomMemberEventContent {
268 is_direct,
269 reason: reason.cloned(),
270 ..RoomMemberEventContent::new(MembershipState::Invite)
271 })
272 .await;
273
274 self.services
275 .timeline
276 .build_and_append_pdu(
277 PduBuilder::state(user_id.to_string(), &content),
278 sender_user,
279 room_id,
280 &state_lock,
281 )
282 .await?;
283
284 drop(state_lock);
285
286 Ok(())
287}