Skip to main content

mas_handlers/
lib.rs

1// Copyright 2025, 2026 Element Creations Ltd.
2// Copyright 2024, 2025 New Vector Ltd.
3// Copyright 2021-2024 The Matrix.org Foundation C.I.C.
4//
5// SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
6// Please see LICENSE files in the repository root for full details.
7
8#![deny(clippy::future_not_send)]
9#![allow(
10    // Some axum handlers need that
11    clippy::unused_async,
12    // Because of how axum handlers work, we sometime have take many arguments
13    clippy::too_many_arguments,
14    // Code generated by tracing::instrument trigger this when returning an `impl Trait`
15    // See https://github.com/tokio-rs/tracing/issues/2613
16    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/// Implement `From<E>` for `RouteError`, for "internal server error" kind of
89/// errors.
90#[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        // Pass the undocumented_oauth2_access parameter through the request extension, as it is
146        // per-listener
147        .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    // All those routes are API-like, with a common CORS layer
215    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                    // Swagger will send this header, so we have to allow it to avoid CORS errors
255                    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    // A sub-router for human-facing routes with error handling
278    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    // A sub-router for API-facing routes with CORS
298    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        // XXX: hard-coded redirect from /account to /account/
357        .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    // Error responses should have an ErrorContext attached to them
483    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
496/// The fallback handler for all routes that don't match anything else.
497///
498/// # Errors
499///
500/// Returns an error if the template rendering fails.
501pub 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    // XXX: this should look at the Accept header and return JSON if requested
510
511    let res = templates.render_not_found(&ctx)?;
512
513    Ok((StatusCode::NOT_FOUND, Html(res)))
514}