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