Skip to main content

mtop_client/dns/
client.rs

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/// Configuration for creating a new `DnsClient` instance.
23#[derive(Debug, Clone)]
24pub struct DnsClientConfig {
25    /// One or more DNS nameservers to use for resolution. These servers will be tried
26    /// in order for each resolution unless `rotate` is set.
27    pub nameservers: Vec<SocketAddr>,
28
29    /// Timeout for each resolution. This timeout is applied to each attempt and so a
30    /// single call to `DnsClient::resolve` may take longer based on the value of `attempts`.
31    pub timeout: Duration,
32
33    /// Number of attempts to make performing a resolution for a single name. Note that
34    /// any response from a DNS server counts as "success". Only timeout or network errors
35    /// will trigger retries.
36    pub attempts: u8,
37
38    /// If true, `nameservers` will be round-robin load balanced for each resolution. If false
39    /// the nameservers are tried in-order for each resolution.
40    pub rotate: bool,
41
42    /// Max number of open sockets or connections to each nameserver. Default is not to keep
43    /// any open socket or connections open.
44    pub pool_max_idle: u64,
45}
46
47impl Default for DnsClientConfig {
48    fn default() -> Self {
49        // Default values picked based on `man 5 resolv.conf` when relevant.
50        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/// Client for performing DNS queries and returning the results.
61///
62/// There is currently only a single non-test implementation because this
63/// trait exists to make testing consumers easier.
64#[async_trait]
65pub trait DnsClient {
66    /// Resolve a single domain name of the provided type, ensuring it is
67    /// fully qualified before making the query.
68    async fn resolve(
69        &self,
70        id: MessageId,
71        name: Name,
72        rtype: RecordType,
73        rclass: RecordClass,
74    ) -> Result<Message, MtopError>;
75}
76
77/// Implementation of a `DnsClient` that uses UDP with TCP fallback.
78///
79/// Supports nameserver rotation, retries, timeouts, and pooling of client
80/// connections. Names are assumed to already be fully qualified, meaning
81/// that they are not combined with a search domain.
82///
83/// Timeouts are handled by the client itself and so callers should _not_
84/// add a timeout on the `resolve` method. Note that timeouts are per-network
85/// operation. This means that a single call to `resolve` make take longer
86/// than the timeout since failed network operations are retried.
87#[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    /// Create a new `DnsClient` that will resolve names using UDP or TCP connections
97    /// and behavior based on a resolv.conf configuration file.
98    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 the UDP response indicates the message was truncated, we discard
135        // it and repeat the query using TCP.
136        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    // Get the index of nameserver that should be used for a query based on if the client has
155    // been configured to roundrobin between nameservers or not.
156    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    // Get an iterator that will visit every nameserver once starting from the provided index.
165    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                        // NoError or a NameError is a conclusive answer. We either have results
197                        // or this is a bad domain. Any other type of response means we have to try
198                        // another server.
199                        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
231/// Connection for unconditionally sending and receiving DNS messages using TCP streams.
232/// Messages are sent with a two byte prefix that indicates the size of the message.
233/// Responses are expected to have the same prefix. The message ID of responses is
234/// checked to ensure it matches the request ID. If it does not, an error is returned.
235pub 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        // Write the message to a local buffer and then send it, prefixed
256        // with the size of the message.
257        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        // Read the prefixed size of the response in big-endian (network)
271        // order and then read exactly that many bytes into our buffer.
272        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
298/// Connection for unconditionally sending and receiving DNS messages using UDP packets.
299/// The message ID of responses is checked to ensure it matches the request ID. If it
300/// does not, the response is discarded and the client will wait for another response
301/// until it gets one with a matching ID.
302pub 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        // We expect this to be a datagram socket so we only do a single write.
327        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        // Resize to our packet size since the .read() call will only read up to
338        // the size of the buffer at most.
339        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
359/// Adapter for reading and writing to a `UdpSocket` using the `AsyncRead` and `AsyncWrite`
360/// traits. This exists to enable easier testing of `UdpConnection` by allowing alternate
361/// implementations of those traits to be used.
362pub(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/// Implementation of `ClientFactory` for creating concrete `UdpConnection` instances
391/// that use a UDP socket.
392#[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/// Implementation of `ClientFactory` for creating concrete `TcpConnection` instances
409/// that use a TCP socket.
410#[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, // message length
486        ]);
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        // Note that the request has the ID of the second response because
571        // we are testing the that first response is discarded due to the ID
572        // not matching.
573        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}