1use std::any::TypeId;
2
3use ruma::{
4 CanonicalJsonValue,
5 api::{
6 auth_scheme::{
7 AccessToken, AccessTokenOptional, AppserviceToken, AppserviceTokenOptional,
8 AuthScheme, NoAccessToken, NoAuthentication,
9 },
10 client::{account::change_password, rtc::transports},
11 error::{ErrorKind, UnknownTokenErrorData},
12 federation::authentication::ServerSignatures,
13 },
14};
15use tuwunel_core::{Err, Error, Result};
16use tuwunel_service::Services;
17
18use super::{Auth, Request, Token, appservice::auth_appservice, server::auth_server};
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub(in crate::router) enum Scheme {
28 None,
29 AccessToken,
30 AccessTokenOptional,
31 AppserviceToken,
32 AppserviceTokenOptional,
33 ServerSignatures,
34}
35
36pub(in crate::router) trait AuthDispatch: AuthScheme {
42 const SCHEME: Scheme;
43
44 fn dispatch(
45 services: &Services,
46 request: &mut Request,
47 json_body: Option<&CanonicalJsonValue>,
48 token: Token,
49 route: TypeId,
50 ) -> impl Future<Output = Result<Auth>> + Send;
51}
52
53impl AuthDispatch for NoAccessToken {
54 const SCHEME: Scheme = Scheme::None;
55
56 async fn dispatch(
57 services: &Services,
58 request: &mut Request,
59 json_body: Option<&CanonicalJsonValue>,
60 token: Token,
61 route: TypeId,
62 ) -> Result<Auth> {
63 <NoAuthentication as AuthDispatch>::dispatch(services, request, json_body, token, route)
64 .await
65 }
66}
67
68impl AuthDispatch for NoAuthentication {
69 const SCHEME: Scheme = Scheme::None;
70
71 async fn dispatch(
72 _services: &Services,
73 _request: &mut Request,
74 _json_body: Option<&CanonicalJsonValue>,
75 token: Token,
76 _route: TypeId,
77 ) -> Result<Auth> {
78 match token {
79 | Token::Invalid | Token::Expired(_) | Token::None => Ok(Auth::default()),
82
83 | Token::User(user) => Ok(Auth {
84 sender_user: Some(user.0),
85 sender_device: Some(user.1),
86 _expires_at: user.2,
87 ..Auth::default()
88 }),
89
90 | Token::Appservice(info) => Ok(Auth {
91 appservice_info: Some(*info),
92 ..Auth::default()
93 }),
94 }
95 }
96}
97
98impl AuthDispatch for AccessToken {
99 const SCHEME: Scheme = Scheme::AccessToken;
100
101 async fn dispatch(
102 services: &Services,
103 request: &mut Request,
104 _json_body: Option<&CanonicalJsonValue>,
105 token: Token,
106 route: TypeId,
107 ) -> Result<Auth> {
108 match token {
109 | Token::Invalid => unknown_token(),
110 | Token::Expired(access_token) => expired_token(services, &access_token).await,
111 | Token::Appservice(info) => Ok(auth_appservice(services, request, info).await?),
112 | Token::User(user) => Ok(Auth {
113 sender_user: Some(user.0),
114 sender_device: Some(user.1),
115 _expires_at: user.2,
116 ..Auth::default()
117 }),
118 | Token::None if allows_missing_access_token(route) => Ok(Auth::default()),
119
120 | Token::None => Err!(Request(MissingToken("Missing access token."))),
121 }
122 }
123}
124
125fn allows_missing_access_token(route: TypeId) -> bool {
131 route == TypeId::of::<transports::v1::Request>()
132 || route == TypeId::of::<change_password::v3::Request>()
133}
134
135impl AuthDispatch for AccessTokenOptional {
136 const SCHEME: Scheme = Scheme::AccessTokenOptional;
137
138 async fn dispatch(
139 services: &Services,
140 _request: &mut Request,
141 _json_body: Option<&CanonicalJsonValue>,
142 token: Token,
143 _route: TypeId,
144 ) -> Result<Auth> {
145 match token {
146 | Token::Invalid => unknown_token(),
147 | Token::Expired(access_token) => expired_token(services, &access_token).await,
148 | Token::User(user) => Ok(Auth {
149 sender_user: Some(user.0),
150 sender_device: Some(user.1),
151 _expires_at: user.2,
152 ..Auth::default()
153 }),
154 | Token::Appservice(info) => Ok(Auth {
155 appservice_info: Some(*info),
156 ..Auth::default()
157 }),
158 | Token::None => Ok(Auth::default()),
159 }
160 }
161}
162
163impl AuthDispatch for AppserviceToken {
164 const SCHEME: Scheme = Scheme::AppserviceToken;
165
166 async fn dispatch(
167 services: &Services,
168 _request: &mut Request,
169 _json_body: Option<&CanonicalJsonValue>,
170 token: Token,
171 _route: TypeId,
172 ) -> Result<Auth> {
173 match token {
174 | Token::Invalid => unknown_token(),
175 | Token::Expired(access_token) => expired_token(services, &access_token).await,
176 | Token::User(_) =>
177 Err!(Request(Unauthorized("Appservice tokens must be used on this endpoint."))),
178 | Token::Appservice(info) => Ok(Auth {
179 appservice_info: Some(*info),
180 ..Auth::default()
181 }),
182 | Token::None => Err!(Request(MissingToken("Missing access token."))),
183 }
184 }
185}
186
187impl AuthDispatch for AppserviceTokenOptional {
188 const SCHEME: Scheme = Scheme::AppserviceTokenOptional;
189
190 async fn dispatch(
191 services: &Services,
192 _request: &mut Request,
193 _json_body: Option<&CanonicalJsonValue>,
194 token: Token,
195 _route: TypeId,
196 ) -> Result<Auth> {
197 match token {
198 | Token::Invalid => unknown_token(),
199 | Token::Expired(access_token) => expired_token(services, &access_token).await,
200 | Token::User(user) => Ok(Auth {
201 sender_user: Some(user.0),
202 sender_device: Some(user.1),
203 _expires_at: user.2,
204 ..Auth::default()
205 }),
206 | Token::Appservice(info) => Ok(Auth {
207 appservice_info: Some(*info),
208 ..Auth::default()
209 }),
210 | Token::None => Ok(Auth::default()),
211 }
212 }
213}
214
215impl AuthDispatch for ServerSignatures {
216 const SCHEME: Scheme = Scheme::ServerSignatures;
217
218 async fn dispatch(
219 services: &Services,
220 request: &mut Request,
221 json_body: Option<&CanonicalJsonValue>,
222 token: Token,
223 _route: TypeId,
224 ) -> Result<Auth> {
225 match token {
226 | Token::Invalid => unknown_token(),
227 | Token::Expired(access_token) => expired_token(services, &access_token).await,
228 | Token::Appservice(_) | Token::User(_) =>
229 Err!(Request(Unauthorized("Server signatures must be used on this endpoint."))),
230 | Token::None => Ok(auth_server(services, request, json_body).await?),
231 }
232 }
233}
234
235fn unknown_token() -> Result<Auth> {
236 Err(Error::BadRequest(
237 ErrorKind::UnknownToken(UnknownTokenErrorData::new()),
238 "Unknown access token.",
239 ))
240}
241
242async fn expired_token(services: &Services, access_token: &str) -> Result<Auth> {
243 services
244 .users
245 .remove_access_token_value(access_token)
246 .await;
247
248 Err(Error::BadRequest(
249 ErrorKind::UnknownToken(UnknownTokenErrorData { soft_logout: true }),
250 "Expired access token.",
251 ))
252}
253
254#[cfg(test)]
255mod tests {
256 use ruma::api::client::discovery::get_supported_versions;
257
258 use super::*;
259
260 #[test]
261 fn dedicated_routes_bypass_access_token_auth() {
262 assert!(allows_missing_access_token(TypeId::of::<transports::v1::Request>()));
263 assert!(allows_missing_access_token(TypeId::of::<change_password::v3::Request>()));
264 assert!(!allows_missing_access_token(TypeId::of::<get_supported_versions::Request>()));
265 }
266}