1use std::{fmt::Debug, net::IpAddr};
2
3use futures::{FutureExt, TryFutureExt};
4use hickory_resolver::{
5 net::{DnsError, NetError},
6 proto::rr::{RData, rdata::SRV},
7};
8use ipaddress::IPAddress;
9use ruma::ServerName;
10use tuwunel_core::{
11 Err, Result, debug, debug_info, debug_warn, err, error, format_array_string, implement,
12 trace, utils::string::to_small_string,
13};
14
15use super::{
16 DestString, FedDest,
17 cache::{CachedDest, CachedOverride, MAX_IPS},
18 fed::{HostString, PortString, add_port_to_hostname, get_ip_with_port},
19};
20
21#[derive(Clone, Debug)]
22pub(crate) struct ActualDest {
23 pub(crate) dest: FedDest,
24 pub(crate) host: DestString,
25}
26
27impl ActualDest {
28 #[inline]
29 pub(crate) fn to_string(&self) -> DestString { self.dest.https_string() }
30}
31
32#[implement(super::Service)]
33#[tracing::instrument(skip_all, level = "debug", name = "resolve")]
34pub(crate) async fn get_actual_dest(&self, server_name: &ServerName) -> Result<ActualDest> {
35 let (CachedDest { dest, host, .. }, _cached) = self
36 .lookup_actual_dest_with_policy(server_name, false)
37 .await?;
38
39 Ok(ActualDest { dest, host })
40}
41
42#[implement(super::Service)]
43#[tracing::instrument(skip_all, level = "debug", name = "resolve")]
44pub(crate) async fn get_actual_dest_allow_self(
45 &self,
46 server_name: &ServerName,
47) -> Result<ActualDest> {
48 let (CachedDest { dest, host, .. }, _cached) = self
49 .lookup_actual_dest_with_policy(server_name, true)
50 .await?;
51
52 Ok(ActualDest { dest, host })
53}
54
55#[implement(super::Service)]
56async fn lookup_actual_dest_with_policy(
57 &self,
58 server_name: &ServerName,
59 allow_self: bool,
60) -> Result<(CachedDest, bool)> {
61 self.validate_self_destination(server_name, allow_self)?;
62
63 if let Ok(result) = self.cache.get_destination(server_name).await {
64 return Ok((result, true));
65 }
66
67 let _dedup = self.resolving.lock(server_name).await;
68 if let Ok(result) = self.cache.get_destination(server_name).await {
69 return Ok((result, true));
70 }
71
72 self.validate_dest_address(server_name)?;
73
74 self.resolve_actual_dest_unchecked(server_name, true)
75 .inspect_ok(|result| self.cache.set_destination(server_name, result))
76 .map_ok(|result| (result, false))
77 .boxed()
78 .await
79}
80
81#[implement(super::Service)]
86pub async fn resolve_actual_dest(&self, dest: &ServerName, cache: bool) -> Result<CachedDest> {
87 self.validate_dest(dest, false)?;
88 self.resolve_actual_dest_unchecked(dest, cache)
89 .await
90}
91
92#[implement(super::Service)]
93#[tracing::instrument(name = "actual", level = "debug", skip(self, cache))]
94async fn resolve_actual_dest_unchecked(
95 &self,
96 dest: &ServerName,
97 cache: bool,
98) -> Result<CachedDest> {
99 let mut host: DestString = dest.as_str().into();
100 let actual_dest = self.actual_dest(dest, cache, &mut host).await?;
101 let actual_host = Self::dest_host(&host);
102
103 debug!("Actual destination: {actual_dest:?} hostname: {actual_host:?}");
104 Ok(CachedDest {
105 dest: actual_dest,
106 host: actual_host.uri_string(),
107 expire: CachedDest::default_expire(),
108 })
109}
110
111#[implement(super::Service)]
112fn dest_host(host: &DestString) -> FedDest {
113 host.parse()
115 .map(FedDest::Literal)
116 .or_else(|_| {
117 host.parse().map(|addr: IpAddr| {
118 FedDest::Named(addr.to_string().into(), FedDest::default_port())
119 })
120 })
121 .unwrap_or_else(|_| {
122 host.find(':').map_or_else(
123 || FedDest::Named(host.as_str().into(), FedDest::default_port()),
124 |pos| {
125 let (host, port) = host.split_at(pos);
126
127 FedDest::Named(
128 host.into(),
129 port.try_into()
130 .unwrap_or_else(|_| FedDest::default_port()),
131 )
132 },
133 )
134 })
135}
136
137#[implement(super::Service)]
138async fn actual_dest(
139 &self,
140 dest: &ServerName,
141 cache: bool,
142 host: &mut DestString,
143) -> Result<FedDest> {
144 match get_ip_with_port(dest.as_str()) {
145 | Some(host_port) => Self::actual_dest_1(host_port),
146 | None if let Some(pos) = dest.as_str().find(':') =>
147 self.actual_dest_2(dest, cache, pos).await,
148 | None => {
149 self.maybe_query_and_cache(dest.as_str(), 8448, true)
150 .await?;
151 self.services.server.check_running()?;
152 match self.request_well_known(dest.as_str()).await? {
153 | Some(delegated) => self.actual_dest_3(host, cache, &delegated).await,
154 | _ => match self.query_srv_record(dest.as_str()).await? {
155 | Some(overrider) => self.actual_dest_4(host, cache, overrider).await,
156 | _ => self.actual_dest_5(dest, cache).await,
157 },
158 }
159 },
160 }
161}
162
163#[implement(super::Service)]
164fn actual_dest_1(host_port: FedDest) -> Result<FedDest> {
165 debug!("1: IP literal with provided or default port");
166 Ok(host_port)
167}
168
169#[implement(super::Service)]
170async fn actual_dest_2(&self, dest: &ServerName, cache: bool, pos: usize) -> Result<FedDest> {
171 debug!("2: Hostname with included port");
172 let (host, port) = dest.as_str().split_at(pos);
173 let port_num = port
174 .trim_start_matches(':')
175 .parse::<u16>()
176 .unwrap_or(8448);
177
178 self.maybe_query_and_cache(host, port_num, cache)
179 .await?;
180
181 let port = port
182 .try_into()
183 .unwrap_or_else(|_| FedDest::default_port());
184
185 Ok(FedDest::Named(host.into(), port))
186}
187
188#[implement(super::Service)]
189async fn actual_dest_3(
190 &self,
191 host: &mut DestString,
192 cache: bool,
193 delegated: &str,
194) -> Result<FedDest> {
195 debug!("3: A .well-known file is available");
196 *host = add_port_to_hostname(delegated).uri_string();
197 match get_ip_with_port(delegated) {
198 | Some(host_and_port) => Self::actual_dest_3_1(host_and_port),
199 | None =>
200 if let Some(pos) = delegated.find(':') {
201 self.actual_dest_3_2(cache, delegated, pos).await
202 } else {
203 trace!("Delegated hostname has no port in this branch");
204 match self.query_srv_record(delegated).await? {
205 | Some(overrider) =>
206 self.actual_dest_3_3(cache, delegated, overrider)
207 .await,
208 | _ => self.actual_dest_3_4(cache, delegated).await,
209 }
210 },
211 }
212}
213
214#[implement(super::Service)]
215fn actual_dest_3_1(host_and_port: FedDest) -> Result<FedDest> {
216 debug!("3.1: IP literal in .well-known file");
217 Ok(host_and_port)
218}
219
220#[implement(super::Service)]
221async fn actual_dest_3_2(&self, cache: bool, delegated: &str, pos: usize) -> Result<FedDest> {
222 debug!("3.2: Hostname with port in .well-known file");
223 let (host, port) = delegated.split_at(pos);
224 let port_num = port
225 .trim_start_matches(':')
226 .parse::<u16>()
227 .unwrap_or(8448);
228
229 self.maybe_query_and_cache(host, port_num, cache)
230 .await?;
231
232 let port = port
233 .try_into()
234 .unwrap_or_else(|_| FedDest::default_port());
235
236 Ok(FedDest::Named(host.into(), port))
237}
238
239#[implement(super::Service)]
240async fn actual_dest_3_3(
241 &self,
242 cache: bool,
243 delegated: &str,
244 overrider: FedDest,
245) -> Result<FedDest> {
246 debug!("3.3: SRV lookup successful");
247 let force_port = overrider.port();
248 self.maybe_query_and_cache_override(
249 delegated,
250 &overrider.hostname(),
251 force_port.unwrap_or(8448),
252 cache,
253 )
254 .await?;
255
256 if let Some(port) = force_port {
257 let port: PortString = format_array_string!(":{port}");
258
259 return Ok(FedDest::Named(delegated.into(), port));
260 }
261
262 Ok(add_port_to_hostname(delegated))
263}
264
265#[implement(super::Service)]
266async fn actual_dest_3_4(&self, cache: bool, delegated: &str) -> Result<FedDest> {
267 debug!("3.4: No SRV records, just use the hostname from .well-known");
268 self.maybe_query_and_cache(delegated, 8448, cache)
269 .await?;
270
271 Ok(add_port_to_hostname(delegated))
272}
273
274#[implement(super::Service)]
275async fn actual_dest_4(&self, host: &str, cache: bool, overrider: FedDest) -> Result<FedDest> {
276 debug!("4: No .well-known; SRV record found");
277 let force_port = overrider.port();
278 self.maybe_query_and_cache_override(
279 host,
280 &overrider.hostname(),
281 force_port.unwrap_or(8448),
282 cache,
283 )
284 .await?;
285
286 if let Some(port) = force_port {
287 let port: PortString = format_array_string!(":{port}");
288
289 return Ok(FedDest::Named(host.into(), port));
290 }
291
292 Ok(add_port_to_hostname(host))
293}
294
295#[implement(super::Service)]
296async fn actual_dest_5(&self, dest: &ServerName, cache: bool) -> Result<FedDest> {
297 debug!("5: No SRV record found");
298 self.maybe_query_and_cache(dest.as_str(), 8448, cache)
299 .await?;
300
301 Ok(add_port_to_hostname(dest.as_str()))
302}
303
304#[implement(super::Service)]
305#[inline]
306async fn maybe_query_and_cache(&self, hostname: &str, port: u16, cache: bool) -> Result {
307 self.maybe_query_and_cache_override(hostname, hostname, port, cache)
308 .await
309}
310
311#[implement(super::Service)]
312#[inline]
313async fn maybe_query_and_cache_override(
314 &self,
315 untername: &str,
316 hostname: &str,
317 port: u16,
318 cache: bool,
319) -> Result {
320 if !cache {
321 return Ok(());
322 }
323
324 if self.cache.has_override(untername).await {
325 return Ok(());
326 }
327
328 self.query_and_cache_override(untername, hostname, port)
329 .await
330}
331
332#[implement(super::Service)]
333#[tracing::instrument(name = "ip", level = "debug", skip(self))]
334async fn query_and_cache_override(
335 &self,
336 untername: &'_ str,
337 hostname: &'_ str,
338 port: u16,
339) -> Result {
340 self.services.server.check_running()?;
341
342 debug!("querying IP for {untername:?} ({hostname:?}:{port})");
343 match self
344 .resolver
345 .resolver
346 .lookup_ip(hostname.to_owned())
347 .await
348 {
349 | Err(e) => Self::handle_resolve_error(&e, hostname),
350 | Ok(override_ip) => {
351 self.cache
352 .set_override(untername, &CachedOverride {
353 ips: override_ip.iter().take(MAX_IPS).collect(),
354 port,
355 expire: CachedOverride::default_expire(),
356 overriding: (hostname != untername)
357 .then_some(hostname.into())
358 .inspect(|_| debug_info!("{untername:?} overridden by {hostname:?}")),
359 });
360
361 Ok(())
362 },
363 }
364}
365
366#[implement(super::Service)]
367#[tracing::instrument(name = "srv", level = "debug", skip(self))]
368async fn query_srv_record(&self, hostname: &'_ str) -> Result<Option<FedDest>> {
369 let hostnames =
370 [format!("_matrix-fed._tcp.{hostname}."), format!("_matrix._tcp.{hostname}.")];
371
372 for hostname in hostnames {
373 self.services.server.check_running()?;
374
375 debug!("querying SRV for {hostname:?}");
376 let hostname = hostname.trim_end_matches('.');
377 match self.resolver.resolver.srv_lookup(hostname).await {
378 | Err(e) => Self::handle_resolve_error(&e, hostname)?,
379 | Ok(result) => {
380 let srv = result
381 .answers()
382 .iter()
383 .find_map(|r| match &r.data {
384 | RData::SRV(srv) => Some(srv),
385 | _ => None,
386 });
387
388 return Ok(srv.map(Self::srv_dest));
389 },
390 }
391 }
392
393 Ok(None)
394}
395
396#[implement(super::Service)]
397fn srv_dest(srv: &SRV) -> FedDest {
398 let host: HostString = to_small_string(&srv.target);
399 let port: PortString = format_array_string!(":{}", srv.port);
400
401 FedDest::Named(host.trim_end_matches('.').into(), port)
402}
403
404#[implement(super::Service)]
405fn handle_resolve_error(e: &NetError, host: &'_ str) -> Result {
406 match e {
411 | NetError::Dns(DnsError::NoRecordsFound(_)) => {
412 debug!(%host, "No DNS records found: {e}");
414 Ok(())
415 },
416 | NetError::Dns(_) => {
417 debug_warn!(%host, "DNS response error: {e}");
418 Ok(())
419 },
420 | NetError::Timeout => Err!(warn!(%host, "DNS {e}")),
421 | NetError::NoConnections => {
422 error!(
423 "Your DNS server is overloaded and has ran out of connections. It is strongly \
424 recommended you remediate this issue to ensure proper federation connectivity."
425 );
426
427 Err!(error!(%host, "DNS error: {e}"))
428 },
429 | _ => Err!(error!(%host, "DNS error: {e}")),
430 }
431}
432
433#[implement(super::Service)]
434fn validate_dest(&self, dest: &ServerName, allow_self: bool) -> Result {
435 self.validate_self_destination(dest, allow_self)?;
436 self.validate_dest_address(dest)
437}
438
439#[implement(super::Service)]
440fn validate_self_destination(&self, dest: &ServerName, allow_self: bool) -> Result {
441 if !allow_self
442 && dest == self.services.server.name
443 && !self.services.server.config.federation_loopback
444 {
445 return Err!("Won't send federation request to ourselves");
446 }
447
448 Ok(())
449}
450
451#[implement(super::Service)]
452fn validate_dest_address(&self, dest: &ServerName) -> Result {
453 if dest.is_ip_literal() || IPAddress::is_valid(dest.host()) {
454 self.validate_dest_ip_literal(dest)?;
455 }
456
457 Ok(())
458}
459
460#[implement(super::Service)]
461fn validate_dest_ip_literal(&self, dest: &ServerName) -> Result {
462 trace!("Destination is an IP literal, checking against IP range denylist.",);
463 debug_assert!(
464 dest.is_ip_literal() || !IPAddress::is_valid(dest.host()),
465 "Destination is not an IP literal."
466 );
467 let ip = IPAddress::parse(dest.host()).map_err(|e| {
468 err!(BadServerResponse(debug_error!("Failed to parse IP literal from string: {e}")))
469 })?;
470
471 self.validate_ip(&ip)?;
472
473 Ok(())
474}
475
476#[implement(super::Service)]
477pub(crate) fn validate_ip(&self, ip: &IPAddress) -> Result {
478 if !self.services.client.valid_cidr_range(ip) {
479 return Err!(BadServerResponse("Not allowed to send requests to this IP"));
480 }
481
482 Ok(())
483}