Skip to main content

matrix_oracle/
client.rs

1//! Resolution for the client-server API
2
3pub 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/// well-known information for the client-server API.
15#[derive(Debug, Clone, Serialize, Deserialize)]
16pub struct ClientWellKnown {
17	/// Information about the homeserver to connect to.
18	#[serde(rename = "m.homeserver")]
19	pub homeserver: HomeserverInfo,
20
21	/// Information about the identity server to connect to.
22	#[serde(rename = "m.identity_server")]
23	pub identity_server: Option<IdentityServerInfo>,
24}
25
26/// Information about the homeserver to connect to.
27#[derive(Debug, Clone, Serialize, Deserialize)]
28pub struct HomeserverInfo {
29	/// The base url to use for client-server API endpoints.
30	base_url: String,
31}
32
33/// Information about the identity server to connect to.
34#[derive(Debug, Clone, Serialize, Deserialize)]
35pub struct IdentityServerInfo {
36	/// The base url to use for identity server API endpoints.
37	base_url: String,
38}
39
40/// Resolver for well-known lookups for the client-server API.
41#[derive(Clone, Debug)]
42pub struct Resolver {
43	/// The HTTP client used to send and receive requests. Should transparently
44	/// handle HTTP caching.
45	http: ClientWithMiddleware,
46}
47
48/// Represents the set of matrix versions a server support. Used exclusively for
49/// validating the contents of a response
50#[allow(dead_code)]
51#[derive(Deserialize)]
52struct Versions {
53	/// List of matrix spec versions the server supports.
54	pub versions: Vec<String>,
55	/// Set of unstable matrix extensions which the server supports
56	#[serde(default)]
57	pub unstable_features: BTreeMap<String, bool>,
58}
59
60impl Resolver {
61	/// Construct a new resolver.
62	#[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	/// Construct a new resolver with the given reqwest client.
72	#[must_use]
73	pub fn with(http: reqwest::Client) -> Self {
74		Self { http: reqwest_middleware::ClientBuilder::new(http).with(cache()).build() }
75	}
76
77	/// Get the base URL for the client-server API with the given name.
78	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		// 3. make a GET request to the well-known endpoint
85		let response = self.http.get(url.join(".well-known/matrix/client")?).send().await?;
86		// a. if the returned status code is 404, then IGNORE
87		if response.status() == StatusCode::NOT_FOUND {
88			return Ok(url);
89		};
90		// b. otherwise, if the status code is not 200, FAIL_PROMPT
91		if response.status() != StatusCode::OK {
92			return Err(UnexpectedStatus(response.status()).into());
93		}
94		// c. parse the response as json
95		let well_known = response.json::<ClientWellKnown>().await?;
96		// d+e.i Extract base_url and parse it as a URL
97		let url = Url::parse(&well_known.homeserver.base_url)?;
98		// e.ii Validate versions endpoint: a non-200 response is a FAIL_ERROR
99		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		// f. if present, validate identity server v2 endpoint
111		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	/// Tests that a 404 response is correctly handled
144	#[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	/// Tests that a present identity server is validated at the current spec
211	/// endpoint (`/_matrix/identity/v2`), not the removed v1 endpoint.
212	#[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		// The identity server MUST be validated at the v2 endpoint.
236		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	/// Per spec: a non-404, non-200 response to the well-known request must
254	/// `FAIL_PROMPT`, even if the body happens to parse as valid JSON.
255	#[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	/// Per spec: a failed `/_matrix/client/versions` validation must
282	/// `FAIL_ERROR`, even if the response body happens to parse as `Versions`.
283	#[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	/// Per spec: only a `200` well-known response is accepted. A
318	/// 2xx-but-not-200 status (e.g. `201`) must `FAIL_PROMPT`, even with a
319	/// parseable body.
320	#[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	/// Per spec: homeserver validation requires a `200` from
355	/// `/_matrix/client/versions`. A 2xx-but-not-200 status must `FAIL_ERROR`,
356	/// even with a parseable body.
357	#[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	/// Per spec: identity-server validation requires a `200` from
392	/// `/_matrix/identity/v2`. A 2xx-but-not-200 status must `FAIL_ERROR`.
393	#[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}