Skip to main content

mas_handlers/oauth2/authorization/
callback.rs

1// Copyright 2024, 2025 New Vector Ltd.
2// Copyright 2022-2024 The Matrix.org Foundation C.I.C.
3//
4// SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
5// Please see LICENSE files in the repository root for full details.
6
7#![allow(clippy::module_name_repetitions)]
8
9use std::collections::HashMap;
10
11use axum::response::{Html, IntoResponse, Redirect, Response};
12use mas_data_model::AuthorizationGrant;
13use mas_i18n::DataLocale;
14use mas_templates::{FormPostContext, Templates};
15use oauth2_types::requests::ResponseMode;
16use serde::Serialize;
17use thiserror::Error;
18use url::Url;
19
20#[derive(Debug, Clone)]
21enum CallbackDestinationMode {
22    Query {
23        existing_params: HashMap<String, String>,
24    },
25    Fragment,
26    FormPost,
27}
28
29#[derive(Debug, Clone)]
30pub struct CallbackDestination {
31    mode: CallbackDestinationMode,
32    safe_redirect_uri: Url,
33    state: Option<String>,
34}
35
36#[derive(Debug, Error)]
37pub enum IntoCallbackDestinationError {
38    #[error("Redirect URI can't have a fragment")]
39    RedirectUriFragmentNotAllowed,
40
41    #[error("Existing query parameters are not valid")]
42    RedirectUriInvalidQueryParams(#[from] serde_urlencoded::de::Error),
43
44    #[error("Requested response_mode is not supported")]
45    UnsupportedResponseMode,
46}
47
48#[derive(Debug, Error)]
49pub enum CallbackDestinationError {
50    #[error("Failed to render the form_post template")]
51    FormPostRender(#[from] mas_templates::TemplateError),
52
53    #[error("Failed to serialize parameters query string")]
54    ParamsSerialization(#[from] serde_urlencoded::ser::Error),
55}
56
57impl TryFrom<&AuthorizationGrant> for CallbackDestination {
58    type Error = IntoCallbackDestinationError;
59
60    fn try_from(value: &AuthorizationGrant) -> Result<Self, Self::Error> {
61        Self::try_new(
62            &value.response_mode,
63            value.redirect_uri.clone(),
64            value.state.clone(),
65        )
66    }
67}
68
69impl CallbackDestination {
70    pub fn try_new(
71        mode: &ResponseMode,
72        mut redirect_uri: Url,
73        state: Option<String>,
74    ) -> Result<Self, IntoCallbackDestinationError> {
75        if redirect_uri.fragment().is_some() {
76            return Err(IntoCallbackDestinationError::RedirectUriFragmentNotAllowed);
77        }
78
79        let mode = match mode {
80            ResponseMode::Query => {
81                let existing_params = redirect_uri
82                    .query()
83                    .map(serde_urlencoded::from_str)
84                    .transpose()?
85                    .unwrap_or_default();
86
87                // Remove the query from the URL
88                redirect_uri.set_query(None);
89
90                CallbackDestinationMode::Query { existing_params }
91            }
92            ResponseMode::Fragment => CallbackDestinationMode::Fragment,
93            ResponseMode::FormPost => CallbackDestinationMode::FormPost,
94            _ => return Err(IntoCallbackDestinationError::UnsupportedResponseMode),
95        };
96
97        Ok(Self {
98            mode,
99            safe_redirect_uri: redirect_uri,
100            state,
101        })
102    }
103
104    pub fn go<T: Serialize + Send + Sync>(
105        self,
106        templates: &Templates,
107        locale: &DataLocale,
108        params: T,
109    ) -> Result<Response, CallbackDestinationError> {
110        #[derive(Serialize)]
111        struct AllParams<'s, T> {
112            #[serde(flatten, skip_serializing_if = "Option::is_none")]
113            existing: Option<&'s HashMap<String, String>>,
114
115            #[serde(skip_serializing_if = "Option::is_none")]
116            state: Option<String>,
117
118            #[serde(flatten)]
119            params: T,
120        }
121
122        let mut redirect_uri = self.safe_redirect_uri;
123        let state = self.state;
124
125        match self.mode {
126            CallbackDestinationMode::Query { existing_params } => {
127                let merged = AllParams {
128                    existing: Some(&existing_params),
129                    state,
130                    params,
131                };
132
133                let new_qs = serde_urlencoded::to_string(merged)?;
134
135                redirect_uri.set_query(Some(&new_qs));
136                if redirect_uri.fragment().is_none() {
137                    // Ensure that the Location header (redirect target)
138                    // includes a URL fragment (#) of some sort.
139                    //
140                    // Any fragment present in the Location header URL that the server redirects to
141                    // (e.g., via a 303 response) will overwrite the client’s existing fragment,
142                    // otherwise the fragment will be preserved across the
143                    // redirect (and may contain sensitive information,
144                    // or confuse the downstream client).
145                    //
146                    // If the redirect_uri already contains a fragment, that fragment will do the
147                    // same job, so we leave it alone — we don't want to mangle the client's
148                    // configured redirect URL by replacing it with a blank fragment.
149                    // Otherwise, set a fragment of empty string (effectively appending `#` to the
150                    // URL).
151                    //
152                    // Browser behaviour is documented as part of the 'location URL' algorithm at
153                    // https://fetch.spec.whatwg.org/commit-snapshots/809904366f33a673a9489b81155ee9e3edd29c12#concept-response-location-url
154                    redirect_uri.set_fragment(Some(""));
155                }
156
157                Ok(Redirect::to(redirect_uri.as_str()).into_response())
158            }
159
160            CallbackDestinationMode::Fragment => {
161                let merged = AllParams {
162                    existing: None,
163                    state,
164                    params,
165                };
166
167                let new_qs = serde_urlencoded::to_string(merged)?;
168
169                redirect_uri.set_fragment(Some(&new_qs));
170
171                Ok(Redirect::to(redirect_uri.as_str()).into_response())
172            }
173
174            CallbackDestinationMode::FormPost => {
175                let merged = AllParams {
176                    existing: None,
177                    state,
178                    params,
179                };
180                let ctx = FormPostContext::new_for_url(redirect_uri, merged).with_language(locale);
181                let rendered = templates.render_form_post(&ctx)?;
182                Ok(Html(rendered).into_response())
183            }
184        }
185    }
186}
187
188#[cfg(test)]
189mod tests {
190    use hyper::{Request, StatusCode};
191    use mas_router::SimpleRoute;
192    use oauth2_types::registration::ClientRegistrationResponse;
193    use sqlx::PgPool;
194
195    use crate::test_utils::{RequestBuilderExt, ResponseExt, TestState, setup};
196
197    /// Test that checks the content of the `Location` header
198    /// in response to an authorization request.
199    ///
200    /// Specifically, we expect to see an empty fragment (`#`)
201    /// at the end of the URL in order to overwrite any fragment
202    /// that the browser might otherwise preserve across the redirect.
203    #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
204    async fn test_query_mode_location_header(pool: PgPool) {
205        setup();
206        let state = TestState::from_pool(pool).await.unwrap();
207
208        // Register an OAuth2 client
209        let request =
210            Request::post(mas_router::OAuth2RegistrationEndpoint::PATH).json(serde_json::json!({
211                "client_uri": "https://example.com/",
212                "redirect_uris": ["https://example.com/callback"],
213                "token_endpoint_auth_method": "none",
214                "response_types": ["code"],
215                "grant_types": ["authorization_code"],
216            }));
217
218        let response = state.request(request).await;
219        response.assert_status(StatusCode::CREATED);
220
221        let registration: ClientRegistrationResponse = response.json();
222        let client_id = registration.client_id;
223
224        // Send an authorization request with response_mode=query and prompt=none.
225        // prompt=none always fails with login_required since there is no session,
226        // which exercises the CallbackDestinationMode::Query path.
227
228        // Build /authorize query parameters
229        let query = url::form_urlencoded::Serializer::new(String::new())
230            .append_pair("response_type", "code")
231            .append_pair("client_id", &client_id)
232            .append_pair("redirect_uri", "https://example.com/callback")
233            .append_pair("scope", "openid")
234            .append_pair("state", "test-state-value")
235            .append_pair("response_mode", "query")
236            .append_pair("prompt", "none")
237            .finish();
238
239        let response = state
240            .request(Request::get(format!("https://example.com/authorize?{query}")).empty())
241            .await;
242
243        response.assert_status(StatusCode::SEE_OTHER);
244
245        // Check the form of the Location redirect
246        response.assert_header_value(
247            hyper::header::LOCATION,
248            "https://example.com/callback?state=test-state-value&error=login_required&error_description=The+Authorization+Server+requires+End-User+authentication.#",
249        );
250    }
251}