1use std::{collections::BTreeMap, mem::take};
2
3use axum::extract::State;
4use base64::{Engine as _, engine::general_purpose};
5use futures::StreamExt;
6use ruma::{
7 CanonicalJsonObject, CanonicalJsonValue, OwnedRoomId, OwnedUserId, RoomId, RoomVersionId,
8 ServerName, UserId,
9 api::{
10 appservice::event::push_events::{self, v1::DeviceLists},
11 error::{ErrorKind, IncompatibleRoomVersionErrorData},
12 federation::membership::{RawStrippedState, create_invite},
13 },
14 events::{
15 AnyStrippedStateEvent, GlobalAccountDataEventType, StateEventType,
16 invite_permission_config::InvitePermission,
17 push_rules::PushRulesEvent,
18 room::member::{MembershipState, RoomMemberEventContent},
19 },
20 push,
21 serde::{JsonObject, Raw},
22};
23use serde::Deserialize;
24use tuwunel_core::{
25 Err, Error, Result, debug_warn, err,
26 matrix::{Event, PduCount, PduEvent, event::gen_event_id},
27 utils,
28 utils::hash::sha256,
29};
30use tuwunel_service::{
31 Services,
32 membership::{
33 StrippedCreateVerdict, dedup_stripped_state, enforce_stripped_create,
34 into_client_stripped, v12_room_ids, without_member,
35 },
36 rooms::state_cache::MembershipUpdate,
37};
38
39use crate::{ClientIp, Ruma};
40
41#[derive(Deserialize)]
43struct ExtractIsDirect {
44 #[serde(default)]
45 is_direct: bool,
46}
47
48#[tracing::instrument(skip_all, fields(%client), name = "invite")]
52pub(crate) async fn create_invite_route(
53 State(services): State<crate::State>,
54 ClientIp(client): ClientIp,
55 mut body: Ruma<create_invite::v2::Request>,
56) -> Result<create_invite::v2::Response> {
57 services
58 .sending
59 .notify_peer_alive(body.origin())
60 .await;
61
62 validate_request(&services, &body).await?;
63
64 let stripped_state = dedup_stripped_state(take(&mut body.body.invite_room_state));
66
67 enforce_stripped_state(&services, &body, &stripped_state).await?;
68
69 let (mut signed_event, invited_user) = parse_and_validate_event(&services, &body).await?;
70
71 sign_event(&services, &mut signed_event, &body.room_version)?;
72
73 let sender = validate_origins(&signed_event, body.origin())?;
74
75 check_invite_permitted(&services, &body, &invited_user, sender).await?;
76
77 let pdu = build_pdu(&body)?;
78
79 let invite_state: Vec<_> = without_member(stripped_state, &invited_user)
81 .filter_map(|state| into_client_stripped(&body.room_id, state))
82 .chain([pdu.to_format()])
83 .collect();
84
85 let _federation_lock = services
88 .event_handler
89 .mutex_federation
90 .lock(&body.room_id)
91 .await;
92
93 record_local_invite(&services, &body, &invited_user, sender, invite_state, &pdu).await?;
94
95 Ok(create_invite::v2::Response {
96 event: services
97 .federation
98 .format_pdu_into(signed_event, Some(&body.room_version))
99 .await,
100 })
101}
102
103async fn validate_request(
104 services: &Services,
105 body: &Ruma<create_invite::v2::Request>,
106) -> Result<()> {
107 services
108 .event_handler
109 .acl_check(body.origin(), &body.room_id)
110 .await?;
111
112 if !services
113 .config
114 .supported_room_version(&body.room_version)
115 {
116 return Err(Error::BadRequest(
117 ErrorKind::IncompatibleRoomVersion(IncompatibleRoomVersionErrorData::new(
118 body.room_version.clone(),
119 )),
120 "Server does not support this room version.",
121 ));
122 }
123
124 if let Some(server) = body.room_id.server_name()
125 && services
126 .config
127 .is_forbidden_remote_server_name(server)
128 {
129 return Err!(Request(Forbidden("Server is banned on this homeserver.")));
130 }
131
132 Ok(())
133}
134
135async fn enforce_stripped_state(
138 services: &Services,
139 body: &Ruma<create_invite::v2::Request>,
140 stripped_state: &[RawStrippedState],
141) -> Result {
142 let verdict = services
143 .membership
144 .validate_stripped_create(stripped_state, &body.room_id, &body.room_version)
145 .await?;
146
147 if verdict != StrippedCreateVerdict::Valid {
148 debug_warn!(
149 ?verdict,
150 room_id = %body.room_id,
151 "MSC4311 invite create-event validation failed",
152 );
153 }
154
155 if enforce_stripped_create(
156 verdict,
157 v12_room_ids(&body.room_version),
158 services
159 .config
160 .enforce_stripped_state_pdu_validation,
161 ) {
162 return Err!(Request(MissingParam(
163 "The invite's m.room.create event is missing or does not validate for this room."
164 )));
165 }
166
167 Ok(())
168}
169
170async fn parse_and_validate_event(
171 services: &Services,
172 body: &Ruma<create_invite::v2::Request>,
173) -> Result<(CanonicalJsonObject, OwnedUserId)> {
174 let signed_event = utils::to_canonical_object(&body.event)
175 .map_err(|_| err!(Request(InvalidParam("Invite event is invalid."))))?;
176
177 let room_id: OwnedRoomId = signed_event
178 .get("room_id")
179 .try_into()
180 .map(RoomId::to_owned)
181 .map_err(|e| err!(Request(InvalidParam("Invalid room_id property: {e}"))))?;
182
183 if body.room_id != room_id {
184 return Err!(Request(InvalidParam("Event room_id does not match the request path.")));
185 }
186
187 let kind: StateEventType = signed_event
188 .get("type")
189 .and_then(CanonicalJsonValue::as_str)
190 .ok_or_else(|| err!(Request(BadJson("Missing type in event."))))?
191 .into();
192
193 if kind != StateEventType::RoomMember {
194 return Err!(Request(InvalidParam("Event must be m.room.member type.")));
195 }
196
197 let invited_user: OwnedUserId = signed_event
198 .get("state_key")
199 .try_into()
200 .map(UserId::to_owned)
201 .map_err(|e| err!(Request(InvalidParam("Invalid state_key property: {e}"))))?;
202
203 if !services.globals.user_is_local(&invited_user) {
204 return Err!(Request(InvalidParam("User does not belong to this homeserver.")));
205 }
206
207 let content: RoomMemberEventContent = signed_event
208 .get("content")
209 .cloned()
210 .map(Into::into)
211 .map(serde_json::from_value)
212 .transpose()
213 .map_err(|e| err!(Request(InvalidParam("Invalid content object in event: {e}"))))?
214 .ok_or_else(|| err!(Request(BadJson("Missing content in event."))))?;
215
216 if content.membership != MembershipState::Invite {
217 return Err!(Request(InvalidParam("Event membership must be invite.")));
218 }
219
220 services
221 .event_handler
222 .acl_check(invited_user.server_name(), &body.room_id)
223 .await?;
224
225 Ok((signed_event, invited_user))
226}
227
228fn sign_event(
229 services: &Services,
230 signed_event: &mut CanonicalJsonObject,
231 room_version: &RoomVersionId,
232) -> Result<()> {
233 services
234 .server_keys
235 .hash_and_sign_event(signed_event, room_version)
236 .map_err(|e| err!(Request(InvalidParam("Failed to sign event: {e}"))))?;
237
238 let event_id = gen_event_id(signed_event, room_version)?;
239 signed_event.insert("event_id".into(), CanonicalJsonValue::String(event_id.to_string()));
240
241 Ok(())
242}
243
244fn validate_origins<'a>(
245 signed_event: &'a CanonicalJsonObject,
246 body_origin: &ServerName,
247) -> Result<&'a UserId> {
248 let origin: Option<&str> = signed_event
249 .get("origin")
250 .and_then(CanonicalJsonValue::as_str);
251
252 let sender: &UserId = signed_event
253 .get("sender")
254 .try_into()
255 .map_err(|e| err!(Request(InvalidParam("Invalid sender property: {e}"))))?;
256
257 if sender.server_name() != body_origin {
258 return Err!(Request(Forbidden("Can only send invites on behalf of your users.")));
259 }
260
261 if origin.is_some_and(|origin| origin != body_origin) {
262 return Err!(Request(Forbidden("Can only send events from your origin.")));
263 }
264
265 Ok(sender)
266}
267
268async fn check_invite_permitted(
269 services: &Services,
270 body: &Ruma<create_invite::v2::Request>,
271 invited_user: &UserId,
272 sender: &UserId,
273) -> Result<()> {
274 if services.metadata.is_banned(&body.room_id).await
275 && !services.admin.user_is_admin(invited_user).await
276 {
277 return Err!(Request(Forbidden("This room is banned on this homeserver.")));
278 }
279
280 if services.config.block_non_admin_invites
281 && !services.admin.user_is_admin(invited_user).await
282 {
283 return Err!(Request(Forbidden("This server does not allow room invites.")));
284 }
285
286 if services
291 .users
292 .invite_permission(sender, invited_user)
293 .await
294 .eq(&InvitePermission::Block)
295 {
296 return Err!(Request(InviteBlocked("{invited_user} has blocked this invite.")));
297 }
298
299 Ok(())
300}
301
302fn build_pdu(body: &Ruma<create_invite::v2::Request>) -> Result<PduEvent> {
303 let mut event: JsonObject = serde_json::from_str(body.event.get())
304 .map_err(|e| err!(Request(BadJson("Invalid invite event PDU: {e}"))))?;
305
306 event.insert("event_id".into(), "$placeholder".into());
307
308 serde_json::from_value(event.into())
309 .map_err(|e| err!(Request(BadJson("Invalid invite event PDU: {e}"))))
310}
311
312async fn record_local_invite(
319 services: &Services,
320 body: &Ruma<create_invite::v2::Request>,
321 invited_user: &UserId,
322 sender: &UserId,
323 invite_state: Vec<Raw<AnyStrippedStateEvent>>,
324 pdu: &PduEvent,
325) -> Result {
326 let state_lock = services.state.mutex.lock(&body.room_id).await;
327
328 if services
329 .state_cache
330 .server_in_room(services.globals.server_name(), &body.room_id)
331 .await
332 {
333 return Ok(());
334 }
335
336 if services
337 .state_accessor
338 .room_state_get_content::<RoomMemberEventContent>(
339 &body.room_id,
340 &StateEventType::RoomMember,
341 invited_user.as_str(),
342 )
343 .await
344 .is_ok_and(|content| content.membership == MembershipState::Ban)
345 {
346 debug_warn!(
347 room_id = %body.room_id,
348 user_id = %invited_user,
349 "Recording invite while local room state shows banned membership.",
350 );
351 }
352
353 let count = services.globals.next_count();
354 services
355 .state_cache
356 .update_membership(MembershipUpdate {
357 room_id: &body.room_id,
358 user_id: invited_user,
359 membership_event: RoomMemberEventContent::new(MembershipState::Invite),
360 sender,
361 last_state: Some(invite_state),
362 invite_via: body.via.clone(),
363 update_joined_count: true,
364 count: PduCount::Normal(*count),
365 })
366 .await?;
367
368 drop(count);
369 drop(state_lock);
370
371 let is_direct = pdu
372 .get_content()
373 .is_ok_and(|content: ExtractIsDirect| content.is_direct);
374
375 services
376 .membership
377 .auto_accept(&body.room_id, invited_user, sender, is_direct);
378
379 notify_pushers(services, invited_user, pdu).await;
380
381 for appservice in services.appservice.read().await.values() {
382 if appservice.is_user_match(invited_user) {
383 services
384 .appservice
385 .send_request(appservice.registration.clone(), push_events::v1::Request {
386 events: vec![pdu.to_format()],
387 txn_id: general_purpose::URL_SAFE_NO_PAD
388 .encode(sha256::hash(pdu.event_id.as_bytes()))
389 .into(),
390 ephemeral: Vec::new(),
391 to_device: Vec::new(),
392 device_lists: DeviceLists::new(),
393 device_one_time_keys_count: BTreeMap::new(),
394 device_unused_fallback_key_types: BTreeMap::new(),
395 })
396 .await
397 .map_err(|_| {
398 err!(BadServerResponse("Failed to notify appservice about incoming invite."))
399 })?;
400 }
401 }
402
403 Ok(())
404}
405
406async fn notify_pushers(services: &Services, invited_user: &UserId, pdu: &PduEvent) {
407 services
408 .pusher
409 .get_pushkeys(invited_user)
410 .map(ToOwned::to_owned)
411 .for_each(async |pushkey| {
412 let Ok(pusher) = services
413 .pusher
414 .get_pusher(invited_user, &pushkey)
415 .await
416 else {
417 return;
418 };
419
420 let ruleset = services
421 .account_data
422 .get_global(invited_user, GlobalAccountDataEventType::PushRules)
423 .await
424 .map_or_else(
425 |_| push::Ruleset::server_default(invited_user),
426 |ev: PushRulesEvent| ev.content.global,
427 );
428
429 services
430 .pusher
431 .send_push_notice(invited_user, &pusher, &ruleset, pdu)
432 .await
433 .ok();
434 })
435 .await;
436}