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