1use std::{fmt::Debug, mem, time::Duration};
7
8use bytes::Bytes;
9use ipaddress::IPAddress;
10use reqwest::{Client, Method, Request, Response, Url};
11use ruma::{
12 ServerName,
13 api::{
14 EndpointError, IncomingResponse, MatrixVersion, OutgoingRequest, OutgoingRequestExt,
15 SupportedVersions,
16 error::{Error as RumaError, ErrorBody},
17 },
18};
19use tokio::time::timeout;
20use tuwunel_core::{
21 Err, Error, Result, debug, debug::INFO_SPAN_LEVEL, debug_error, debug_warn, err, implement,
22 trace,
23};
24
25use super::{
26 ShouldAttempt,
27 peer::classify_error,
28 scheme::{FedAuth, FedPath},
29};
30use crate::{client::read_response_capped, resolver::actual::ActualDest};
31
32#[implement(super::Service)]
38#[tracing::instrument(skip_all, name = "request", level = "debug")]
39pub async fn execute<T>(&self, dest: &ServerName, request: T) -> Result<T::IncomingResponse>
40where
41 T: OutgoingRequest + Debug + Send,
42 T::Authentication: FedAuth,
43 T::PathBuilder: FedPath,
44{
45 let client = &self.services.client.federation;
46 self.execute_on(client, dest, request).await
47}
48
49#[implement(super::Service)]
56#[tracing::instrument(skip_all, name = "keys", level = "debug")]
57pub async fn execute_keys<T>(&self, dest: &ServerName, request: T) -> Result<T::IncomingResponse>
58where
59 T: OutgoingRequest + Debug + Send,
60 T::Authentication: FedAuth,
61 T::PathBuilder: FedPath,
62{
63 if matches!(self.should_attempt(dest).await, ShouldAttempt::No { .. }) {
64 return Err!("{dest} is in federation backoff; skipping key lookup");
65 }
66
67 let timeout_dur = Duration::from_secs(
68 self.services
69 .server
70 .config
71 .federation_keys_timeout,
72 );
73
74 let client = &self.services.client.federation;
75
76 match timeout(timeout_dur, self.execute_uncounted(client, dest, request)).await {
77 | Ok(result) => result,
78 | Err(_elapsed) => Err!("{dest} key lookup exceeded {}s", timeout_dur.as_secs()),
79 }
80}
81
82#[implement(super::Service)]
87#[tracing::instrument(skip_all, name = "synapse", level = "debug")]
88pub async fn execute_synapse<T>(
89 &self,
90 dest: &ServerName,
91 request: T,
92) -> Result<T::IncomingResponse>
93where
94 T: OutgoingRequest + Debug + Send,
95 T::Authentication: FedAuth,
96 T::PathBuilder: FedPath,
97{
98 let client = &self.services.client.synapse;
99 self.execute_on(client, dest, request).await
100}
101
102#[implement(super::Service)]
108pub async fn execute_on<T>(
109 &self,
110 client: &Client,
111 dest: &ServerName,
112 request: T,
113) -> Result<T::IncomingResponse>
114where
115 T: OutgoingRequest + Send,
116 T::Authentication: FedAuth,
117 T::PathBuilder: FedPath,
118{
119 let result = self
120 .execute_uncounted(client, dest, request)
121 .await;
122
123 match &result {
124 | Ok(_) => self.record_success(dest).await,
125 | Err(error) =>
126 if let Some(class) = classify_error(error) {
127 self.record_failure(dest, class);
128 },
129 }
130
131 result
132}
133
134#[implement(super::Service)]
139pub(super) async fn execute_on_allow_self<T>(
140 &self,
141 client: &Client,
142 dest: &ServerName,
143 request: T,
144) -> Result<T::IncomingResponse>
145where
146 T: OutgoingRequest + Send,
147 T::Authentication: FedAuth,
148 T::PathBuilder: FedPath,
149{
150 let result = self
151 .execute_uncounted_allow_self(client, dest, request)
152 .await;
153
154 match &result {
155 | Ok(_) => self.record_success(dest).await,
156 | Err(error) =>
157 if let Some(class) = classify_error(error) {
158 self.record_failure(dest, class);
159 },
160 }
161
162 result
163}
164
165#[implement(super::Service)]
170#[tracing::instrument(
171 name = "fed",
172 level = INFO_SPAN_LEVEL,
173 skip(self, client, request),
174)]
175pub(super) async fn execute_uncounted<T>(
176 &self,
177 client: &Client,
178 dest: &ServerName,
179 request: T,
180) -> Result<T::IncomingResponse>
181where
182 T: OutgoingRequest + Send,
183 T::Authentication: FedAuth,
184 T::PathBuilder: FedPath,
185{
186 self.validate_request_destination(dest)?;
187 let actual = self
188 .services
189 .resolver
190 .get_actual_dest(dest)
191 .await?;
192 let request = self.prepare(&actual, dest, request)?;
193
194 self.perform::<T>(&actual, dest, request, client)
195 .await
196}
197
198#[implement(super::Service)]
203#[tracing::instrument(name = "fed", level = "debug", skip(self, client, request))]
204pub(super) async fn execute_uncounted_allow_self<T>(
205 &self,
206 client: &Client,
207 dest: &ServerName,
208 request: T,
209) -> Result<T::IncomingResponse>
210where
211 T: OutgoingRequest + Send,
212 T::Authentication: FedAuth,
213 T::PathBuilder: FedPath,
214{
215 self.validate_request_destination(dest)?;
216 let actual = self
217 .services
218 .resolver
219 .get_actual_dest_allow_self(dest)
220 .await?;
221 let request = self.prepare(&actual, dest, request)?;
222
223 self.perform::<T>(&actual, dest, request, client)
224 .await
225}
226
227#[implement(super::Service)]
228fn validate_request_destination(&self, dest: &ServerName) -> Result {
229 if !self.services.server.config.allow_federation {
230 return Err!(Config("allow_federation", "Federation is disabled."));
231 }
232
233 if self
234 .services
235 .server
236 .config
237 .is_forbidden_remote_server_name(dest)
238 {
239 return Err!(Request(Forbidden(debug_warn!("Federation with {dest} is not allowed."))));
240 }
241
242 Ok(())
243}
244
245#[implement(super::Service)]
246async fn perform<T>(
247 &self,
248 actual: &ActualDest,
249 dest: &ServerName,
250 request: Request,
251 client: &Client,
252) -> Result<T::IncomingResponse>
253where
254 T: OutgoingRequest + Send,
255 T::Authentication: FedAuth,
256 T::PathBuilder: FedPath,
257{
258 let url = request.url().clone();
259 let method = request.method().clone();
260
261 debug!(?method, ?url, "Sending request");
262 let limit = self.services.server.config.max_response_size;
263
264 match client.execute(request).await {
265 | Ok(response) => handle_response::<T>(actual, dest, &method, &url, response, limit)
266 .await
267 .inspect_err(|error| self.evict_misrouted(dest, actual, error)),
268 | Err(error) => Err(self
269 .handle_error(dest, actual, &method, &url, error)
270 .expect_err("always returns error")),
271 }
272}
273
274#[implement(super::Service)]
275fn prepare<T>(&self, actual: &ActualDest, dest: &ServerName, request: T) -> Result<Request>
276where
277 T: OutgoingRequest + Send,
278 T::Authentication: FedAuth,
279 T::PathBuilder: FedPath,
280{
281 let request = self.to_http_request::<T>(actual, dest, request)?;
282 let request = Request::try_from(request)?;
283 self.validate_url(request.url())?;
284 self.services.server.check_running()?;
285
286 Ok(request)
287}
288
289#[implement(super::Service)]
290fn validate_url(&self, url: &Url) -> Result {
291 if let Some(url_host) = url.host_str()
292 && let Ok(ip) = IPAddress::parse(url_host)
293 {
294 trace!("Checking request URL IP {ip:?}");
295 self.services.resolver.validate_ip(&ip)?;
296 }
297
298 Ok(())
299}
300
301async fn handle_response<T>(
302 actual: &ActualDest,
303 dest: &ServerName,
304 method: &Method,
305 url: &Url,
306 response: Response,
307 limit: usize,
308) -> Result<T::IncomingResponse>
309where
310 T: OutgoingRequest + Send,
311 T::Authentication: FedAuth,
312 T::PathBuilder: FedPath,
313{
314 let response = into_http_response(dest, actual, method, url, response, limit).await?;
315
316 T::IncomingResponse::try_from_http_response(response)
317 .map_err(|e| err!(BadServerResponse("Server returned bad 200 response: {e:?}")))
318}
319
320async fn into_http_response(
321 dest: &ServerName,
322 actual: &ActualDest,
323 method: &Method,
324 url: &Url,
325 mut response: Response,
326 limit: usize,
327) -> Result<http::Response<Bytes>> {
328 let status = response.status();
329 trace!(
330 ?status, ?method,
331 request_url = ?url,
332 response_url = ?response.url(),
333 "Received response from {}",
334 actual.to_string(),
335 );
336
337 let mut http_response_builder = http::Response::builder()
338 .status(status)
339 .version(response.version());
340
341 mem::swap(
342 response.headers_mut(),
343 http_response_builder
344 .headers_mut()
345 .expect("http::response::Builder is usable"),
346 );
347
348 trace!("Waiting for response body...");
350 let body = read_response_capped(response, limit).await?;
351
352 let http_response = http_response_builder
353 .body(body)
354 .expect("reqwest body is valid http body");
355
356 debug!("Got {status:?} for {method} {url}");
357 if !status.is_success() {
358 return Err(Error::Federation(
359 dest.to_owned(),
360 RumaError::from_http_response(http_response),
361 ));
362 }
363
364 Ok(http_response)
365}
366
367#[implement(super::Service)]
368fn handle_error(
369 &self,
370 dest: &ServerName,
371 actual: &ActualDest,
372 method: &Method,
373 url: &Url,
374 mut e: reqwest::Error,
375) -> Result {
376 if e.is_timeout() || e.is_connect() {
377 e = e.without_url();
378 debug_warn!("{e:?}");
379 } else if e.is_redirect() {
380 debug_error!(
381 method = ?method,
382 url = ?url,
383 final_url = ?e.url(),
384 "Redirect loop {}: {}",
385 actual.host,
386 e,
387 );
388 } else {
389 debug_error!("{e:?}");
390 }
391
392 self.evict_route(dest, actual);
393
394 Err(e.into())
395}
396
397#[implement(super::Service)]
400fn evict_misrouted(&self, dest: &ServerName, actual: &ActualDest, error: &Error) {
401 let Error::Federation(_, response) = error else {
402 return;
403 };
404
405 if matches!(response.body, ErrorBody::NotJson { .. }) {
406 self.evict_route(dest, actual);
407 }
408}
409
410#[implement(super::Service)]
413fn evict_route(&self, dest: &ServerName, actual: &ActualDest) {
414 self.services.resolver.cache.del_destination(dest);
415 self.services
416 .resolver
417 .cache
418 .del_override(&actual.dest.hostname());
419}
420
421#[implement(super::Service)]
422fn to_http_request<T>(
423 &self,
424 actual: &ActualDest,
425 dest: &ServerName,
426 request: T,
427) -> Result<http::Request<Vec<u8>>>
428where
429 T: OutgoingRequest + Send,
430 T::Authentication: FedAuth,
431 T::PathBuilder: FedPath,
432{
433 const VERSIONS: [MatrixVersion; 1] = [MatrixVersion::V1_11];
434 let supported = SupportedVersions {
435 versions: VERSIONS.into(),
436 features: Default::default(),
437 };
438
439 let auth = T::Authentication::input(
440 self.services.server.name.clone(),
441 dest.to_owned(),
442 self.services.server_keys.keypair(),
443 );
444 let path = T::PathBuilder::input(&supported);
445
446 request
447 .try_into_http_request::<Vec<u8>>(actual.to_string().as_str(), auth, path)
448 .map_err(|e| err!(BadServerResponse("Invalid destination: {e:?}")))
449}