Skip to main content

tuwunel_api/router/
auth.rs

1mod appservice;
2mod dispatch;
3mod server;
4mod uiaa;
5
6use std::{any::TypeId, fmt::Debug, time::SystemTime};
7
8use axum::RequestPartsExt;
9use axum_extra::{
10	TypedHeader,
11	headers::{Authorization, authorization::Bearer},
12};
13use futures::{
14	FutureExt, TryFutureExt,
15	future::{select_ok, try_join},
16	pin_mut,
17};
18use ruma::{
19	CanonicalJsonValue, OwnedDeviceId, OwnedServerName, OwnedUserId,
20	api::client::{
21		directory::get_public_rooms,
22		knock::knock_room,
23		membership::{
24			ban_user, invite_user, join_room_by_id, join_room_by_id_or_alias, kick_user,
25			unban_user,
26		},
27		profile::{delete_profile_field, get_profile, get_profile_field, set_profile_field},
28		room::{create_room, upgrade_room},
29		session::{logout, logout_all},
30	},
31};
32use tuwunel_core::{Err, Result, is_less_than, smallstr::SmallString};
33use tuwunel_service::{Services, appservice::RegistrationInfo};
34
35pub(super) use self::dispatch::AuthDispatch;
36use self::dispatch::Scheme;
37pub(crate) use self::uiaa::auth_uiaa;
38use super::request::Request;
39
40type AccessToken = SmallString<[u8; 32]>;
41
42pub(super) enum Token {
43	Appservice(Box<RegistrationInfo>),
44	User((OwnedUserId, OwnedDeviceId, Option<SystemTime>)),
45	Expired(AccessToken),
46	Invalid,
47	None,
48}
49
50#[derive(Debug, Default)]
51pub(super) struct Auth {
52	pub(super) origin: Option<OwnedServerName>,
53	pub(super) sender_user: Option<OwnedUserId>,
54	pub(super) sender_device: Option<OwnedDeviceId>,
55	pub(super) appservice_info: Option<RegistrationInfo>,
56	pub(super) _expires_at: Option<SystemTime>,
57}
58
59#[tracing::instrument(level = "trace", skip_all, err(level = "debug"))]
60pub(super) async fn auth<A: AuthDispatch>(
61	services: &Services,
62	request: &mut Request,
63	json_body: Option<&CanonicalJsonValue>,
64	route: TypeId,
65) -> Result<Auth> {
66	let bearer: Option<TypedHeader<Authorization<Bearer>>> =
67		request.parts.extract().await.unwrap_or(None);
68
69	let access_token = match &bearer {
70		| Some(TypedHeader(Authorization(bearer))) => Some(bearer.token()),
71		| None => request.query.access_token.as_deref(),
72	};
73
74	let token = match find_token(services, access_token).await? {
75		| Token::User((_, _, expires_at))
76			if expires_at.is_some_and(is_less_than!(SystemTime::now())) =>
77			Token::Expired(access_token.unwrap_or_default().into()),
78
79		| token => token,
80	};
81
82	if A::SCHEME == Scheme::None {
83		check_auth_still_required(services, &token, route)?;
84	}
85
86	let auth = A::dispatch(services, request, json_body, token, route).await?;
87
88	try_join(
89		locked_account_check(services, &auth, route),
90		suspended_account_check(services, &auth, route),
91	)
92	.await?;
93
94	Ok(auth)
95}
96
97/// MSC3939: 401 `M_USER_LOCKED` for locked accounts; logout endpoints
98/// bypass. `soft_logout: true` is emitted by ruma for this errcode.
99#[inline(never)]
100async fn locked_account_check(services: &Services, auth: &Auth, route: TypeId) -> Result {
101	let Some(user_id) = auth.sender_user.as_deref() else {
102		return Ok(());
103	};
104
105	let is_logout = route == TypeId::of::<logout::v3::Request>()
106		|| route == TypeId::of::<logout_all::v3::Request>();
107
108	if is_logout {
109		return Ok(());
110	}
111
112	services.users.locked_check(user_id).await
113}
114
115/// MSC3823: 403 `M_USER_SUSPENDED` on membership, room create/upgrade, and
116/// profile routes. Companion checks: self-redaction and self-leave carve-outs
117/// in the /send, /redact, and /state handlers; propagation in the profile
118/// service.
119#[inline(never)]
120async fn suspended_account_check(services: &Services, auth: &Auth, route: TypeId) -> Result {
121	let Some(user_id) = auth.sender_user.as_deref() else {
122		return Ok(());
123	};
124
125	let blocked = route == TypeId::of::<join_room_by_id::v3::Request>()
126		|| route == TypeId::of::<join_room_by_id_or_alias::v3::Request>()
127		|| route == TypeId::of::<invite_user::v3::Request>()
128		|| route == TypeId::of::<knock_room::v3::Request>()
129		|| route == TypeId::of::<kick_user::v3::Request>()
130		|| route == TypeId::of::<ban_user::v3::Request>()
131		|| route == TypeId::of::<unban_user::v3::Request>()
132		|| route == TypeId::of::<create_room::v3::Request>()
133		|| route == TypeId::of::<upgrade_room::v3::Request>()
134		|| route == TypeId::of::<set_profile_field::v3::Request>()
135		|| route == TypeId::of::<delete_profile_field::v3::Request>();
136
137	if !blocked || !services.users.is_suspended(user_id).await {
138		return Ok(());
139	}
140
141	Err!(Request(UserSuspended("Account is suspended.")))
142}
143
144#[inline(never)]
145fn check_auth_still_required(services: &Services, token: &Token, route: TypeId) -> Result {
146	let is_profile = route == TypeId::of::<get_profile::v3::Request>()
147		|| route == TypeId::of::<get_profile_field::v3::Request>();
148
149	let is_public_rooms = route == TypeId::of::<get_public_rooms::v3::Request>();
150
151	if (is_profile
152		&& services
153			.server
154			.config
155			.require_auth_for_profile_requests)
156		|| (is_public_rooms
157			&& !services
158				.server
159				.config
160				.allow_public_room_directory_without_auth)
161	{
162		match token {
163			| Token::Appservice(_) | Token::User(_) => Ok(()),
164			| Token::None | Token::Expired(_) | Token::Invalid =>
165				Err!(Request(MissingToken("Missing or invalid access token."))),
166		}
167	} else {
168		Ok(())
169	}
170}
171
172async fn find_token(services: &Services, token: Option<&str>) -> Result<Token> {
173	let Some(token) = token else {
174		return Ok(Token::None);
175	};
176
177	let user_token = services
178		.users
179		.find_from_token(token)
180		.map_ok(Token::User);
181
182	let appservice_token = services
183		.appservice
184		.find_from_access_token(token)
185		.map_ok(Box::new)
186		.map_ok(Token::Appservice);
187
188	pin_mut!(user_token, appservice_token);
189	match select_ok([user_token.left_future(), appservice_token.right_future()]).await {
190		| Err(e) if !e.is_not_found() => Err(e),
191		| Ok((token, _)) => Ok(token),
192		| _ => Ok(Token::Invalid),
193	}
194}