Skip to main content

tuwunel_api/client/session/
sso.rs

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/// Grant phase query string.
56#[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
81/// Path attribute for the grant-session cookie. The set and removal cookies
82/// must use the same value: a cookie is only replaced (and thus removed) when
83/// its name, domain and path all match (RFC 6265 5.3), and a Set-Cookie
84/// without an explicit path defaults to the request-URI's directory.
85fn 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/// # `GET /_matrix/client/v3/login/sso/redirect`
148///
149/// A web-based Matrix client should instruct the user’s browser to navigate to
150/// this endpoint in order to log in via SSO.
151#[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/// # `GET /_matrix/client/v3/login/sso/redirect/{idpId}`
186///
187/// This endpoint is the same as /login/sso/redirect, though with an IdP ID from
188/// the original identity_providers array to inform the server of which IdP the
189/// client/user would like to continue with.
190#[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				// Base wins on key collision so extras cannot disable CSRF/PKCE.
265				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				// Present in LDAP is an existing user to provision, not a new registration.
457				| 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	// Find the UIAA session by its ID. SECURITY: Ensure the user authenticating via
623	// SSO is the owner of the UIAA session
624	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	// MSC4312 m.oauth flow → mark OAuth.
632	let has_oauth_flow = uiaainfo
633		.flows
634		.iter()
635		.any(|f| f.stages.contains(&AuthType::OAuth));
636
637	// Mark the completed step based on the UIAA session's flow.
638	if has_oauth_flow && !uiaainfo.completed.contains(&AuthType::OAuth) {
639		// Grant 10-minute bypass for cross-signing key replacement (like Synapse).
640		services
641			.users
642			.allow_cross_signing_replacement(&user_id);
643
644		uiaainfo.completed.push(AuthType::OAuth);
645	}
646
647	// Legacy m.login.sso flow → mark Sso.
648	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	// Redirect back to the fallback page to render the success HTML
662	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	// log in conduit admin channel if a non-guest user registered
724	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(&notice).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		// Keep the candidate on other errors because account creation retries the gate.
920		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}