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 ruma::{OwnedUserId, UserId};
10use serde::Deserialize;
11use serde_json::json;
12use tuwunel_core::{
13 Err, Result, err,
14 smallstr::SmallString,
15 utils::{self, hash, html::escape as html_escape},
16};
17use tuwunel_service::{Services, users::Register};
18use url::Url;
19
20use super::{
21 account::{
22 ACCOUNT_HEAD, account_error_response, account_html_response, account_redirect_response,
23 },
24 url_encode,
25};
26use crate::ClientIp;
27
28type AccountAction = SmallString<[u8; 32]>;
29type DeviceId = SmallString<[u8; 24]>;
30
31const LOGIN_TOKEN_LENGTH: usize = 32;
32
33#[derive(Debug, Default, Deserialize)]
34struct NativeQuery {
35 oidc_req_id: Option<String>,
36 user_code: Option<String>,
37 action: Option<AccountAction>,
38 device_id: Option<DeviceId>,
39 view: Option<String>,
40}
41
42#[derive(Debug, Deserialize)]
43pub(crate) struct NativeSubmit {
44 #[serde(default)]
45 oidc_req_id: Option<String>,
46 #[serde(default)]
47 user_code: Option<String>,
48 #[serde(default)]
49 action: Option<AccountAction>,
50 #[serde(default)]
51 device_id: Option<DeviceId>,
52 #[serde(default)]
53 mode: Option<String>,
54 username: String,
55 password: String,
56 #[serde(default)]
57 registration_token: Option<String>,
58 #[serde(default)]
59 accept_terms: Option<String>,
60}
61
62#[derive(Clone, Copy)]
63enum Flow<'a> {
64 Account {
65 action: &'a str,
66 device_id: &'a str,
67 },
68 Authorization(&'a str),
69 Device(&'a str),
70}
71
72pub(crate) async fn native_get_route(
75 State(services): State<crate::State>,
76 request: Request,
77) -> Response {
78 if let Err(e) = require_native(&services) {
79 return account_error_response(&e);
80 }
81
82 let params: NativeQuery =
83 match serde_html_form::from_str(request.uri().query().unwrap_or_default()) {
84 | Ok(params) => params,
85 | Err(e) => return account_error_response(&e.into()),
86 };
87
88 let context = match parse_flow(
89 params.oidc_req_id.as_deref(),
90 params.user_code.as_deref(),
91 params.action.as_deref(),
92 params.device_id.as_deref(),
93 ) {
94 | Ok(context) => context,
95 | Err(e) => return account_error_response(&e),
96 };
97
98 let view = params.view.as_deref().unwrap_or("login");
99
100 account_html_response(StatusCode::OK, render_page(&services, view, context, None).await)
101}
102
103fn parse_flow<'a>(
104 oidc_req_id: Option<&'a str>,
105 user_code: Option<&'a str>,
106 action: Option<&'a str>,
107 device_id: Option<&'a str>,
108) -> Result<Flow<'a>> {
109 match (
110 oidc_req_id.filter(|value| !value.is_empty()),
111 user_code.filter(|value| !value.is_empty()),
112 action.filter(|value| !value.is_empty()),
113 ) {
114 | (Some(req_id), None, None) => Ok(Flow::Authorization(req_id)),
115 | (None, Some(user_code), None) => Ok(Flow::Device(user_code)),
116 | (None, None, Some(action)) => Ok(Flow::Account {
117 action,
118 device_id: device_id.unwrap_or_default(),
119 }),
120 | _ => Err!(Request(InvalidParam(
121 "Exactly one OIDC request ID, user code, or account action is required."
122 ))),
123 }
124}
125
126pub(crate) async fn native_submit_route(
129 State(services): State<crate::State>,
130 ClientIp(client): ClientIp,
131 Form(body): Form<NativeSubmit>,
132) -> Response {
133 match native_submit(&services, client, &body).await {
134 | Ok(response) => response,
135 | Err(e) => {
136 let context = match parse_flow(
137 body.oidc_req_id.as_deref(),
138 body.user_code.as_deref(),
139 body.action.as_deref(),
140 body.device_id.as_deref(),
141 ) {
142 | Ok(context) => context,
143 | Err(context_error) => return account_error_response(&context_error),
144 };
145
146 let view = match (context, body.mode.as_deref()) {
147 | (Flow::Authorization(_), Some("register")) => "register",
148 | _ => "login",
149 };
150
151 let msg = e.sanitized_message();
152 let html = render_page(&services, view, context, Some(&msg)).await;
153
154 account_html_response(e.status_code(), html)
155 },
156 }
157}
158
159async fn native_submit(
160 services: &Services,
161 client: IpAddr,
162 body: &NativeSubmit,
163) -> Result<Response> {
164 require_native(services)?;
165 services.oauth.check_device_rate_limit(client)?;
167 services.oauth.check_rate_limit(client)?;
168
169 let context = parse_flow(
170 body.oidc_req_id.as_deref(),
171 body.user_code.as_deref(),
172 body.action.as_deref(),
173 body.device_id.as_deref(),
174 )?;
175
176 let user_id = match (context, body.mode.as_deref()) {
177 | (Flow::Authorization(_), Some("register")) => do_register(services, body).await?,
178 | _ => verify_credentials(services, &body.username, &body.password).await?,
179 };
180
181 let token = utils::random_string(LOGIN_TOKEN_LENGTH);
182 let _expires_in = services
183 .users
184 .create_login_token(&user_id, &token);
185
186 let redirect = complete_redirect(services, context, &token)?;
187
188 Ok(account_redirect_response(redirect))
189}
190
191async fn verify_credentials(
194 services: &Services,
195 username: &str,
196 password: &str,
197) -> Result<OwnedUserId> {
198 let invalid = || err!(Request(Forbidden("Invalid username or password.")));
199 let server_name = &services.config.server_name;
200
201 let user_id = UserId::parse_with_server_name(username, server_name).map_err(|_| invalid())?;
202
203 if !services.globals.user_is_local(&user_id) {
204 return Err(invalid());
205 }
206
207 let (user_id, hash) = match services.users.password_hash(&user_id).await {
210 | Ok(hash) => (user_id, hash),
211 | Err(_) => {
212 let lowercased = UserId::parse_with_server_name(username.to_lowercase(), server_name)
213 .map_err(|_| invalid())?;
214
215 let hash = services
216 .users
217 .password_hash(&lowercased)
218 .await
219 .map_err(|_| invalid())?;
220
221 (lowercased, hash)
222 },
223 };
224
225 if services
227 .users
228 .origin(&user_id)
229 .await
230 .is_ok_and(|origin| origin != "password")
231 {
232 return Err(invalid());
233 }
234
235 if hash.is_empty() {
236 return Err(invalid());
237 }
238
239 hash::verify_password(password, &hash).map_err(|_| invalid())?;
240
241 Ok(user_id)
242}
243
244async fn do_register(services: &Services, body: &NativeSubmit) -> Result<OwnedUserId> {
245 if !services.config.allow_registration {
246 return Err!(Request(Forbidden("Registration is disabled on this server.")));
247 }
248
249 let username = body.username.trim().to_lowercase();
250 if username.is_empty() {
251 return Err!(Request(InvalidUsername("A username is required.")));
252 }
253
254 if body.password.is_empty() {
255 return Err!(Request(InvalidParam("A password is required.")));
256 }
257
258 let token_required = services.registration_tokens.is_enabled().await;
261 let smtp = &services.config.smtp;
262 let email_required = smtp.connection_uri.is_some()
263 && (smtp.require_email_for_registration
264 || (token_required && smtp.require_email_for_token_registration));
265
266 if email_required {
267 return Err!(Request(Forbidden(
268 "This server requires an email to register, which this page cannot collect."
269 )));
270 }
271
272 if services
273 .config
274 .forbidden_usernames
275 .is_match(&username)
276 {
277 return Err!(Request(Forbidden("That username is not allowed.")));
278 }
279
280 let user_id = UserId::parse_with_server_name(&username, &services.config.server_name)
281 .map_err(|_| err!(Request(InvalidUsername("That username is not valid."))))?;
282
283 user_id.validate_strict().map_err(|_| {
284 err!(Request(InvalidUsername("That username contains disallowed characters.")))
285 })?;
286
287 if services
288 .appservice
289 .is_exclusive_user_id(&user_id)
290 .await
291 {
292 return Err!(Request(Exclusive("That username is reserved by an appservice.")));
293 }
294
295 if services.users.exists(&user_id).await {
296 return Err!(Request(UserInUse("That username is taken.")));
297 }
298
299 if !services.config.registration_terms.is_empty()
302 && body.accept_terms.as_deref() != Some("on")
303 {
304 return Err!(Request(Forbidden("You must accept the terms to register.")));
305 }
306
307 if token_required {
308 let token = body
309 .registration_token
310 .as_deref()
311 .unwrap_or_default();
312
313 services
314 .registration_tokens
315 .try_consume(token)
316 .await?;
317 }
318
319 services
320 .users
321 .full_register(Register {
322 user_id: Some(&user_id),
323 password: Some(&body.password),
324 grant_first_user_admin: true,
325 ..Default::default()
326 })
327 .await?;
328
329 record_accepted_terms(services, &user_id).await?;
330
331 Ok(user_id)
332}
333
334async fn record_accepted_terms(services: &Services, user_id: &UserId) -> Result {
335 let accepted: Vec<String> = services
336 .config
337 .registration_terms
338 .values()
339 .flat_map(|policy| policy.translations.values())
340 .map(|translation| translation.url.to_string())
341 .collect();
342
343 if accepted.is_empty() {
344 return Ok(());
345 }
346
347 let event_type = "m.accepted_terms";
348 let event = json!({
349 "type": event_type,
350 "content": { "accepted": accepted },
351 });
352
353 services
354 .account_data
355 .update(None, user_id, event_type.into(), &event)
356 .await
357}
358
359fn complete_redirect(services: &Services, flow: Flow<'_>, login_token: &str) -> Result<Redirect> {
362 let issuer = services.oauth.get_server()?.issuer_url()?;
363 let base = issuer.trim_end_matches('/');
364
365 let url = match flow {
366 | Flow::Device(user_code) =>
367 Url::parse_with_params(&format!("{base}/_tuwunel/oidc/device_callback"), [
368 ("user_code", user_code),
369 ("loginToken", login_token),
370 ]),
371 | Flow::Authorization(req_id) =>
372 Url::parse_with_params(&format!("{base}/_tuwunel/oidc/_complete"), [
373 ("oidc_req_id", req_id),
374 ("loginToken", login_token),
375 ]),
376 | Flow::Account { action, device_id } =>
377 Url::parse_with_params(&format!("{base}/_tuwunel/oidc/account_callback"), [
378 ("action", action),
379 ("device_id", device_id),
380 ("loginToken", login_token),
381 ]),
382 }
383 .map_err(|_| err!(error!("Failed to build completion URL")))?;
384
385 Ok(Redirect::to(url.as_str()))
386}
387
388fn require_native(services: &Services) -> Result {
389 services.oauth.get_server()?;
390
391 services
392 .config
393 .oidc_native_auth
394 .then_some(())
395 .ok_or_else(|| err!(Request(NotFound("Native authentication is not enabled"))))
396}
397
398async fn render_page(
399 services: &Services,
400 view: &str,
401 context: Flow<'_>,
402 error: Option<&str>,
403) -> String {
404 let registration_enabled = services.config.allow_registration;
405
406 match (context, view) {
407 | (Flow::Authorization(req_id), "register") if registration_enabled =>
408 render_register(services, req_id, error).await,
409 | _ => render_login(context, error, registration_enabled),
410 }
411}
412
413fn render_login(context: Flow<'_>, error: Option<&str>, show_register: bool) -> String {
414 let (context_fields, register_link) = match context {
415 | Flow::Device(user_code) => {
416 let context_fields = format!(
417 r#"<input type="hidden" name="user_code" value="{}">"#,
418 html_escape(user_code),
419 );
420
421 (context_fields, String::new())
422 },
423 | Flow::Account { action, device_id } => {
424 let context_fields = format!(
425 concat!(
426 r#"<input type="hidden" name="action" value="{}">"#,
427 "\n\t\t\t",
428 r#"<input type="hidden" name="device_id" value="{}">"#,
429 ),
430 html_escape(action),
431 html_escape(device_id),
432 );
433
434 (context_fields, String::new())
435 },
436 | Flow::Authorization(req_id) => {
437 let context_fields = format!(
438 r#"<input type="hidden" name="oidc_req_id" value="{}">"#,
439 html_escape(req_id),
440 );
441
442 let register_link = show_register
443 .then(|| {
444 format!(
445 r#"<p class="nav">No account? <a href="/_tuwunel/oidc/native?oidc_req_id={}&view=register">Create one</a>.</p>"#,
446 url_encode(req_id),
447 )
448 })
449 .unwrap_or_default();
450
451 (context_fields, register_link)
452 },
453 };
454
455 LOGIN_HTML
456 .replace("{register_link}", ®ister_link)
457 .replace("{error}", &error_block(error))
458 .replace("{context_fields}", &context_fields)
460}
461
462async fn render_register(services: &Services, req_id: &str, error: Option<&str>) -> String {
463 let token_field = services
464 .registration_tokens
465 .is_enabled()
466 .await
467 .then_some(TOKEN_FIELD)
468 .unwrap_or_default();
469
470 REGISTER_HTML
471 .replace("{token_field}", token_field)
472 .replace("{req_id_enc}", &url_encode(req_id))
473 .replace("{terms}", &terms_block(services))
474 .replace("{error}", &error_block(error))
475 .replace("{req_id}", &html_escape(req_id))
477}
478
479fn error_block(error: Option<&str>) -> String {
480 error
481 .map(|msg| format!(r#"<p class="err">{}</p>"#, html_escape(msg)))
482 .unwrap_or_default()
483}
484
485fn terms_block(services: &Services) -> String {
486 let policies = &services.config.registration_terms;
487 if policies.is_empty() {
488 return String::new();
489 }
490
491 let links = policies
492 .values()
493 .filter_map(|policy| {
494 policy
495 .translations
496 .get("en")
497 .or_else(|| policy.translations.values().next())
498 })
499 .fold(String::new(), |mut links, translation| {
500 write!(
501 links,
502 r#"<li><a href="{}" target="_blank" rel="noopener noreferrer">{}</a></li>"#,
503 html_escape(translation.url.as_str()),
504 html_escape(&translation.name),
505 )
506 .ok();
507
508 links
509 });
510
511 format!(
512 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>"#
513 )
514}
515
516static LOGIN_HTML: &str = const_format!(
517 r#"
518<!DOCTYPE html>
519<html lang="en">
520 <head>
521 {ACCOUNT_HEAD}
522 <title>Sign In</title>
523 </head>
524 <body>
525 <h1>Sign In</h1>
526 {{error}}
527 <form method="POST" action="/_tuwunel/oidc/native">
528 {{context_fields}}
529 <input type="hidden" name="mode" value="login">
530 <label>
531 Username
532 <input type="text" name="username" autocomplete="username" autofocus required>
533 </label>
534 <label>
535 Password
536 <input type="password" name="password" autocomplete="current-password" required>
537 </label>
538 <button type="submit">Sign in</button>
539 </form>
540 {{register_link}}
541 </body>
542</html>"#
543);
544
545static REGISTER_HTML: &str = const_format!(
546 r#"
547<!DOCTYPE html>
548<html lang="en">
549 <head>
550 {ACCOUNT_HEAD}
551 <title>Create Account</title>
552 </head>
553 <body>
554 <h1>Create Account</h1>
555 {{error}}
556 <form method="POST" action="/_tuwunel/oidc/native">
557 <input type="hidden" name="oidc_req_id" value="{{req_id}}">
558 <input type="hidden" name="mode" value="register">
559 <label>
560 Username
561 <input type="text" name="username" autocomplete="username" autofocus required>
562 </label>
563 <label>
564 Password
565 <input type="password" name="password" autocomplete="new-password" required>
566 </label>
567 {{token_field}}
568 {{terms}}
569 <button type="submit">Create account</button>
570 </form>
571 <p class="nav">Have an account? <a href="/_tuwunel/oidc/native?oidc_req_id={{req_id_enc}}&view=login">Sign in</a>.</p>
572 </body>
573</html>"#
574);
575
576static TOKEN_FIELD: &str = r#"<label>
577 Registration token
578 <input type="text" name="registration_token" autocomplete="off" required>
579 </label>"#;
580
581#[cfg(test)]
582mod tests {
583 use super::{Flow, error_block, parse_flow, render_login};
584
585 #[test]
586 fn login_page_has_form_and_hidden_req_id() {
587 let html = render_login(Flow::Authorization("REQ123"), None, false);
588
589 assert!(html.contains(r#"action="/_tuwunel/oidc/native""#));
590 assert!(html.contains(r#"name="oidc_req_id" value="REQ123""#));
591 assert!(html.contains(r#"name="username""#));
592 assert!(html.contains(r#"name="password""#));
593 assert!(!html.contains("view=register"));
594 }
595
596 #[test]
597 fn login_page_links_to_register_when_enabled() {
598 let html = render_login(Flow::Authorization("REQ123"), None, true);
599
600 assert!(html.contains("oidc_req_id=REQ123&view=register"));
601 }
602
603 #[test]
604 fn login_page_escapes_error_and_req_id() {
605 let html =
606 render_login(Flow::Authorization("a<b>c"), Some("<script>alert(1)</script>"), false);
607
608 assert!(!html.contains("<script>"));
609 assert!(html.contains("<script>"));
610 assert!(!html.contains("a<b>c"));
611 assert!(html.contains("a<b>c"));
612 }
613
614 #[test]
615 fn login_page_does_not_expand_smuggled_placeholder() {
616 let html = render_login(Flow::Authorization("{error}"), Some("BOOM"), false);
618
619 assert_eq!(html.matches("BOOM").count(), 1);
620 assert!(html.contains(r#"value="{error}""#));
621 }
622
623 #[test]
624 fn device_login_page_has_only_hidden_user_code() {
625 let html = render_login(Flow::Device("BCDF-GHJK"), None, true);
626
627 assert!(html.contains(r#"name="user_code" value="BCDF-GHJK""#));
628 assert!(!html.contains(r#"name="oidc_req_id""#));
629 assert!(!html.contains("view=register"));
630 }
631
632 #[test]
633 fn device_login_page_escapes_and_does_not_expand_context() {
634 let html = render_login(Flow::Device("a<{error}>"), Some("BOOM"), true);
635
636 assert_eq!(html.matches("BOOM").count(), 1);
637 assert!(!html.contains("a<{error}>"));
638 assert!(html.contains(r#"value="a<{error}>""#));
639 }
640
641 #[test]
642 fn account_login_page_has_hidden_action_and_device_id() {
643 let context = Flow::Account {
644 action: "org.matrix.sessions_list",
645 device_id: "",
646 };
647
648 let html = render_login(context, None, true);
649
650 assert!(html.contains(r#"name="action" value="org.matrix.sessions_list""#));
651 assert!(html.contains(r#"name="device_id" value="""#));
652 assert!(!html.contains(r#"name="oidc_req_id""#));
653 assert!(!html.contains(r#"name="user_code""#));
654 assert!(!html.contains("view=register"));
655 }
656
657 #[test]
658 fn account_login_page_escapes_and_does_not_expand_context() {
659 let context = Flow::Account {
660 action: "a<{error}>",
661 device_id: "b<{error}>",
662 };
663
664 let html = render_login(context, Some("BOOM"), true);
665
666 assert_eq!(html.matches("BOOM").count(), 1);
667 assert!(!html.contains("a<{error}>"));
668 assert!(!html.contains("b<{error}>"));
669 assert!(html.contains(r#"name="action" value="a<{error}>""#));
670 assert!(html.contains(r#"name="device_id" value="b<{error}>""#));
671 }
672
673 #[test]
674 fn flow_requires_exactly_one_nonempty_value() {
675 assert!(matches!(
676 parse_flow(Some("REQ123"), None, None, None),
677 Ok(Flow::Authorization("REQ123"))
678 ));
679
680 assert!(matches!(
681 parse_flow(None, Some("BCDF-GHJK"), None, None),
682 Ok(Flow::Device("BCDF-GHJK"))
683 ));
684
685 assert!(matches!(
686 parse_flow(None, None, Some("org.matrix.sessions_list"), None),
687 Ok(Flow::Account {
688 action: "org.matrix.sessions_list",
689 device_id: "",
690 })
691 ));
692
693 assert!(matches!(
694 parse_flow(None, None, Some("org.matrix.session_view"), Some("DEVICE")),
695 Ok(Flow::Account {
696 action: "org.matrix.session_view",
697 device_id: "DEVICE",
698 })
699 ));
700
701 assert!(matches!(
702 parse_flow(None, None, Some("org.matrix.sessions_list"), Some("")),
703 Ok(Flow::Account {
704 action: "org.matrix.sessions_list",
705 device_id: "",
706 })
707 ));
708
709 assert!(parse_flow(None, None, None, None).is_err());
710 assert!(parse_flow(None, None, None, Some("DEVICE")).is_err());
711 assert!(parse_flow(Some(""), None, None, None).is_err());
712 assert!(parse_flow(None, Some(""), None, None).is_err());
713 assert!(parse_flow(None, None, Some(""), None).is_err());
714 assert!(parse_flow(None, None, Some(""), Some("DEVICE")).is_err());
715 assert!(parse_flow(Some("REQ123"), Some("BCDF-GHJK"), None, None).is_err());
716 assert!(
717 parse_flow(Some("REQ123"), None, Some("org.matrix.sessions_list"), None).is_err()
718 );
719
720 assert!(
721 parse_flow(None, Some("BCDF-GHJK"), Some("org.matrix.sessions_list"), None).is_err()
722 );
723
724 assert!(
725 parse_flow(
726 Some("REQ123"),
727 Some("BCDF-GHJK"),
728 Some("org.matrix.sessions_list"),
729 None,
730 )
731 .is_err()
732 );
733 }
734
735 #[test]
736 fn error_block_renders_only_when_present() {
737 let block = error_block(None);
738
739 assert!(block.is_empty(), "{block:?}");
740 assert!(error_block(Some("oops")).contains(r#"class="err""#));
741 }
742}