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)]
660mod tests {
661 use super::*;
662
663 fn query(id: u16, name: &str, qtype: u16) -> Vec<u8> {
665 let mut out = Vec::new();
666 out.extend_from_slice(&id.to_be_bytes());
667 out.extend_from_slice(&FLAG_RD.to_be_bytes());
668 out.extend_from_slice(&1u16.to_be_bytes());
669 out.extend_from_slice(&0u16.to_be_bytes());
670 out.extend_from_slice(&0u16.to_be_bytes());
671 out.extend_from_slice(&0u16.to_be_bytes());
672 for label in name.split('.') {
673 out.push(label.len() as u8);
674 out.extend_from_slice(label.as_bytes());
675 }
676 out.push(0);
677 out.extend_from_slice(&qtype.to_be_bytes());
678 out.extend_from_slice(&CLASS_IN.to_be_bytes());
679 out
680 }
681
682 fn rcode(resp: &[u8]) -> u16 {
683 u16::from_be_bytes([resp[2], resp[3]]) & 0x000F
684 }
685
686 fn ancount(resp: &[u8]) -> u16 {
687 u16::from_be_bytes([resp[6], resp[7]])
688 }
689
690 fn rdata(resp: &[u8]) -> Vec<u8> {
692 let q = parse_question(resp).expect("response echoes the question");
693 let rdlen = usize::from(u16::from_be_bytes([resp[q.end + 10], resp[q.end + 11]]));
694 resp[q.end + 12..q.end + 12 + rdlen].to_vec()
695 }
696
697 fn cfg() -> ResolverConfig {
699 ResolverConfig::for_bind("localhost", "::".parse().unwrap())
700 }
701
702 #[test]
703 fn a_query_under_tld_answers_loopback() {
704 let resp = handle_query(&query(0x1234, "myapp.localhost", TYPE_A), &cfg()).unwrap();
705 assert_eq!(&resp[0..2], &0x1234u16.to_be_bytes());
706 assert_eq!(rcode(&resp), RCODE_NOERROR);
707 assert_eq!(ancount(&resp), 1);
708 assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
709 let flags = u16::from_be_bytes([resp[2], resp[3]]);
711 assert_eq!(flags & FLAG_QR, FLAG_QR);
712 assert_eq!(flags & FLAG_AA, FLAG_AA);
713 assert_eq!(flags & FLAG_RD, FLAG_RD);
714 }
715
716 #[test]
717 fn a_query_answers_multi_level_names() {
718 let resp = handle_query(
720 &query(1, "core.fix-refs.entiredb.localhost", TYPE_A),
721 &cfg(),
722 )
723 .unwrap();
724 assert_eq!(rcode(&resp), RCODE_NOERROR);
725 assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
726 }
727
728 #[test]
729 fn tld_apex_resolves() {
730 let resp = handle_query(&query(1, "localhost", TYPE_A), &cfg()).unwrap();
731 assert_eq!(rcode(&resp), RCODE_NOERROR);
732 assert_eq!(ancount(&resp), 1);
733 }
734
735 #[test]
736 fn matching_is_case_insensitive() {
737 let resp = handle_query(&query(1, "MyApp.LOCALHOST", TYPE_A), &cfg()).unwrap();
738 assert_eq!(rcode(&resp), RCODE_NOERROR);
739 assert_eq!(ancount(&resp), 1);
740 }
741
742 #[test]
743 fn aaaa_query_answers_ipv6_loopback() {
744 let resp = handle_query(&query(1, "myapp.localhost", TYPE_AAAA), &cfg()).unwrap();
745 assert_eq!(rcode(&resp), RCODE_NOERROR);
746 assert_eq!(ancount(&resp), 1);
747 assert_eq!(rdata(&resp), Ipv6Addr::LOCALHOST.octets().to_vec());
748 }
749
750 #[test]
751 fn lan_mode_answers_lan_ip_and_nodata_for_aaaa() {
752 let cfg = ResolverConfig::lan("local", Ipv4Addr::new(192, 168, 1, 42));
753 let a = handle_query(&query(1, "myapp.local", TYPE_A), &cfg).unwrap();
754 assert_eq!(rdata(&a), vec![192, 168, 1, 42]);
755
756 let aaaa = handle_query(&query(1, "myapp.local", TYPE_AAAA), &cfg).unwrap();
759 assert_eq!(rcode(&aaaa), RCODE_NOERROR);
760 assert_eq!(ancount(&aaaa), 0);
761 }
762
763 #[test]
764 fn name_outside_tld_is_refused_not_nxdomain() {
765 let resp = handle_query(&query(1, "example.com", TYPE_A), &cfg()).unwrap();
766 assert_eq!(rcode(&resp), RCODE_REFUSED);
769 assert_eq!(ancount(&resp), 0);
770 assert_eq!(u16::from_be_bytes([resp[2], resp[3]]) & FLAG_AA, 0);
772 }
773
774 #[test]
775 fn tld_suffix_without_label_boundary_is_refused() {
776 let resp = handle_query(&query(1, "notlocalhost", TYPE_A), &cfg()).unwrap();
778 assert_eq!(rcode(&resp), RCODE_REFUSED);
779 }
780
781 #[test]
782 fn lan_mode_serves_ipv4_whatever_proxy_host_says() {
783 let cfg = ResolverConfig::lan("local", Ipv4Addr::new(192, 168, 1, 42));
787 assert_eq!(cfg.ipv4, Some(Ipv4Addr::new(192, 168, 1, 42)));
788 assert_eq!(cfg.ipv6, None);
789 let fallback = ResolverConfig::loopback("local");
791 assert_eq!(fallback.ipv4, Some(Ipv4Addr::LOCALHOST));
792 assert_eq!(fallback.ipv6, None);
793 }
794
795 #[test]
796 fn answers_name_only_addresses_the_proxy_listens_on() {
797 use std::net::IpAddr;
798
799 let v6 = ResolverConfig::for_bind("test", IpAddr::V6(Ipv6Addr::LOCALHOST));
802 assert_eq!(v6.ipv4, None);
803 assert_eq!(v6.ipv6, Some(Ipv6Addr::LOCALHOST));
804 let a = handle_query(&query(1, "x.test", TYPE_A), &v6).unwrap();
805 assert_eq!(rcode(&a), RCODE_NOERROR);
806 assert_eq!(ancount(&a), 0);
807 let aaaa = handle_query(&query(1, "x.test", TYPE_AAAA), &v6).unwrap();
808 assert_eq!(rdata(&aaaa), Ipv6Addr::LOCALHOST.octets().to_vec());
809
810 let specific_v4 = ResolverConfig::for_bind("test", "192.168.1.5".parse().unwrap());
813 assert_eq!(specific_v4.ipv4, Some(Ipv4Addr::new(192, 168, 1, 5)));
814 assert_eq!(specific_v4.ipv6, None);
815 let specific_v6 = ResolverConfig::for_bind("test", "fd00::1".parse().unwrap());
816 assert_eq!(specific_v6.ipv4, None);
817 assert_eq!(specific_v6.ipv6, Some("fd00::1".parse().unwrap()));
818
819 let any_v4 = ResolverConfig::for_bind("test", "0.0.0.0".parse().unwrap());
822 assert_eq!(any_v4.ipv4, Some(Ipv4Addr::LOCALHOST));
823 assert_eq!(any_v4.ipv6, None);
824 let any_v6 = ResolverConfig::for_bind("test", "::".parse().unwrap());
825 assert_eq!(any_v6.ipv4, Some(Ipv4Addr::LOCALHOST));
826 assert_eq!(any_v6.ipv6, Some(Ipv6Addr::LOCALHOST));
827 }
828
829 #[test]
830 fn aaaa_is_nodata_unless_the_proxy_listens_on_ipv6() {
831 let v4_only = ResolverConfig::loopback("localhost");
834 let resp = handle_query(&query(1, "myapp.localhost", TYPE_AAAA), &v4_only).unwrap();
835 assert_eq!(rcode(&resp), RCODE_NOERROR);
836 assert_eq!(ancount(&resp), 0);
837 let a = handle_query(&query(1, "myapp.localhost", TYPE_A), &v4_only).unwrap();
839 assert_eq!(ancount(&a), 1);
840 }
841
842 #[test]
843 fn unsupported_record_type_under_tld_is_nodata() {
844 const TYPE_MX: u16 = 15;
845 let resp = handle_query(&query(1, "myapp.localhost", TYPE_MX), &cfg()).unwrap();
846 assert_eq!(rcode(&resp), RCODE_NOERROR);
847 assert_eq!(ancount(&resp), 0);
848 }
849
850 #[test]
851 fn non_internet_class_is_refused() {
852 let mut q = query(1, "myapp.localhost", TYPE_A);
853 let len = q.len();
854 q[len - 2..].copy_from_slice(&3u16.to_be_bytes()); let resp = handle_query(&q, &cfg()).unwrap();
856 assert_eq!(rcode(&resp), RCODE_REFUSED);
857 }
858
859 #[test]
860 fn malformed_and_unsupported_messages() {
861 assert!(handle_query(&[0u8; 4], &cfg()).is_none());
863 let mut resp_msg = query(1, "myapp.localhost", TYPE_A);
865 resp_msg[2] |= 0x80;
866 assert!(handle_query(&resp_msg, &cfg()).is_none());
867 let q = query(1, "myapp.localhost", TYPE_A);
869 let resp = handle_query(&q[..16], &cfg()).unwrap();
870 assert_eq!(rcode(&resp), RCODE_FORMERR);
871 let mut upd = query(1, "myapp.localhost", TYPE_A);
873 upd[2] |= 5 << 3;
874 let resp = handle_query(&upd, &cfg()).unwrap();
875 assert_eq!(rcode(&resp), RCODE_NOTIMP);
876 }
877
878 #[test]
879 fn compression_pointer_in_question_is_rejected() {
880 let mut q = query(1, "myapp.localhost", TYPE_A);
881 q[12] = 0xC0;
882 let resp = handle_query(&q, &cfg()).unwrap();
883 assert_eq!(rcode(&resp), RCODE_FORMERR);
884 }
885
886 #[tokio::test]
887 async fn an_idle_tcp_client_is_dropped_rather_than_held() {
888 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
892 let addr = listener.local_addr().unwrap();
893 let idle = std::time::Duration::from_millis(50);
894
895 let server = tokio::spawn(async move {
896 let (stream, _) = listener.accept().await.unwrap();
897 serve_tcp_conn(stream, &std::sync::RwLock::new(cfg()), idle).await
898 });
899
900 let mut client = tokio::net::TcpStream::connect(addr).await.unwrap();
901 client.write_all(&16u16.to_be_bytes()).await.unwrap();
903
904 let err = tokio::time::timeout(std::time::Duration::from_secs(5), server)
905 .await
906 .expect("the handler should give up on its own")
907 .unwrap()
908 .expect_err("an idle connection is an error, not a clean close");
909 assert_eq!(err.kind(), std::io::ErrorKind::TimedOut);
910
911 let mut buf = [0u8; 1];
913 let n = tokio::time::timeout(std::time::Duration::from_secs(5), client.read(&mut buf))
914 .await
915 .expect("the connection is already closed")
916 .unwrap();
917 assert_eq!(n, 0);
918 }
919
920 #[test]
921 fn the_tcp_connection_cap_is_bounded() {
922 assert!((8..=1024).contains(&MAX_TCP_CONNECTIONS));
925 }
926
927 #[tokio::test]
928 async fn serves_over_udp_and_tcp() {
929 let cancel = tokio_util::sync::CancellationToken::new();
930 let (tx, rx) = tokio::sync::oneshot::channel();
931 let probe = TcpListener::bind("127.0.0.1:0").await.unwrap();
934 let addr = probe.local_addr().unwrap();
935 drop(probe);
936
937 let task = tokio::spawn({
938 let cancel = cancel.clone();
939 async move { serve(cfg(), addr, tx, cancel).await }
940 });
941 rx.await.unwrap().expect("resolver binds");
942
943 let q = query(0x4242, "deep.nested.myapp.localhost", TYPE_A);
944
945 let sock = UdpSocket::bind("127.0.0.1:0").await.unwrap();
946 sock.send_to(&q, addr).await.unwrap();
947 let mut buf = [0u8; 512];
948 let (n, _) =
949 tokio::time::timeout(std::time::Duration::from_secs(5), sock.recv_from(&mut buf))
950 .await
951 .expect("udp reply arrives")
952 .unwrap();
953 assert_eq!(rdata(&buf[..n]), vec![127, 0, 0, 1]);
954
955 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
956 stream
957 .write_all(&(q.len() as u16).to_be_bytes())
958 .await
959 .unwrap();
960 stream.write_all(&q).await.unwrap();
961 let mut len_buf = [0u8; 2];
962 stream.read_exact(&mut len_buf).await.unwrap();
963 let mut resp = vec![0u8; usize::from(u16::from_be_bytes(len_buf))];
964 stream.read_exact(&mut resp).await.unwrap();
965 assert_eq!(rdata(&resp), vec![127, 0, 0, 1]);
966 stream
968 .write_all(&(q.len() as u16).to_be_bytes())
969 .await
970 .unwrap();
971 stream.write_all(&q).await.unwrap();
972 stream.read_exact(&mut len_buf).await.unwrap();
973 let mut resp2 = vec![0u8; usize::from(u16::from_be_bytes(len_buf))];
974 stream.read_exact(&mut resp2).await.unwrap();
975 assert_eq!(rcode(&resp2), RCODE_NOERROR);
976 drop(stream);
977
978 cancel.cancel();
979 tokio::time::timeout(std::time::Duration::from_secs(5), task)
980 .await
981 .expect("the resolver did not stop within 5s of cancellation")
982 .expect("the resolver task panicked")
983 .expect("the resolver returned an error");
984 }
985}