1use std::{fmt::Write, net::IpAddr};
2
3use axum::{
4 extract::{Form, Request, State},
5 response::{Redirect, Response},
6};
7use const_str::format as const_format;
8use http::StatusCode;
9use itertools::Either::{Left, Right};
10use ruma::{OwnedUserId, UserId};
11use serde::Deserialize;
12use serde_json::json;
13use tuwunel_core::{
14 Err, Error, Result,
15 config::IdentityProvider,
16 err,
17 smallstr::SmallString,
18 utils::{self, BoolExt, hash::verify_password, html::escape as html_escape},
19};
20use tuwunel_service::{Services, users::Register};
21use url::Url;
22
23use super::{
24 account::{
25 ACCOUNT_HEAD, account_error_response, account_html_response, account_redirect_response,
26 },
27 authorization_sso_url, url_encode,
28};
29use crate::ClientIp;
30
31type AccountAction = SmallString<[u8; 32]>;
32type DeviceId = SmallString<[u8; 24]>;
33type IdpId = SmallString<[u8; 32]>;
34type ProviderChoice<'a> = (&'a str, &'a str);
35
36const LOGIN_TOKEN_LENGTH: usize = 32;
37
38#[derive(Debug, Default, Deserialize)]
39struct NativeQuery {
40 oidc_req_id: Option<String>,
41 idp_id: Option<IdpId>,
42 user_code: Option<String>,
43 action: Option<AccountAction>,
44 device_id: Option<DeviceId>,
45 view: Option<String>,
46}
47
48#[derive(Debug, Deserialize)]
49pub(crate) struct NativeSubmit {
50 #[serde(default)]
51 oidc_req_id: Option<String>,
52 #[serde(default)]
53 user_code: Option<String>,
54 #[serde(default)]
55 action: Option<AccountAction>,
56 #[serde(default)]
57 device_id: Option<DeviceId>,
58 #[serde(default)]
59 mode: Option<String>,
60 username: String,
61 password: String,
62 #[serde(default)]
63 registration_token: Option<String>,
64 #[serde(default)]
65 accept_terms: Option<String>,
66}
67
68#[derive(Clone, Copy)]
69enum Flow<'a> {
70 Account {
71 action: &'a str,
72 device_id: &'a str,
73 },
74 Authorization(&'a str),
75 Device(&'a str),
76}
77
78pub(crate) async fn native_get_route(
85 State(services): State<crate::State>,
86 ClientIp(client): ClientIp,
87 request: Request,
88) -> Response {
89 if let Err(e) = require_native(&services) {
90 return account_error_response(&e);
91 }
92
93 let params: NativeQuery =
94 match serde_html_form::from_str(request.uri().query().unwrap_or_default()) {
95 | Ok(params) => params,
96 | Err(e) => return account_error_response(&e.into()),
97 };
98
99 let context = match parse_flow(
100 params.oidc_req_id.as_deref(),
101 params.user_code.as_deref(),
102 params.action.as_deref(),
103 params.device_id.as_deref(),
104 ) {
105 | Ok(context) => context,
106 | Err(e) => return account_error_response(&e),
107 };
108
109 if let Some(idp_id) = params.idp_id.as_deref() {
110 return provider_redirect(&services, client, context, idp_id)
111 .await
112 .map_or_else(|e| account_error_response(&e), account_redirect_response);
113 }
114
115 let view = params.view.as_deref().unwrap_or("login");
116
117 render_page(&services, view, context, None)
118 .await
119 .map(|html| account_html_response(StatusCode::OK, html))
120 .unwrap_or_else(|e| account_error_response(&e))
121}
122
123fn parse_flow<'a>(
124 oidc_req_id: Option<&'a str>,
125 user_code: Option<&'a str>,
126 action: Option<&'a str>,
127 device_id: Option<&'a str>,
128) -> Result<Flow<'a>> {
129 match (
130 oidc_req_id.filter(|value| !value.is_empty()),
131 user_code.filter(|value| !value.is_empty()),
132 action.filter(|value| !value.is_empty()),
133 ) {
134 | (Some(req_id), None, None) => Ok(Flow::Authorization(req_id)),
135 | (None, Some(user_code), None) => Ok(Flow::Device(user_code)),
136 | (None, None, Some(action)) => Ok(Flow::Account {
137 action,
138 device_id: device_id.unwrap_or_default(),
139 }),
140 | _ => Err!(Request(InvalidParam(
141 "Exactly one OIDC request ID, user code, or account action is required."
142 ))),
143 }
144}
145
146async fn provider_redirect(
152 services: &Services,
153 client: IpAddr,
154 context: Flow<'_>,
155 idp_id: &str,
156) -> Result<Redirect> {
157 let Flow::Authorization(req_id) = context else {
158 return Err!(Request(InvalidParam(
159 "Provider selection requires an authorization request"
160 )));
161 };
162
163 services.oauth.check_rate_limit(client)?;
164
165 let oidc = services.oauth.get_server()?;
166 let provider_id = services
167 .oauth
168 .providers
169 .find_config(idp_id)
170 .map_err(|_| err!(Request(InvalidParam("Unrecognized identity provider"))))?
171 .id();
172
173 let sso_url = authorization_sso_url(&oidc.issuer_url()?, provider_id, req_id)?;
174
175 oidc.bind_auth_request_to_provider(req_id, provider_id)
176 .await?;
177
178 Ok(Redirect::temporary(sso_url.as_str()))
179}
180
181pub(crate) async fn native_submit_route(
184 State(services): State<crate::State>,
185 ClientIp(client): ClientIp,
186 Form(body): Form<NativeSubmit>,
187) -> Response {
188 match native_submit(&services, client, &body).await {
189 | Ok(response) => response,
190 | Err(e) => render_submit_error(&services, &body, &e).await,
191 }
192}
193
194async fn native_submit(
195 services: &Services,
196 client: IpAddr,
197 body: &NativeSubmit,
198) -> Result<Response> {
199 require_native(services)?;
200 services.oauth.check_device_rate_limit(client)?;
202 services.oauth.check_rate_limit(client)?;
203
204 let context = parse_flow(
205 body.oidc_req_id.as_deref(),
206 body.user_code.as_deref(),
207 body.action.as_deref(),
208 body.device_id.as_deref(),
209 )?;
210
211 let user_id = match context {
212 | Flow::Authorization(req_id) => authenticate_local(services, req_id, body).await?,
213 | _ => verify_credentials(services, &body.username, &body.password).await?,
214 };
215
216 let token = utils::random_string(LOGIN_TOKEN_LENGTH);
217 let _expires_in = services
218 .users
219 .create_login_token(&user_id, &token);
220
221 let redirect = complete_redirect(services, context, &token)?;
222
223 Ok(account_redirect_response(redirect))
224}
225
226async fn render_submit_error(
231 services: &Services,
232 body: &NativeSubmit,
233 error: &Error,
234) -> Response {
235 let context = match parse_flow(
236 body.oidc_req_id.as_deref(),
237 body.user_code.as_deref(),
238 body.action.as_deref(),
239 body.device_id.as_deref(),
240 ) {
241 | Ok(context) => context,
242 | Err(e) => return account_error_response(&e),
243 };
244
245 let view = match (context, body.mode.as_deref()) {
246 | (Flow::Authorization(_), Some("register")) => "register",
247 | _ => "login",
248 };
249
250 let msg = error.sanitized_message();
251
252 render_page(services, view, context, Some(&msg))
253 .await
254 .map(|html| account_html_response(error.status_code(), html))
255 .unwrap_or_else(|e| account_error_response(&e))
256}
257
258async fn authenticate_local(
264 services: &Services,
265 req_id: &str,
266 body: &NativeSubmit,
267) -> Result<OwnedUserId> {
268 let oidc = services.oauth.get_server()?;
269
270 oidc.check_local_auth_request(req_id).await?;
271
272 if body.mode.as_deref() == Some("register") {
273 return do_register(services, req_id, body).await;
274 }
275
276 let user_id = verify_credentials(services, &body.username, &body.password).await?;
277
278 oidc.bind_auth_request_to_local(req_id).await?;
279
280 Ok(user_id)
281}
282
283async fn verify_credentials(
286 services: &Services,
287 username: &str,
288 password: &str,
289) -> Result<OwnedUserId> {
290 let invalid = || err!(Request(Forbidden("Invalid username or password.")));
291 let server_name = &services.config.server_name;
292
293 let user_id = UserId::parse_with_server_name(username, server_name).map_err(|_| invalid())?;
294
295 if !services.globals.user_is_local(&user_id) {
296 return Err(invalid());
297 }
298
299 let reservation = services
302 .login_ratelimit
303 .reserve_login_attempt(&user_id)?;
304
305 let (user_id, hash) = match services.users.password_hash(&user_id).await {
309 | Ok(hash) => (user_id, hash),
310 | Err(_) => {
311 let lowercased = UserId::parse_with_server_name(username.to_lowercase(), server_name)
312 .map_err(|_| invalid())?;
313
314 let hash = services
315 .users
316 .password_hash(&lowercased)
317 .await
318 .map_err(|_| invalid())?;
319
320 (lowercased, hash)
321 },
322 };
323
324 let unchecked = hash.is_empty()
327 || services
328 .users
329 .origin(&user_id)
330 .await
331 .is_ok_and(|origin| origin != "password");
332
333 if unchecked {
334 services
335 .login_ratelimit
336 .refund_login_attempt(reservation)?;
337
338 return Err(invalid());
339 }
340
341 verify_password(password, &hash).map_err(|_| invalid())?;
342
343 services
344 .login_ratelimit
345 .record_login(reservation)?;
346
347 Ok(user_id)
348}
349
350async fn do_register(
351 services: &Services,
352 req_id: &str,
353 body: &NativeSubmit,
354) -> Result<OwnedUserId> {
355 if !services.config.allow_registration {
356 return Err!(Request(Forbidden("Registration is disabled on this server.")));
357 }
358
359 let username = body.username.trim().to_lowercase();
360 if username.is_empty() {
361 return Err!(Request(InvalidUsername("A username is required.")));
362 }
363
364 if body.password.is_empty() {
365 return Err!(Request(InvalidParam("A password is required.")));
366 }
367
368 let token_required = services.registration_tokens.is_enabled().await;
371 let smtp = &services.config.smtp;
372 let email_required = smtp.connection_uri.is_some()
373 && (smtp.require_email_for_registration
374 || (token_required && smtp.require_email_for_token_registration));
375
376 if email_required {
377 return Err!(Request(Forbidden(
378 "This server requires an email to register, which this page cannot collect."
379 )));
380 }
381
382 if services
383 .config
384 .forbidden_usernames
385 .is_match(&username)
386 {
387 return Err!(Request(Forbidden("That username is not allowed.")));
388 }
389
390 let user_id = UserId::parse_with_server_name(&username, &services.config.server_name)
391 .map_err(|_| err!(Request(InvalidUsername("That username is not valid."))))?;
392
393 user_id.validate_strict().map_err(|_| {
394 err!(Request(InvalidUsername("That username contains disallowed characters.")))
395 })?;
396
397 if services
398 .appservice
399 .is_exclusive_user_id(&user_id)
400 .await
401 {
402 return Err!(Request(Exclusive("That username is reserved by an appservice.")));
403 }
404
405 if services.users.exists(&user_id).await {
406 return Err!(Request(UserInUse("That username is taken.")));
407 }
408
409 services.users.check_creation(&user_id).await?;
410
411 if !services.config.registration_terms.is_empty()
414 && body.accept_terms.as_deref() != Some("on")
415 {
416 return Err!(Request(Forbidden("You must accept the terms to register.")));
417 }
418
419 let token = body
420 .registration_token
421 .as_deref()
422 .unwrap_or_default();
423
424 if token_required {
426 services
427 .registration_tokens
428 .is_token_valid(token)
429 .await?;
430 }
431
432 services
434 .oauth
435 .get_server()?
436 .bind_auth_request_to_local(req_id)
437 .await?;
438
439 if token_required {
440 services
441 .registration_tokens
442 .try_consume(token)
443 .await?;
444 }
445
446 services
447 .users
448 .full_register(Register {
449 user_id: Some(&user_id),
450 password: Some(&body.password),
451 grant_first_user_admin: true,
452 ..Default::default()
453 })
454 .await?;
455
456 record_accepted_terms(services, &user_id).await?;
457
458 Ok(user_id)
459}
460
461async fn record_accepted_terms(services: &Services, user_id: &UserId) -> Result {
462 let accepted: Vec<String> = services
463 .config
464 .registration_terms
465 .values()
466 .flat_map(|policy| policy.translations.values())
467 .map(|translation| translation.url.to_string())
468 .collect();
469
470 if accepted.is_empty() {
471 return Ok(());
472 }
473
474 let event_type = "m.accepted_terms";
475 let event = json!({
476 "type": event_type,
477 "content": { "accepted": accepted },
478 });
479
480 services
481 .account_data
482 .update(None, user_id, event_type.into(), &event)
483 .await
484}
485
486fn complete_redirect(services: &Services, flow: Flow<'_>, login_token: &str) -> Result<Redirect> {
489 let issuer = services.oauth.get_server()?.issuer_url()?;
490 let base = issuer.trim_end_matches('/');
491
492 let url = match flow {
493 | Flow::Device(user_code) =>
494 Url::parse_with_params(&format!("{base}/_tuwunel/oidc/device_callback"), [
495 ("user_code", user_code),
496 ("loginToken", login_token),
497 ]),
498 | Flow::Authorization(req_id) =>
499 Url::parse_with_params(&format!("{base}/_tuwunel/oidc/_complete"), [
500 ("oidc_req_id", req_id),
501 ("loginToken", login_token),
502 ]),
503 | Flow::Account { action, device_id } =>
504 Url::parse_with_params(&format!("{base}/_tuwunel/oidc/account_callback"), [
505 ("action", action),
506 ("device_id", device_id),
507 ("loginToken", login_token),
508 ]),
509 }
510 .map_err(|_| err!(error!("Failed to build completion URL")))?;
511
512 Ok(Redirect::to(url.as_str()))
513}
514
515fn require_native(services: &Services) -> Result {
516 services.oauth.get_server()?;
517
518 services
519 .config
520 .oidc_native_auth
521 .then_some(())
522 .ok_or_else(|| err!(Request(NotFound("Native authentication is not enabled"))))
523}
524
525async fn render_page(
530 services: &Services,
531 view: &str,
532 context: Flow<'_>,
533 error: Option<&str>,
534) -> Result<String> {
535 let registration_enabled = services.config.allow_registration;
536 let Flow::Authorization(req_id) = context else {
537 return Ok(render_login(context, error, registration_enabled, ""));
538 };
539
540 let bound = services
541 .oauth
542 .get_server()?
543 .peek_auth_request(req_id)
544 .await?
545 .idp_id;
546
547 let page = match bound.as_deref() {
548 | None if view == "register" && registration_enabled =>
549 render_register(services, req_id, error).await,
550
551 | Some(idp_id) => {
552 let provider = services.oauth.providers.find_config(idp_id)?;
553
554 render_bound(req_id, provider_choice(provider), error)
555 },
556
557 | None => {
558 let sso_options =
559 render_sso_options("Or sign in with", req_id, sso_choices(services));
560
561 render_login(context, error, registration_enabled, &sso_options)
562 },
563 };
564
565 Ok(page)
566}
567
568fn render_login(
569 context: Flow<'_>,
570 error: Option<&str>,
571 show_register: bool,
572 sso_options: &str,
573) -> String {
574 let (context_fields, register_link) = match context {
575 | Flow::Device(user_code) => {
576 let context_fields = format!(
577 r#"<input type="hidden" name="user_code" value="{}">"#,
578 html_escape(user_code),
579 );
580
581 (context_fields, String::new())
582 },
583 | Flow::Account { action, device_id } => {
584 let context_fields = format!(
585 concat!(
586 r#"<input type="hidden" name="action" value="{}">"#,
587 "\n\t\t\t",
588 r#"<input type="hidden" name="device_id" value="{}">"#,
589 ),
590 html_escape(action),
591 html_escape(device_id),
592 );
593
594 (context_fields, String::new())
595 },
596 | Flow::Authorization(req_id) => {
597 let context_fields = format!(
598 r#"<input type="hidden" name="oidc_req_id" value="{}">"#,
599 html_escape(req_id),
600 );
601
602 let register_link = show_register
603 .then(|| {
604 format!(
605 r#"<p class="auth-nav">New to this server? <a href="/_tuwunel/oidc/native?oidc_req_id={}&view=register">Create an account</a></p>"#,
606 url_encode(req_id),
607 )
608 })
609 .unwrap_or_default();
610
611 (context_fields, register_link)
612 },
613 };
614
615 LOGIN_HTML
616 .replace("{register_link}", ®ister_link)
617 .replace("{sso_options}", sso_options)
618 .replace("{error}", &error_block(error))
619 .replace("{context_fields}", &context_fields)
621}
622
623async fn render_register(services: &Services, req_id: &str, error: Option<&str>) -> String {
624 let token_field = services
625 .registration_tokens
626 .is_enabled()
627 .await
628 .then_some(TOKEN_FIELD)
629 .unwrap_or_default();
630
631 REGISTER_HTML
632 .replace("{token_field}", token_field)
633 .replace("{req_id_enc}", &url_encode(req_id))
634 .replace("{terms}", &terms_block(services))
635 .replace("{error}", &error_block(error))
636 .replace("{req_id}", &html_escape(req_id))
638}
639
640fn provider_choice(provider: &IdentityProvider) -> ProviderChoice<'_> {
641 (provider.id(), provider.display_name())
642}
643
644fn render_bound(req_id: &str, provider: ProviderChoice<'_>, error: Option<&str>) -> String {
649 let sso_options = render_sso_options("Continue with", req_id, [provider]);
650
651 BOUND_HTML
652 .replace("{sso_options}", &sso_options)
653 .replace("{error}", &error_block(error))
654}
655
656fn sso_choices(services: &Services) -> impl Iterator<Item = ProviderChoice<'_>> {
661 let listed = services
662 .config
663 .identity_provider
664 .values()
665 .map(provider_choice);
666
667 let single = services
668 .oauth
669 .providers
670 .find_default_config()
671 .map(|provider| (provider.id(), "Single sign-on"));
672
673 match services.config.lists_identity_providers() {
674 | true => Left(listed),
675 | false => Right(single.into_iter()),
676 }
677}
678
679fn render_sso_options<'a, I>(heading: &str, req_id: &str, providers: I) -> String
684where
685 I: IntoIterator<Item = ProviderChoice<'a>>,
686{
687 let req_id = url_encode(req_id);
688 let options = providers
689 .into_iter()
690 .map(|(id, name)| {
691 let name = html_escape(name)
692 .replace('{', "{")
693 .replace('}', "}");
694
695 (url_encode(id), name)
696 })
697 .fold(String::new(), |mut out, (id, name)| {
698 write!(
699 out,
700 r#"<li><a href="/_tuwunel/oidc/native?oidc_req_id={req_id}&idp_id={id}">{name}</a></li>"#,
701 )
702 .ok();
703
704 out
705 });
706
707 options
708 .is_empty()
709 .is_false()
710 .then(|| {
711 format!(
712 r#"<section class="sso-options"><h2>{heading}</h2><ul>{options}</ul></section>"#
713 )
714 })
715 .unwrap_or_default()
716}
717
718fn error_block(error: Option<&str>) -> String {
719 error
720 .map(|msg| format!(r#"<p class="err">{}</p>"#, html_escape(msg)))
721 .unwrap_or_default()
722}
723
724fn terms_block(services: &Services) -> String {
725 let policies = &services.config.registration_terms;
726 if policies.is_empty() {
727 return String::new();
728 }
729
730 let links = policies
731 .values()
732 .filter_map(|policy| {
733 policy
734 .translations
735 .get("en")
736 .or_else(|| policy.translations.values().next())
737 })
738 .fold(String::new(), |mut links, translation| {
739 write!(
740 links,
741 r#"<li><a href="{}" target="_blank" rel="noopener noreferrer">{}</a></li>"#,
742 html_escape(translation.url.as_str()),
743 html_escape(&translation.name),
744 )
745 .ok();
746
747 links
748 });
749
750 format!(
751 r#"<fieldset class="terms"><legend>Terms</legend><ul>{links}</ul><label><input type="checkbox" name="accept_terms" value="on" required> I accept the terms above.</label></fieldset>"#
752 )
753}
754
755static LOGIN_HTML: &str = const_format!(
756 r#"
757<!DOCTYPE html>
758<html lang="en">
759 <head>
760 {ACCOUNT_HEAD}
761 <title>Sign in · Tuwunel</title>
762 </head>
763 <body class="auth-page">
764 <main class="auth-card" aria-labelledby="auth-title">
765 <h1 id="auth-title">Sign in</h1>
766 <p class="auth-description">Sign in to your Tuwunel account.</p>
767 {{error}}
768 <form class="auth-form" method="POST" action="/_tuwunel/oidc/native">
769 {{context_fields}}
770 <input type="hidden" name="mode" value="login">
771 <label for="auth-username">Username</label>
772 <input id="auth-username" type="text" name="username" autocomplete="username" autofocus required>
773 <label for="auth-password">Password</label>
774 <input id="auth-password" type="password" name="password" autocomplete="current-password" required>
775 <button type="submit">Sign in</button>
776 </form>
777 {{sso_options}}
778 {{register_link}}
779 </main>
780 </body>
781</html>"#
782);
783
784static REGISTER_HTML: &str = const_format!(
785 r#"
786<!DOCTYPE html>
787<html lang="en">
788 <head>
789 {ACCOUNT_HEAD}
790 <title>Create account · Tuwunel</title>
791 </head>
792 <body class="auth-page">
793 <main class="auth-card" aria-labelledby="auth-title">
794 <h1 id="auth-title">Create account</h1>
795 <p class="auth-description">Set up your account on this homeserver.</p>
796 {{error}}
797 <form class="auth-form" method="POST" action="/_tuwunel/oidc/native">
798 <input type="hidden" name="oidc_req_id" value="{{req_id}}">
799 <input type="hidden" name="mode" value="register">
800 <label for="auth-username">Username</label>
801 <input id="auth-username" type="text" name="username" autocomplete="username" autofocus required>
802 <label for="auth-password">Password</label>
803 <input id="auth-password" type="password" name="password" autocomplete="new-password" required>
804 {{token_field}}
805 {{terms}}
806 <button type="submit">Create account</button>
807 </form>
808 <p class="auth-nav">Already have an account? <a href="/_tuwunel/oidc/native?oidc_req_id={{req_id_enc}}&view=login">Sign in</a></p>
809 </main>
810 </body>
811</html>"#
812);
813
814static BOUND_HTML: &str = const_format!(
815 r#"
816<!DOCTYPE html>
817<html lang="en">
818 <head>
819 {ACCOUNT_HEAD}
820 <title>Continue signing in · Tuwunel</title>
821 </head>
822 <body class="auth-page">
823 <main class="auth-card" aria-labelledby="auth-title">
824 <h1 id="auth-title">Continue signing in</h1>
825 <p class="auth-description">Finish signing in with the provider you chose.</p>
826 {{error}}
827 {{sso_options}}
828 </main>
829 </body>
830</html>"#
831);
832
833static TOKEN_FIELD: &str = r#"<label for="auth-token">Registration token</label>
834 <input id="auth-token" type="text" name="registration_token" autocomplete="off" placeholder="Enter your token" required>"#;
835
836#[cfg(test)]
837mod tests {
838 use super::{Flow, error_block, parse_flow, render_bound, render_login, render_sso_options};
839
840 #[test]
841 fn login_page_has_form_and_hidden_req_id() {
842 let html = render_login(Flow::Authorization("REQ123"), None, false, "");
843
844 assert!(html.contains(r#"action="/_tuwunel/oidc/native""#));
845 assert!(html.contains(r#"name="oidc_req_id" value="REQ123""#));
846 assert!(html.contains(r#"name="username""#));
847 assert!(html.contains(r#"name="password""#));
848 assert!(!html.contains("view=register"));
849 }
850
851 #[test]
852 fn login_page_links_to_register_when_enabled() {
853 let html = render_login(Flow::Authorization("REQ123"), None, true, "");
854
855 assert!(html.contains("oidc_req_id=REQ123&view=register"));
856 }
857
858 #[test]
859 fn login_page_offers_each_provider_with_a_bound_request() {
860 let providers =
861 [("first/provider", "First provider"), ("second", "Second {error} <provider>")];
862
863 let options = render_sso_options("Or sign in with", "REQ123", providers);
864 let html = render_login(Flow::Authorization("REQ123"), None, false, &options);
865
866 assert!(html.contains("oidc_req_id=REQ123&idp_id=first%2Fprovider"));
867 assert!(html.contains("oidc_req_id=REQ123&idp_id=second"));
868 assert!(html.contains("Second {error} <provider>"));
869 assert!(!html.contains("<provider>"));
870 assert!(html.contains(r#"name="password""#));
871 }
872
873 #[test]
874 fn bound_page_offers_only_its_provider() {
875 let html = render_bound("REQ123", ("first", "First <provider>"), Some("Already chosen"));
876
877 assert!(html.contains("oidc_req_id=REQ123&idp_id=first"));
878 assert!(html.contains("First <provider>"));
879 assert!(html.contains("Already chosen"));
880 assert!(!html.contains(r#"name="password""#));
881 assert!(!html.contains("{sso_options}"));
882 }
883
884 #[test]
885 fn login_page_escapes_error_and_req_id() {
886 let html = render_login(
887 Flow::Authorization("a<b>c"),
888 Some("<script>alert(1)</script>"),
889 false,
890 "",
891 );
892
893 assert!(!html.contains("<script>"));
894 assert!(html.contains("<script>"));
895 assert!(!html.contains("a<b>c"));
896 assert!(html.contains("a<b>c"));
897 }
898
899 #[test]
900 fn login_page_does_not_expand_smuggled_placeholder() {
901 let html = render_login(Flow::Authorization("{error}"), Some("BOOM"), false, "");
903
904 assert_eq!(html.matches("BOOM").count(), 1);
905 assert!(html.contains(r#"value="{error}""#));
906 }
907
908 #[test]
909 fn device_login_page_has_only_hidden_user_code() {
910 let html = render_login(Flow::Device("BCDF-GHJK"), None, true, "");
911
912 assert!(html.contains(r#"name="user_code" value="BCDF-GHJK""#));
913 assert!(!html.contains(r#"name="oidc_req_id""#));
914 assert!(!html.contains("view=register"));
915 }
916
917 #[test]
918 fn device_login_page_escapes_and_does_not_expand_context() {
919 let html = render_login(Flow::Device("a<{error}>"), Some("BOOM"), true, "");
920
921 assert_eq!(html.matches("BOOM").count(), 1);
922 assert!(!html.contains("a<{error}>"));
923 assert!(html.contains(r#"value="a<{error}>""#));
924 }
925
926 #[test]
927 fn account_login_page_has_hidden_action_and_device_id() {
928 let context = Flow::Account {
929 action: "org.matrix.sessions_list",
930 device_id: "",
931 };
932
933 let html = render_login(context, None, true, "");
934
935 assert!(html.contains(r#"name="action" value="org.matrix.sessions_list""#));
936 assert!(html.contains(r#"name="device_id" value="""#));
937 assert!(!html.contains(r#"name="oidc_req_id""#));
938 assert!(!html.contains(r#"name="user_code""#));
939 assert!(!html.contains("view=register"));
940 }
941
942 #[test]
943 fn account_login_page_escapes_and_does_not_expand_context() {
944 let context = Flow::Account {
945 action: "a<{error}>",
946 device_id: "b<{error}>",
947 };
948
949 let html = render_login(context, Some("BOOM"), true, "");
950
951 assert_eq!(html.matches("BOOM").count(), 1);
952 assert!(!html.contains("a<{error}>"));
953 assert!(!html.contains("b<{error}>"));
954 assert!(html.contains(r#"name="action" value="a<{error}>""#));
955 assert!(html.contains(r#"name="device_id" value="b<{error}>""#));
956 }
957
958 #[test]
959 fn flow_requires_exactly_one_nonempty_value() {
960 assert!(matches!(
961 parse_flow(Some("REQ123"), None, None, None),
962 Ok(Flow::Authorization("REQ123"))
963 ));
964
965 assert!(matches!(
966 parse_flow(None, Some("BCDF-GHJK"), None, None),
967 Ok(Flow::Device("BCDF-GHJK"))
968 ));
969
970 assert!(matches!(
971 parse_flow(None, None, Some("org.matrix.sessions_list"), None),
972 Ok(Flow::Account {
973 action: "org.matrix.sessions_list",
974 device_id: "",
975 })
976 ));
977
978 assert!(matches!(
979 parse_flow(None, None, Some("org.matrix.session_view"), Some("DEVICE")),
980 Ok(Flow::Account {
981 action: "org.matrix.session_view",
982 device_id: "DEVICE",
983 })
984 ));
985
986 assert!(matches!(
987 parse_flow(None, None, Some("org.matrix.sessions_list"), Some("")),
988 Ok(Flow::Account {
989 action: "org.matrix.sessions_list",
990 device_id: "",
991 })
992 ));
993
994 assert!(parse_flow(None, None, None, None).is_err());
995 assert!(parse_flow(None, None, None, Some("DEVICE")).is_err());
996 assert!(parse_flow(Some(""), None, None, None).is_err());
997 assert!(parse_flow(None, Some(""), None, None).is_err());
998 assert!(parse_flow(None, None, Some(""), None).is_err());
999 assert!(parse_flow(None, None, Some(""), Some("DEVICE")).is_err());
1000 assert!(parse_flow(Some("REQ123"), Some("BCDF-GHJK"), None, None).is_err());
1001 assert!(
1002 parse_flow(Some("REQ123"), None, Some("org.matrix.sessions_list"), None).is_err()
1003 );
1004
1005 assert!(
1006 parse_flow(None, Some("BCDF-GHJK"), Some("org.matrix.sessions_list"), None).is_err()
1007 );
1008
1009 assert!(
1010 parse_flow(
1011 Some("REQ123"),
1012 Some("BCDF-GHJK"),
1013 Some("org.matrix.sessions_list"),
1014 None,
1015 )
1016 .is_err()
1017 );
1018 }
1019
1020 #[test]
1021 fn error_block_renders_only_when_present() {
1022 let block = error_block(None);
1023
1024 assert!(block.is_empty(), "{block:?}");
1025 assert!(error_block(Some("oops")).contains(r#"class="err""#));
1026 }
1027}