1use std::collections::{BTreeSet, HashMap};
20
21pub use language_tags;
22use language_tags::LanguageTag;
23use matrix_sdk_base::deserialized_responses::PrivOwnedStr;
24use oauth2::{AsyncHttpClient, ClientId, HttpClientError, RequestTokenError};
25use ruma::{
26 SecondsSinceUnixEpoch,
27 api::client::discovery::get_authorization_server_metadata::v1::{GrantType, ResponseType},
28 serde::{Raw, StringEnum},
29};
30use serde::{Deserialize, Serialize, ser::SerializeMap};
31use url::Url;
32
33use super::{
34 OAuthHttpClient,
35 error::OAuthClientRegistrationError,
36 http_client::{check_http_response_json_content_type, check_http_response_status_code},
37};
38
39#[tracing::instrument(skip_all, fields(registration_endpoint))]
53pub(super) async fn register_client(
54 http_client: &OAuthHttpClient,
55 registration_endpoint: &Url,
56 client_metadata: &Raw<ClientMetadata>,
57) -> Result<ClientRegistrationResponse, OAuthClientRegistrationError> {
58 tracing::debug!("Registering client...");
59
60 let body =
61 serde_json::to_vec(client_metadata).map_err(OAuthClientRegistrationError::IntoJson)?;
62 let request = http::Request::post(registration_endpoint.as_str())
63 .header(http::header::CONTENT_TYPE, mime::APPLICATION_JSON.to_string())
64 .body(body)
65 .map_err(|err| RequestTokenError::Request(HttpClientError::Http(err)))?;
66
67 let response = http_client.call(request).await.map_err(RequestTokenError::Request)?;
68
69 check_http_response_status_code(&response)?;
70 check_http_response_json_content_type(&response)?;
71
72 let response = serde_json::from_slice(&response.into_body())
73 .map_err(OAuthClientRegistrationError::FromJson)?;
74
75 Ok(response)
76}
77
78#[derive(Debug, Clone, Deserialize)]
82pub struct ClientRegistrationResponse {
83 pub client_id: ClientId,
85
86 pub client_id_issued_at: Option<SecondsSinceUnixEpoch>,
88}
89
90#[derive(Debug, Clone, Serialize)]
103#[serde(into = "ClientMetadataSerializeHelper")]
104pub struct ClientMetadata {
105 pub application_type: ApplicationType,
107
108 pub grant_types: Vec<OAuthGrantType>,
112
113 pub client_uri: Localized<Url>,
115
116 pub client_name: Option<Localized<String>>,
118
119 pub logo_uri: Option<Localized<Url>>,
121
122 pub policy_uri: Option<Localized<Url>>,
125
126 pub tos_uri: Option<Localized<Url>>,
129}
130
131impl ClientMetadata {
132 pub fn new(
134 application_type: ApplicationType,
135 grant_types: Vec<OAuthGrantType>,
136 client_uri: Localized<Url>,
137 ) -> Self {
138 Self {
139 application_type,
140 grant_types,
141 client_uri,
142 client_name: None,
143 logo_uri: None,
144 policy_uri: None,
145 tos_uri: None,
146 }
147 }
148}
149
150#[derive(Debug, Clone)]
156#[non_exhaustive]
157pub enum OAuthGrantType {
158 AuthorizationCode {
165 redirect_uris: Vec<Url>,
167 },
168
169 DeviceCode,
176}
177
178#[derive(Clone, StringEnum)]
180#[ruma_enum(rename_all = "lowercase")]
181#[non_exhaustive]
182pub enum ApplicationType {
183 Web,
188
189 Native,
193
194 #[doc(hidden)]
195 _Custom(PrivOwnedStr),
196}
197
198#[derive(Debug, Clone, PartialEq, Eq)]
202pub struct Localized<T> {
203 non_localized: T,
204 localized: HashMap<LanguageTag, T>,
205}
206
207impl<T> Localized<T> {
208 pub fn new(non_localized: T, localized: impl IntoIterator<Item = (LanguageTag, T)>) -> Self {
211 Self { non_localized, localized: localized.into_iter().collect() }
212 }
213
214 pub fn non_localized(&self) -> &T {
216 &self.non_localized
217 }
218
219 pub fn get(&self, language: Option<&LanguageTag>) -> Option<&T> {
221 match language {
222 Some(lang) => self.localized.get(lang),
223 None => Some(&self.non_localized),
224 }
225 }
226}
227
228impl<T> From<(T, HashMap<LanguageTag, T>)> for Localized<T> {
229 fn from(t: (T, HashMap<LanguageTag, T>)) -> Self {
230 Localized { non_localized: t.0, localized: t.1 }
231 }
232}
233
234#[derive(Serialize)]
235struct ClientMetadataSerializeHelper {
236 #[serde(skip_serializing_if = "Vec::is_empty")]
237 redirect_uris: Vec<Url>,
238 token_endpoint_auth_method: &'static str,
239 grant_types: BTreeSet<GrantType>,
240 #[serde(skip_serializing_if = "Vec::is_empty")]
241 response_types: Vec<ResponseType>,
242 application_type: ApplicationType,
243 #[serde(flatten)]
244 localized: ClientMetadataLocalizedFields,
245}
246
247impl From<ClientMetadata> for ClientMetadataSerializeHelper {
248 fn from(value: ClientMetadata) -> Self {
249 let ClientMetadata {
250 application_type,
251 grant_types: oauth_grant_types,
252 client_uri,
253 client_name,
254 logo_uri,
255 policy_uri,
256 tos_uri,
257 } = value;
258
259 let mut redirect_uris = None;
260 let mut response_types = None;
261 let mut grant_types = BTreeSet::new();
262
263 grant_types.insert(GrantType::RefreshToken);
265
266 for oauth_grant_type in oauth_grant_types {
267 match oauth_grant_type {
268 OAuthGrantType::AuthorizationCode { redirect_uris: uris } => {
269 redirect_uris = Some(uris);
270 response_types = Some(vec![ResponseType::Code]);
271 grant_types.insert(GrantType::AuthorizationCode);
272 }
273 OAuthGrantType::DeviceCode => {
274 grant_types.insert(GrantType::DeviceCode);
275 }
276 }
277 }
278
279 ClientMetadataSerializeHelper {
280 redirect_uris: redirect_uris.unwrap_or_default(),
281 token_endpoint_auth_method: "none",
283 grant_types,
284 response_types: response_types.unwrap_or_default(),
285 application_type,
286 localized: ClientMetadataLocalizedFields {
287 client_uri,
288 client_name,
289 logo_uri,
290 policy_uri,
291 tos_uri,
292 },
293 }
294 }
295}
296
297struct ClientMetadataLocalizedFields {
302 client_uri: Localized<Url>,
303 client_name: Option<Localized<String>>,
304 logo_uri: Option<Localized<Url>>,
305 policy_uri: Option<Localized<Url>>,
306 tos_uri: Option<Localized<Url>>,
307}
308
309impl Serialize for ClientMetadataLocalizedFields {
310 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
311 where
312 S: serde::Serializer,
313 {
314 fn serialize_localized_into_map<M: SerializeMap, T: Serialize>(
315 map: &mut M,
316 field_name: &str,
317 value: &Localized<T>,
318 ) -> Result<(), M::Error> {
319 map.serialize_entry(field_name, &value.non_localized)?;
320
321 for (lang, localized) in &value.localized {
322 map.serialize_entry(&format!("{field_name}#{lang}"), localized)?;
323 }
324
325 Ok(())
326 }
327
328 let mut map = serializer.serialize_map(None)?;
329
330 serialize_localized_into_map(&mut map, "client_uri", &self.client_uri)?;
331
332 if let Some(client_name) = &self.client_name {
333 serialize_localized_into_map(&mut map, "client_name", client_name)?;
334 }
335
336 if let Some(logo_uri) = &self.logo_uri {
337 serialize_localized_into_map(&mut map, "logo_uri", logo_uri)?;
338 }
339
340 if let Some(policy_uri) = &self.policy_uri {
341 serialize_localized_into_map(&mut map, "policy_uri", policy_uri)?;
342 }
343
344 if let Some(tos_uri) = &self.tos_uri {
345 serialize_localized_into_map(&mut map, "tos_uri", tos_uri)?;
346 }
347
348 map.end()
349 }
350}
351
352#[cfg(test)]
353mod tests {
354 use language_tags::LanguageTag;
355 use serde_json::json;
356 use url::Url;
357
358 use super::{ApplicationType, ClientMetadata, Localized, OAuthGrantType};
359
360 #[test]
361 fn test_serialize_minimal_client_metadata() {
362 let metadata = ClientMetadata::new(
363 ApplicationType::Native,
364 vec![OAuthGrantType::AuthorizationCode {
365 redirect_uris: vec![Url::parse("http://127.0.0.1/").unwrap()],
366 }],
367 Localized::new(
368 Url::parse("https://github.com/matrix-org/matrix-rust-sdk").unwrap(),
369 [],
370 ),
371 );
372
373 assert_eq!(
374 serde_json::to_value(metadata).unwrap(),
375 json!({
376 "application_type": "native",
377 "grant_types": ["authorization_code", "refresh_token"],
378 "response_types": ["code"],
379 "token_endpoint_auth_method": "none",
380 "redirect_uris": ["http://127.0.0.1/"],
381 "client_uri": "https://github.com/matrix-org/matrix-rust-sdk",
382 }),
383 );
384 }
385
386 #[test]
387 fn test_serialize_full_client_metadata() {
388 let lang_fr = LanguageTag::parse("fr").unwrap();
389 let lang_mas = LanguageTag::parse("mas").unwrap();
390
391 let mut metadata = ClientMetadata::new(
392 ApplicationType::Web,
393 vec![
394 OAuthGrantType::AuthorizationCode {
395 redirect_uris: vec![
396 Url::parse("http://127.0.0.1/").unwrap(),
397 Url::parse("http://[::1]/").unwrap(),
398 ],
399 },
400 OAuthGrantType::DeviceCode,
401 ],
402 Localized::new(
403 Url::parse("https://example.org/matrix-client").unwrap(),
404 [
405 (lang_fr.clone(), Url::parse("https://example.org/fr/matrix-client").unwrap()),
406 (
407 lang_mas.clone(),
408 Url::parse("https://example.org/mas/matrix-client").unwrap(),
409 ),
410 ],
411 ),
412 );
413
414 metadata.client_name = Some(Localized::new(
415 "My Matrix client".to_owned(),
416 [(lang_fr.clone(), "Mon client Matrix".to_owned())],
417 ));
418 metadata.logo_uri =
419 Some(Localized::new(Url::parse("https://example.org/logo.svg").unwrap(), []));
420 metadata.policy_uri = Some(Localized::new(
421 Url::parse("https://example.org/policy").unwrap(),
422 [
423 (lang_fr.clone(), Url::parse("https://example.org/fr/policy").unwrap()),
424 (lang_mas.clone(), Url::parse("https://example.org/mas/policy").unwrap()),
425 ],
426 ));
427 metadata.tos_uri = Some(Localized::new(
428 Url::parse("https://example.org/tos").unwrap(),
429 [
430 (lang_fr, Url::parse("https://example.org/fr/tos").unwrap()),
431 (lang_mas, Url::parse("https://example.org/mas/tos").unwrap()),
432 ],
433 ));
434
435 assert_eq!(
436 serde_json::to_value(metadata).unwrap(),
437 json!({
438 "application_type": "web",
439 "grant_types": [
440 "authorization_code",
441 "refresh_token",
442 "urn:ietf:params:oauth:grant-type:device_code",
443 ],
444 "response_types": ["code"],
445 "token_endpoint_auth_method": "none",
446 "redirect_uris": ["http://127.0.0.1/", "http://[::1]/"],
447 "client_uri": "https://example.org/matrix-client",
448 "client_uri#fr": "https://example.org/fr/matrix-client",
449 "client_uri#mas": "https://example.org/mas/matrix-client",
450 "client_name": "My Matrix client",
451 "client_name#fr": "Mon client Matrix",
452 "logo_uri": "https://example.org/logo.svg",
453 "policy_uri": "https://example.org/policy",
454 "policy_uri#fr": "https://example.org/fr/policy",
455 "policy_uri#mas": "https://example.org/mas/policy",
456 "tos_uri": "https://example.org/tos",
457 "tos_uri#fr": "https://example.org/fr/tos",
458 "tos_uri#mas": "https://example.org/mas/tos",
459 }),
460 );
461 }
462}