Skip to main content

tuwunel_service/resolver/
actual.rs

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/// Returns: `actual_destination`, host header
82/// Implemented according to the specification at <https://matrix.org/docs/spec/server_server/r0.1.4#resolving-server-names>
83/// Numbers in comments below refer to bullet points in linked section of
84/// specification
85#[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	// Preserve an unspecified port on an IP address.
114	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	// `NetError::Dns(_)` covers responses returned by the remote side (NXDOMAIN,
407	// SERVFAIL, REFUSED, ...) only seen with verbose-logging. Local-origin failures
408	// (Timeout, NoConnections, Io, ...) keep their warn/error level so an operator
409	// notices when their own resolver is unhealthy.
410	match e {
411		| NetError::Dns(DnsError::NoRecordsFound(_)) => {
412			// Raise to debug_warn if we can find out the result wasn't from cache
413			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}