1use crate::core::MtopError;
2use crate::dns::core::{RecordClass, RecordType};
3use crate::dns::message::{Flags, Message, MessageId, Question, ResponseCode};
4use crate::dns::name::Name;
5use crate::net::tcp_connect;
6use crate::pool::{ClientFactory, ClientPool, ClientPoolConfig};
7use crate::timeout::Timeout;
8use async_trait::async_trait;
9use std::fmt;
10use std::io::{self, Cursor, Error};
11use std::net::{IpAddr, Ipv4Addr, SocketAddr};
12use std::pin::Pin;
13use std::sync::atomic::{AtomicUsize, Ordering};
14use std::task::{Context, Poll};
15use std::time::Duration;
16use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader, BufWriter, ReadBuf};
17use tokio::net::UdpSocket;
18
19const DEFAULT_NAMESERVER: SocketAddr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 53);
20const DEFAULT_MESSAGE_BUFFER: usize = 512;
21
22#[derive(Debug, Clone)]
24pub struct DnsClientConfig {
25 pub nameservers: Vec<SocketAddr>,
28
29 pub timeout: Duration,
32
33 pub attempts: u8,
37
38 pub rotate: bool,
41
42 pub pool_max_idle: u64,
45}
46
47impl Default for DnsClientConfig {
48 fn default() -> Self {
49 Self {
51 nameservers: vec![DEFAULT_NAMESERVER],
52 timeout: Duration::from_secs(5),
53 attempts: 2,
54 rotate: false,
55 pool_max_idle: 0,
56 }
57 }
58}
59
60#[async_trait]
65pub trait DnsClient {
66 async fn resolve(
69 &self,
70 id: MessageId,
71 name: Name,
72 rtype: RecordType,
73 rclass: RecordClass,
74 ) -> Result<Message, MtopError>;
75}
76
77#[derive(Debug)]
88pub struct DefaultDnsClient {
89 config: DnsClientConfig,
90 server_idx: AtomicUsize,
91 udp_pool: ClientPool<SocketAddr, UdpConnection>,
92 tcp_pool: ClientPool<SocketAddr, TcpConnection>,
93}
94
95impl DefaultDnsClient {
96 pub fn new<U, T>(config: DnsClientConfig, udp_factory: U, tcp_factory: T) -> Self
99 where
100 U: ClientFactory<SocketAddr, UdpConnection> + Send + Sync + 'static,
101 T: ClientFactory<SocketAddr, TcpConnection> + Send + Sync + 'static,
102 {
103 let udp_config = ClientPoolConfig {
104 name: "dns-udp".to_owned(),
105 max_idle: config.pool_max_idle,
106 };
107
108 let tcp_config = ClientPoolConfig {
109 name: "dns-tcp".to_owned(),
110 max_idle: config.pool_max_idle,
111 };
112
113 Self {
114 config,
115 server_idx: AtomicUsize::new(0),
116 udp_pool: ClientPool::new(udp_config, udp_factory),
117 tcp_pool: ClientPool::new(tcp_config, tcp_factory),
118 }
119 }
120
121 async fn exchange(&self, msg: &Message, server: &SocketAddr) -> Result<Message, MtopError> {
122 let res = async {
123 let mut conn = self.udp_pool.get(server).await?;
124 let res = conn.exchange(msg).await;
125 if res.is_ok() {
126 self.udp_pool.put(conn).await;
127 }
128
129 res
130 }
131 .timeout(self.config.timeout, format!("client.exchange udp://{}", server))
132 .await?;
133
134 if res.flags().is_truncated() {
137 tracing::debug!(message = "UDP response truncated, retrying with TCP", flags = ?res.flags(), server = %server);
138 async {
139 let mut conn = self.tcp_pool.get(server).await?;
140 let res = conn.exchange(msg).await;
141 if res.is_ok() {
142 self.tcp_pool.put(conn).await;
143 }
144
145 res
146 }
147 .timeout(self.config.timeout, format!("client.exchange tcp://{}", server))
148 .await
149 } else {
150 Ok(res)
151 }
152 }
153
154 fn starting_idx(&self) -> usize {
157 if self.config.rotate {
158 self.server_idx.fetch_add(1, Ordering::Relaxed)
159 } else {
160 0
161 }
162 }
163
164 fn nameserver_iterator(&self, idx: usize) -> impl Iterator<Item = &SocketAddr> {
166 self.config
167 .nameservers
168 .iter()
169 .cycle()
170 .skip(idx)
171 .take(self.config.nameservers.len())
172 }
173}
174
175#[async_trait]
176impl DnsClient for DefaultDnsClient {
177 async fn resolve(
178 &self,
179 id: MessageId,
180 name: Name,
181 rtype: RecordType,
182 rclass: RecordClass,
183 ) -> Result<Message, MtopError> {
184 let full = name.to_fqdn();
185 let flags = Flags::default().set_recursion_desired();
186 let question = Question::new(full.clone(), rtype).set_qclass(rclass);
187 let message = Message::new(id, flags).add_question(question);
188
189 let start = self.starting_idx();
190
191 let mut errors = Vec::new();
192 for attempt in 0..self.config.attempts {
193 for server in self.nameserver_iterator(start) {
194 match self.exchange(&message, server).await {
195 Ok(v) => {
196 let rc = v.flags().get_response_code();
200 if rc == ResponseCode::NoError || rc == ResponseCode::NameError {
201 return Ok(v);
202 }
203
204 tracing::debug!(message = "unsuitable response from nameserver, trying next one", server = %server, attempt = attempt + 1, max_attempts = self.config.attempts, response_code = ?rc);
205 errors.push(rc.to_string());
206 }
207 Err(e) => {
208 tracing::debug!(message = "nameserver failed, trying next one", server = %server, attempt = attempt + 1, max_attempts = self.config.attempts, err = %e);
209 errors.push(e.to_string());
210 }
211 }
212 }
213
214 if attempt + 1 < self.config.attempts {
215 tracing::debug!(
216 message = "all nameservers failed, retrying",
217 attempt = attempt + 1,
218 max_attempts = self.config.attempts
219 );
220 }
221 }
222
223 Err(MtopError::runtime(format!(
224 "no nameservers returned suitable responses for names {}: {}",
225 full,
226 errors.join("; "),
227 )))
228 }
229}
230
231pub struct TcpConnection {
236 read: BufReader<Box<dyn AsyncRead + Send + Sync + Unpin>>,
237 write: BufWriter<Box<dyn AsyncWrite + Send + Sync + Unpin>>,
238 buffer: Vec<u8>,
239}
240
241impl TcpConnection {
242 pub fn new<R, W>(read: R, write: W) -> Self
243 where
244 R: AsyncRead + Unpin + Sync + Send + 'static,
245 W: AsyncWrite + Unpin + Sync + Send + 'static,
246 {
247 Self {
248 read: BufReader::new(Box::new(read)),
249 write: BufWriter::new(Box::new(write)),
250 buffer: Vec::with_capacity(DEFAULT_MESSAGE_BUFFER),
251 }
252 }
253
254 pub async fn exchange(&mut self, msg: &Message) -> Result<Message, MtopError> {
255 self.buffer.clear();
258 msg.write_network_bytes(&mut self.buffer)?;
259 assert!(
260 self.buffer.len() < usize::from(u16::MAX),
261 "message size of {} exceeds maximum of {}",
262 self.buffer.len(),
263 u16::MAX
264 );
265
266 self.write.write_u16(u16::try_from(self.buffer.len()).unwrap()).await?;
267 self.write.write_all(&self.buffer).await?;
268 self.write.flush().await?;
269
270 let sz = self.read.read_u16().await?;
273 self.buffer.clear();
274 self.buffer.resize(usize::from(sz), 0);
275 self.read.read_exact(&mut self.buffer).await?;
276
277 let mut cur = Cursor::new(&self.buffer);
278 let res = Message::read_network_bytes(&mut cur)?;
279
280 if res.id() == msg.id() {
281 Ok(res)
282 } else {
283 Err(MtopError::runtime(format!(
284 "unexpected DNS MessageId; expected {}, got {}",
285 msg.id(),
286 res.id()
287 )))
288 }
289 }
290}
291
292impl fmt::Debug for TcpConnection {
293 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
294 write!(f, "TcpConnection {{ ... }}")
295 }
296}
297
298pub struct UdpConnection {
303 read: Box<dyn AsyncRead + Send + Sync + Unpin>,
304 write: Box<dyn AsyncWrite + Send + Sync + Unpin>,
305 buffer: Vec<u8>,
306 packet_size: usize,
307}
308
309impl UdpConnection {
310 pub fn new<R, W>(read: R, write: W) -> Self
311 where
312 R: AsyncRead + Unpin + Sync + Send + 'static,
313 W: AsyncWrite + Unpin + Sync + Send + 'static,
314 {
315 Self {
316 read: Box::new(read),
317 write: Box::new(write),
318 buffer: Vec::with_capacity(DEFAULT_MESSAGE_BUFFER),
319 packet_size: DEFAULT_MESSAGE_BUFFER,
320 }
321 }
322
323 pub async fn exchange(&mut self, msg: &Message) -> Result<Message, MtopError> {
324 self.buffer.clear();
325 msg.write_network_bytes(&mut self.buffer)?;
326 let n = self.write.write(&self.buffer).await?;
328 if n != self.buffer.len() {
329 return Err(MtopError::runtime(format!(
330 "short write to UDP socket. expected {}, got {}",
331 self.buffer.len(),
332 n
333 )));
334 }
335 self.write.flush().await?;
336
337 self.buffer.clear();
340 self.buffer.resize(self.packet_size, 0);
341
342 loop {
343 let n = self.read.read(&mut self.buffer).await?;
344 let cur = Cursor::new(&self.buffer[0..n]);
345 let res = Message::read_network_bytes(cur)?;
346 if res.id() == msg.id() {
347 return Ok(res);
348 }
349 }
350 }
351}
352
353impl fmt::Debug for UdpConnection {
354 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
355 write!(f, "UdpConnection {{ ... }}")
356 }
357}
358
359pub(crate) struct SocketAdapter(UdpSocket);
363
364impl SocketAdapter {
365 pub(crate) fn new(sock: UdpSocket) -> Self {
366 Self(sock)
367 }
368}
369
370impl AsyncRead for SocketAdapter {
371 fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
372 self.0.poll_recv(cx, buf)
373 }
374}
375
376impl AsyncWrite for SocketAdapter {
377 fn poll_write(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Error>> {
378 self.0.poll_send(cx, buf)
379 }
380
381 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
382 Poll::Ready(Ok(()))
383 }
384
385 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
386 Poll::Ready(Ok(()))
387 }
388}
389
390#[derive(Debug, Clone, Default)]
393pub struct UdpConnectionFactory;
394
395#[async_trait]
396impl ClientFactory<SocketAddr, UdpConnection> for UdpConnectionFactory {
397 async fn make(&self, address: &SocketAddr) -> Result<UdpConnection, MtopError> {
398 let local = if address.is_ipv4() { "0.0.0.0:0" } else { "[::]:0" };
399 let sock = UdpSocket::bind(local).await?;
400 sock.connect(address).await?;
401
402 let adapter = SocketAdapter::new(sock);
403 let (read, write) = tokio::io::split(adapter);
404 Ok(UdpConnection::new(read, write))
405 }
406}
407
408#[derive(Debug, Clone, Default)]
411pub struct TcpConnectionFactory;
412
413#[async_trait]
414impl ClientFactory<SocketAddr, TcpConnection> for TcpConnectionFactory {
415 async fn make(&self, address: &SocketAddr) -> Result<TcpConnection, MtopError> {
416 let (read, write) = tcp_connect(address).await?;
417 Ok(TcpConnection::new(read, write))
418 }
419}
420
421#[cfg(test)]
422mod test {
423 use super::{DefaultDnsClient, DnsClient, DnsClientConfig, TcpConnection, UdpConnection};
424 use crate::core::ErrorKind;
425 use crate::dns::core::{RecordClass, RecordType};
426 use crate::dns::message::{Flags, Message, MessageId, Question, Record, ResponseCode};
427 use crate::dns::name::Name;
428 use crate::dns::rdata::{RecordData, RecordDataA};
429 use crate::dns::test::{
430 TestPooledTcpClientFactory, TestPooledUdpClientFactory, TestTcpSocket, TestUdpSocket,
431 TestUnpooledTcpClientFactory, TestUnpooledUdpClientFactory,
432 };
433 use std::collections::HashMap;
434 use std::io::Cursor;
435 use std::net::{Ipv4Addr, SocketAddr};
436 use std::str::FromStr;
437
438 fn new_request(id: MessageId) -> Message {
439 let flags = Flags::default().set_query().set_recursion_desired();
440 let question = Question::new(Name::from_str("example.com.").unwrap(), RecordType::A);
441 Message::new(id, flags).add_question(question)
442 }
443
444 fn new_empty_response(id: MessageId) -> Message {
445 let flags = Flags::default()
446 .set_response()
447 .set_recursion_desired()
448 .set_recursion_available();
449 let question = Question::new(Name::from_str("example.com.").unwrap(), RecordType::A);
450 Message::new(id, flags).add_question(question)
451 }
452
453 fn new_response(id: MessageId) -> Message {
454 let response = new_empty_response(id);
455 let answer = Record::new(
456 Name::from_str("example.com.").unwrap(),
457 RecordType::A,
458 RecordClass::INET,
459 300,
460 RecordData::A(RecordDataA::new(Ipv4Addr::LOCALHOST)),
461 );
462
463 response.add_answer(answer)
464 }
465
466 #[tokio::test]
467 async fn test_tcp_client_eof_reading_length() {
468 let write = Vec::new();
469 let read = Cursor::new(Vec::new());
470
471 let id = MessageId::from(123);
472 let request = new_request(id);
473
474 let mut client = TcpConnection::new(read, write);
475
476 let res = client.exchange(&request).await;
477 let err = res.unwrap_err();
478 assert_eq!(ErrorKind::IO, err.kind());
479 }
480
481 #[tokio::test]
482 async fn test_tcp_client_eof_reading_message() {
483 let write = Vec::new();
484 let read = Cursor::new(vec![
485 0, 200, ]);
487
488 let id = MessageId::from(123);
489 let request = new_request(id);
490
491 let mut client = TcpConnection::new(read, write);
492
493 let res = client.exchange(&request).await;
494 let err = res.unwrap_err();
495 assert_eq!(ErrorKind::IO, err.kind());
496 }
497
498 #[tokio::test]
499 async fn test_tcp_client_id_mismatch() {
500 let response_id = MessageId::from(456);
501 let response = new_response(response_id);
502
503 let request_id = MessageId::from(123);
504 let request = new_request(request_id);
505
506 let sock = TestTcpSocket::new(vec![response]);
507 let (read, write) = tokio::io::split(sock);
508 let mut client = TcpConnection::new(read, write);
509
510 let result = client.exchange(&request).await;
511 let err = result.unwrap_err();
512 assert_eq!(ErrorKind::Runtime, err.kind());
513 }
514
515 #[tokio::test]
516 async fn test_tcp_client_single_message() {
517 let id = MessageId::from(123);
518 let response = new_response(id);
519 let request = new_request(id);
520
521 let sock = TestTcpSocket::new(vec![response.clone()]);
522 let (read, write) = tokio::io::split(sock);
523 let mut client = TcpConnection::new(read, write);
524
525 let result = client.exchange(&request).await.unwrap();
526 assert_eq!(response, result);
527 }
528
529 #[tokio::test]
530 async fn test_tcp_client_multiple_message() {
531 let id1 = MessageId::from(123);
532 let response1 = new_response(id1);
533 let request1 = new_request(id1);
534 let id2 = MessageId::from(456);
535 let response2 = new_response(id2);
536 let request2 = new_request(id2);
537
538 let sock = TestTcpSocket::new(vec![response2.clone(), response1.clone()]);
539 let (read, write) = tokio::io::split(sock);
540 let mut client = TcpConnection::new(read, write);
541
542 let result1 = client.exchange(&request1).await.unwrap();
543 assert_eq!(response1, result1);
544
545 let result2 = client.exchange(&request2).await.unwrap();
546 assert_eq!(response2, result2);
547 }
548
549 #[tokio::test]
550 async fn test_udp_client_success() {
551 let id = MessageId::from(123);
552 let response = new_response(id);
553 let request = new_request(id);
554
555 let sock = TestUdpSocket::new(vec![response.clone()]);
556 let (read, write) = tokio::io::split(sock);
557 let mut client = UdpConnection::new(read, write);
558
559 let result = client.exchange(&request).await.unwrap();
560 assert_eq!(response, result);
561 }
562
563 #[tokio::test]
564 async fn test_udp_client_one_id_mismatch() {
565 let id1 = MessageId::from(456);
566 let response1 = new_response(id1);
567 let id2 = MessageId::from(123);
568 let response2 = new_response(id2);
569
570 let request = new_request(id2);
574
575 let sock = TestUdpSocket::new(vec![response2.clone(), response1.clone()]);
576 let (read, write) = tokio::io::split(sock);
577 let mut client = UdpConnection::new(read, write);
578
579 let result = client.exchange(&request).await.unwrap();
580 assert_eq!(response2, result);
581 }
582
583 #[tokio::test]
584 async fn test_default_dns_client_resolve_name_error() {
585 let id = MessageId::from(123);
586 let name = Name::from_str("example.com.").unwrap();
587 let server = "127.0.0.1:53".parse().unwrap();
588
589 let udp_response = new_empty_response(id);
590 let flags = udp_response.flags().set_response_code(ResponseCode::NameError);
591 let udp_response = udp_response.set_flags(flags);
592
593 let mut udp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
594 udp_mapping.entry(server).or_default().push(udp_response);
595 let udp_factory = TestUnpooledUdpClientFactory::new(udp_mapping);
596 let tcp_factory = TestUnpooledTcpClientFactory::new(HashMap::new());
597
598 let cfg = DnsClientConfig::default();
599 let client = DefaultDnsClient::new(cfg, udp_factory, tcp_factory);
600 let result = client.resolve(id, name, RecordType::A, RecordClass::INET).await.unwrap();
601
602 assert_eq!(ResponseCode::NameError, result.flags().get_response_code());
603 assert!(result.answers().is_empty());
604 }
605
606 #[tokio::test]
607 async fn test_default_dns_client_resolve_success() {
608 let id = MessageId::from(123);
609 let name = Name::from_str("example.com.").unwrap();
610 let server = "127.0.0.1:53".parse().unwrap();
611
612 let udp_response = new_response(id);
613 let mut udp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
614 udp_mapping.entry(server).or_default().push(udp_response.clone());
615 let udp_factory = TestUnpooledUdpClientFactory::new(udp_mapping);
616 let tcp_factory = TestUnpooledTcpClientFactory::new(HashMap::new());
617
618 let cfg = DnsClientConfig::default();
619 let client = DefaultDnsClient::new(cfg, udp_factory, tcp_factory);
620 let result = client.resolve(id, name, RecordType::A, RecordClass::INET).await.unwrap();
621
622 assert_eq!(udp_response, result);
623 }
624
625 #[tokio::test]
626 async fn test_default_dns_client_resolve_one_error() {
627 let id = MessageId::from(123);
628 let name = Name::from_str("example.com.").unwrap();
629 let server = "127.0.0.1:53".parse().unwrap();
630
631 let udp_response1 = new_empty_response(id);
632 let flags = udp_response1.flags().set_response_code(ResponseCode::ServerFailure);
633 let udp_response1 = udp_response1.set_flags(flags);
634 let udp_response2 = new_response(id);
635
636 let mut udp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
637 let entry = udp_mapping.entry(server).or_default();
638 entry.push(udp_response2.clone());
639 entry.push(udp_response1);
640
641 let udp_factory = TestUnpooledUdpClientFactory::new(udp_mapping);
642 let tcp_factory = TestUnpooledTcpClientFactory::new(HashMap::new());
643
644 let cfg = DnsClientConfig::default();
645 let client = DefaultDnsClient::new(cfg, udp_factory, tcp_factory);
646 let result = client.resolve(id, name, RecordType::A, RecordClass::INET).await.unwrap();
647
648 assert_eq!(udp_response2, result);
649 }
650
651 #[tokio::test]
652 async fn test_default_dns_client_resolve_all_errors() {
653 let id = MessageId::from(123);
654 let name = Name::from_str("example.com.").unwrap();
655 let server = "127.0.0.1:53".parse().unwrap();
656
657 let udp_response1 = new_empty_response(id);
658 let flags = udp_response1.flags().set_response_code(ResponseCode::ServerFailure);
659 let udp_response1 = udp_response1.set_flags(flags);
660
661 let udp_response2 = new_empty_response(id);
662 let flags = udp_response2.flags().set_response_code(ResponseCode::ServerFailure);
663 let udp_response2 = udp_response2.set_flags(flags);
664
665 let mut udp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
666 let entry = udp_mapping.entry(server).or_default();
667 entry.push(udp_response2.clone());
668 entry.push(udp_response1);
669
670 let udp_factory = TestUnpooledUdpClientFactory::new(udp_mapping);
671 let tcp_factory = TestUnpooledTcpClientFactory::new(HashMap::new());
672
673 let cfg = DnsClientConfig::default();
674 let client = DefaultDnsClient::new(cfg, udp_factory, tcp_factory);
675 let err = client.resolve(id, name, RecordType::A, RecordClass::INET).await.unwrap_err();
676
677 assert_eq!(ErrorKind::Runtime, err.kind());
678 }
679
680 #[tokio::test]
681 async fn test_default_dns_client_resolve_one_bad_server() {
682 let id = MessageId::from(123);
683 let name = Name::from_str("example.com.").unwrap();
684 let server1 = "127.0.0.1:53".parse().unwrap();
685 let server2 = "127.0.0.2:53".parse().unwrap();
686
687 let udp_response1 = new_empty_response(id);
688 let flags = udp_response1.flags().set_response_code(ResponseCode::ServerFailure);
689 let udp_response1 = udp_response1.set_flags(flags);
690 let udp_response2 = new_response(id);
691
692 let mut udp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
693 udp_mapping.entry(server1).or_default().push(udp_response1);
694 udp_mapping.entry(server2).or_default().push(udp_response2.clone());
695
696 let udp_factory = TestUnpooledUdpClientFactory::new(udp_mapping);
697 let tcp_factory = TestUnpooledTcpClientFactory::new(HashMap::new());
698
699 let cfg = DnsClientConfig {
700 nameservers: vec![server1, server2],
701 ..Default::default()
702 };
703 let client = DefaultDnsClient::new(cfg, udp_factory, tcp_factory);
704 let result = client.resolve(id, name, RecordType::A, RecordClass::INET).await.unwrap();
705
706 assert_eq!(udp_response2, result);
707 }
708
709 #[tokio::test]
710 async fn test_default_dns_client_resolve_udp_truncation() {
711 let id = MessageId::from(123);
712 let name = Name::from_str("example.com.").unwrap();
713 let server = "127.0.0.1:53".parse().unwrap();
714
715 let udp_response = new_empty_response(id);
716 let flags = udp_response.flags().set_truncated();
717 let udp_response = udp_response.set_flags(flags);
718 let tcp_response = new_response(id);
719
720 let mut udp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
721 udp_mapping.entry(server).or_default().push(udp_response);
722
723 let mut tcp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
724 tcp_mapping.entry(server).or_default().push(tcp_response.clone());
725
726 let udp_factory = TestUnpooledUdpClientFactory::new(udp_mapping);
727 let tcp_factory = TestUnpooledTcpClientFactory::new(tcp_mapping);
728
729 let cfg = DnsClientConfig::default();
730 let client = DefaultDnsClient::new(cfg, udp_factory, tcp_factory);
731 let result = client.resolve(id, name, RecordType::A, RecordClass::INET).await.unwrap();
732
733 assert_eq!(tcp_response, result);
734 }
735
736 #[tokio::test]
737 async fn test_default_dns_client_resolve_tcp_reuse() {
738 let id = MessageId::from(123);
739 let name = Name::from_str("example.com.").unwrap();
740 let server = "127.0.0.1:53".parse().unwrap();
741
742 let udp_response = new_empty_response(id);
743 let flags = udp_response.flags().set_truncated();
744 let udp_response = udp_response.set_flags(flags);
745 let tcp_response = new_response(id);
746
747 let mut udp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
748 let udp_entry = udp_mapping.entry(server).or_default();
749 udp_entry.push(udp_response.clone());
750 udp_entry.push(udp_response.clone());
751
752 let mut tcp_mapping: HashMap<SocketAddr, Vec<Message>> = HashMap::new();
753 let tcp_entry = tcp_mapping.entry(server).or_default();
754 tcp_entry.push(tcp_response.clone());
755 tcp_entry.push(tcp_response.clone());
756
757 let udp_factory = TestPooledUdpClientFactory::new(udp_mapping);
758 let tcp_factory = TestPooledTcpClientFactory::new(tcp_mapping);
759
760 let cfg = DnsClientConfig {
761 pool_max_idle: 1,
762 ..Default::default()
763 };
764 let client = DefaultDnsClient::new(cfg, udp_factory, tcp_factory);
765
766 let result1 = client
767 .resolve(id, name.clone(), RecordType::A, RecordClass::INET)
768 .await
769 .unwrap();
770 assert_eq!(tcp_response, result1);
771
772 let result2 = client.resolve(id, name, RecordType::A, RecordClass::INET).await.unwrap();
773 assert_eq!(tcp_response, result2);
774 }
775}