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#[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#[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}