1use std::{borrow::Cow, fmt, marker::PhantomData};
6
7use bytes::BufMut;
8use ruma_common::{
9 api::{
10 EndpointError, OutgoingBody, OutgoingBodyJson, OutgoingResponse,
11 error::{Error as MatrixError, ErrorResponseBody, IntoHttpError, StandardErrorBody},
12 },
13 serde::StringEnum,
14};
15use serde::{Deserialize, Deserializer, Serialize, de};
16use serde_json::{from_slice as from_json_slice, value::RawValue as RawJsonValue};
17
18use crate::PrivOwnedStr;
19
20mod auth_data;
21mod auth_params;
22pub mod get_uiaa_fallback_page;
23
24pub use self::{auth_data::*, auth_params::*};
25
26#[doc = include_str!(concat!(env!("CARGO_MANIFEST_DIR"), "/src/doc/string_enum.md"))]
28#[derive(Clone, StringEnum)]
29#[non_exhaustive]
30pub enum AuthType {
31 #[ruma_enum(rename = "m.login.password")]
33 Password,
34
35 #[ruma_enum(rename = "m.login.recaptcha")]
37 ReCaptcha,
38
39 #[ruma_enum(rename = "m.login.email.identity")]
41 EmailIdentity,
42
43 #[ruma_enum(rename = "m.login.msisdn")]
45 Msisdn,
46
47 #[ruma_enum(rename = "m.login.sso")]
49 Sso,
50
51 #[ruma_enum(rename = "m.login.dummy")]
53 Dummy,
54
55 #[ruma_enum(rename = "m.login.registration_token")]
57 RegistrationToken,
58
59 #[ruma_enum(rename = "m.login.terms")]
63 Terms,
64
65 #[ruma_enum(rename = "m.oauth", alias = "org.matrix.cross_signing_reset")]
70 OAuth,
71
72 #[doc(hidden)]
73 _Custom(PrivOwnedStr),
74}
75
76#[derive(Clone, Debug, Deserialize, Serialize, OutgoingBodyJson)]
79#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
80pub struct UiaaInfo {
81 pub flows: Vec<AuthFlow>,
83
84 #[serde(default, skip_serializing_if = "Vec::is_empty")]
86 pub completed: Vec<AuthType>,
87
88 #[serde(skip_serializing_if = "Option::is_none")]
92 pub params: Option<Box<RawJsonValue>>,
93
94 #[serde(skip_serializing_if = "Option::is_none")]
96 pub session: Option<String>,
97
98 #[serde(flatten, skip_serializing_if = "Option::is_none")]
100 pub auth_error: Option<Box<StandardErrorBody>>,
101}
102
103impl UiaaInfo {
104 pub fn new(flows: Vec<AuthFlow>) -> Self {
106 Self { flows, completed: Vec::new(), params: None, session: None, auth_error: None }
107 }
108
109 pub fn params<'a, T: Deserialize<'a>>(
127 &'a self,
128 auth_type: &AuthType,
129 ) -> Result<Option<T>, serde_json::Error> {
130 struct AuthTypeVisitor<'b, T> {
131 auth_type: &'b AuthType,
132 _phantom: PhantomData<T>,
133 }
134
135 impl<'de, T> de::Visitor<'de> for AuthTypeVisitor<'_, T>
136 where
137 T: Deserialize<'de>,
138 {
139 type Value = Option<T>;
140
141 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
142 formatter.write_str("a key-value map")
143 }
144
145 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
146 where
147 A: de::MapAccess<'de>,
148 {
149 let mut params = None;
150
151 while let Some(key) = map.next_key::<Cow<'de, str>>()? {
152 if AuthType::from(key) == *self.auth_type {
153 params = Some(map.next_value()?);
154 } else {
155 map.next_value::<de::IgnoredAny>()?;
156 }
157 }
158
159 Ok(params)
160 }
161 }
162
163 let Some(params) = &self.params else {
164 return Ok(None);
165 };
166
167 let mut deserializer = serde_json::Deserializer::from_str(params.get());
168 deserializer.deserialize_map(AuthTypeVisitor { auth_type, _phantom: PhantomData })
169 }
170}
171
172#[derive(Clone, Debug, Default, Deserialize, Serialize)]
174#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
175pub struct AuthFlow {
176 #[serde(default, skip_serializing_if = "Vec::is_empty")]
178 pub stages: Vec<AuthType>,
179}
180
181impl AuthFlow {
182 pub fn new(stages: Vec<AuthType>) -> Self {
186 Self { stages }
187 }
188}
189
190#[derive(Clone, Debug)]
192#[allow(clippy::exhaustive_enums)]
193pub enum UiaaResponse {
194 AuthResponse(UiaaInfo),
196
197 MatrixError(MatrixError),
199}
200
201impl fmt::Display for UiaaResponse {
202 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
203 match self {
204 Self::AuthResponse(_) => write!(f, "User-Interactive Authentication required."),
205 Self::MatrixError(err) => write!(f, "{err}"),
206 }
207 }
208}
209
210impl From<MatrixError> for UiaaResponse {
211 fn from(error: MatrixError) -> Self {
212 Self::MatrixError(error)
213 }
214}
215
216impl EndpointError for UiaaResponse {
217 fn from_http_response(response: http::Response<&[u8]>) -> Self {
218 if response.status() == http::StatusCode::UNAUTHORIZED
219 && let Ok(uiaa_info) = from_json_slice(response.body())
220 {
221 return Self::AuthResponse(uiaa_info);
222 }
223
224 Self::MatrixError(MatrixError::from_http_response(response))
225 }
226}
227
228impl std::error::Error for UiaaResponse {}
229
230impl OutgoingResponse for UiaaResponse {
231 type Body = ResponseBody;
232
233 fn try_into_http_response_inner(self) -> Result<http::Response<Self::Body>, IntoHttpError> {
234 Ok(match self {
235 UiaaResponse::AuthResponse(authentication_info) => http::Response::builder()
236 .header(http::header::CONTENT_TYPE, ruma_common::http_headers::APPLICATION_JSON)
237 .status(http::StatusCode::UNAUTHORIZED)
238 .body(ResponseBody::AuthResponse(authentication_info))?,
239 UiaaResponse::MatrixError(error) => {
240 let (parts, body) = error.try_into_http_response_inner()?.into_parts();
241 http::Response::from_parts(parts, ResponseBody::MatrixError(body))
242 }
243 })
244 }
245}
246
247#[doc(hidden)]
248#[allow(clippy::exhaustive_enums)]
249pub enum ResponseBody {
250 AuthResponse(UiaaInfo),
251 MatrixError(ErrorResponseBody),
252}
253
254impl OutgoingBody for ResponseBody {
255 type Error = IntoHttpError;
256
257 fn try_into_buf<T: Default + BufMut + AsRef<[u8]>>(self) -> Result<T, Self::Error> {
258 Ok(match self {
259 Self::AuthResponse(uiaa_info) => uiaa_info.try_into_buf()?,
260 Self::MatrixError(error_response_body) => error_response_body.try_into_buf()?,
261 })
262 }
263}
264
265#[cfg(test)]
266mod tests {
267 use assert_matches2::{assert_let, assert_matches};
268 use ruma_common::serde::JsonObject;
269 use serde_json::{from_value as from_json_value, json};
270
271 use super::{AuthType, LoginTermsParams, OAuthParams, UiaaInfo};
272
273 #[test]
274 fn uiaa_info_params() {
275 let json = json!({
276 "flows": [{
277 "stages": ["m.login.terms", "m.login.email.identity", "local.custom.stage"],
278 }],
279 "params": {
280 "local.custom.stage": {
281 "foo": "bar",
282 },
283 "m.login.terms": {
284 "policies": {
285 "privacy": {
286 "en-US": {
287 "name": "Privacy Policy",
288 "url": "http://matrix.local/en-US/privacy",
289 },
290 "fr-FR": {
291 "name": "Politique de confidentialité",
292 "url": "http://matrix.local/fr-FR/privacy",
293 },
294 "version": "1",
295 },
296 },
297 }
298 },
299 "session": "abcdef",
300 });
301
302 let info = from_json_value::<UiaaInfo>(json).unwrap();
303
304 assert_matches!(info.params::<JsonObject>(&AuthType::EmailIdentity), Ok(None));
305 assert_matches!(
306 info.params::<JsonObject>(&AuthType::from("local.custom.stage")),
307 Ok(Some(_))
308 );
309
310 assert_let!(Ok(Some(params)) = info.params::<LoginTermsParams>(&AuthType::Terms));
311 assert_eq!(params.policies.len(), 1);
312
313 let policy = params.policies.get("privacy").unwrap();
314 assert_eq!(policy.version, "1");
315 assert_eq!(policy.translations.len(), 2);
316 let translation = policy.translations.get("en-US").unwrap();
317 assert_eq!(translation.name, "Privacy Policy");
318 assert_eq!(translation.url, "http://matrix.local/en-US/privacy");
319 let translation = policy.translations.get("fr-FR").unwrap();
320 assert_eq!(translation.name, "Politique de confidentialité");
321 assert_eq!(translation.url, "http://matrix.local/fr-FR/privacy");
322 }
323
324 #[test]
325 fn uiaa_info_oauth_params() {
326 let url = "http://auth.matrix.local/reset";
327 let stable_json = json!({
328 "flows": [{
329 "stages": ["m.oauth"],
330 }],
331 "params": {
332 "m.oauth": {
333 "url": url,
334 }
335 },
336 "session": "abcdef",
337 });
338 let unstable_json = json!({
339 "flows": [{
340 "stages": ["org.matrix.cross_signing_reset"],
341 }],
342 "params": {
343 "org.matrix.cross_signing_reset": {
344 "url": url,
345 }
346 },
347 "session": "abcdef",
348 });
349
350 let info = from_json_value::<UiaaInfo>(stable_json).unwrap();
351 assert_let!(Ok(Some(params)) = info.params::<OAuthParams>(&AuthType::OAuth));
352 assert_eq!(params.url, url);
353
354 let info = from_json_value::<UiaaInfo>(unstable_json).unwrap();
355 assert_let!(Ok(Some(params)) = info.params::<OAuthParams>(&AuthType::OAuth));
356 assert_eq!(params.url, url);
357 }
358}