1use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
16use std::sync::Arc;
17
18use tokio::io::{AsyncReadExt, AsyncWriteExt};
19use tokio::net::{TcpListener, UdpSocket};
20
21pub const DEFAULT_DNS_PORT: u16 = 15353;
26
27const MAX_UDP_PAYLOAD: usize = 512;
29
30const MAX_TCP_MESSAGE: usize = 65535;
32
33const TTL: u32 = 60;
40
41const ACCEPT_ERROR_BACKOFF: std::time::Duration = std::time::Duration::from_millis(100);
46
47const TCP_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
52
53const REFUSAL_LOG_INTERVAL: std::time::Duration = std::time::Duration::from_secs(60);
55
56static REFUSED_TCP: crate::proxy::LogThrottle = crate::proxy::LogThrottle::new();
58
59const MAX_TCP_CONNECTIONS: usize = 64;
65
66const TCP_CONNECTION_LIFETIME: std::time::Duration = std::time::Duration::from_secs(60);
73
74const TYPE_A: u16 = 1;
75const TYPE_AAAA: u16 = 28;
76const CLASS_IN: u16 = 1;
77
78const RCODE_NOERROR: u16 = 0;
79const RCODE_FORMERR: u16 = 1;
80const RCODE_NOTIMP: u16 = 4;
81const RCODE_REFUSED: u16 = 5;
82
83const FLAG_QR: u16 = 0x8000;
84const FLAG_AA: u16 = 0x0400;
85const FLAG_TC: u16 = 0x0200;
86const FLAG_RD: u16 = 0x0100;
87
88#[derive(Clone, Debug, PartialEq, Eq)]
90pub struct ResolverConfig {
91 pub tld: String,
93 pub ipv4: Option<Ipv4Addr>,
99 pub ipv6: Option<Ipv6Addr>,
106}
107
108impl ResolverConfig {
109 pub fn loopback(tld: impl Into<String>) -> Self {
111 Self {
112 tld: tld.into(),
113 ipv4: Some(Ipv4Addr::LOCALHOST),
114 ipv6: None,
115 }
116 }
117
118 pub fn lan(tld: impl Into<String>, ip: Ipv4Addr) -> Self {
120 Self {
121 tld: tld.into(),
122 ipv4: Some(ip),
123 ipv6: None,
124 }
125 }
126
127 pub fn for_bind(tld: impl Into<String>, bind_ip: std::net::IpAddr) -> Self {
138 let tld = tld.into();
139 match bind_ip {
140 std::net::IpAddr::V4(ip) if ip.is_unspecified() => Self {
141 tld,
142 ipv4: Some(Ipv4Addr::LOCALHOST),
143 ipv6: None,
144 },
145 std::net::IpAddr::V4(ip) => Self {
146 tld,
147 ipv4: Some(ip),
148 ipv6: None,
149 },
150 std::net::IpAddr::V6(ip) if ip.is_unspecified() => Self {
151 tld,
152 ipv4: Some(Ipv4Addr::LOCALHOST),
153 ipv6: Some(Ipv6Addr::LOCALHOST),
154 },
155 std::net::IpAddr::V6(ip) => Self {
156 tld,
157 ipv4: None,
158 ipv6: Some(ip),
159 },
160 }
161 }
162
163 fn owns(&self, name: &str) -> bool {
169 super::owns_name(&self.tld, name)
172 }
173}
174
175#[derive(Debug, PartialEq, Eq)]
177struct Question {
178 name: String,
179 qtype: u16,
180 qclass: u16,
181 end: usize,
183}
184
185fn parse_question(msg: &[u8]) -> Option<Question> {
190 let mut pos = 12;
191 let mut name = String::new();
192 loop {
193 let len = *msg.get(pos)? as usize;
194 pos += 1;
195 if len == 0 {
196 break;
197 }
198 if len & 0xC0 != 0 {
201 return None;
202 }
203 let label = msg.get(pos..pos + len)?;
204 pos += len;
205 if !name.is_empty() {
206 name.push('.');
207 }
208 name.push_str(&String::from_utf8_lossy(label));
209 if name.len() > 255 {
210 return None;
211 }
212 }
213 let qtype = u16::from_be_bytes([*msg.get(pos)?, *msg.get(pos + 1)?]);
214 let qclass = u16::from_be_bytes([*msg.get(pos + 2)?, *msg.get(pos + 3)?]);
215 Some(Question {
216 name,
217 qtype,
218 qclass,
219 end: pos + 4,
220 })
221}
222
223fn header_only(id: u16, flags: u16, rcode: u16) -> Vec<u8> {
225 let mut out = Vec::with_capacity(12);
226 out.extend_from_slice(&id.to_be_bytes());
227 out.extend_from_slice(&(flags | rcode).to_be_bytes());
228 out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&0u16.to_be_bytes()); out
233}
234
235pub fn handle_query(query: &[u8], cfg: &ResolverConfig) -> Option<Vec<u8>> {
240 if query.len() < 12 {
241 return None;
242 }
243 let id = u16::from_be_bytes([query[0], query[1]]);
244 let req_flags = u16::from_be_bytes([query[2], query[3]]);
245 if req_flags & FLAG_QR != 0 {
246 return None;
248 }
249 let opcode = req_flags & 0x7800;
250 let qdcount = u16::from_be_bytes([query[4], query[5]]);
251
252 let base_flags = FLAG_QR | opcode | (req_flags & FLAG_RD);
255
256 if opcode != 0 {
258 return Some(header_only(id, base_flags, RCODE_NOTIMP));
259 }
260 if qdcount != 1 {
261 return Some(header_only(id, base_flags, RCODE_FORMERR));
262 }
263 let Some(q) = parse_question(query) else {
264 return Some(header_only(id, base_flags, RCODE_FORMERR));
265 };
266
267 let owned = q.qclass == CLASS_IN && cfg.owns(&q.name);
268 let answer = if !owned {
269 None
270 } else {
271 match q.qtype {
272 TYPE_A => cfg.ipv4.map(|ip| ip.octets().to_vec()),
273 TYPE_AAAA => cfg.ipv6.map(|ip| ip.octets().to_vec()),
274 _ => None,
275 }
276 };
277
278 let rcode = if owned { RCODE_NOERROR } else { RCODE_REFUSED };
286 let ancount: u16 = u16::from(answer.is_some());
287
288 let mut out = Vec::with_capacity(query.len() + 32);
289 out.extend_from_slice(&id.to_be_bytes());
290 let aa = if owned { FLAG_AA } else { 0 };
292 out.extend_from_slice(&(base_flags | aa | rcode).to_be_bytes());
293 out.extend_from_slice(&1u16.to_be_bytes()); out.extend_from_slice(&ancount.to_be_bytes());
295 out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&0u16.to_be_bytes()); out.extend_from_slice(&query[12..q.end]); if let Some(rdata) = answer {
300 out.extend_from_slice(&[0xC0, 0x0C]);
303 out.extend_from_slice(&q.qtype.to_be_bytes());
304 out.extend_from_slice(&CLASS_IN.to_be_bytes());
305 out.extend_from_slice(&TTL.to_be_bytes());
306 out.extend_from_slice(&(rdata.len() as u16).to_be_bytes());
307 out.extend_from_slice(&rdata);
308 }
309
310 Some(out)
311}
312
313fn truncate_for_udp(mut resp: Vec<u8>) -> Vec<u8> {
316 if resp.len() <= MAX_UDP_PAYLOAD {
317 return resp;
318 }
319 let flags = u16::from_be_bytes([resp[2], resp[3]]) | FLAG_TC;
320 resp[2..4].copy_from_slice(&flags.to_be_bytes());
321 resp[6..8].copy_from_slice(&0u16.to_be_bytes());
323 resp.truncate(MAX_UDP_PAYLOAD);
324 resp
325}
326
327static ACTIVE_CONFIG: std::sync::RwLock<Option<Arc<std::sync::RwLock<ResolverConfig>>>> =
336 std::sync::RwLock::new(None);
337
338pub fn update_lan_ip(ip: Ipv4Addr) {
342 let cfg = match ACTIVE_CONFIG.read() {
343 Ok(active) => active.clone(),
344 Err(e) => {
345 log::warn!("Could not read the active DNS resolver config: {e}");
346 return;
347 }
348 };
349 let Some(cfg) = cfg else {
350 return;
351 };
352 match cfg.write() {
353 Ok(mut cfg) if cfg.ipv4 != Some(ip) => {
354 log::info!("DNS resolver now answering *.{} with {ip}", cfg.tld);
355 cfg.ipv4 = Some(ip);
356 }
357 Ok(_) => {}
358 Err(e) => log::warn!("Could not update the DNS resolver address: {e}"),
359 }
360}
361
362pub async fn serve(
367 cfg: ResolverConfig,
368 addr: SocketAddr,
369 bind_tx: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
370 cancel: tokio_util::sync::CancellationToken,
371) -> crate::Result<()> {
372 let udp = match UdpSocket::bind(addr).await {
373 Ok(s) => s,
374 Err(e) => {
375 let msg = format!("DNS resolver failed to bind UDP {addr}: {e}");
376 let _ = bind_tx.send(Err(msg.clone()));
377 miette::bail!("{msg}");
378 }
379 };
380 let tcp = match TcpListener::bind(addr).await {
381 Ok(l) => l,
382 Err(e) => {
383 let msg = format!("DNS resolver failed to bind TCP {addr}: {e}");
384 let _ = bind_tx.send(Err(msg.clone()));
385 miette::bail!("{msg}");
386 }
387 };
388 let _ = bind_tx.send(Ok(()));
389 {
390 let answers = [
391 cfg.ipv4.map(|ip| ip.to_string()),
392 cfg.ipv6.map(|ip| ip.to_string()),
393 ]
394 .into_iter()
395 .flatten()
396 .collect::<Vec<_>>()
397 .join(", ");
398 log::info!(
399 "DNS resolver listening on {addr} (udp+tcp), answering *.{} with {answers}",
400 cfg.tld,
401 );
402 }
403
404 let cfg = Arc::new(std::sync::RwLock::new(cfg));
407 match ACTIVE_CONFIG.write() {
408 Ok(mut active) => *active = Some(Arc::clone(&cfg)),
409 Err(e) => log::warn!("Could not publish the DNS resolver config: {e}"),
410 }
411
412 fn answer(cfg: &std::sync::RwLock<ResolverConfig>, query: &[u8]) -> Option<Vec<u8>> {
414 match cfg.read() {
415 Ok(cfg) => handle_query(query, &cfg),
416 Err(e) => {
417 log::warn!("DNS resolver config lock poisoned: {e}");
418 None
419 }
420 }
421 }
422
423 let mut buf = vec![0u8; MAX_UDP_PAYLOAD];
424 let mut conns: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
425 loop {
426 while conns.try_join_next().is_some() {}
427 tokio::select! {
428 recv = udp.recv_from(&mut buf) => {
429 let (len, peer) = match recv {
430 Ok(v) => v,
431 Err(e) => {
432 log::debug!("DNS UDP receive error: {e}");
433 tokio::select! {
440 _ = tokio::time::sleep(ACCEPT_ERROR_BACKOFF) => continue,
441 _ = cancel.cancelled() => {
442 log::info!("DNS resolver shutting down");
443 break;
444 }
445 }
446 }
447 };
448 if let Some(resp) = answer(&cfg, &buf[..len])
449 && let Err(e) = udp.send_to(&truncate_for_udp(resp), peer).await
450 {
451 log::debug!("DNS UDP send error to {peer}: {e}");
452 }
453 }
454 accept = tcp.accept() => {
455 let (stream, peer) = match accept {
456 Ok(v) => v,
457 Err(e) => {
458 log::debug!("DNS TCP accept error: {e}");
459 tokio::select! {
472 _ = tokio::time::sleep(ACCEPT_ERROR_BACKOFF) => continue,
473 _ = cancel.cancelled() => {
474 log::info!("DNS resolver shutting down");
475 break;
476 }
477 }
478 }
479 };
480 while conns.try_join_next().is_some() {}
489 if conns.len() >= MAX_TCP_CONNECTIONS {
490 if let Some(suppressed) = REFUSED_TCP.allow(REFUSAL_LOG_INTERVAL) {
494 log::warn!(
495 "DNS resolver refused a TCP connection from {peer}: \
496 {MAX_TCP_CONNECTIONS} already in flight \
497 ({suppressed} similar refusals since the last message)"
498 );
499 }
500 drop(stream);
501 continue;
502 }
503 let cfg = Arc::clone(&cfg);
504 conns.spawn(async move {
505 match tokio::time::timeout(
506 TCP_CONNECTION_LIFETIME,
507 serve_tcp_conn(stream, &cfg, TCP_IDLE_TIMEOUT),
508 )
509 .await
510 {
511 Ok(Ok(())) => {}
512 Ok(Err(e)) => log::debug!("DNS TCP connection from {peer} ended: {e}"),
513 Err(_) => log::debug!(
514 "DNS TCP connection from {peer} closed after \
515 {TCP_CONNECTION_LIFETIME:?}"
516 ),
517 }
518 });
519 }
520 _ = cancel.cancelled() => {
521 log::info!("DNS resolver shutting down");
522 break;
523 }
524 }
525 }
526 conns.abort_all();
527 if let Ok(mut active) = ACTIVE_CONFIG.write()
529 && active.as_ref().is_some_and(|c| Arc::ptr_eq(c, &cfg))
530 {
531 *active = None;
532 }
533 Ok(())
534}
535
536async fn serve_tcp_conn(
541 mut stream: tokio::net::TcpStream,
542 cfg: &std::sync::RwLock<ResolverConfig>,
543 idle: std::time::Duration,
544) -> std::io::Result<()> {
545 async fn read_exact_timeout(
547 stream: &mut tokio::net::TcpStream,
548 buf: &mut [u8],
549 idle: std::time::Duration,
550 ) -> std::io::Result<()> {
551 tokio::time::timeout(idle, stream.read_exact(buf))
552 .await
553 .map_err(|_| {
554 std::io::Error::new(std::io::ErrorKind::TimedOut, "idle DNS connection")
555 })??;
556 Ok(())
557 }
558
559 loop {
560 let mut len_buf = [0u8; 2];
561 match read_exact_timeout(&mut stream, &mut len_buf, idle).await {
562 Ok(()) => {}
563 Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(()),
565 Err(e) => return Err(e),
566 }
567 let len = usize::from(u16::from_be_bytes(len_buf));
568 if len == 0 || len > MAX_TCP_MESSAGE {
569 return Ok(());
570 }
571 let mut msg = vec![0u8; len];
572 read_exact_timeout(&mut stream, &mut msg, idle).await?;
574 let Some(resp) = (match cfg.read() {
575 Ok(cfg) => handle_query(&msg, &cfg),
576 Err(_) => None,
577 }) else {
578 continue;
579 };
580 let Ok(resp_len) = u16::try_from(resp.len()) else {
589 log::debug!(
590 "DNS reply of {} bytes cannot be framed over TCP; closing the connection",
591 resp.len()
592 );
593 return Ok(());
594 };
595 tokio::time::timeout(idle, async {
599 stream.write_all(&resp_len.to_be_bytes()).await?;
600 stream.write_all(&resp).await?;
601 stream.flush().await
602 })
603 .await
604 .map_err(|_| {
605 std::io::Error::new(
606 std::io::ErrorKind::TimedOut,
607 "DNS client not reading replies",
608 )
609 })??;
610 }
611}
612
613pub fn config_from_settings(
619 s: &crate::settings::Settings,
620 lan_ip: Option<Ipv4Addr>,
621) -> ResolverConfig {
622 let lan_enabled = s.proxy.lan || !s.proxy.lan_ip.is_empty();
623 let tld = crate::proxy::effective_tld(s).to_string();
627 if lan_enabled {
628 return match lan_ip {
630 Some(ip) => ResolverConfig::lan(tld, ip),
631 None => ResolverConfig::loopback(tld),
632 };
633 }
634 let bind_ip = s
638 .proxy
639 .host
640 .parse()
641 .unwrap_or(std::net::IpAddr::V4(Ipv4Addr::LOCALHOST));
642 ResolverConfig::for_bind(tld, bind_ip)
643}
644
645pub fn dns_port(s: &crate::settings::Settings) -> u16 {
647 u16::try_from(s.proxy.dns_port)
648 .ok()
649 .filter(|&p| p > 0)
650 .unwrap_or_else(|| {
651 log::warn!(
652 "proxy.dns_port {} is out of valid port range (1-65535), using {DEFAULT_DNS_PORT}",
653 s.proxy.dns_port
654 );
655 DEFAULT_DNS_PORT
656 })
657}
658
659#[cfg(test)]
667pub(crate) async fn free_udp_and_tcp_addr() -> SocketAddr {
668 for _ in 0..50 {
669 let udp = tokio::net::UdpSocket::bind("127.0.0.1:0")
670 .await
671 .expect("bind UDP to an OS-chosen port");
672 let addr = udp.local_addr().expect("UDP local address");
673 if tokio::net::TcpListener::bind(addr).await.is_ok() {
674 return addr;
675 }
676 }
677 panic!("no loopback port free for both UDP and TCP after 50 tries");
678}
679
680#[cfg(test)]
681mod tests {
682 use super::*;
683
684 #[tokio::test]
685 async fn free_udp_and_tcp_addr_is_bindable_by_both() {
686 let taken = |r: &std::io::Result<_>| matches!(r, Err(e) if e.kind() == std::io::ErrorKind::AddrInUse);
692 for _ in 0..100 {
693 let addr = free_udp_and_tcp_addr().await;
694 let udp = tokio::net::UdpSocket::bind(addr).await.map(drop);
695 assert!(udp.is_ok() || taken(&udp), "UDP {addr}: {udp:?}");
696 let tcp = tokio::net::TcpListener::bind(addr).await.map(drop);
697 assert!(tcp.is_ok() || taken(&tcp), "TCP {addr}: {tcp:?}");
698 }
699 }
700
701 fn query(id: u16, name: &str, qtype: u16) -> Vec<u8> {
703 let mut out = Vec::new();
704 out.extend_from_slice(&id.to_be_bytes());
705 out.extend_from_slice(&FLAG_RD.to_be_bytes());
706 out.extend_from_slice(&1u16.to_be_bytes());
707 out.extend_from_slice(&0u16.to_be_bytes());
708 out.extend_from_slice(&0u16.to_be_bytes());
709 out.extend_from_slice(&0u16.to_be_bytes());
710 for label in name.split('.') {
711 out.push(label.len() as u8);
712 out.extend_from_slice(label.as_bytes());
713 }
714 out.push(0);
715 out.extend_from_slice(&qtype.to_be_bytes());
716 out.extend_from_slice(&CLASS_IN.to_be_bytes());
717 out
718 }
719
720 fn rcode(resp: &[u8]) -> u16 {
721 u16::from_be_bytes([resp[2], resp[3]]) & 0x000F
722 }
723
724 fn ancount(resp: &[u8]) -> u16 {
725 u16::from_be_bytes([resp[6], resp[7]])
726 }
727
728 fn rdata(resp: &[u8]) -> Vec<u8> {
730 let q = parse_question(resp).expect("response echoes the question");
731 let rdlen = usize::from(u16::from_be_bytes([resp[q.end + 10], resp[q.end + 11]]));
732 resp[q.end + 12..q.end + 12 + rdlen].to_vec()
733 }
734
735 fn cfg() -> ResolverConfig {
737 ResolverConfig::for_bind("localhost", "::".parse().unwrap())
738 }
739
740 #[test]
741 fn a_query_under_tld_answers_loopback() {
742 let resp = handle_query(&query(0x1234, "myapp.localhost", TYPE_A), &cfg()).unwrap();
743 assert_eq!(&resp[0..2], &0x1234u16.to_be_bytes());
744 assert_eq!(rcode(&resp), RCODE_NOERROR);
745 assert_eq!(ancount(&resp), 1);
746 assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
747 let flags = u16::from_be_bytes([resp[2], resp[3]]);
749 assert_eq!(flags & FLAG_QR, FLAG_QR);
750 assert_eq!(flags & FLAG_AA, FLAG_AA);
751 assert_eq!(flags & FLAG_RD, FLAG_RD);
752 }
753
754 #[test]
755 fn a_query_answers_multi_level_names() {
756 let resp = handle_query(
758 &query(1, "core.fix-refs.entiredb.localhost", TYPE_A),
759 &cfg(),
760 )
761 .unwrap();
762 assert_eq!(rcode(&resp), RCODE_NOERROR);
763 assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
764 }
765
766 #[test]
767 fn tld_apex_resolves() {
768 let resp = handle_query(&query(1, "localhost", TYPE_A), &cfg()).unwrap();
769 assert_eq!(rcode(&resp), RCODE_NOERROR);
770 assert_eq!(ancount(&resp), 1);
771 }
772
773 #[test]
774 fn matching_is_case_insensitive() {
775 let resp = handle_query(&query(1, "MyApp.LOCALHOST", TYPE_A), &cfg()).unwrap();
776 assert_eq!(rcode(&resp), RCODE_NOERROR);
777 assert_eq!(ancount(&resp), 1);
778 }
779
780 #[test]
781 fn aaaa_query_answers_ipv6_loopback() {
782 let resp = handle_query(&query(1, "myapp.localhost", TYPE_AAAA), &cfg()).unwrap();
783 assert_eq!(rcode(&resp), RCODE_NOERROR);
784 assert_eq!(ancount(&resp), 1);
785 assert_eq!(rdata(&resp), Ipv6Addr::LOCALHOST.octets().to_vec());
786 }
787
788 #[test]
789 fn lan_mode_answers_lan_ip_and_nodata_for_aaaa() {
790 let cfg = ResolverConfig::lan("local", Ipv4Addr::new(192, 168, 1, 42));
791 let a = handle_query(&query(1, "myapp.local", TYPE_A), &cfg).unwrap();
792 assert_eq!(rdata(&a), vec![192, 168, 1, 42]);
793
794 let aaaa = handle_query(&query(1, "myapp.local", TYPE_AAAA), &cfg).unwrap();
797 assert_eq!(rcode(&aaaa), RCODE_NOERROR);
798 assert_eq!(ancount(&aaaa), 0);
799 }
800
801 #[test]
802 fn name_outside_tld_is_refused_not_nxdomain() {
803 let resp = handle_query(&query(1, "example.com", TYPE_A), &cfg()).unwrap();
804 assert_eq!(rcode(&resp), RCODE_REFUSED);
807 assert_eq!(ancount(&resp), 0);
808 assert_eq!(u16::from_be_bytes([resp[2], resp[3]]) & FLAG_AA, 0);
810 }
811
812 #[test]
813 fn tld_suffix_without_label_boundary_is_refused() {
814 let resp = handle_query(&query(1, "notlocalhost", TYPE_A), &cfg()).unwrap();
816 assert_eq!(rcode(&resp), RCODE_REFUSED);
817 }
818
819 #[test]
820 fn lan_mode_serves_ipv4_whatever_proxy_host_says() {
821 let cfg = ResolverConfig::lan("local", Ipv4Addr::new(192, 168, 1, 42));
825 assert_eq!(cfg.ipv4, Some(Ipv4Addr::new(192, 168, 1, 42)));
826 assert_eq!(cfg.ipv6, None);
827 let fallback = ResolverConfig::loopback("local");
829 assert_eq!(fallback.ipv4, Some(Ipv4Addr::LOCALHOST));
830 assert_eq!(fallback.ipv6, None);
831 }
832
833 #[test]
834 fn answers_name_only_addresses_the_proxy_listens_on() {
835 use std::net::IpAddr;
836
837 let v6 = ResolverConfig::for_bind("test", IpAddr::V6(Ipv6Addr::LOCALHOST));
840 assert_eq!(v6.ipv4, None);
841 assert_eq!(v6.ipv6, Some(Ipv6Addr::LOCALHOST));
842 let a = handle_query(&query(1, "x.test", TYPE_A), &v6).unwrap();
843 assert_eq!(rcode(&a), RCODE_NOERROR);
844 assert_eq!(ancount(&a), 0);
845 let aaaa = handle_query(&query(1, "x.test", TYPE_AAAA), &v6).unwrap();
846 assert_eq!(rdata(&aaaa), Ipv6Addr::LOCALHOST.octets().to_vec());
847
848 let specific_v4 = ResolverConfig::for_bind("test", "192.168.1.5".parse().unwrap());
851 assert_eq!(specific_v4.ipv4, Some(Ipv4Addr::new(192, 168, 1, 5)));
852 assert_eq!(specific_v4.ipv6, None);
853 let specific_v6 = ResolverConfig::for_bind("test", "fd00::1".parse().unwrap());
854 assert_eq!(specific_v6.ipv4, None);
855 assert_eq!(specific_v6.ipv6, Some("fd00::1".parse().unwrap()));
856
857 let any_v4 = ResolverConfig::for_bind("test", "0.0.0.0".parse().unwrap());
860 assert_eq!(any_v4.ipv4, Some(Ipv4Addr::LOCALHOST));
861 assert_eq!(any_v4.ipv6, None);
862 let any_v6 = ResolverConfig::for_bind("test", "::".parse().unwrap());
863 assert_eq!(any_v6.ipv4, Some(Ipv4Addr::LOCALHOST));
864 assert_eq!(any_v6.ipv6, Some(Ipv6Addr::LOCALHOST));
865 }
866
867 #[test]
868 fn aaaa_is_nodata_unless_the_proxy_listens_on_ipv6() {
869 let v4_only = ResolverConfig::loopback("localhost");
872 let resp = handle_query(&query(1, "myapp.localhost", TYPE_AAAA), &v4_only).unwrap();
873 assert_eq!(rcode(&resp), RCODE_NOERROR);
874 assert_eq!(ancount(&resp), 0);
875 let a = handle_query(&query(1, "myapp.localhost", TYPE_A), &v4_only).unwrap();
877 assert_eq!(ancount(&a), 1);
878 }
879
880 #[test]
881 fn unsupported_record_type_under_tld_is_nodata() {
882 const TYPE_MX: u16 = 15;
883 let resp = handle_query(&query(1, "myapp.localhost", TYPE_MX), &cfg()).unwrap();
884 assert_eq!(rcode(&resp), RCODE_NOERROR);
885 assert_eq!(ancount(&resp), 0);
886 }
887
888 #[test]
889 fn non_internet_class_is_refused() {
890 let mut q = query(1, "myapp.localhost", TYPE_A);
891 let len = q.len();
892 q[len - 2..].copy_from_slice(&3u16.to_be_bytes()); let resp = handle_query(&q, &cfg()).unwrap();
894 assert_eq!(rcode(&resp), RCODE_REFUSED);
895 }
896
897 #[test]
898 fn malformed_and_unsupported_messages() {
899 assert!(handle_query(&[0u8; 4], &cfg()).is_none());
901 let mut resp_msg = query(1, "myapp.localhost", TYPE_A);
903 resp_msg[2] |= 0x80;
904 assert!(handle_query(&resp_msg, &cfg()).is_none());
905 let q = query(1, "myapp.localhost", TYPE_A);
907 let resp = handle_query(&q[..16], &cfg()).unwrap();
908 assert_eq!(rcode(&resp), RCODE_FORMERR);
909 let mut upd = query(1, "myapp.localhost", TYPE_A);
911 upd[2] |= 5 << 3;
912 let resp = handle_query(&upd, &cfg()).unwrap();
913 assert_eq!(rcode(&resp), RCODE_NOTIMP);
914 }
915
916 #[test]
917 fn compression_pointer_in_question_is_rejected() {
918 let mut q = query(1, "myapp.localhost", TYPE_A);
919 q[12] = 0xC0;
920 let resp = handle_query(&q, &cfg()).unwrap();
921 assert_eq!(rcode(&resp), RCODE_FORMERR);
922 }
923
924 #[tokio::test]
925 async fn an_idle_tcp_client_is_dropped_rather_than_held() {
926 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
930 let addr = listener.local_addr().unwrap();
931 let idle = std::time::Duration::from_millis(50);
932
933 let server = tokio::spawn(async move {
934 let (stream, _) = listener.accept().await.unwrap();
935 serve_tcp_conn(stream, &std::sync::RwLock::new(cfg()), idle).await
936 });
937
938 let mut client = tokio::net::TcpStream::connect(addr).await.unwrap();
939 client.write_all(&16u16.to_be_bytes()).await.unwrap();
941
942 let err = tokio::time::timeout(std::time::Duration::from_secs(5), server)
943 .await
944 .expect("the handler should give up on its own")
945 .unwrap()
946 .expect_err("an idle connection is an error, not a clean close");
947 assert_eq!(err.kind(), std::io::ErrorKind::TimedOut);
948
949 let mut buf = [0u8; 1];
951 let n = tokio::time::timeout(std::time::Duration::from_secs(5), client.read(&mut buf))
952 .await
953 .expect("the connection is already closed")
954 .unwrap();
955 assert_eq!(n, 0);
956 }
957
958 #[test]
959 fn the_tcp_connection_cap_is_bounded() {
960 assert!((8..=1024).contains(&MAX_TCP_CONNECTIONS));
963 }
964
965 #[tokio::test]
966 async fn serves_over_udp_and_tcp() {
967 let cancel = tokio_util::sync::CancellationToken::new();
968 let (tx, rx) = tokio::sync::oneshot::channel();
969 let addr = super::free_udp_and_tcp_addr().await;
971
972 let task = tokio::spawn({
973 let cancel = cancel.clone();
974 async move { serve(cfg(), addr, tx, cancel).await }
975 });
976 rx.await.unwrap().expect("resolver binds");
977
978 let q = query(0x4242, "deep.nested.myapp.localhost", TYPE_A);
979
980 let sock = UdpSocket::bind("127.0.0.1:0").await.unwrap();
981 sock.send_to(&q, addr).await.unwrap();
982 let mut buf = [0u8; 512];
983 let (n, _) =
984 tokio::time::timeout(std::time::Duration::from_secs(5), sock.recv_from(&mut buf))
985 .await
986 .expect("udp reply arrives")
987 .unwrap();
988 assert_eq!(rdata(&buf[..n]), vec![127, 0, 0, 1]);
989
990 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
991 stream
992 .write_all(&(q.len() as u16).to_be_bytes())
993 .await
994 .unwrap();
995 stream.write_all(&q).await.unwrap();
996 let mut len_buf = [0u8; 2];
997 stream.read_exact(&mut len_buf).await.unwrap();
998 let mut resp = vec![0u8; usize::from(u16::from_be_bytes(len_buf))];
999 stream.read_exact(&mut resp).await.unwrap();
1000 assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
1001 stream
1003 .write_all(&(q.len() as u16).to_be_bytes())
1004 .await
1005 .unwrap();
1006 stream.write_all(&q).await.unwrap();
1007 stream.read_exact(&mut len_buf).await.unwrap();
1008 let mut resp2 = vec![0u8; usize::from(u16::from_be_bytes(len_buf))];
1009 stream.read_exact(&mut resp2).await.unwrap();
1010 assert_eq!(rcode(&resp2), RCODE_NOERROR);
1011 drop(stream);
1012
1013 cancel.cancel();
1014 tokio::time::timeout(std::time::Duration::from_secs(5), task)
1015 .await
1016 .expect("the resolver did not stop within 5s of cancellation")
1017 .expect("the resolver task panicked")
1018 .expect("the resolver returned an error");
1019 }
1020}