1#![deny(clippy::future_not_send)]
9#![allow(
10 clippy::unused_async,
12 clippy::too_many_arguments,
14 clippy::let_with_type_underscore,
17)]
18
19use std::{
20 convert::Infallible,
21 sync::{Arc, LazyLock},
22 time::Duration,
23};
24
25use axum::{
26 Extension, Router,
27 extract::{FromRef, FromRequestParts, OriginalUri, RawQuery, State},
28 http::Method,
29 response::{Html, IntoResponse},
30 routing::{get, post},
31};
32use headers::HeaderName;
33use hyper::{
34 StatusCode, Version,
35 header::{
36 ACCEPT, ACCEPT_LANGUAGE, AUTHORIZATION, CONTENT_LANGUAGE, CONTENT_LENGTH, CONTENT_TYPE,
37 X_FRAME_OPTIONS,
38 },
39};
40use mas_axum_utils::{InternalError, cookies::CookieJar};
41use mas_data_model::SiteConfig;
42use mas_http::CorsLayerExt;
43use mas_keystore::{Encrypter, Keystore};
44use mas_matrix::HomeserverConnection;
45use mas_policy::Policy;
46use mas_router::{Route, UrlBuilder};
47use mas_storage::{BoxRepository, BoxRepositoryFactory};
48use mas_templates::{ErrorContext, NotFoundContext, TemplateContext, Templates};
49use opentelemetry::metrics::Meter;
50use sqlx::PgPool;
51use tower::util::AndThenLayer;
52use tower_http::{
53 cors::{Any, CorsLayer},
54 set_header::SetResponseHeaderLayer,
55};
56
57use self::{graphql::ExtraRouterParameters, passwords::PasswordManager};
58
59mod admin;
60mod compat;
61mod graphql;
62mod health;
63mod oauth2;
64pub mod passwords;
65pub mod upstream_oauth2;
66mod views;
67
68mod activity_tracker;
69mod captcha;
70#[cfg(test)]
71mod cleanup_tests;
72mod client_ip;
73mod preferred_language;
74mod rate_limit;
75mod session;
76#[cfg(test)]
77mod test_utils;
78
79static METER: LazyLock<Meter> = LazyLock::new(|| {
80 let scope = opentelemetry::InstrumentationScope::builder(env!("CARGO_PKG_NAME"))
81 .with_version(env!("CARGO_PKG_VERSION"))
82 .with_schema_url(opentelemetry_semantic_conventions::SCHEMA_URL)
83 .build();
84
85 opentelemetry::global::meter_with_scope(scope)
86});
87
88#[macro_export]
91macro_rules! impl_from_error_for_route {
92 ($route_error:ty : $error:ty) => {
93 impl From<$error> for $route_error {
94 fn from(e: $error) -> Self {
95 Self::Internal(Box::new(e))
96 }
97 }
98 };
99 ($error:ty) => {
100 impl_from_error_for_route!(self::RouteError: $error);
101 };
102}
103
104pub use mas_axum_utils::{ErrorWrapper, cookies::CookieManager};
105use mas_data_model::{BoxClock, BoxRng};
106
107pub use self::{
108 activity_tracker::{ActivityTracker, Bound as BoundActivityTracker},
109 admin::router as admin_api_router,
110 client_ip::ClientIp,
111 graphql::{
112 GraphQLOperation, Schema as GraphQLSchema, schema as graphql_schema,
113 schema_builder as graphql_schema_builder,
114 },
115 preferred_language::PreferredLanguage,
116 rate_limit::{Limiter, RequesterFingerprint},
117 upstream_oauth2::cache::MetadataCache,
118};
119
120pub fn healthcheck_router<S>() -> Router<S>
121where
122 S: Clone + Send + Sync + 'static,
123 PgPool: FromRef<S>,
124{
125 Router::new().route(mas_router::Healthcheck::route(), get(self::health::get))
126}
127
128pub fn graphql_router<S>(undocumented_oauth2_access: bool) -> Router<S>
129where
130 S: Clone + Send + Sync + 'static,
131 graphql::Schema: FromRef<S>,
132 BoundActivityTracker: FromRequestParts<S>,
133 BoxRepository: FromRequestParts<S>,
134 BoxClock: FromRequestParts<S>,
135 Encrypter: FromRef<S>,
136 CookieJar: FromRequestParts<S>,
137 Limiter: FromRef<S>,
138 RequesterFingerprint: FromRequestParts<S>,
139{
140 Router::new()
141 .route(
142 mas_router::GraphQL::route(),
143 get(self::graphql::get).post(self::graphql::post),
144 )
145 .layer(Extension(ExtraRouterParameters {
148 undocumented_oauth2_access,
149 }))
150 .layer(
151 CorsLayer::new()
152 .allow_origin(Any)
153 .allow_methods(Any)
154 .allow_otel_headers([
155 AUTHORIZATION,
156 ACCEPT,
157 ACCEPT_LANGUAGE,
158 CONTENT_LANGUAGE,
159 CONTENT_TYPE,
160 ]),
161 )
162}
163
164pub fn discovery_router<S>() -> Router<S>
165where
166 S: Clone + Send + Sync + 'static,
167 Keystore: FromRef<S>,
168 SiteConfig: FromRef<S>,
169 UrlBuilder: FromRef<S>,
170 BoxClock: FromRequestParts<S>,
171 BoxRng: FromRequestParts<S>,
172{
173 Router::new()
174 .route(
175 mas_router::OidcConfiguration::route(),
176 get(self::oauth2::discovery::get),
177 )
178 .route(
179 mas_router::Webfinger::route(),
180 get(self::oauth2::webfinger::get),
181 )
182 .layer(
183 CorsLayer::new()
184 .allow_origin(Any)
185 .allow_methods(Any)
186 .allow_otel_headers([
187 AUTHORIZATION,
188 ACCEPT,
189 ACCEPT_LANGUAGE,
190 CONTENT_LANGUAGE,
191 CONTENT_TYPE,
192 ])
193 .max_age(Duration::from_hours(1)),
194 )
195}
196
197pub fn api_router<S>() -> Router<S>
198where
199 S: Clone + Send + Sync + 'static,
200 Keystore: FromRef<S>,
201 UrlBuilder: FromRef<S>,
202 BoxRepository: FromRequestParts<S>,
203 ActivityTracker: FromRequestParts<S>,
204 BoundActivityTracker: FromRequestParts<S>,
205 Encrypter: FromRef<S>,
206 reqwest::Client: FromRef<S>,
207 SiteConfig: FromRef<S>,
208 Templates: FromRef<S>,
209 Arc<dyn HomeserverConnection>: FromRef<S>,
210 BoxClock: FromRequestParts<S>,
211 BoxRng: FromRequestParts<S>,
212 Policy: FromRequestParts<S>,
213{
214 Router::new()
216 .route(
217 mas_router::OAuth2Keys::route(),
218 get(self::oauth2::keys::get),
219 )
220 .route(
221 mas_router::OidcUserinfo::route(),
222 get(self::oauth2::userinfo::get).post(self::oauth2::userinfo::get),
223 )
224 .route(
225 mas_router::OAuth2Introspection::route(),
226 post(self::oauth2::introspection::post),
227 )
228 .route(
229 mas_router::OAuth2Revocation::route(),
230 post(self::oauth2::revoke::post),
231 )
232 .route(
233 mas_router::OAuth2TokenEndpoint::route(),
234 post(self::oauth2::token::post),
235 )
236 .route(
237 mas_router::OAuth2RegistrationEndpoint::route(),
238 post(self::oauth2::registration::post),
239 )
240 .route(
241 mas_router::OAuth2DeviceAuthorizationEndpoint::route(),
242 post(self::oauth2::device::authorize::post),
243 )
244 .layer(
245 CorsLayer::new()
246 .allow_origin(Any)
247 .allow_methods(Any)
248 .allow_otel_headers([
249 AUTHORIZATION,
250 ACCEPT,
251 ACCEPT_LANGUAGE,
252 CONTENT_LANGUAGE,
253 CONTENT_TYPE,
254 HeaderName::from_static("x-requested-with"),
256 ])
257 .max_age(Duration::from_hours(1)),
258 )
259}
260
261pub fn compat_router<S>(templates: Templates) -> Router<S>
262where
263 S: Clone + Send + Sync + 'static,
264 UrlBuilder: FromRef<S>,
265 SiteConfig: FromRef<S>,
266 Arc<dyn HomeserverConnection>: FromRef<S>,
267 PasswordManager: FromRef<S>,
268 Limiter: FromRef<S>,
269 BoxRepositoryFactory: FromRef<S>,
270 BoundActivityTracker: FromRequestParts<S>,
271 RequesterFingerprint: FromRequestParts<S>,
272 BoxRepository: FromRequestParts<S>,
273 BoxClock: FromRequestParts<S>,
274 BoxRng: FromRequestParts<S>,
275 Policy: FromRequestParts<S>,
276{
277 let human_router = Router::new()
279 .route(
280 mas_router::CompatLoginSsoRedirect::route(),
281 get(self::compat::login_sso_redirect::get),
282 )
283 .route(
284 mas_router::CompatLoginSsoRedirectIdp::route(),
285 get(self::compat::login_sso_redirect::get),
286 )
287 .route(
288 mas_router::CompatLoginSsoRedirectSlash::route(),
289 get(self::compat::login_sso_redirect::get),
290 )
291 .layer(AndThenLayer::new(
292 async move |response: axum::response::Response| {
293 Ok::<_, Infallible>(recover_error(&templates, response))
294 },
295 ));
296
297 let api_router = Router::new()
299 .route(
300 mas_router::CompatLogin::route(),
301 get(self::compat::login::get).post(self::compat::login::post),
302 )
303 .route(
304 mas_router::CompatLogout::route(),
305 post(self::compat::logout::post),
306 )
307 .route(
308 mas_router::CompatLogoutAll::route(),
309 post(self::compat::logout_all::post),
310 )
311 .route(
312 mas_router::CompatRefresh::route(),
313 post(self::compat::refresh::post),
314 )
315 .layer(
316 CorsLayer::new()
317 .allow_origin(Any)
318 .allow_methods(Any)
319 .allow_otel_headers([
320 AUTHORIZATION,
321 ACCEPT,
322 ACCEPT_LANGUAGE,
323 CONTENT_LANGUAGE,
324 CONTENT_TYPE,
325 HeaderName::from_static("x-requested-with"),
326 ])
327 .max_age(Duration::from_hours(1)),
328 );
329
330 Router::new().merge(human_router).merge(api_router)
331}
332
333pub fn human_router<S>(templates: Templates) -> Router<S>
334where
335 S: Clone + Send + Sync + 'static,
336 UrlBuilder: FromRef<S>,
337 PreferredLanguage: FromRequestParts<S>,
338 BoxRepository: FromRequestParts<S>,
339 CookieJar: FromRequestParts<S>,
340 BoundActivityTracker: FromRequestParts<S>,
341 RequesterFingerprint: FromRequestParts<S>,
342 Encrypter: FromRef<S>,
343 Templates: FromRef<S>,
344 Keystore: FromRef<S>,
345 PasswordManager: FromRef<S>,
346 MetadataCache: FromRef<S>,
347 SiteConfig: FromRef<S>,
348 Limiter: FromRef<S>,
349 reqwest::Client: FromRef<S>,
350 Arc<dyn HomeserverConnection>: FromRef<S>,
351 BoxClock: FromRequestParts<S>,
352 BoxRng: FromRequestParts<S>,
353 Policy: FromRequestParts<S>,
354{
355 Router::new()
356 .route(
358 "/account",
359 get(
360 async |State(url_builder): State<UrlBuilder>, RawQuery(query): RawQuery| {
361 let prefix = url_builder.prefix().unwrap_or_default();
362 let route = mas_router::Account::route();
363 let destination = if let Some(query) = query {
364 format!("{prefix}{route}?{query}")
365 } else {
366 format!("{prefix}{route}")
367 };
368
369 axum::response::Redirect::to(&destination)
370 },
371 ),
372 )
373 .route(mas_router::Account::route(), get(self::views::app::get))
374 .route(
375 mas_router::AccountWildcard::route(),
376 get(self::views::app::get),
377 )
378 .route(
379 mas_router::AccountRecoveryFinish::route(),
380 get(self::views::app::get_anonymous),
381 )
382 .route(
383 mas_router::ChangePasswordDiscovery::route(),
384 get(async |State(url_builder): State<UrlBuilder>| {
385 url_builder.redirect(&mas_router::AccountPasswordChange)
386 }),
387 )
388 .route(mas_router::Index::route(), get(self::views::index::get))
389 .route(
390 mas_router::Login::route(),
391 get(self::views::login::get).post(self::views::login::post),
392 )
393 .route(mas_router::Logout::route(), post(self::views::logout::post))
394 .route(
395 mas_router::Register::route(),
396 get(self::views::register::get),
397 )
398 .route(
399 mas_router::PasswordRegister::route(),
400 get(self::views::register::password::get).post(self::views::register::password::post),
401 )
402 .route(
403 mas_router::RegisterVerifyEmail::route(),
404 get(self::views::register::steps::verify_email::get)
405 .post(self::views::register::steps::verify_email::post),
406 )
407 .route(
408 mas_router::RegisterToken::route(),
409 get(self::views::register::steps::registration_token::get)
410 .post(self::views::register::steps::registration_token::post),
411 )
412 .route(
413 mas_router::RegisterDisplayName::route(),
414 get(self::views::register::steps::display_name::get)
415 .post(self::views::register::steps::display_name::post),
416 )
417 .route(
418 mas_router::RegisterFinish::route(),
419 get(self::views::register::steps::finish::get),
420 )
421 .route(
422 mas_router::AccountRecoveryStart::route(),
423 get(self::views::recovery::start::get).post(self::views::recovery::start::post),
424 )
425 .route(
426 mas_router::AccountRecoveryProgress::route(),
427 get(self::views::recovery::progress::get).post(self::views::recovery::progress::post),
428 )
429 .route(
430 mas_router::OAuth2AuthorizationEndpoint::route(),
431 get(self::oauth2::authorization::get),
432 )
433 .route(
434 mas_router::Consent::route(),
435 get(self::oauth2::authorization::consent::get)
436 .post(self::oauth2::authorization::consent::post),
437 )
438 .route(
439 mas_router::CompatLoginSsoComplete::route(),
440 get(self::compat::login_sso_complete::get).post(self::compat::login_sso_complete::post),
441 )
442 .route(
443 mas_router::UpstreamOAuth2Authorize::route(),
444 get(self::upstream_oauth2::authorize::get),
445 )
446 .route(
447 mas_router::UpstreamOAuth2Callback::route(),
448 get(self::upstream_oauth2::callback::handler)
449 .post(self::upstream_oauth2::callback::handler),
450 )
451 .route(
452 mas_router::UpstreamOAuth2Link::route(),
453 get(self::upstream_oauth2::link::get).post(self::upstream_oauth2::link::post),
454 )
455 .route(
456 mas_router::UpstreamOAuth2BackchannelLogout::route(),
457 post(self::upstream_oauth2::backchannel_logout::post),
458 )
459 .route(
460 mas_router::DeviceCodeLink::route(),
461 get(self::oauth2::device::link::get).post(self::oauth2::device::link::post),
462 )
463 .route(
464 mas_router::DeviceCodeConsent::route(),
465 get(self::oauth2::device::consent::get).post(self::oauth2::device::consent::post),
466 )
467 .layer(AndThenLayer::new(
468 async move |response: axum::response::Response| {
469 Ok::<_, Infallible>(recover_error(&templates, response))
470 },
471 ))
472 .layer(SetResponseHeaderLayer::if_not_present(
473 X_FRAME_OPTIONS,
474 http::HeaderValue::from_static("DENY"),
475 ))
476}
477
478fn recover_error(
479 templates: &Templates,
480 response: axum::response::Response,
481) -> axum::response::Response {
482 let ext = response.extensions().get::<ErrorContext>();
484 if let Some(ctx) = ext
485 && let Ok(res) = templates.render_error(ctx)
486 {
487 let (mut parts, _original_body) = response.into_parts();
488 parts.headers.remove(CONTENT_TYPE);
489 parts.headers.remove(CONTENT_LENGTH);
490 return (parts, Html(res)).into_response();
491 }
492
493 response
494}
495
496pub async fn fallback(
502 State(templates): State<Templates>,
503 OriginalUri(uri): OriginalUri,
504 method: Method,
505 version: Version,
506 PreferredLanguage(locale): PreferredLanguage,
507) -> Result<impl IntoResponse, InternalError> {
508 let ctx = NotFoundContext::new(&method, version, &uri).with_language(locale);
509 let res = templates.render_not_found(&ctx)?;
512
513 Ok((StatusCode::NOT_FOUND, Html(res)))
514}