1use std::{
2 collections::BTreeMap,
3 ops::ControlFlow,
4 sync::{Arc, RwLock},
5};
6
7use futures::{TryStreamExt, pin_mut};
8use ruma::{
9 CanonicalJsonValue, DeviceId, OwnedDeviceId, OwnedUserId, UserId,
10 api::{
11 client::uiaa::{
12 AuthData, AuthType, EmailIdentity, Password, ThirdpartyIdCredentials, UiaaInfo,
13 UserIdentifier,
14 },
15 error::{ErrorKind, StandardErrorBody},
16 },
17};
18use tuwunel_core::{
19 Err, Result, err, error, extract, implement,
20 utils::{self, BoolExt, hash::verify_password, string::EMPTY},
21};
22use tuwunel_database::{Deserialized, Json, Map};
23
24use crate::users::is_password_hash;
25
26pub struct Service {
27 userdevicesessionid_uiaarequest: RwLock<RequestMap>,
28 db: Data,
29 services: Arc<crate::services::OnceServices>,
30}
31
32struct Data {
33 userdevicesessionid_uiaainfo: Arc<Map>,
34}
35
36type RequestMap = BTreeMap<RequestKey, CanonicalJsonValue>;
37type RequestKey = (OwnedUserId, OwnedDeviceId, String);
38
39pub const SESSION_ID_LENGTH: usize = 32;
40
41#[derive(Clone, Copy)]
42enum EmailIdentityMode {
43 Validate,
44 Claim,
45}
46
47impl crate::Service for Service {
48 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
49 Ok(Arc::new(Self {
50 userdevicesessionid_uiaarequest: RwLock::new(RequestMap::new()),
51 db: Data {
52 userdevicesessionid_uiaainfo: args.db["userdevicesessionid_uiaainfo"].clone(),
53 },
54 services: args.services.clone(),
55 }))
56 }
57
58 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
59}
60
61#[implement(Service)]
63pub fn create(
64 &self,
65 user_id: &UserId,
66 device_id: &DeviceId,
67 uiaainfo: &UiaaInfo,
68 json_body: &CanonicalJsonValue,
69) {
70 let session = uiaainfo
73 .session
74 .as_ref()
75 .expect("session should be set");
76
77 self.set_uiaa_request(user_id, device_id, session, json_body);
78
79 self.update_uiaa_session(user_id, device_id, session, Some(uiaainfo));
80}
81
82#[implement(Service)]
87pub async fn try_auth(
88 &self,
89 user_id: &UserId,
90 device_id: &DeviceId,
91 auth: &AuthData,
92 uiaainfo: &UiaaInfo,
93) -> Result<(bool, UiaaInfo)> {
94 self.try_auth_inner(user_id, device_id, auth, uiaainfo, EmailIdentityMode::Validate)
95 .await
96}
97
98#[implement(Service)]
103pub async fn try_auth_registration(
104 &self,
105 user_id: &UserId,
106 device_id: &DeviceId,
107 auth: &AuthData,
108 uiaainfo: &UiaaInfo,
109) -> Result<(bool, UiaaInfo)> {
110 self.try_auth_inner(user_id, device_id, auth, uiaainfo, EmailIdentityMode::Claim)
111 .await
112}
113
114#[implement(Service)]
115async fn try_auth_inner(
116 &self,
117 user_id: &UserId,
118 device_id: &DeviceId,
119 auth: &AuthData,
120 uiaainfo: &UiaaInfo,
121 email_identity_mode: EmailIdentityMode,
122) -> Result<(bool, UiaaInfo)> {
123 let mut uiaainfo = if let Some(session) = auth.session() {
124 self.get_uiaa_session(user_id, device_id, session)
125 .await?
126 } else {
127 uiaainfo.clone()
128 };
129
130 if uiaainfo.session.is_none() {
131 uiaainfo.session = Some(utils::random_string(SESSION_ID_LENGTH));
132 }
133
134 match auth {
135 | AuthData::Password(password) => {
137 if let ControlFlow::Break(authed) = self
138 .verify_password(user_id, &mut uiaainfo, password)
139 .await?
140 {
141 return Ok((authed, uiaainfo));
142 }
143 },
144 | AuthData::RegistrationToken(t) => {
145 let token = t.token.trim();
146 if self
147 .services
148 .registration_tokens
149 .try_consume(token)
150 .await
151 .is_ok()
152 {
153 uiaainfo
154 .completed
155 .push(AuthType::RegistrationToken);
156 } else {
157 uiaainfo.auth_error = Some(Box::new(StandardErrorBody {
158 kind: ErrorKind::forbidden(),
159 message: "Invalid registration token.".to_owned(),
160 }));
161
162 return Ok((false, uiaainfo));
163 }
164 },
165 | AuthData::FallbackAcknowledgement(_session) => {
166 },
169 | AuthData::OAuth(_) => {
170 if !uiaainfo.completed.contains(&AuthType::OAuth) {
173 if self
174 .services
175 .users
176 .can_replace_cross_signing_keys(user_id)
177 .await
178 {
179 uiaainfo.completed.push(AuthType::OAuth);
180 } else {
181 uiaainfo.auth_error = Some(Box::new(StandardErrorBody {
182 kind: ErrorKind::forbidden(),
183 message: "OAuth cross-signing reset not approved for this session."
184 .to_owned(),
185 }));
186
187 return Ok((false, uiaainfo));
188 }
189 }
190 },
191 | AuthData::Dummy(_) => {
192 uiaainfo.completed.push(AuthType::Dummy);
193 },
194 | AuthData::Terms(_) => {
195 uiaainfo.completed.push(AuthType::Terms);
197 },
198 | AuthData::EmailIdentity(EmailIdentity { thirdparty_id_creds, .. }) => {
199 let validated = self
201 .authenticate_email_identity(
202 user_id,
203 device_id,
204 &uiaainfo,
205 thirdparty_id_creds,
206 email_identity_mode,
207 )
208 .await?;
209
210 if !validated {
211 uiaainfo.auth_error = Some(Box::new(StandardErrorBody {
212 kind: ErrorKind::forbidden(),
213 message: "Email address has not been validated.".to_owned(),
214 }));
215
216 return Ok((false, uiaainfo));
217 }
218
219 uiaainfo.completed.push(AuthType::EmailIdentity);
220 },
221 | auth => error!("AuthData type not supported: {auth:?}"),
222 }
223
224 let mut completed = false;
226 'flows: for flow in &mut uiaainfo.flows {
227 for stage in &flow.stages {
228 if !uiaainfo.completed.contains(stage) {
229 continue 'flows;
230 }
231 }
232 completed = true;
234 }
235
236 let session = uiaainfo
237 .session
238 .as_ref()
239 .expect("session is always set");
240
241 if matches!(email_identity_mode, EmailIdentityMode::Claim)
242 && !matches!(auth, AuthData::EmailIdentity(_))
243 && uiaainfo
244 .completed
245 .contains(&AuthType::EmailIdentity)
246 {
247 let claim = (user_id.to_owned(), device_id.to_owned(), session.as_str().into());
248
249 if !self
250 .services
251 .threepid
252 .refresh_claim(&claim)
253 .await?
254 {
255 uiaainfo
256 .completed
257 .retain(|stage| stage != &AuthType::EmailIdentity);
258
259 uiaainfo.auth_error = Some(Box::new(StandardErrorBody {
260 kind: ErrorKind::forbidden(),
261 message: "Email address has not been validated.".to_owned(),
262 }));
263
264 self.update_uiaa_session(user_id, device_id, session, Some(&uiaainfo));
265
266 return Ok((false, uiaainfo));
267 }
268 }
269
270 if !completed {
271 self.update_uiaa_session(user_id, device_id, session, Some(&uiaainfo));
272
273 return Ok((false, uiaainfo));
274 }
275
276 let retain_session = matches!(email_identity_mode, EmailIdentityMode::Claim)
278 && uiaainfo
279 .completed
280 .contains(&AuthType::EmailIdentity);
281
282 self.update_uiaa_session(user_id, device_id, session, retain_session.then_some(&uiaainfo));
283
284 Ok((true, uiaainfo))
285}
286
287#[implement(Service)]
288async fn authenticate_email_identity(
289 &self,
290 user_id: &UserId,
291 device_id: &DeviceId,
292 uiaainfo: &UiaaInfo,
293 creds: &ThirdpartyIdCredentials,
294 mode: EmailIdentityMode,
295) -> Result<bool> {
296 match mode {
297 | EmailIdentityMode::Validate => Ok(self
298 .services
299 .threepid
300 .session_validated(creds.sid.as_str(), creds.client_secret.as_str())
301 .await),
302 | EmailIdentityMode::Claim => {
303 let session = uiaainfo
304 .session
305 .as_ref()
306 .expect("session is always set");
307
308 let claim = (user_id.to_owned(), device_id.to_owned(), session.as_str().into());
309
310 self.services
311 .threepid
312 .claim_validated(creds.sid.as_str(), creds.client_secret.as_str(), claim)
313 .await
314 },
315 }
316}
317
318#[implement(Service)]
319async fn verify_password(
320 &self,
321 user_id: &UserId,
322 uiaainfo: &mut UiaaInfo,
323 password: &Password,
324) -> Result<ControlFlow<bool>> {
325 let Password { identifier, password, user, .. } = password;
326
327 let username = extract!(identifier, x in Some(UserIdentifier::Matrix(ruma::api::client::uiaa::MatrixUserIdentifier { user: x, .. })))
328 .or_else(|| cfg!(feature = "element_hacks").and(user.as_ref()))
329 .ok_or(err!(Request(Unrecognized("Identifier type not recognized."))))?;
330
331 let user_id_from_username =
332 UserId::parse_with_server_name(username.clone(), self.services.globals.server_name())
333 .map_err(|_| err!(Request(InvalidParam("User ID is invalid."))))?;
334
335 if user_id.localpart() != user_id_from_username.localpart() {
337 return Err!(Request(Forbidden("User ID and access token mismatch.")));
338 }
339
340 let user_id = user_id_from_username;
341
342 let reservation = self
344 .services
345 .login_ratelimit
346 .reserve_login_attempt(&user_id)?;
347
348 let verified = self
351 .services
352 .users
353 .password_hash(&user_id)
354 .await
355 .ok()
356 .filter(|hash| is_password_hash(hash))
357 .map(|hash| verify_password(password, &hash).is_ok());
358
359 #[cfg(feature = "ldap")]
362 let verified = if verified != Some(true)
363 && self.services.server.config.ldap.enable
364 && self
365 .services
366 .users
367 .origin(&user_id)
368 .await
369 .is_ok_and(|origin| origin == "ldap")
370 && let Ok(dns) = self.services.users.search_ldap(&user_id).await
371 && let Some((user_dn, _is_admin)) = dns.first()
372 {
373 let bound = self
374 .services
375 .users
376 .auth_ldap(user_dn, password)
377 .await
378 .is_ok();
379
380 Some(bound)
381 } else {
382 verified
383 };
384
385 if verified != Some(false) {
388 self.services
389 .login_ratelimit
390 .refund_login_attempt(reservation)?;
391 }
392
393 if verified != Some(true) {
394 uiaainfo.auth_error = Some(Box::new(StandardErrorBody {
395 kind: ErrorKind::forbidden(),
396 message: "Invalid username or password.".to_owned(),
397 }));
398
399 return Ok(ControlFlow::Break(false));
400 }
401
402 uiaainfo.completed.push(AuthType::Password);
403
404 Ok(ControlFlow::Continue(()))
405}
406
407#[implement(Service)]
408fn set_uiaa_request(
409 &self,
410 user_id: &UserId,
411 device_id: &DeviceId,
412 session: &str,
413 request: &CanonicalJsonValue,
414) {
415 let key = (user_id.to_owned(), device_id.to_owned(), session.to_owned());
416
417 self.userdevicesessionid_uiaarequest
418 .write()
419 .expect("locked for writing")
420 .insert(key, request.to_owned());
421}
422
423#[implement(Service)]
424pub fn get_uiaa_request(
425 &self,
426 user_id: &UserId,
427 device_id: Option<&DeviceId>,
428 session: &str,
429) -> Option<CanonicalJsonValue> {
430 let device_id = device_id.unwrap_or_else(|| EMPTY.into());
431 let key = (user_id.to_owned(), device_id.to_owned(), session.to_owned());
432
433 self.userdevicesessionid_uiaarequest
434 .read()
435 .expect("locked for reading")
436 .get(&key)
437 .cloned()
438}
439
440#[implement(Service)]
441pub fn update_uiaa_session(
442 &self,
443 user_id: &UserId,
444 device_id: &DeviceId,
445 session: &str,
446 uiaainfo: Option<&UiaaInfo>,
447) {
448 let key = (user_id, device_id, session);
449
450 if let Some(uiaainfo) = uiaainfo {
451 self.db
452 .userdevicesessionid_uiaainfo
453 .put(key, Json(uiaainfo));
454 } else {
455 self.db.userdevicesessionid_uiaainfo.del(key);
456 }
457}
458
459#[implement(Service)]
460async fn get_uiaa_session(
461 &self,
462 user_id: &UserId,
463 device_id: &DeviceId,
464 session: &str,
465) -> Result<UiaaInfo> {
466 let key = (user_id, device_id, session);
467
468 self.db
469 .userdevicesessionid_uiaainfo
470 .qry(&key)
471 .await
472 .deserialized()
473 .map_err(|_| err!(Request(Forbidden("UIAA session does not exist."))))
474}
475
476#[implement(Service)]
477pub async fn get_uiaa_session_by_session_id(
478 &self,
479 session_id: &str,
480) -> Option<(OwnedUserId, OwnedDeviceId, UiaaInfo)> {
481 let stream = self
483 .db
484 .userdevicesessionid_uiaainfo
485 .keys::<(OwnedUserId, OwnedDeviceId, String)>();
486
487 pin_mut!(stream);
488 while let Ok(Some((user_id, device_id, session))) = stream.try_next().await {
489 if session == session_id {
490 if let Ok(uiaainfo) = self
492 .get_uiaa_session(&user_id, &device_id, session_id)
493 .await
494 {
495 return Some((user_id, device_id, uiaainfo));
496 }
497 }
498 }
499
500 None
501}