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