1mod uiaa;
2
3use std::{borrow::Cow, collections::BTreeMap, net::IpAddr, time::Duration};
4
5use axum::extract::State;
6use axum_extra::extract::cookie::{Cookie, CookieJar, SameSite};
7use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD as b64};
8use futures::{FutureExt, TryFutureExt, future::try_join};
9use reqwest::header::{CONTENT_TYPE, HeaderValue};
10use ruma::{
11 Mxc, OwnedMxcUri, OwnedUserId, ServerName, UserId,
12 api::{
13 client::{
14 session::{SsoRedirectAction, sso_callback, sso_login, sso_login_with_provider},
15 uiaa::AuthType,
16 },
17 error::ErrorKind,
18 },
19};
20use serde::{Deserialize, Serialize};
21use serde_json::Value as JsonValue;
22use tuwunel_core::{
23 Err, Result, at,
24 config::IdentityProvider,
25 debug::INFO_SPAN_LEVEL,
26 debug_info, debug_warn, err, info, is_not_equal_to,
27 itertools::Itertools,
28 utils,
29 utils::{
30 OptionExt,
31 content_disposition::make_content_disposition,
32 hash::sha256,
33 result::{FlatOk, LogErr},
34 string::{EMPTY, truncate_deterministic},
35 timepoint_from_now, timepoint_has_passed,
36 },
37 warn,
38};
39use tuwunel_service::{
40 Services,
41 client::read_response_capped,
42 media::MXC_LENGTH,
43 oauth::{
44 CODE_VERIFIER_LENGTH, Provider, SESSION_ID_LENGTH, Session, TokenResponse, UserInfo,
45 unique_id_sub,
46 },
47 users::{PASSWORD_SENTINEL, Register},
48};
49use url::Url;
50
51pub(crate) use self::uiaa::{sso_complete_js_route, sso_css_route, sso_fallback_route};
52use super::TOKEN_LENGTH;
53use crate::{ClientIp, Ruma};
54
55#[derive(Debug, Serialize)]
57struct GrantQuery<'a> {
58 client_id: &'a str,
59 state: &'a str,
60 nonce: &'a str,
61 scope: &'a str,
62 response_type: &'a str,
63 access_type: &'a str,
64 code_challenge_method: &'a str,
65 code_challenge: &'a str,
66 redirect_uri: Option<&'a str>,
67 #[serde(skip_serializing_if = "Option::is_none")]
68 prompt: Option<&'a str>,
69}
70
71#[derive(Debug, Deserialize, Serialize)]
72struct GrantCookie<'a> {
73 client_id: Cow<'a, str>,
74 state: Cow<'a, str>,
75 nonce: Cow<'a, str>,
76 redirect_uri: Cow<'a, str>,
77}
78
79static GRANT_SESSION_COOKIE: &str = "tuwunel_grant_session";
80
81fn grant_session_cookie_path(callback_url: Option<&Url>) -> &str {
86 callback_url.map(Url::path).unwrap_or("/")
87}
88
89fn decode_apple_userinfo_from_id_token(session: &Session) -> Result<UserInfo> {
90 let id_token = session.id_token.as_deref().ok_or_else(|| {
91 err!(Request(Unauthorized("Missing Apple id_token in token response.")))
92 })?;
93
94 let payload_b64 = id_token
95 .split('.')
96 .nth(1)
97 .ok_or_else(|| err!(Request(Unauthorized("Apple id_token is malformed."))))?;
98
99 let payload = b64
100 .decode(payload_b64)
101 .map_err(|_| err!(Request(Unauthorized("Apple id_token payload is invalid base64."))))?;
102
103 let payload: JsonValue = serde_json::from_slice(&payload)
104 .map_err(|_| err!(Request(Unauthorized("Apple id_token payload is not valid JSON."))))?;
105
106 let sub = payload
107 .get("sub")
108 .and_then(JsonValue::as_str)
109 .ok_or_else(|| {
110 err!(Request(Unauthorized("Apple id_token missing required sub claim.")))
111 })?;
112
113 let email = payload
114 .get("email")
115 .and_then(JsonValue::as_str)
116 .map(ToOwned::to_owned);
117
118 let preferred_username = email
119 .as_deref()
120 .and_then(|value| value.split_once('@'))
121 .map(at!(0))
122 .map(ToOwned::to_owned);
123
124 Ok(UserInfo {
125 sub: sub.to_owned(),
126 preferred_username: preferred_username.clone(),
127 username: preferred_username,
128 nickname: None,
129 name: payload
130 .get("name")
131 .and_then(JsonValue::as_str)
132 .map(ToOwned::to_owned),
133 given_name: payload
134 .get("given_name")
135 .and_then(JsonValue::as_str)
136 .map(ToOwned::to_owned),
137 family_name: payload
138 .get("family_name")
139 .and_then(JsonValue::as_str)
140 .map(ToOwned::to_owned),
141 email,
142 avatar_url: None,
143 picture: None,
144 })
145}
146
147#[tracing::instrument(
152 name = "sso_login",
153 level = "debug",
154 skip_all,
155 fields(%client),
156)]
157pub(crate) async fn sso_login_route(
158 State(services): State<crate::State>,
159 ClientIp(client): ClientIp,
160 body: Ruma<sso_login::v3::Request>,
161) -> Result<sso_login::v3::Response> {
162 if services.config.sso_custom_providers_page {
163 return Err!(Request(NotImplemented(
164 "sso_custom_providers_page has been enabled but this URL has not been overridden \
165 with any custom page listing the available providers..."
166 )));
167 }
168
169 let redirect_url = body.body.redirect_url;
170 let action = body.body.action;
171 let default_idp_id = services
172 .oauth
173 .providers
174 .get_default_id()
175 .unwrap_or_default();
176
177 handle_sso_login(&services, &client, default_idp_id, redirect_url, None, action)
178 .map_ok(|response| sso_login::v3::Response {
179 location: response.location,
180 cookie: response.cookie,
181 })
182 .await
183}
184
185#[tracing::instrument(
191 name = "sso_login_with_provider",
192 level = "info",
193 skip_all,
194 ret(level = "debug")
195 fields(
196 %client,
197 idp_id = body.body.idp_id,
198 ),
199)]
200pub(crate) async fn sso_login_with_provider_route(
201 State(services): State<crate::State>,
202 ClientIp(client): ClientIp,
203 body: Ruma<sso_login_with_provider::v3::Request>,
204) -> Result<sso_login_with_provider::v3::Response> {
205 let idp_id = body.body.idp_id;
206 let redirect_url = body.body.redirect_url;
207 let login_token = body.body.login_token;
208 let action = body.body.action;
209
210 handle_sso_login(&services, &client, idp_id, redirect_url, login_token, action).await
211}
212
213async fn handle_sso_login(
214 services: &Services,
215 _client: &IpAddr,
216 idp_id: String,
217 redirect_url: String,
218 login_token: Option<String>,
219 action: Option<SsoRedirectAction>,
220) -> Result<sso_login_with_provider::v3::Response> {
221 let redirect_url: Url = redirect_url.parse().map_err(|e| {
222 err!(Request(InvalidParam(debug_warn!(
223 ?e,
224 ?redirect_url,
225 "Failed to parse redirect_url.",
226 ))))
227 })?;
228
229 let provider = services.oauth.providers.get(&idp_id).await?;
230 let sess_id = utils::random_string(SESSION_ID_LENGTH);
231 let query_nonce = utils::random_string(CODE_VERIFIER_LENGTH);
232 let cookie_nonce = utils::random_string(CODE_VERIFIER_LENGTH);
233 let code_verifier = utils::random_string(CODE_VERIFIER_LENGTH);
234 let code_challenge = b64.encode(sha256::hash(code_verifier.as_bytes()));
235 let callback_uri = provider.callback_url.as_ref().map(Url::as_str);
236 let scope = provider.scope.iter().join(" ");
237 let prompt = action
238 .filter(|_| provider.forward_action_prompt)
239 .and_then(|action| matches!(action, SsoRedirectAction::Register).then_some("create"));
240
241 let query = GrantQuery {
242 client_id: &provider.client_id,
243 state: &sess_id,
244 nonce: &query_nonce,
245 access_type: "online",
246 response_type: "code",
247 code_challenge_method: "S256",
248 code_challenge: &code_challenge,
249 redirect_uri: callback_uri,
250 prompt,
251 scope: scope
252 .is_empty()
253 .then_some("openid email profile")
254 .unwrap_or(scope.as_str()),
255 };
256
257 let location = provider
258 .authorization_url
259 .clone()
260 .map(|mut location| {
261 let query = serde_html_form::to_string(&query).ok();
262 location.set_query(query.as_deref());
263 if !provider.extra_authorization_parameters.is_empty() {
264 let merged: BTreeMap<String, String> = provider
266 .extra_authorization_parameters
267 .clone()
268 .into_iter()
269 .chain(
270 location
271 .query_pairs()
272 .map(|(k, v)| (k.into_owned(), v.into_owned())),
273 )
274 .collect();
275
276 location.set_query(None);
277 location.query_pairs_mut().extend_pairs(&merged);
278 }
279 location
280 })
281 .ok_or_else(|| {
282 err!(Config("authorization_url", "Missing required IdentityProvider config"))
283 })?;
284
285 let cookie_val = GrantCookie {
286 client_id: query.client_id.into(),
287 state: query.state.into(),
288 nonce: cookie_nonce.as_str().into(),
289 redirect_uri: redirect_url.as_str().into(),
290 };
291
292 let cookie_path = grant_session_cookie_path(provider.callback_url.as_ref());
293
294 let cookie_max_age = provider
295 .grant_session_duration
296 .map(Duration::from_secs)
297 .expect("Defaulted to Some value during configure_idp()")
298 .try_into()
299 .expect("std::time::Duration to time::Duration conversion failure");
300
301 let cookie = Cookie::build((GRANT_SESSION_COOKIE, serde_html_form::to_string(&cookie_val)?))
302 .path(cookie_path)
303 .max_age(cookie_max_age)
304 .same_site(SameSite::None)
305 .secure(true)
306 .http_only(true)
307 .build()
308 .to_string()
309 .into();
310
311 let session = Session {
312 idp_id: Some(idp_id),
313 sess_id: Some(sess_id.clone()),
314 redirect_url: Some(redirect_url),
315 code_verifier: Some(code_verifier),
316 query_nonce: Some(query_nonce),
317 cookie_nonce: Some(cookie_nonce),
318 authorize_expires_at: provider
319 .grant_session_duration
320 .map(Duration::from_secs)
321 .map(timepoint_from_now)
322 .transpose()?,
323
324 user_id: login_token
325 .as_deref()
326 .map_async(|token| services.users.find_from_login_token(token))
327 .map(FlatOk::flat_ok)
328 .await,
329
330 ..Default::default()
331 };
332
333 services.oauth.sessions.put(&session).await;
334
335 Ok(sso_login_with_provider::v3::Response {
336 location: location.into(),
337 cookie: Some(cookie),
338 })
339}
340
341#[tracing::instrument(
342 name = "sso_callback"
343 level = "debug",
344 skip_all,
345 fields(
346 %client,
347 cookie = ?body.cookie,
348 body = ?body.body,
349 ),
350)]
351pub(crate) async fn sso_callback_route(
352 State(services): State<crate::State>,
353 ClientIp(client): ClientIp,
354 body: Ruma<sso_callback::unstable::Request>,
355) -> Result<sso_callback::unstable::Response> {
356 let sess_id = body
357 .body
358 .state
359 .as_deref()
360 .ok_or_else(|| err!(Request(Forbidden("Missing sess_id in callback."))))?;
361
362 let code = body
363 .body
364 .code
365 .as_deref()
366 .ok_or_else(|| err!(Request(Forbidden("Missing code in callback."))))?;
367
368 let session = services
369 .oauth
370 .sessions
371 .get(sess_id)
372 .map_err(|_| err!(Request(Forbidden("Invalid state in callback"))));
373
374 let provider = services
375 .oauth
376 .providers
377 .get(body.body.idp_id.as_str());
378
379 let (provider, session) = try_join(provider, session).await.log_err()?;
380 let idp_id = provider.id();
381
382 if session.sess_id.as_deref() != Some(sess_id) {
383 return Err!(Request(Unauthorized("Session ID {sess_id:?} not recognized.")));
384 }
385
386 if session.idp_id.as_deref() != Some(idp_id) {
387 return Err!(Request(Unauthorized(
388 "Identity Provider {idp_id:?} session not recognized."
389 )));
390 }
391
392 if session
393 .authorize_expires_at
394 .is_some_and(timepoint_has_passed)
395 {
396 return Err!(Request(Unauthorized("Authorization grant session has expired.")));
397 }
398
399 if provider.check_cookie {
400 validate_session_cookie(&body.cookie, &provider, &session, sess_id)?;
401 }
402
403 let token_response = services
404 .oauth
405 .request_token((&provider, &session), code)
406 .await?;
407
408 let session = apply_token_response(session, token_response)?;
409
410 let userinfo = services
411 .oauth
412 .request_userinfo((&provider, &session))
413 .await
414 .or_else(|error| {
415 if provider.brand != "appleoidc" {
416 return Err(error);
417 }
418
419 debug_warn!(
420 ?error,
421 idp_id = provider.id(),
422 "Failed to fetch Apple userinfo endpoint; falling back to id_token claims.",
423 );
424
425 decode_apple_userinfo_from_id_token(&session).map_err(|decode_error| {
426 debug_warn!(
427 ?decode_error,
428 idp_id = provider.id(),
429 "Failed to decode Apple id_token fallback.",
430 );
431 error
432 })
433 })?;
434
435 let unique_id = unique_id_sub((&provider, &userinfo.sub))?;
436
437 let complete_identity = async |old_user_id: Option<OwnedUserId>| {
438 let session = Session {
439 user_info: Some(userinfo.clone()),
440 ..session
441 };
442
443 let user_id = match (session.user_id, old_user_id) {
444 | (Some(user_id), ..) | (None, Some(user_id)) => user_id,
445 | (None, None) => decide_user_id(&services, &provider, &userinfo, &unique_id).await?,
446 };
447
448 let session = Session {
449 user_id: Some(user_id.clone()),
450 ..session
451 };
452
453 if !services.users.exists(&user_id).await {
454 let origin = match provider.registration {
455 | true => "sso",
456 | false if ldap_user_exists(&services, &user_id).await => "ldap",
458 | false =>
459 return Err!(Request(Forbidden(
460 "Registration from this provider is disabled"
461 ))),
462 };
463
464 register_user(&services, &provider, &session, &userinfo, &user_id, origin).await?;
465 }
466
467 Ok((session, user_id))
468 };
469
470 let (session, user_id, old_sess_id) = services
471 .oauth
472 .sessions
473 .commit_identity_session(&unique_id, complete_identity)
474 .await?;
475
476 if let Some(old_sess_id) = old_sess_id
477 .as_deref()
478 .filter(is_not_equal_to!(&sess_id))
479 {
480 services.oauth.sessions.delete(old_sess_id).await;
481 }
482
483 if !services.users.is_active_local(&user_id).await {
484 return Err!(Request(UserDeactivated("This user has been deactivated.")));
485 }
486
487 let cookie = Cookie::build((GRANT_SESSION_COOKIE, EMPTY))
488 .path(grant_session_cookie_path(provider.callback_url.as_ref()))
489 .removal()
490 .build()
491 .to_string()
492 .into();
493
494 if let Some(redirect_url) = session
495 .redirect_url
496 .as_ref()
497 .filter(|url| url.scheme() == "uiaa")
498 {
499 return handle_uiaa(&services, &user_id, cookie, redirect_url).await;
500 }
501
502 let next_idp_url = chain_next_idp_url(&services, &provider, &session, idp_id);
503
504 let location = finalize_login_redirect(&services, &session, next_idp_url, &user_id)?;
505
506 Ok(sso_callback::unstable::Response { location, cookie: Some(cookie) })
507}
508
509fn validate_session_cookie(
510 cookies: &CookieJar,
511 provider: &Provider,
512 session: &Session,
513 sess_id: &str,
514) -> Result {
515 let client_id = &provider.client_id;
516 let cookie = cookies
517 .get(GRANT_SESSION_COOKIE)
518 .map(Cookie::value)
519 .map(serde_html_form::from_str::<GrantCookie<'_>>)
520 .transpose()?
521 .ok_or_else(|| err!(Request(Unauthorized("Missing cookie {GRANT_SESSION_COOKIE:?}"))))?;
522
523 if cookie.client_id.as_ref() != client_id.as_str() {
524 return Err!(Request(Unauthorized("Client ID {client_id:?} cookie mismatch.")));
525 }
526
527 if Some(cookie.nonce.as_ref()) != session.cookie_nonce.as_deref() {
528 return Err!(Request(Unauthorized("Cookie nonce does not match session state.")));
529 }
530
531 if cookie.state.as_ref() != sess_id {
532 return Err!(Request(Unauthorized("Session ID {sess_id:?} cookie mismatch.")));
533 }
534
535 Ok(())
536}
537
538fn apply_token_response(session: Session, token: TokenResponse) -> Result<Session> {
539 let expires_at = token
540 .expires_in
541 .map(Duration::from_secs)
542 .map(timepoint_from_now)
543 .transpose()?;
544
545 let refresh_token_expires_at = token
546 .refresh_token_expires_in
547 .map(Duration::from_secs)
548 .map(timepoint_from_now)
549 .transpose()?;
550
551 Ok(Session {
552 scope: token.scope,
553 token_type: token.token_type,
554 access_token: token.access_token,
555 id_token: token.id_token,
556 expires_at,
557 refresh_token: token.refresh_token,
558 refresh_token_expires_at,
559 ..session
560 })
561}
562
563fn chain_next_idp_url(
564 services: &Services,
565 provider: &Provider,
566 session: &Session,
567 idp_id: &str,
568) -> Option<Url> {
569 services
570 .config
571 .identity_provider
572 .values()
573 .filter(|idp| idp.default || services.config.single_sso)
574 .skip_while(|idp| idp.id() != idp_id)
575 .nth(1)
576 .map(IdentityProvider::id)
577 .and_then(|next_idp| {
578 provider.callback_url.clone().map(|mut url| {
579 let path = format!("/_matrix/client/v3/login/sso/redirect/{next_idp}");
580 url.set_path(&path);
581
582 if let Some(redirect_url) = session.redirect_url.as_ref() {
583 url.query_pairs_mut()
584 .append_pair("redirectUrl", redirect_url.as_str());
585 }
586
587 url
588 })
589 })
590}
591
592fn finalize_login_redirect(
593 services: &Services,
594 session: &Session,
595 next_idp_url: Option<Url>,
596 user_id: &UserId,
597) -> Result<String> {
598 let login_token = utils::random_string(TOKEN_LENGTH);
599 let _login_token_expires_in = services
600 .users
601 .create_login_token(user_id, &login_token);
602
603 let location = next_idp_url
604 .or_else(|| session.redirect_url.clone())
605 .ok_or_else(|| err!(Request(InvalidParam("Missing redirect URL in session data"))))?
606 .query_pairs_mut()
607 .append_pair("loginToken", &login_token)
608 .finish()
609 .to_string();
610
611 Ok(location)
612}
613
614async fn handle_uiaa(
615 services: &Services,
616 user_id: &UserId,
617 cookie: Cow<'static, str>,
618 redirect_url: &Url,
619) -> Result<sso_callback::unstable::Response> {
620 let uiaa_session_id = redirect_url.path();
621
622 let (user_id, device_id, mut uiaainfo) = services
625 .uiaa
626 .get_uiaa_session_by_session_id(uiaa_session_id)
627 .await
628 .filter(|(db_user_id, ..)| user_id.eq(db_user_id))
629 .ok_or_else(|| err!(Request(Forbidden("UIAA session not found."))))?;
630
631 let has_oauth_flow = uiaainfo
633 .flows
634 .iter()
635 .any(|f| f.stages.contains(&AuthType::OAuth));
636
637 if has_oauth_flow && !uiaainfo.completed.contains(&AuthType::OAuth) {
639 services
641 .users
642 .allow_cross_signing_replacement(&user_id);
643
644 uiaainfo.completed.push(AuthType::OAuth);
645 }
646
647 let has_sso_flow = uiaainfo
649 .flows
650 .iter()
651 .any(|f| f.stages.contains(&AuthType::Sso));
652
653 if has_sso_flow && !uiaainfo.completed.contains(&AuthType::Sso) {
654 uiaainfo.completed.push(AuthType::Sso);
655 }
656
657 services
658 .uiaa
659 .update_uiaa_session(&user_id, &device_id, uiaa_session_id, Some(&uiaainfo));
660
661 let location =
663 format!("/_matrix/client/v3/auth/m.login.sso/fallback/web?session={uiaa_session_id}");
664
665 Ok(sso_callback::unstable::Response { location, cookie: Some(cookie) })
666}
667
668async fn ldap_user_exists(services: &Services, user_id: &UserId) -> bool {
669 cfg!(feature = "ldap")
670 && services.config.ldap.enable
671 && services
672 .users
673 .search_ldap(user_id)
674 .await
675 .log_err()
676 .is_ok_and(|dns| !dns.is_empty())
677}
678
679#[tracing::instrument(
680 name = "register",
681 level = INFO_SPAN_LEVEL,
682 skip_all,
683 fields(user_id, userinfo)
684)]
685async fn register_user(
686 services: &Services,
687 provider: &Provider,
688 session: &Session,
689 userinfo: &UserInfo,
690 user_id: &UserId,
691 origin: &str,
692) -> Result {
693 debug_info!(%user_id, "Creating new user account...");
694
695 services
696 .users
697 .full_register(Register {
698 user_id: Some(user_id),
699 password: Some(PASSWORD_SENTINEL),
700 origin: Some(origin),
701 displayname: userinfo.name.as_deref(),
702 grant_first_user_admin: true,
703 ..Default::default()
704 })
705 .await?;
706
707 if let Some(avatar_url) = userinfo
708 .avatar_url
709 .as_deref()
710 .or(userinfo.picture.as_deref())
711 {
712 set_avatar(services, provider, session, userinfo, user_id, avatar_url)
713 .await
714 .ok();
715 }
716
717 let idp_id = provider.id();
718 let idp_name = provider
719 .name
720 .as_deref()
721 .unwrap_or(provider.brand.as_str());
722
723 let notice =
725 format!("New user \"{user_id}\" registered on this server via {idp_name} ({idp_id})");
726
727 info!("{notice}");
728 services.admin.notify(¬ice).await;
729
730 Ok(())
731}
732
733#[tracing::instrument(level = "debug", skip_all, fields(user_id, avatar_url))]
734async fn set_avatar(
735 services: &Services,
736 _provider: &Provider,
737 _session: &Session,
738 _userinfo: &UserInfo,
739 user_id: &UserId,
740 avatar_url: &str,
741) -> Result {
742 use reqwest::Response;
743
744 let response = services
745 .client
746 .default
747 .get(avatar_url)
748 .send()
749 .await
750 .and_then(Response::error_for_status)?;
751
752 let content_type = response
753 .headers()
754 .get(CONTENT_TYPE)
755 .map(HeaderValue::to_str)
756 .flat_ok()
757 .map(ToOwned::to_owned);
758
759 let mxc = Mxc {
760 server_name: services.globals.server_name(),
761 media_id: &utils::random_string(MXC_LENGTH),
762 };
763
764 let content_disposition = make_content_disposition(None, content_type.as_deref(), None);
765 let limit = services.server.config.max_response_size;
766 let bytes = read_response_capped(response, limit).await?;
767 services
768 .media
769 .create(&mxc, Some(user_id), Some(&content_disposition), content_type.as_deref(), &bytes)
770 .await?;
771
772 let mxc_uri: OwnedMxcUri = mxc.to_string().into();
773 services
774 .profile
775 .set_avatar_url(user_id, Some(&mxc_uri), None)
776 .await?;
777
778 Ok(())
779}
780
781#[tracing::instrument(
782 level = "debug",
783 ret(level = "debug")
784 skip_all,
785 fields(user),
786)]
787async fn decide_user_id(
788 services: &Services,
789 provider: &Provider,
790 userinfo: &UserInfo,
791 unique_id: &str,
792) -> Result<OwnedUserId> {
793 if let Some(user_id) = services
794 .oauth
795 .sessions
796 .find_user_association_pending(provider.id(), userinfo)
797 {
798 debug_info!(
799 provider = ?provider.id(),
800 ?user_id,
801 ?userinfo,
802 "Matched pending association"
803 );
804
805 return Ok(user_id);
806 }
807
808 let explicit = |claim: &str| provider.userid_claims.contains(claim);
809
810 let allowed = |claim: &str| provider.userid_claims.is_empty() || explicit(claim);
811
812 let choices = [
813 explicit("sub")
814 .then_some(userinfo.sub.as_str())
815 .map(str::to_lowercase),
816 userinfo
817 .preferred_username
818 .as_deref()
819 .map(str::to_lowercase)
820 .filter(|_| allowed("preferred_username")),
821 userinfo
822 .username
823 .as_deref()
824 .map(str::to_lowercase)
825 .filter(|_| allowed("username")),
826 userinfo
827 .nickname
828 .as_deref()
829 .map(str::to_lowercase)
830 .filter(|_| allowed("nickname")),
831 provider
832 .brand
833 .eq(&"github")
834 .then_some(userinfo.sub.as_str())
835 .map(str::to_lowercase)
836 .filter(|_| allowed("login")),
837 userinfo
838 .email
839 .as_deref()
840 .and_then(|email| email.split_once('@'))
841 .map(at!(0))
842 .map(str::to_lowercase)
843 .filter(|_| allowed("email")),
844 ];
845
846 for choice in choices.into_iter().flatten() {
847 if let Some(user_id) = try_user_id(services, provider, &choice, false).await {
848 return Ok(user_id);
849 }
850 }
851
852 let length = Some(15..23);
853 let unique_id = truncate_deterministic(unique_id, length).to_lowercase();
854 if let Some(user_id) = try_user_id(services, provider, &unique_id, true).await {
855 return Ok(user_id);
856 }
857
858 Err!(Request(UserInUse("User ID is not available.")))
859}
860
861#[tracing::instrument(level = "debug", skip_all, fields(username))]
862async fn try_user_id(
863 services: &Services,
864 provider: &Provider,
865 username: &str,
866 unique_id: bool,
867) -> Option<OwnedUserId> {
868 let server_name = services.globals.server_name();
869 let user_id = parse_user_id(server_name, username)
870 .inspect_err(|e| warn!(?username, "Username invalid: {e}"))
871 .ok()?;
872
873 if services
874 .config
875 .forbidden_usernames
876 .is_match(username)
877 {
878 warn!(?username, "Username forbidden.");
879 return None;
880 }
881
882 if services.users.exists(&user_id).await {
883 if provider.trusted {
884 info!(
885 ?username,
886 provider = ?provider.brand,
887 "Authorizing trusted provider access to existing account."
888 );
889
890 return Some(user_id);
891 }
892
893 if services
894 .users
895 .origin(&user_id)
896 .await
897 .ok()
898 .is_none_or(|origin| origin != "sso")
899 {
900 debug_warn!(?username, "Existing username has non-sso origin.");
901 return None;
902 }
903
904 if !unique_id {
905 debug_warn!(?username, "Username exists.");
906 return None;
907 }
908 } else {
909 if unique_id && !provider.unique_id_fallbacks {
910 debug_warn!(
911 ?username,
912 provider = ?provider.brand,
913 "Unique ID fallbacks disabled.",
914 );
915
916 return None;
917 }
918
919 if services
921 .users
922 .check_creation(&user_id)
923 .await
924 .is_err_and(|error| error.kind() == ErrorKind::UserInUse)
925 {
926 return None;
927 }
928 }
929
930 Some(user_id)
931}
932
933fn parse_user_id(server_name: &ServerName, username: &str) -> Result<OwnedUserId> {
934 match UserId::parse_with_server_name(username, server_name) {
935 | Err(e) => {
936 Err!(Request(InvalidUsername(debug_error!("Username {username} is not valid: {e}"))))
937 },
938 | Ok(user_id) => match user_id.validate_strict() {
939 | Ok(()) => Ok(user_id),
940 | Err(e) => Err!(Request(InvalidUsername(debug_error!(
941 "Username {username} contains disallowed characters or spaces: {e}"
942 )))),
943 },
944 }
945}
946
947#[cfg(test)]
948mod tests {
949 use serde_json::json;
950
951 use super::*;
952
953 fn apple_session_with_claims(claims: &serde_json::Value) -> Session {
954 let payload = b64.encode(serde_json::to_vec(claims).expect("serialize claims"));
955
956 Session {
957 id_token: Some(format!("header.{payload}.signature")),
958 ..Default::default()
959 }
960 }
961
962 #[test]
963 fn decode_apple_userinfo_from_id_token_extracts_expected_claims() {
964 let session = apple_session_with_claims(&json!({
965 "sub": "apple-user-123",
966 "email": "alice@example.com",
967 "name": "Alice Example",
968 "given_name": "Alice",
969 "family_name": "Example"
970 }));
971
972 let userinfo =
973 decode_apple_userinfo_from_id_token(&session).expect("decode Apple id_token claims");
974
975 assert_eq!(userinfo.sub, "apple-user-123");
976 assert_eq!(userinfo.email.as_deref(), Some("alice@example.com"));
977 assert_eq!(userinfo.preferred_username.as_deref(), Some("alice"));
978 assert_eq!(userinfo.username.as_deref(), Some("alice"));
979 assert_eq!(userinfo.name.as_deref(), Some("Alice Example"));
980 assert_eq!(userinfo.given_name.as_deref(), Some("Alice"));
981 assert_eq!(userinfo.family_name.as_deref(), Some("Example"));
982 }
983
984 #[test]
985 fn decode_apple_userinfo_from_id_token_requires_sub_claim() {
986 let session = apple_session_with_claims(&json!({
987 "email": "alice@example.com"
988 }));
989
990 let error = decode_apple_userinfo_from_id_token(&session)
991 .expect_err("missing sub claim should fail");
992
993 let message = format!("{error}");
994 assert!(message.contains("sub claim"), "unexpected error: {message}");
995 }
996
997 #[test]
998 fn decode_apple_userinfo_from_id_token_requires_id_token() {
999 let session = Session::default();
1000
1001 let error = decode_apple_userinfo_from_id_token(&session)
1002 .expect_err("missing id_token should fail");
1003
1004 let message = format!("{error}");
1005 assert!(message.contains("Missing Apple id_token"), "unexpected error: {message}");
1006 }
1007
1008 #[test]
1009 fn decode_apple_userinfo_from_id_token_rejects_invalid_payload() {
1010 let session = Session {
1011 id_token: Some("header.!.signature".to_owned()),
1012 ..Default::default()
1013 };
1014
1015 let error = decode_apple_userinfo_from_id_token(&session)
1016 .expect_err("invalid id_token payload should fail");
1017
1018 let message = format!("{error}");
1019 assert!(message.contains("invalid base64"), "unexpected error: {message}");
1020 }
1021
1022 #[test]
1023 fn grant_session_cookie_path_uses_the_callback_path() {
1024 let callback = Url::parse(
1025 "https://matrix.example/_matrix/client/unstable/login/sso/callback/tuwunel_abc",
1026 )
1027 .expect("valid URL");
1028
1029 assert_eq!(
1030 grant_session_cookie_path(Some(&callback)),
1031 "/_matrix/client/unstable/login/sso/callback/tuwunel_abc"
1032 );
1033 assert_eq!(grant_session_cookie_path(None), "/");
1034 }
1035
1036 #[test]
1037 fn grant_session_removal_cookie_path_matches_the_set_cookie_path() {
1038 let callback = Url::parse(
1039 "https://matrix.example/_matrix/client/unstable/login/sso/callback/tuwunel_abc",
1040 )
1041 .expect("valid URL");
1042
1043 let set = Cookie::build((GRANT_SESSION_COOKIE, "grant"))
1044 .path(grant_session_cookie_path(Some(&callback)))
1045 .build();
1046
1047 let removal = Cookie::build((GRANT_SESSION_COOKIE, EMPTY))
1048 .path(grant_session_cookie_path(Some(&callback)))
1049 .removal()
1050 .build();
1051
1052 assert_eq!(set.path(), removal.path());
1053 assert!(
1054 removal
1055 .to_string()
1056 .contains("Path=/_matrix/client/unstable/login/sso/callback/tuwunel_abc"),
1057 "unexpected removal cookie: {removal}"
1058 );
1059 }
1060}