1pub mod error;
4
5use std::collections::BTreeMap;
6
7use reqwest::{StatusCode, Url};
8use reqwest_middleware::ClientWithMiddleware;
9use serde::{Deserialize, Serialize};
10
11use self::error::{Error, FailError, UnexpectedStatus};
12use crate::cache;
13
14#[derive(Debug, Clone, Serialize, Deserialize)]
16pub struct ClientWellKnown {
17 #[serde(rename = "m.homeserver")]
19 pub homeserver: HomeserverInfo,
20
21 #[serde(rename = "m.identity_server")]
23 pub identity_server: Option<IdentityServerInfo>,
24}
25
26#[derive(Debug, Clone, Serialize, Deserialize)]
28pub struct HomeserverInfo {
29 base_url: String,
31}
32
33#[derive(Debug, Clone, Serialize, Deserialize)]
35pub struct IdentityServerInfo {
36 base_url: String,
38}
39
40#[derive(Clone, Debug)]
42pub struct Resolver {
43 http: ClientWithMiddleware,
46}
47
48#[allow(dead_code)]
51#[derive(Deserialize)]
52struct Versions {
53 pub versions: Vec<String>,
55 #[serde(default)]
57 pub unstable_features: BTreeMap<String, bool>,
58}
59
60impl Resolver {
61 #[must_use]
63 pub fn new() -> Self {
64 Self {
65 http: reqwest_middleware::ClientBuilder::new(reqwest::Client::new())
66 .with(cache())
67 .build(),
68 }
69 }
70
71 #[must_use]
73 pub fn with(http: reqwest::Client) -> Self {
74 Self { http: reqwest_middleware::ClientBuilder::new(http).with(cache()).build() }
75 }
76
77 pub async fn resolve(&self, name: &str) -> Result<Url, Error> {
79 #[cfg(not(test))]
80 let url = Url::parse(&format!("https://{}", name))?;
81 #[cfg(test)]
82 let url = Url::parse(&format!("http://{}", name))?;
83
84 let response = self.http.get(url.join(".well-known/matrix/client")?).send().await?;
86 if response.status() == StatusCode::NOT_FOUND {
88 return Ok(url);
89 };
90 if response.status() != StatusCode::OK {
92 return Err(UnexpectedStatus(response.status()).into());
93 }
94 let well_known = response.json::<ClientWellKnown>().await?;
96 let url = Url::parse(&well_known.homeserver.base_url)?;
98 let versions = self
100 .http
101 .get(url.join("_matrix/client/versions")?)
102 .send()
103 .await
104 .map_err(FailError::Http)?;
105 if versions.status() != StatusCode::OK {
106 return Err(Error::Fail(UnexpectedStatus(versions.status()).into()));
107 }
108 versions.json::<Versions>().await.map_err(|e| FailError::Http(e.into()))?;
109
110 if let Some(identity) = well_known.identity_server {
112 let url = Url::parse(&identity.base_url)?;
113 let result: Result<_, FailError> = async {
114 let response = self.http.get(url.join("_matrix/identity/v2")?).send().await?;
115 if response.status() != StatusCode::OK {
116 return Err(UnexpectedStatus(response.status()).into());
117 }
118 Ok(())
119 }
120 .await;
121 result?;
122 }
123
124 Ok(url)
125 }
126}
127
128impl Default for Resolver {
129 fn default() -> Self {
130 Self { http: ClientWithMiddleware::from(reqwest::Client::new()) }
131 }
132}
133
134#[cfg(test)]
135mod tests {
136 use wiremock::{
137 Mock, MockServer, ResponseTemplate,
138 matchers::{method, path},
139 };
140
141 use super::{Resolver, error::Error};
142
143 #[tokio::test]
145 async fn not_found() -> Result<(), Box<dyn std::error::Error>> {
146 let mock_server = MockServer::start().await;
147
148 Mock::given(method("GET"))
149 .and(path("/.well-known/matrix/client"))
150 .respond_with(ResponseTemplate::new(404))
151 .expect(1)
152 .mount(&mock_server)
153 .await;
154
155 let http =
156 reqwest::Client::builder().resolve("example.test", *mock_server.address()).build()?;
157 let resolver = Resolver::with(http);
158 let url =
159 resolver.resolve(&format!("example.test:{}", mock_server.address().port())).await?;
160
161 assert_eq!(
162 format!("http://example.test:{}/", mock_server.address().port()),
163 url.to_string()
164 );
165 Ok(())
166 }
167
168 #[tokio::test]
169 async fn resolve() -> Result<(), Box<dyn std::error::Error>> {
170 let mock_server = MockServer::start().await;
171
172 let port = mock_server.address().port();
173
174 Mock::given(method("GET"))
175 .and(path("/.well-known/matrix/client"))
176 .respond_with(ResponseTemplate::new(200).set_body_raw(
177 format!(
178 r#"{{"m.homeserver": {{"base_url": "http://destination.test:{}"}} }}"#,
179 port
180 ),
181 "application/json",
182 ))
183 .expect(1)
184 .mount(&mock_server)
185 .await;
186
187 Mock::given(method("GET"))
188 .and(path("/_matrix/client/versions"))
189 .respond_with(
190 ResponseTemplate::new(200)
191 .set_body_raw(r#"{"versions":["r0.0.1"]}"#, "application/json"),
192 )
193 .expect(1)
194 .mount(&mock_server)
195 .await;
196
197 let http = reqwest::Client::builder()
198 .resolve("example.test", *mock_server.address())
199 .resolve("destination.test", *mock_server.address())
200 .build()?;
201 let resolver = Resolver::with(http);
202
203 let url =
204 resolver.resolve(&format!("example.test:{}", mock_server.address().port())).await?;
205
206 assert_eq!(url.to_string(), format!("http://destination.test:{}/", port));
207 Ok(())
208 }
209
210 #[tokio::test]
213 async fn resolve_with_identity_server() -> Result<(), Box<dyn std::error::Error>> {
214 let mock_server = MockServer::start().await;
215 let port = mock_server.address().port();
216
217 Mock::given(method("GET"))
218 .and(path("/.well-known/matrix/client"))
219 .respond_with(ResponseTemplate::new(200).set_body_raw(
220 format!(
221 r#"{{"m.homeserver":{{"base_url":"http://destination.test:{port}"}},"m.identity_server":{{"base_url":"http://destination.test:{port}"}}}}"#
222 ),
223 "application/json",
224 ))
225 .mount(&mock_server)
226 .await;
227 Mock::given(method("GET"))
228 .and(path("/_matrix/client/versions"))
229 .respond_with(
230 ResponseTemplate::new(200)
231 .set_body_raw(r#"{"versions":["v1.11"]}"#, "application/json"),
232 )
233 .mount(&mock_server)
234 .await;
235 Mock::given(method("GET"))
237 .and(path("/_matrix/identity/v2"))
238 .respond_with(ResponseTemplate::new(200).set_body_raw("{}", "application/json"))
239 .expect(1)
240 .mount(&mock_server)
241 .await;
242
243 let http = reqwest::Client::builder()
244 .resolve("example.test", *mock_server.address())
245 .resolve("destination.test", *mock_server.address())
246 .build()?;
247 let resolver = Resolver::with(http);
248 let url = resolver.resolve(&format!("example.test:{port}")).await?;
249 assert_eq!(url.to_string(), format!("http://destination.test:{port}/"));
250 Ok(())
251 }
252
253 #[tokio::test]
256 async fn well_known_non_200_fails_prompt() -> Result<(), Box<dyn std::error::Error>> {
257 let mock_server = MockServer::start().await;
258 let port = mock_server.address().port();
259
260 Mock::given(method("GET"))
261 .and(path("/.well-known/matrix/client"))
262 .respond_with(ResponseTemplate::new(500).set_body_raw(
263 format!(r#"{{"m.homeserver":{{"base_url":"http://destination.test:{port}"}}}}"#),
264 "application/json",
265 ))
266 .expect(1)
267 .mount(&mock_server)
268 .await;
269
270 let http = reqwest::Client::builder()
271 .resolve("example.test", *mock_server.address())
272 .resolve("destination.test", *mock_server.address())
273 .build()?;
274 let resolver = Resolver::with(http);
275 let result = resolver.resolve(&format!("example.test:{port}")).await;
276
277 assert!(matches!(result, Err(Error::Prompt(_))), "expected FAIL_PROMPT, got {result:?}");
278 Ok(())
279 }
280
281 #[tokio::test]
284 async fn versions_non_200_fails_error() -> Result<(), Box<dyn std::error::Error>> {
285 let mock_server = MockServer::start().await;
286 let port = mock_server.address().port();
287
288 Mock::given(method("GET"))
289 .and(path("/.well-known/matrix/client"))
290 .respond_with(ResponseTemplate::new(200).set_body_raw(
291 format!(r#"{{"m.homeserver":{{"base_url":"http://destination.test:{port}"}}}}"#),
292 "application/json",
293 ))
294 .mount(&mock_server)
295 .await;
296 Mock::given(method("GET"))
297 .and(path("/_matrix/client/versions"))
298 .respond_with(
299 ResponseTemplate::new(500)
300 .set_body_raw(r#"{"versions":["v1.11"]}"#, "application/json"),
301 )
302 .expect(1)
303 .mount(&mock_server)
304 .await;
305
306 let http = reqwest::Client::builder()
307 .resolve("example.test", *mock_server.address())
308 .resolve("destination.test", *mock_server.address())
309 .build()?;
310 let resolver = Resolver::with(http);
311 let result = resolver.resolve(&format!("example.test:{port}")).await;
312
313 assert!(matches!(result, Err(Error::Fail(_))), "expected FAIL_ERROR, got {result:?}");
314 Ok(())
315 }
316
317 #[tokio::test]
321 async fn well_known_2xx_non_200_fails_prompt() -> Result<(), Box<dyn std::error::Error>> {
322 let mock_server = MockServer::start().await;
323 let port = mock_server.address().port();
324
325 Mock::given(method("GET"))
326 .and(path("/.well-known/matrix/client"))
327 .respond_with(ResponseTemplate::new(201).set_body_raw(
328 format!(r#"{{"m.homeserver":{{"base_url":"http://destination.test:{port}"}}}}"#),
329 "application/json",
330 ))
331 .expect(1)
332 .mount(&mock_server)
333 .await;
334 Mock::given(method("GET"))
335 .and(path("/_matrix/client/versions"))
336 .respond_with(
337 ResponseTemplate::new(200)
338 .set_body_raw(r#"{"versions":["v1.11"]}"#, "application/json"),
339 )
340 .mount(&mock_server)
341 .await;
342
343 let http = reqwest::Client::builder()
344 .resolve("example.test", *mock_server.address())
345 .resolve("destination.test", *mock_server.address())
346 .build()?;
347 let resolver = Resolver::with(http);
348 let result = resolver.resolve(&format!("example.test:{port}")).await;
349
350 assert!(matches!(result, Err(Error::Prompt(_))), "expected FAIL_PROMPT, got {result:?}");
351 Ok(())
352 }
353
354 #[tokio::test]
358 async fn versions_2xx_non_200_fails_error() -> Result<(), Box<dyn std::error::Error>> {
359 let mock_server = MockServer::start().await;
360 let port = mock_server.address().port();
361
362 Mock::given(method("GET"))
363 .and(path("/.well-known/matrix/client"))
364 .respond_with(ResponseTemplate::new(200).set_body_raw(
365 format!(r#"{{"m.homeserver":{{"base_url":"http://destination.test:{port}"}}}}"#),
366 "application/json",
367 ))
368 .mount(&mock_server)
369 .await;
370 Mock::given(method("GET"))
371 .and(path("/_matrix/client/versions"))
372 .respond_with(
373 ResponseTemplate::new(201)
374 .set_body_raw(r#"{"versions":["v1.11"]}"#, "application/json"),
375 )
376 .expect(1)
377 .mount(&mock_server)
378 .await;
379
380 let http = reqwest::Client::builder()
381 .resolve("example.test", *mock_server.address())
382 .resolve("destination.test", *mock_server.address())
383 .build()?;
384 let resolver = Resolver::with(http);
385 let result = resolver.resolve(&format!("example.test:{port}")).await;
386
387 assert!(matches!(result, Err(Error::Fail(_))), "expected FAIL_ERROR, got {result:?}");
388 Ok(())
389 }
390
391 #[tokio::test]
394 async fn identity_2xx_non_200_fails_error() -> Result<(), Box<dyn std::error::Error>> {
395 let mock_server = MockServer::start().await;
396 let port = mock_server.address().port();
397
398 Mock::given(method("GET"))
399 .and(path("/.well-known/matrix/client"))
400 .respond_with(ResponseTemplate::new(200).set_body_raw(
401 format!(
402 r#"{{"m.homeserver":{{"base_url":"http://destination.test:{port}"}},"m.identity_server":{{"base_url":"http://destination.test:{port}"}}}}"#
403 ),
404 "application/json",
405 ))
406 .mount(&mock_server)
407 .await;
408 Mock::given(method("GET"))
409 .and(path("/_matrix/client/versions"))
410 .respond_with(
411 ResponseTemplate::new(200)
412 .set_body_raw(r#"{"versions":["v1.11"]}"#, "application/json"),
413 )
414 .mount(&mock_server)
415 .await;
416 Mock::given(method("GET"))
417 .and(path("/_matrix/identity/v2"))
418 .respond_with(ResponseTemplate::new(204))
419 .expect(1)
420 .mount(&mock_server)
421 .await;
422
423 let http = reqwest::Client::builder()
424 .resolve("example.test", *mock_server.address())
425 .resolve("destination.test", *mock_server.address())
426 .build()?;
427 let resolver = Resolver::with(http);
428 let result = resolver.resolve(&format!("example.test:{port}")).await;
429
430 assert!(matches!(result, Err(Error::Fail(_))), "expected FAIL_ERROR, got {result:?}");
431 Ok(())
432 }
433}