1use std::net::{IpAddr, SocketAddr};
16
17use async_trait::async_trait;
18use hickory_resolver::{
19 TokioResolver,
20 config::{ConnectionConfig, NameServerConfig, ResolverConfig, ResolverOpts},
21 net::runtime::TokioRuntimeProvider,
22 proto::rr::{Name, RData, rdata::TXT},
23};
24
25use crate::config::DnsConfig;
26
27#[async_trait]
32pub trait Resolver: Send + Sync {
33 async fn reverse(&self, ip: IpAddr) -> Result<Vec<String>, String>;
35
36 async fn forward(&self, name: &str) -> Result<Vec<IpAddr>, String>;
38
39 async fn txt(&self, name: &str) -> Result<Vec<String>, String>;
45}
46
47pub struct HickoryResolver {
49 inner: TokioResolver,
50}
51
52impl HickoryResolver {
53 pub fn from_system() -> anyhow::Result<Self> {
56 Self::build(None, |_| {})
57 }
58
59 pub fn from_system_uncached() -> anyhow::Result<Self> {
66 Self::build(None, |options| options.cache_size = 0)
67 }
68
69 pub fn from_address(addr: SocketAddr) -> anyhow::Result<Self> {
72 Self::build(Some(addr), |_| {})
73 }
74
75 pub fn from_address_uncached(addr: SocketAddr) -> anyhow::Result<Self> {
79 Self::build(Some(addr), |options| options.cache_size = 0)
80 }
81
82 fn build(
86 addr: Option<SocketAddr>,
87 configure: impl FnOnce(&mut ResolverOpts),
88 ) -> anyhow::Result<Self> {
89 let mut builder = match addr {
90 None => TokioResolver::builder_tokio()
91 .map_err(|error| anyhow::anyhow!("reading system resolver config: {error}"))?,
92 Some(addr) => {
93 let mut udp = ConnectionConfig::udp();
94 udp.port = addr.port();
95 let mut tcp = ConnectionConfig::tcp();
96 tcp.port = addr.port();
97 let name_server = NameServerConfig::new(addr.ip(), true, vec![udp, tcp]);
98 let config = ResolverConfig::from_parts(None, vec![], vec![name_server]);
99 TokioResolver::builder_with_config(config, TokioRuntimeProvider::default())
100 }
101 };
102 configure(builder.options_mut());
103 let inner = builder
104 .build()
105 .map_err(|error| anyhow::anyhow!("building resolver: {error}"))?;
106 Ok(Self { inner })
107 }
108}
109
110#[async_trait]
111impl Resolver for HickoryResolver {
112 async fn reverse(&self, ip: IpAddr) -> Result<Vec<String>, String> {
126 let lookup = match self.inner.reverse_lookup(Name::from(ip)).await {
128 Ok(lookup) => lookup,
129 Err(error) if error.is_no_records_found() => return Ok(Vec::new()),
130 Err(error) => return Err(error.to_string()),
131 };
132
133 Ok(lookup
134 .answers()
135 .iter()
136 .filter_map(|record| match &record.data {
137 RData::PTR(ptr) => Some(strip_root(&ptr.to_string())),
138 _ => None,
139 })
140 .collect())
141 }
142
143 async fn forward(&self, name: &str) -> Result<Vec<IpAddr>, String> {
144 let lookup = match self.inner.lookup_ip(name).await {
145 Ok(lookup) => lookup,
146 Err(error) if error.is_no_records_found() => return Ok(Vec::new()),
147 Err(error) => return Err(error.to_string()),
148 };
149 Ok(lookup.iter().collect())
150 }
151
152 async fn txt(&self, name: &str) -> Result<Vec<String>, String> {
153 let lookup = match self.inner.txt_lookup(name).await {
154 Ok(lookup) => lookup,
155 Err(error) if error.is_no_records_found() => return Ok(Vec::new()),
156 Err(error) => return Err(error.to_string()),
157 };
158
159 Ok(lookup
160 .answers()
161 .iter()
162 .filter_map(|record| match &record.data {
163 RData::TXT(txt) => Some(join_character_strings(txt)),
164 _ => None,
165 })
166 .collect())
167 }
168}
169
170fn join_character_strings(txt: &TXT) -> String {
174 let bytes: Vec<u8> = txt
175 .txt_data
176 .iter()
177 .flat_map(|chunk| chunk.iter().copied())
178 .collect();
179 String::from_utf8_lossy(&bytes).into_owned()
180}
181
182pub(crate) fn strip_root(name: &str) -> String {
185 name.strip_suffix('.').unwrap_or(name).to_string()
186}
187
188pub fn resolver_addr(dns: &DnsConfig) -> anyhow::Result<Option<SocketAddr>> {
200 match dns.resolver.as_deref() {
201 None | Some("") => Ok(None),
202 Some(addr) => addr.parse::<SocketAddr>().map(Some).map_err(|error| {
203 anyhow::anyhow!("dns.resolver {addr:?} is not a valid address: {error}")
204 }),
205 }
206}
207
208pub(crate) async fn connect(
226 resolver: &dyn Resolver,
227 host: &str,
228 port: u16,
229) -> Result<tokio::net::TcpStream, String> {
230 let ips = match host.parse::<IpAddr>() {
231 Ok(ip) => vec![ip],
232 Err(_) => resolver.forward(host).await?,
233 };
234 if ips.is_empty() {
235 return Err(format!("no address found for {host}"));
236 }
237
238 let mut last_error = None;
239 for ip in ips {
240 match tokio::net::TcpStream::connect((ip, port)).await {
241 Ok(stream) => return Ok(stream),
242 Err(error) => last_error = Some(format!("{ip}: {error}")),
243 }
244 }
245 Err(last_error.unwrap_or_else(|| format!("no address found for {host}")))
246}
247
248#[cfg(test)]
249mod tests {
250 use super::*;
251
252 #[test]
253 fn strip_root_removes_the_trailing_dot() {
254 assert_eq!(strip_root("host.example.com."), "host.example.com");
255 assert_eq!(strip_root("host.example.com"), "host.example.com");
256 }
257
258 #[test]
262 fn txt_character_strings_are_concatenated() {
263 let single = TXT::new(vec!["one-piece".to_string()]);
264 assert_eq!(join_character_strings(&single), "one-piece");
265
266 let split = TXT::new(vec!["first".to_string(), "second".to_string()]);
267 assert_eq!(join_character_strings(&split), "firstsecond");
268
269 assert_eq!(join_character_strings(&TXT::new(vec![])), "");
271 }
272
273 #[test]
276 fn non_utf8_txt_data_is_lossy_rather_than_fatal() {
277 let raw = TXT::from_bytes(vec![&[0xff, 0xfe]]);
278 let joined = join_character_strings(&raw);
279 assert!(!joined.is_empty());
280 assert_ne!(joined, "expected-digest");
281 }
282
283 #[test]
287 fn both_constructors_read_the_same_system_configuration() {
288 assert_eq!(
289 HickoryResolver::from_system().is_ok(),
290 HickoryResolver::from_system_uncached().is_ok(),
291 );
292 }
293
294 #[test]
298 fn from_address_builds_without_reading_system_configuration() {
299 let addr: SocketAddr = "127.0.0.1:5300".parse().unwrap();
300 assert!(HickoryResolver::from_address(addr).is_ok());
301 assert!(HickoryResolver::from_address_uncached(addr).is_ok());
302 }
303
304 struct UnreachableResolver;
307
308 #[async_trait]
309 impl Resolver for UnreachableResolver {
310 async fn reverse(&self, _ip: IpAddr) -> Result<Vec<String>, String> {
311 unreachable!("connect never looks up PTR records")
312 }
313 async fn forward(&self, _name: &str) -> Result<Vec<IpAddr>, String> {
314 unreachable!("a literal IP must short-circuit before this is called")
315 }
316 async fn txt(&self, _name: &str) -> Result<Vec<String>, String> {
317 unreachable!("connect never looks up TXT records")
318 }
319 }
320
321 #[tokio::test]
322 async fn connect_short_circuits_a_literal_ip() {
323 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
324 let port = listener.local_addr().unwrap().port();
325 tokio::spawn(async move {
326 let _ = listener.accept().await;
327 });
328
329 assert!(
330 connect(&UnreachableResolver, "127.0.0.1", port)
331 .await
332 .is_ok()
333 );
334 }
335
336 struct StubForward(Vec<IpAddr>);
337
338 #[async_trait]
339 impl Resolver for StubForward {
340 async fn reverse(&self, _ip: IpAddr) -> Result<Vec<String>, String> {
341 unreachable!()
342 }
343 async fn forward(&self, _name: &str) -> Result<Vec<IpAddr>, String> {
344 Ok(self.0.clone())
345 }
346 async fn txt(&self, _name: &str) -> Result<Vec<String>, String> {
347 unreachable!()
348 }
349 }
350
351 #[tokio::test]
352 async fn connect_errors_when_forward_is_empty() {
353 let error = connect(&StubForward(vec![]), "example.com", 1234)
354 .await
355 .unwrap_err();
356 assert!(error.contains("example.com"), "{error}");
357 }
358
359 #[tokio::test]
366 async fn connect_falls_back_past_an_unreachable_first_address() {
367 let listener = tokio::net::TcpListener::bind("127.0.0.2:0").await.unwrap();
371 let port = listener.local_addr().unwrap().port();
372 tokio::spawn(async move {
373 let _ = listener.accept().await;
374 });
375
376 let unreachable_first = "127.0.0.1".parse().unwrap();
377 let reachable_second = "127.0.0.2".parse().unwrap();
378 let resolver = StubForward(vec![unreachable_first, reachable_second]);
379
380 assert!(connect(&resolver, "example.com", port).await.is_ok());
381 }
382
383 #[tokio::test]
384 async fn connect_errors_when_every_address_refuses() {
385 let port = {
387 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
388 listener.local_addr().unwrap().port()
389 };
390 let resolver = StubForward(vec!["127.0.0.1".parse().unwrap()]);
391 let error = connect(&resolver, "example.com", port).await.unwrap_err();
392 assert!(error.contains("127.0.0.1"), "{error}");
393 }
394
395 #[test]
396 fn resolver_addr_is_none_when_unset() {
397 assert!(resolver_addr(&DnsConfig::default()).unwrap().is_none());
398 }
399
400 #[test]
404 fn resolver_addr_treats_an_empty_string_as_unset() {
405 let dns = DnsConfig {
406 resolver: Some(String::new()),
407 };
408 assert!(resolver_addr(&dns).unwrap().is_none());
409 }
410
411 #[test]
412 fn resolver_addr_parses_a_valid_socket_address() {
413 let dns = DnsConfig {
414 resolver: Some("10.60.0.2:53".to_string()),
415 };
416 assert_eq!(
417 resolver_addr(&dns).unwrap(),
418 Some("10.60.0.2:53".parse().unwrap())
419 );
420 }
421
422 #[test]
423 fn resolver_addr_rejects_a_hostname_without_a_port() {
424 let dns = DnsConfig {
425 resolver: Some("not-an-address".to_string()),
426 };
427 let error = resolver_addr(&dns).unwrap_err().to_string();
428 assert!(error.contains("not-an-address"), "{error}");
429 }
430
431 mod loopback {
439 use super::*;
440 use hickory_proto::op::{Message, MessageType, OpCode, ResponseCode};
441 use hickory_proto::rr::rdata::{A, PTR};
442 use hickory_proto::rr::{DNSClass, Record, RecordType};
443 use hickory_proto::serialize::binary::{BinDecodable, BinEncodable};
444 use tokio::net::UdpSocket;
445
446 async fn spawn(answers: Vec<(&'static str, RecordType, RData)>) -> SocketAddr {
450 let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
451 let addr = socket.local_addr().unwrap();
452
453 tokio::spawn(async move {
454 let mut buffer = vec![0u8; 4096];
455 loop {
456 let Ok((read, peer)) = socket.recv_from(&mut buffer).await else {
457 return;
458 };
459 let Ok(request) = Message::from_bytes(&buffer[..read]) else {
460 continue;
461 };
462 let query = request.queries.first().cloned();
463
464 let mut response = Message::response(request.id, OpCode::Query);
465 response.metadata.message_type = MessageType::Response;
466 response.metadata.authoritative = true;
467 response.metadata.response_code = ResponseCode::NoError;
468 if let Some(query) = &query {
469 response.queries.push(query.clone());
470 for (name, record_type, data) in &answers {
471 let name = Name::from_utf8(name).unwrap();
472 if query.name() == &name && query.query_type() == *record_type {
473 let mut record = Record::from_rdata(name, 60, data.clone());
474 record.dns_class = DNSClass::IN;
475 response.answers.push(record);
476 }
477 }
478 }
479 let bytes = response.to_bytes().unwrap();
480 let _ = socket.send_to(&bytes, peer).await;
481 }
482 });
483
484 addr
485 }
486
487 #[tokio::test]
490 async fn reverse_returns_the_ptr_name_without_its_root_dot() {
491 let addr = spawn(vec![(
494 "10.2.0.192.in-addr.arpa.",
495 RecordType::PTR,
496 RData::PTR(PTR(Name::from_utf8("host.example.com.").unwrap())),
497 )])
498 .await;
499
500 let names = HickoryResolver::from_address(addr)
501 .unwrap()
502 .reverse("192.0.2.10".parse().unwrap())
503 .await
504 .unwrap();
505 assert_eq!(names, vec!["host.example.com".to_string()]);
506 }
507
508 #[tokio::test]
509 async fn forward_returns_the_addresses() {
510 let addr = spawn(vec![(
511 "host.example.com.",
512 RecordType::A,
513 RData::A(A("192.0.2.10".parse().unwrap())),
514 )])
515 .await;
516
517 let addresses = HickoryResolver::from_address(addr)
518 .unwrap()
519 .forward("host.example.com.")
520 .await
521 .unwrap();
522 assert!(
523 addresses.contains(&"192.0.2.10".parse::<IpAddr>().unwrap()),
524 "{addresses:?}"
525 );
526 }
527
528 #[tokio::test]
531 async fn txt_returns_the_concatenated_value() {
532 let addr = spawn(vec![(
533 "_acme-challenge.example.com.",
534 RecordType::TXT,
535 RData::TXT(TXT::new(vec!["first".to_string(), "second".to_string()])),
536 )])
537 .await;
538
539 let values = HickoryResolver::from_address_uncached(addr)
540 .unwrap()
541 .txt("_acme-challenge.example.com.")
542 .await
543 .unwrap();
544 assert_eq!(values, vec!["firstsecond".to_string()]);
545 }
546
547 #[tokio::test]
551 async fn an_empty_answer_is_no_records_rather_than_an_error() {
552 let addr = spawn(Vec::new()).await;
553 let resolver = HickoryResolver::from_address_uncached(addr).unwrap();
554
555 assert_eq!(
556 resolver
557 .reverse("192.0.2.10".parse().unwrap())
558 .await
559 .unwrap(),
560 Vec::<String>::new()
561 );
562 assert_eq!(
563 resolver.forward("nothing.example.com.").await.unwrap(),
564 Vec::<IpAddr>::new()
565 );
566 assert_eq!(
567 resolver.txt("nothing.example.com.").await.unwrap(),
568 Vec::<String>::new()
569 );
570 }
571 }
572}