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 .status(http::StatusCode::UNAUTHORIZED)
237 .body(ResponseBody::AuthResponse(authentication_info))?,
238 UiaaResponse::MatrixError(error) => {
239 let (parts, body) = error.try_into_http_response_inner()?.into_parts();
240 http::Response::from_parts(parts, ResponseBody::MatrixError(body))
241 }
242 })
243 }
244}
245
246#[doc(hidden)]
247#[allow(clippy::exhaustive_enums)]
248pub enum ResponseBody {
249 AuthResponse(UiaaInfo),
250 MatrixError(ErrorResponseBody),
251}
252
253impl OutgoingBody for ResponseBody {
254 type Error = IntoHttpError;
255
256 fn content_type(&self) -> Option<http::HeaderValue> {
257 Some(ruma_common::http_headers::APPLICATION_JSON)
258 }
259
260 fn try_into_buf<T: Default + BufMut + AsRef<[u8]>>(self) -> Result<T, Self::Error> {
261 Ok(match self {
262 Self::AuthResponse(uiaa_info) => uiaa_info.try_into_buf()?,
263 Self::MatrixError(error_response_body) => error_response_body.try_into_buf()?,
264 })
265 }
266}
267
268#[cfg(test)]
269mod tests {
270 use assert_matches::assert_matches;
271 use ruma_common::serde::JsonObject;
272 use serde_json::{from_value as from_json_value, json};
273 use strass::assert_let;
274
275 use super::{AuthType, LoginTermsParams, OAuthParams, UiaaInfo};
276
277 #[test]
278 fn uiaa_info_params() {
279 let json = json!({
280 "flows": [{
281 "stages": ["m.login.terms", "m.login.email.identity", "local.custom.stage"],
282 }],
283 "params": {
284 "local.custom.stage": {
285 "foo": "bar",
286 },
287 "m.login.terms": {
288 "policies": {
289 "privacy": {
290 "en-US": {
291 "name": "Privacy Policy",
292 "url": "http://matrix.local/en-US/privacy",
293 },
294 "fr-FR": {
295 "name": "Politique de confidentialité",
296 "url": "http://matrix.local/fr-FR/privacy",
297 },
298 "version": "1",
299 },
300 },
301 }
302 },
303 "session": "abcdef",
304 });
305
306 let info = from_json_value::<UiaaInfo>(json).unwrap();
307
308 assert_matches!(info.params::<JsonObject>(&AuthType::EmailIdentity), Ok(None));
309 assert_matches!(
310 info.params::<JsonObject>(&AuthType::from("local.custom.stage")),
311 Ok(Some(_))
312 );
313
314 assert_let!(Ok(Some(params)) = info.params::<LoginTermsParams>(&AuthType::Terms));
315 assert_eq!(params.policies.len(), 1);
316
317 let policy = params.policies.get("privacy").unwrap();
318 assert_eq!(policy.version, "1");
319 assert_eq!(policy.translations.len(), 2);
320 let translation = policy.translations.get("en-US").unwrap();
321 assert_eq!(translation.name, "Privacy Policy");
322 assert_eq!(translation.url, "http://matrix.local/en-US/privacy");
323 let translation = policy.translations.get("fr-FR").unwrap();
324 assert_eq!(translation.name, "Politique de confidentialité");
325 assert_eq!(translation.url, "http://matrix.local/fr-FR/privacy");
326 }
327
328 #[test]
329 fn uiaa_info_oauth_params() {
330 let url = "http://auth.matrix.local/reset";
331 let stable_json = json!({
332 "flows": [{
333 "stages": ["m.oauth"],
334 }],
335 "params": {
336 "m.oauth": {
337 "url": url,
338 }
339 },
340 "session": "abcdef",
341 });
342 let unstable_json = json!({
343 "flows": [{
344 "stages": ["org.matrix.cross_signing_reset"],
345 }],
346 "params": {
347 "org.matrix.cross_signing_reset": {
348 "url": url,
349 }
350 },
351 "session": "abcdef",
352 });
353
354 let info = from_json_value::<UiaaInfo>(stable_json).unwrap();
355 assert_let!(Ok(Some(params)) = info.params::<OAuthParams>(&AuthType::OAuth));
356 assert_eq!(params.url, url);
357
358 let info = from_json_value::<UiaaInfo>(unstable_json).unwrap();
359 assert_let!(Ok(Some(params)) = info.params::<OAuthParams>(&AuthType::OAuth));
360 assert_eq!(params.url, url);
361 }
362}