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