Skip to main content

hickory_net/xfer/
dns_multiplexer.rs

1// Copyright 2015-2023 Benjamin Fry <benjaminfry@me.com>
2//
3// Licensed under the Apache License, Version 2.0, <LICENSE-APACHE or
4// https://apache.org/licenses/LICENSE-2.0> or the MIT license <LICENSE-MIT or
5// https://opensource.org/licenses/MIT>, at your option. This file may not be
6// copied, modified, or distributed except according to those terms.
7
8//! `DnsMultiplexer` and associated types implement the state machines for sending DNS messages while using the underlying streams.
9
10use core::{
11    pin::Pin,
12    task::{Context, Poll},
13    time::Duration,
14};
15use std::collections::{HashMap, hash_map::Entry};
16use std::io;
17
18use futures_channel::mpsc;
19use futures_util::{
20    FutureExt,
21    future::BoxFuture,
22    stream::{Stream, StreamExt},
23};
24use rand::RngExt;
25use tracing::debug;
26
27use super::{
28    BufDnsStreamHandle, DnsClientStream, DnsRequestSender, DnsResponseStream, ignore_send,
29};
30use crate::proto::op::{DnsRequest, DnsResponse, SerialMessage};
31#[cfg(feature = "__dnssec")]
32use crate::proto::rr::{TSigVerifier, TSigner};
33use crate::{DnsStreamHandle, error::NetError, runtime::Time};
34
35struct ActiveRequest {
36    // the completion is the channel for a response to the original request
37    completion: mpsc::Sender<Result<DnsResponse, NetError>>,
38    request_id: u16,
39    timeout: BoxFuture<'static, ()>,
40    #[cfg(feature = "__dnssec")]
41    verifier: Option<TSigVerifier>,
42}
43
44impl ActiveRequest {
45    fn new(
46        completion: mpsc::Sender<Result<DnsResponse, NetError>>,
47        request_id: u16,
48        timeout: BoxFuture<'static, ()>,
49        #[cfg(feature = "__dnssec")] verifier: Option<TSigVerifier>,
50    ) -> Self {
51        Self {
52            completion,
53            request_id,
54            // request,
55            timeout,
56            #[cfg(feature = "__dnssec")]
57            verifier,
58        }
59    }
60
61    /// polls the timeout and converts the error
62    fn poll_timeout(&mut self, cx: &mut Context<'_>) -> Poll<()> {
63        self.timeout.poll_unpin(cx)
64    }
65
66    /// Returns true of the other side canceled the request
67    fn is_canceled(&self) -> bool {
68        self.completion.is_closed()
69    }
70
71    /// the request id of the message that was sent
72    fn request_id(&self) -> u16 {
73        self.request_id
74    }
75
76    /// Sends an error
77    fn complete_with_error(mut self, error: NetError) {
78        ignore_send(self.completion.try_send(Err(error)));
79    }
80}
81
82/// A DNS Client implemented over futures-rs.
83///
84/// This Client is generic and capable of wrapping UDP, TCP, and other underlying DNS protocol
85///  implementations. This should be used for underlying protocols that do not natively support
86///  multiplexed sessions.
87#[must_use = "futures do nothing unless polled"]
88pub struct DnsMultiplexer<S> {
89    stream: S,
90    timeout_duration: Duration,
91    stream_handle: BufDnsStreamHandle,
92    active_requests: HashMap<u16, ActiveRequest>,
93    max_active_requests: usize,
94    #[cfg(feature = "__dnssec")]
95    signer: Option<TSigner>,
96    is_shutdown: bool,
97}
98
99impl<S: DnsClientStream> DnsMultiplexer<S> {
100    /// Spawns a new DnsMultiplexer Stream.
101    ///
102    /// This uses a default timeout of 5 seconds for all requests unless changed with
103    /// [`Self::with_timeout()`]. At most 32 in-flight requests are allowed unless
104    /// changed with [`Self::with_max_active_requests()`].
105    ///
106    /// # Arguments
107    ///
108    /// * `stream` - A stream of bytes that can be used to send/receive DNS messages
109    ///   (see TcpClientStream or UdpClientStream)
110    /// * `stream_handle` - The handle for the `stream` on which bytes can be sent/received.
111    pub fn new(stream: S, stream_handle: BufDnsStreamHandle) -> Self {
112        Self {
113            stream,
114            timeout_duration: Duration::from_secs(5),
115            stream_handle,
116            active_requests: HashMap::default(),
117            max_active_requests: 32,
118            #[cfg(feature = "__dnssec")]
119            signer: None,
120            is_shutdown: false,
121        }
122    }
123
124    /// Change the default timeout of the DnsMultiplexer stream.
125    pub fn with_timeout(mut self, timeout: Duration) -> Self {
126        self.timeout_duration = timeout;
127        self
128    }
129
130    /// Set the maximum number of active (in-flight) requests.
131    ///
132    /// This limits how many DNS queries can be simultaneously pending on this
133    /// multiplexed connection. When the limit is reached, new requests will
134    /// return [`NetError::Busy`].
135    pub fn with_max_active_requests(mut self, max: usize) -> Self {
136        self.max_active_requests = max;
137        self
138    }
139
140    /// Specify an optional signer to TSIG authenticate requests.
141    #[cfg(feature = "__dnssec")]
142    pub fn with_signer(mut self, signer: TSigner) -> Self {
143        self.signer = Some(signer);
144        self
145    }
146
147    /// loop over active_requests and remove cancelled requests
148    ///  this should free up space if we already had 4096 active requests
149    fn drop_cancelled(&mut self, cx: &mut Context<'_>) {
150        let mut canceled = HashMap::<u16, NetError>::new();
151        for (&id, active_req) in &mut self.active_requests {
152            if active_req.is_canceled() {
153                canceled.insert(id, NetError::from("requestor canceled"));
154            }
155
156            // check for timeouts...
157            match active_req.poll_timeout(cx) {
158                Poll::Ready(()) => {
159                    debug!("request timed out: {}", id);
160                    canceled.insert(id, NetError::Timeout);
161                }
162                Poll::Pending => (),
163            }
164        }
165
166        // drop all the canceled requests
167        for (id, error) in canceled {
168            if let Some(active_request) = self.active_requests.remove(&id) {
169                // complete the request, it's failed...
170                active_request.complete_with_error(error);
171            }
172        }
173    }
174
175    /// creates random query_id, validates against all active queries
176    fn next_random_query_id(&self) -> Result<u16, NetError> {
177        let mut rand = rand::rng();
178
179        for _ in 0..100 {
180            let id: u16 = rand.random(); // the range is [0 ... u16::max]
181
182            if !self.active_requests.contains_key(&id) {
183                return Ok(id);
184            }
185        }
186
187        Err(NetError::from(
188            "id space exhausted, consider filing an issue",
189        ))
190    }
191
192    /// Closes all outstanding completes with a closed stream error
193    fn stream_closed_close_all(&mut self, error: NetError) {
194        debug!(%error, addr = %self.stream.name_server_addr());
195        for (_, active_request) in self.active_requests.drain() {
196            // complete the request, it's failed...
197            active_request.complete_with_error(error.clone());
198        }
199    }
200}
201
202impl<S: DnsClientStream> DnsRequestSender for DnsMultiplexer<S> {
203    fn send_message(&mut self, request: DnsRequest) -> DnsResponseStream {
204        if self.is_shutdown {
205            panic!("can not send messages after stream is shutdown")
206        }
207
208        if self.active_requests.len() >= self.max_active_requests {
209            return NetError::Busy.into();
210        }
211
212        let query_id = match self.next_random_query_id() {
213            Ok(id) => id,
214            Err(e) => return e.into(),
215        };
216
217        let (mut request, _) = request.into_parts();
218        request.metadata.id = query_id;
219
220        #[cfg(feature = "__dnssec")]
221        let mut verifier = None;
222        #[cfg(feature = "__dnssec")]
223        if let Some(signer) = &self.signer {
224            if signer.should_sign_message(&request) {
225                match request.finalize(signer, S::Time::current_time()) {
226                    Ok(answer_verifier) => verifier = answer_verifier,
227                    Err(e) => {
228                        debug!("could not sign message: {}", e);
229                        return NetError::from(e).into();
230                    }
231                }
232            }
233        }
234
235        // store a Timeout for this message before sending
236        let timeout = S::Time::delay_for(self.timeout_duration);
237
238        let (complete, receiver) = mpsc::channel(QUERY_RESPONSE_BUFFER_SIZE);
239
240        // send the message
241        let active_request = ActiveRequest::new(
242            complete,
243            request.id,
244            timeout,
245            #[cfg(feature = "__dnssec")]
246            verifier,
247        );
248
249        match request.to_vec() {
250            Ok(buffer) => {
251                debug!(id = %active_request.request_id(), "sending message");
252                let serial_message = SerialMessage::new(buffer, self.stream.name_server_addr());
253
254                debug!(
255                    "final message: {}",
256                    serial_message
257                        .to_message()
258                        .expect("bizarre we just made this message")
259                );
260
261                // add to the map -after- the client send b/c we don't want to put it in the map if
262                //  we ended up returning an error from the send.
263                match self.stream_handle.send(serial_message) {
264                    Ok(()) => self
265                        .active_requests
266                        .insert(active_request.request_id(), active_request),
267                    Err(err) => return err.into(),
268                };
269            }
270            Err(error) => {
271                debug!(
272                    id = %active_request.request_id(),
273                    %error,
274                    "error message"
275                );
276                // complete with the error, don't add to the map of active requests
277                return NetError::from(error).into();
278            }
279        }
280
281        receiver.into()
282    }
283
284    fn shutdown(&mut self) {
285        self.is_shutdown = true;
286    }
287
288    fn is_shutdown(&self) -> bool {
289        self.is_shutdown
290    }
291}
292
293impl<S: DnsClientStream> Stream for DnsMultiplexer<S> {
294    type Item = Result<(), NetError>;
295
296    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
297        // Always drop the cancelled queries first
298        self.drop_cancelled(cx);
299
300        if self.is_shutdown && self.active_requests.is_empty() {
301            debug!("stream is done: {}", self.stream.name_server_addr());
302            return Poll::Ready(None);
303        }
304
305        // Collect all inbound requests, max 100 at a time for QoS
306        //   by having a max we will guarantee that the client can't be DOSed in this loop
307        // TODO: make the QoS configurable
308        let mut messages_received = 0;
309        for i in 0..QOS_MAX_RECEIVE_MSGS {
310            match self.stream.poll_next_unpin(cx) {
311                Poll::Ready(Some(Ok(buffer))) => {
312                    messages_received = i;
313
314                    //   deserialize or log decode_error
315                    match DnsResponse::from_buffer(buffer.into_parts().0) {
316                        Ok(response) => match self.active_requests.entry(response.id) {
317                            Entry::Occupied(mut request_entry) => {
318                                // send the response, complete the request...
319                                let active_request = request_entry.get_mut();
320                                #[cfg(feature = "__dnssec")]
321                                if let Some(verifier) = &mut active_request.verifier {
322                                    ignore_send(
323                                        active_request.completion.try_send(
324                                            verifier
325                                                .verify(response.as_buffer())
326                                                .map_err(NetError::from),
327                                        ),
328                                    );
329                                } else {
330                                    ignore_send(active_request.completion.try_send(Ok(response)));
331                                }
332                                #[cfg(not(feature = "__dnssec"))]
333                                ignore_send(active_request.completion.try_send(Ok(response)));
334                            }
335                            Entry::Vacant(..) => debug!("unexpected request_id: {}", response.id),
336                        },
337                        // TODO: return src address for diagnostics
338                        Err(error) => debug!(%error, "error decoding message"),
339                    }
340                }
341                Poll::Ready(err) => {
342                    let err = match err {
343                        Some(Err(e)) => e,
344                        None => NetError::from(io::Error::new(
345                            io::ErrorKind::UnexpectedEof,
346                            "stream closed",
347                        )),
348                        _ => unreachable!(),
349                    };
350
351                    self.stream_closed_close_all(err);
352                    self.is_shutdown = true;
353                    return Poll::Ready(None);
354                }
355                Poll::Pending => break,
356            }
357        }
358
359        // If still active, then if the qos (for _ in 0..100 loop) limit
360        // was hit then "yield". This'll make sure that the future is
361        // woken up immediately on the next turn of the event loop.
362        if messages_received == QOS_MAX_RECEIVE_MSGS {
363            // FIXME: this was a task::current().notify(); is this right?
364            cx.waker().wake_by_ref();
365        }
366
367        // Finally, return not ready to keep the 'driver task' alive.
368        Poll::Pending
369    }
370}
371
372const QOS_MAX_RECEIVE_MSGS: usize = 100; // max number of messages to receive from the UDP socket
373
374/// Buffer size for per-query response channels.
375///
376/// Each outgoing DNS query gets its own channel to receive responses. Standard
377/// DNS queries receive exactly one response so a small buffer is sufficient.
378const QUERY_RESPONSE_BUFFER_SIZE: usize = 8;
379
380#[cfg(test)]
381mod test {
382    use core::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
383
384    use futures_util::{
385        future::{self, BoxFuture},
386        ready,
387        stream::TryStreamExt,
388    };
389    use test_support::subscribe;
390
391    use super::*;
392    use crate::proto::op::{DnsRequestOptions, Message, Query};
393    use crate::proto::rr::rdata::{NS, SOA};
394    use crate::proto::rr::{DNSClass, Name, RData, Record, RecordType};
395    use crate::proto::serialize::binary::BinEncodable;
396    use crate::xfer::{DnsClientStream, StreamReceiver};
397
398    struct MockClientStream {
399        messages: Vec<Message>,
400        addr: SocketAddr,
401        id: Option<u16>,
402        receiver: Option<StreamReceiver>,
403    }
404
405    impl MockClientStream {
406        fn new(
407            mut messages: Vec<Message>,
408            addr: SocketAddr,
409        ) -> BoxFuture<'static, Result<Self, NetError>> {
410            messages.reverse(); // so we can pop() and get messages in order
411            Box::pin(future::ok(Self {
412                messages,
413                addr,
414                id: None,
415                receiver: None,
416            }))
417        }
418    }
419
420    impl Stream for MockClientStream {
421        type Item = Result<SerialMessage, NetError>;
422
423        fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
424            let id = if let Some(id) = self.id {
425                id
426            } else {
427                let serial = ready!(
428                    self.receiver
429                        .as_mut()
430                        .expect("should only be polled after receiver has been set")
431                        .poll_next_unpin(cx)
432                );
433                let message = serial.unwrap().to_message().unwrap();
434                self.id = Some(message.id);
435                message.id
436            };
437
438            if let Some(mut message) = self.messages.pop() {
439                message.metadata.id = id;
440                Poll::Ready(Some(Ok(SerialMessage::new(
441                    message.to_bytes().unwrap(),
442                    self.addr,
443                ))))
444            } else {
445                Poll::Pending
446            }
447        }
448    }
449
450    impl DnsClientStream for MockClientStream {
451        type Time = crate::runtime::TokioTime;
452
453        fn name_server_addr(&self) -> SocketAddr {
454            self.addr
455        }
456    }
457
458    async fn get_mocked_multiplexer(
459        mock_response: Vec<Message>,
460    ) -> DnsMultiplexer<MockClientStream> {
461        let addr = SocketAddr::from(([127, 0, 0, 1], 1234));
462        let mock_response = MockClientStream::new(mock_response, addr).await.unwrap();
463        let (handler, receiver) = BufDnsStreamHandle::new(addr);
464        let mut multiplexer =
465            DnsMultiplexer::new(mock_response, handler).with_timeout(Duration::from_millis(100));
466
467        multiplexer.stream.receiver = Some(receiver); // so it can get the correct request id
468
469        multiplexer
470    }
471
472    fn a_query_answer() -> (DnsRequest, Vec<Message>) {
473        let name = Name::from_ascii("www.example.com.").unwrap();
474
475        let mut request = Message::query();
476        request.metadata.recursion_desired = true;
477        request.add_query({
478            let mut q = Query::query(name.clone(), RecordType::A);
479            q.set_query_class(DNSClass::IN);
480            q
481        });
482
483        let mut response = request.clone().into_response();
484        response.add_answer(Record::from_rdata(
485            name,
486            86400,
487            RData::A(Ipv4Addr::new(93, 184, 215, 14).into()),
488        ));
489        (
490            DnsRequest::new(request, DnsRequestOptions::default()),
491            vec![response],
492        )
493    }
494
495    fn axfr_query() -> Message {
496        let name = Name::from_ascii("example.com.").unwrap();
497
498        let mut msg = Message::query();
499        msg.metadata.recursion_desired = true;
500        msg.add_query({
501            let mut query = Query::query(name, RecordType::AXFR);
502            query.set_query_class(DNSClass::IN);
503            query
504        });
505        msg
506    }
507
508    fn axfr_response() -> Vec<Record> {
509        let origin = Name::from_ascii("example.com.").unwrap();
510        let soa = Record::from_rdata(
511            origin.clone(),
512            3600,
513            RData::SOA(SOA::new(
514                Name::parse("sns.dns.icann.org.", None).unwrap(),
515                Name::parse("noc.dns.icann.org.", None).unwrap(),
516                2015082403,
517                7200,
518                3600,
519                1209600,
520                3600,
521            )),
522        );
523
524        vec![
525            soa.clone(),
526            Record::from_rdata(
527                origin.clone(),
528                86400,
529                RData::NS(NS(Name::parse("a.iana-servers.net.", None).unwrap())),
530            ),
531            Record::from_rdata(
532                origin.clone(),
533                86400,
534                RData::NS(NS(Name::parse("b.iana-servers.net.", None).unwrap())),
535            ),
536            Record::from_rdata(
537                origin.clone(),
538                86400,
539                RData::A(Ipv4Addr::new(93, 184, 215, 14).into()),
540            ),
541            Record::from_rdata(
542                origin,
543                86400,
544                RData::AAAA(
545                    Ipv6Addr::new(
546                        0x2606, 0x2800, 0x21f, 0xcb07, 0x6820, 0x80da, 0xaf6b, 0x8b2c,
547                    )
548                    .into(),
549                ),
550            ),
551            soa,
552        ]
553    }
554
555    fn axfr_query_answer() -> (DnsRequest, Vec<Message>) {
556        let msg = axfr_query();
557
558        let mut response = msg.clone().into_response();
559        response.insert_answers(axfr_response());
560        (
561            DnsRequest::new(msg, DnsRequestOptions::default()),
562            vec![response],
563        )
564    }
565
566    fn axfr_query_answer_multi() -> (DnsRequest, Vec<Message>) {
567        let base = axfr_query();
568
569        let query = base.clone();
570        let mut rr = axfr_response();
571        let rr2 = rr.split_off(3);
572        let mut msg1 = base.clone().into_response();
573        msg1.insert_answers(rr);
574        let mut msg2 = base.into_response();
575        msg2.insert_answers(rr2);
576        (
577            DnsRequest::new(query, DnsRequestOptions::default()),
578            vec![msg1, msg2],
579        )
580    }
581
582    #[tokio::test]
583    async fn test_multiplexer_a() {
584        subscribe();
585        let (query, answer) = a_query_answer();
586        let mut multiplexer = get_mocked_multiplexer(answer).await;
587        let response = multiplexer.send_message(query);
588        let response = tokio::select! {
589            _ = multiplexer.next() => {
590                // polling multiplexer to make it run
591                panic!("should never end")
592            },
593            r = response.try_collect::<Vec<_>>() => r.unwrap(),
594        };
595        assert_eq!(response.len(), 1);
596    }
597
598    #[tokio::test]
599    async fn test_multiplexer_axfr() {
600        subscribe();
601        let (query, answer) = axfr_query_answer();
602        let mut multiplexer = get_mocked_multiplexer(answer).await;
603        let response = multiplexer.send_message(query);
604        let response = tokio::select! {
605            _ = multiplexer.next() => {
606                // polling multiplexer to make it run
607                panic!("should never end")
608            },
609            r = response.try_collect::<Vec<_>>() => r.unwrap(),
610        };
611        assert_eq!(response.len(), 1);
612        assert_eq!(response[0].answers.len(), axfr_response().len());
613    }
614
615    #[tokio::test]
616    async fn test_multiplexer_axfr_multi() {
617        subscribe();
618        let (query, answer) = axfr_query_answer_multi();
619        let mut multiplexer = get_mocked_multiplexer(answer).await;
620        let response = multiplexer.send_message(query);
621        let response = tokio::select! {
622            _ = multiplexer.next() => {
623                // polling multiplexer to make it run
624                panic!("should never end")
625            },
626            r = response.try_collect::<Vec<_>>() => r.unwrap(),
627        };
628        assert_eq!(response.len(), 2);
629        assert_eq!(
630            response.iter().map(|m| m.answers.len()).sum::<usize>(),
631            axfr_response().len()
632        );
633    }
634}