mas_handlers/oauth2/authorization/
callback.rs1#![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 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 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 #[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 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 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 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}