Skip to main content

tuwunel_service/resolver/
dns.rs

1use std::{
2	io,
3	io::ErrorKind::PermissionDenied,
4	net::{IpAddr, SocketAddr},
5	sync::Arc,
6	time::Duration,
7};
8
9use futures::FutureExt;
10use hickory_resolver::{
11	TokioResolver,
12	config::{LookupIpStrategy, NameServerConfig, ProtocolConfig, ResolverConfig, ResolverOpts},
13	lookup_ip::LookupIp,
14	net::runtime::TokioRuntimeProvider,
15	system_conf::read_system_conf,
16};
17use ipaddress::IPAddress;
18use reqwest::dns::{Addrs, Name, Resolve, Resolving};
19use tuwunel_core::{Result, Server, config::proxy::ProxyHosts, err, trace};
20
21use super::cache::{Cache, CachedOverride};
22use crate::client::ipaddress_from_std;
23
24pub struct Resolver {
25	pub(crate) resolver: Arc<TokioResolver>,
26	pub(crate) passthru: Arc<Passthru>,
27	pub(crate) hooked: Arc<Hooked>,
28	server: Arc<Server>,
29}
30
31/// Filters destination DNS answers through the configured CIDR denylist.
32///
33/// Proxy endpoint names are exempt because their addresses describe the
34/// transport hop rather than the request destination.
35pub(crate) struct Validating<R> {
36	inner: Arc<R>,
37	denylist: Arc<[IPAddress]>,
38	proxy_hosts: ProxyHosts,
39}
40
41pub(crate) struct Hooked {
42	resolver: Arc<TokioResolver>,
43	passthru: Arc<Passthru>,
44	cache: Arc<Cache>,
45	server: Arc<Server>,
46}
47
48pub(crate) struct Passthru {
49	resolver: Arc<TokioResolver>,
50	server: Arc<Server>,
51}
52
53type ResolvingResult = Result<Addrs, Box<dyn std::error::Error + Send + Sync>>;
54
55impl Resolver {
56	pub(super) fn build(server: &Arc<Server>, cache: Arc<Cache>) -> Result<Arc<Self>> {
57		let config = &server.config;
58
59		// Create the primary resolver.
60		let (conf, mut opts) = Self::configure(server)?;
61		opts.negative_min_ttl = Some(Duration::from_secs(config.dns_min_ttl_nxdomain));
62		opts.positive_min_ttl = Some(Duration::from_secs(config.dns_min_ttl));
63		opts.cache_size = config.dns_cache_entries.into();
64		let resolver = Self::create(server, conf.clone(), opts.clone())?;
65
66		// Create the passthru resolver with modified options.
67		let (conf, mut opts) = (conf, opts);
68		opts.negative_min_ttl = Some(Duration::ZERO);
69		opts.positive_min_ttl = Some(Duration::ZERO);
70		opts.cache_size = ResolverOpts::default().cache_size;
71		let passthru = Arc::new(Passthru {
72			resolver: Self::create(server, conf, opts)?,
73			server: server.clone(),
74		});
75
76		Ok(Arc::new(Self {
77			hooked: Arc::new(Hooked {
78				resolver: resolver.clone(),
79				passthru: passthru.clone(),
80				server: server.clone(),
81				cache,
82			}),
83			server: server.clone(),
84			passthru,
85			resolver,
86		}))
87	}
88
89	fn create(
90		server: &Arc<Server>,
91		conf: ResolverConfig,
92		opts: ResolverOpts,
93	) -> Result<Arc<TokioResolver>> {
94		let mut builder =
95			TokioResolver::builder_with_config(conf, TokioRuntimeProvider::default());
96		*builder.options_mut() = Self::configure_opts(server, opts);
97
98		builder
99			.build()
100			.map(Arc::new)
101			.map_err(|e| err!(error!("Failed to build DNS resolver: {e}")))
102	}
103
104	fn configure(server: &Arc<Server>) -> Result<(ResolverConfig, ResolverOpts)> {
105		let config = &server.config;
106
107		// hickory's android system_conf panics in ndk-context without a JVM.
108		#[cfg(target_os = "android")]
109		if config.dns_servers.is_empty() {
110			return Err(err!(Config(
111				"dns_servers",
112				"The system resolver requires a JVM on Android; set dns_servers to your \
113				 upstream nameservers instead."
114			)));
115		}
116
117		let (base_conf, opts) = if config.dns_servers.is_empty() {
118			read_system_conf().map_err(|e| {
119				err!(error!("Failed to configure DNS resolver from `/etc/resolv.conf': {e}"))
120			})?
121		} else {
122			(Self::configure_custom(&config.dns_servers)?, ResolverOpts::default())
123		};
124
125		let name_servers = base_conf
126			.name_servers()
127			.iter()
128			.cloned()
129			.map(|mut ns| {
130				ns.trust_negative_responses = !config.query_all_nameservers;
131				if config.query_over_tcp_only {
132					ns.connections
133						.retain(|conn| matches!(conn.protocol, ProtocolConfig::Tcp));
134				}
135
136				ns
137			})
138			.collect();
139
140		let conf = ResolverConfig::from_parts(
141			base_conf.domain().cloned(),
142			base_conf.search().to_vec(),
143			name_servers,
144		);
145
146		Ok((conf, opts))
147	}
148
149	fn configure_custom(servers: &[String]) -> Result<ResolverConfig> {
150		let name_servers = servers
151			.iter()
152			.map(String::as_str)
153			.map(Self::parse_nameserver)
154			.collect::<Result<_>>()?;
155
156		Ok(ResolverConfig::from_parts(None, vec![], name_servers))
157	}
158
159	pub(super) fn parse_nameserver(server: &str) -> Result<NameServerConfig> {
160		let (ip, port) = server
161			.parse::<SocketAddr>()
162			.map(|addr| (addr.ip(), addr.port()))
163			.or_else(|_| server.parse::<IpAddr>().map(|ip| (ip, 53)))
164			.map_err(|e| {
165				err!(Config(
166					"dns_servers",
167					"{server:?} is not an IP address or socket address: {e}"
168				))
169			})?;
170
171		let mut conf = NameServerConfig::udp_and_tcp(ip);
172		for connection in &mut conf.connections {
173			connection.port = port;
174		}
175
176		Ok(conf)
177	}
178
179	#[expect(clippy::as_conversions)]
180	fn configure_opts(server: &Arc<Server>, mut opts: ResolverOpts) -> ResolverOpts {
181		let config = &server.config;
182
183		opts.negative_max_ttl = Some(Duration::from_hours(720));
184		opts.positive_max_ttl = Some(Duration::from_hours(168));
185		opts.timeout = Duration::from_secs(config.dns_timeout);
186		opts.attempts = config.dns_attempts as usize;
187		opts.try_tcp_on_error = config.dns_tcp_fallback;
188		opts.num_concurrent_reqs = 1;
189		opts.edns0 = true;
190		opts.case_randomization = config.dns_case_randomization;
191		opts.preserve_intermediates = true;
192		opts.ip_strategy = match config.ip_lookup_strategy {
193			| 1 => LookupIpStrategy::Ipv4Only,
194			| 2 => LookupIpStrategy::Ipv6Only,
195			| 3 => LookupIpStrategy::Ipv4AndIpv6,
196			| 4 => LookupIpStrategy::Ipv6thenIpv4,
197			| _ => LookupIpStrategy::Ipv4thenIpv6,
198		};
199
200		opts
201	}
202
203	/// Clear the in-memory hickory-dns caches
204	#[inline]
205	pub fn clear_cache(&self) { self.resolver.clear_cache(); }
206}
207
208impl<R: Resolve + 'static> Validating<R> {
209	pub(crate) fn new(
210		inner: Arc<R>,
211		denylist: Arc<[IPAddress]>,
212		proxy_hosts: ProxyHosts,
213	) -> Arc<Self> {
214		Arc::new(Self { inner, denylist, proxy_hosts })
215	}
216}
217
218impl<R: Resolve + 'static> Resolve for Validating<R> {
219	fn resolve(&self, name: Name) -> Resolving {
220		if self
221			.proxy_hosts
222			.iter()
223			.any(|host| host.eq_ignore_ascii_case(name.as_str()))
224		{
225			return self.inner.resolve(name);
226		}
227
228		validate_addrs(self.inner.clone(), self.denylist.clone(), name).boxed()
229	}
230}
231
232async fn validate_addrs<R: Resolve + 'static>(
233	inner: Arc<R>,
234	denylist: Arc<[IPAddress]>,
235	name: Name,
236) -> ResolvingResult {
237	let addrs = inner.resolve(name).await?;
238
239	let mut filtered = addrs
240		.filter(move |sa| {
241			let ip = ipaddress_from_std(sa.ip());
242			!denylist.iter().any(|cidr| cidr.includes(&ip))
243		})
244		.peekable();
245
246	if filtered.peek().is_none() {
247		return Err(Box::new(io::Error::new(
248			PermissionDenied,
249			"All resolved addresses are denied by ip_range_denylist",
250		)));
251	}
252
253	Ok(Box::new(filtered))
254}
255
256impl Resolve for Resolver {
257	fn resolve(&self, name: Name) -> Resolving {
258		let resolver = if self
259			.server
260			.config
261			.dns_passthru_domains
262			.is_match(name.as_str())
263		{
264			trace!(?name, "matched to passthru resolver");
265			&self.passthru.resolver
266		} else {
267			trace!(?name, "using primary resolver");
268			&self.resolver
269		};
270
271		resolve_to_reqwest(self.server.clone(), resolver.clone(), name).boxed()
272	}
273}
274
275impl Resolve for Hooked {
276	fn resolve(&self, name: Name) -> Resolving {
277		let resolver = if self
278			.server
279			.config
280			.dns_passthru_domains
281			.is_match(name.as_str())
282		{
283			trace!(?name, "matched to passthru resolver");
284			&self.passthru.resolver
285		} else {
286			trace!(?name, "using hooked resolver");
287			&self.resolver
288		};
289
290		hooked_resolve(self.cache.clone(), self.server.clone(), resolver.clone(), name).boxed()
291	}
292}
293
294impl Resolve for Passthru {
295	fn resolve(&self, name: Name) -> Resolving {
296		trace!(?name, "using passthru resolver");
297		resolve_to_reqwest(self.server.clone(), self.resolver.clone(), name).boxed()
298	}
299}
300
301#[tracing::instrument(
302	level = "debug",
303	skip_all,
304	fields(name = ?name.as_str())
305)]
306async fn hooked_resolve(
307	cache: Arc<Cache>,
308	server: Arc<Server>,
309	resolver: Arc<TokioResolver>,
310	name: Name,
311) -> Result<Addrs, Box<dyn std::error::Error + Send + Sync>> {
312	match cache.get_override(name.as_str()).await {
313		| Ok(cached) if cached.valid() => cached_to_reqwest(cached),
314		| Ok(CachedOverride { overriding, .. }) if overriding.is_some() =>
315			resolve_to_reqwest(
316				server,
317				resolver,
318				overriding
319					.as_deref()
320					.map(str::parse)
321					.expect("overriding is set for this record")
322					.expect("overriding is a valid internet name"),
323			)
324			.boxed()
325			.await,
326
327		| _ =>
328			resolve_to_reqwest(server, resolver, name)
329				.boxed()
330				.await,
331	}
332}
333
334async fn resolve_to_reqwest(
335	server: Arc<Server>,
336	resolver: Arc<TokioResolver>,
337	name: Name,
338) -> ResolvingResult {
339	use std::io::ErrorKind::Interrupted;
340
341	let handle_shutdown = || Box::new(io::Error::new(Interrupted, "Server shutting down"));
342
343	let handle_results = |results: LookupIp| -> Addrs {
344		let addrs = results
345			.iter()
346			.map(|ip| SocketAddr::new(ip, 0))
347			.collect::<Vec<_>>()
348			.into_iter();
349
350		Box::new(addrs)
351	};
352
353	tokio::select! {
354		results = resolver.lookup_ip(name.as_str()) => Ok(handle_results(results?)),
355		() = server.until_shutdown() => Err(handle_shutdown()),
356	}
357}
358
359fn cached_to_reqwest(cached: CachedOverride) -> ResolvingResult {
360	let addrs = cached
361		.ips
362		.into_iter()
363		.map(move |ip| SocketAddr::new(ip, cached.port));
364
365	Ok(Box::new(addrs))
366}