1use 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 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 timeout,
56 #[cfg(feature = "__dnssec")]
57 verifier,
58 }
59 }
60
61 fn poll_timeout(&mut self, cx: &mut Context<'_>) -> Poll<()> {
63 self.timeout.poll_unpin(cx)
64 }
65
66 fn is_canceled(&self) -> bool {
68 self.completion.is_closed()
69 }
70
71 fn request_id(&self) -> u16 {
73 self.request_id
74 }
75
76 fn complete_with_error(mut self, error: NetError) {
78 ignore_send(self.completion.try_send(Err(error)));
79 }
80}
81
82#[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 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 pub fn with_timeout(mut self, timeout: Duration) -> Self {
126 self.timeout_duration = timeout;
127 self
128 }
129
130 pub fn with_max_active_requests(mut self, max: usize) -> Self {
136 self.max_active_requests = max;
137 self
138 }
139
140 #[cfg(feature = "__dnssec")]
142 pub fn with_signer(mut self, signer: TSigner) -> Self {
143 self.signer = Some(signer);
144 self
145 }
146
147 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 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 for (id, error) in canceled {
168 if let Some(active_request) = self.active_requests.remove(&id) {
169 active_request.complete_with_error(error);
171 }
172 }
173 }
174
175 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(); 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 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 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 let timeout = S::Time::delay_for(self.timeout_duration);
237
238 let (complete, receiver) = mpsc::channel(QUERY_RESPONSE_BUFFER_SIZE);
239
240 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 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 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 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 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 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 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 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 messages_received == QOS_MAX_RECEIVE_MSGS {
363 cx.waker().wake_by_ref();
365 }
366
367 Poll::Pending
369 }
370}
371
372const QOS_MAX_RECEIVE_MSGS: usize = 100; const 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(); 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); 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 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 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 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}