1use core::fmt::{self, Display};
9use core::net::SocketAddr;
10use core::pin::Pin;
11use core::task::{Context, Poll};
12use core::time::Duration;
13use std::collections::HashSet;
14use std::sync::Arc;
15
16use futures_util::{FutureExt, Stream, StreamExt, pin_mut, stream::FuturesUnordered};
17use tracing::{debug, trace, warn};
18
19use crate::error::NetError;
20use crate::proto::op::{DEFAULT_RETRY_FLOOR, DnsRequest, DnsResponse, Message, SerialMessage};
21#[cfg(feature = "__dnssec")]
22use crate::proto::rr::TSigner;
23use crate::runtime::{DnsUdpSocket, RuntimeProvider, Spawn, Time};
24use crate::udp::MAX_RECEIVE_BUFFER_SIZE;
25use crate::udp::udp_stream::NextRandomUdpSocket;
26use crate::xfer::{DnsExchange, DnsRequestSender, DnsResponseStream};
27
28#[must_use = "futures do nothing unless polled"]
34pub struct UdpClientStream<P> {
35 name_server: SocketAddr,
36 timeout: Duration,
37 is_shutdown: bool,
38 #[cfg(feature = "__dnssec")]
39 signer: Option<TSigner>,
40 bind_addr: Option<SocketAddr>,
41 avoid_local_ports: Arc<HashSet<u16>>,
42 os_port_selection: bool,
43 provider: P,
44 max_retries: u8,
45 retry_interval_floor: Duration,
46}
47
48impl<P: RuntimeProvider> UdpClientStream<P> {
49 pub fn builder(name_server: SocketAddr, provider: P) -> UdpClientStreamBuilder<P> {
51 UdpClientStreamBuilder {
52 name_server,
53 timeout: None,
54 #[cfg(feature = "__dnssec")]
55 signer: None,
56 bind_addr: None,
57 avoid_local_ports: Arc::default(),
58 os_port_selection: false,
59 provider,
60 max_retries: 3,
61 retry_interval_floor: DEFAULT_RETRY_FLOOR,
64 }
65 }
66}
67
68impl<P> Display for UdpClientStream<P> {
69 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> Result<(), fmt::Error> {
70 write!(formatter, "UDP({})", self.name_server)
71 }
72}
73
74impl<P: RuntimeProvider> DnsRequestSender for UdpClientStream<P> {
75 fn send_message(&mut self, request: DnsRequest) -> DnsResponseStream {
76 if self.is_shutdown {
77 panic!("can not send messages after stream is shutdown")
78 }
79
80 let retry_interval_time = request.options().retry_interval;
81 let request = UdpRequest::new(request, self);
82
83 let max_retries = self.max_retries;
84 let retry_interval = if retry_interval_time < self.retry_interval_floor {
85 self.retry_interval_floor
86 } else {
87 retry_interval_time
88 };
89
90 P::Timer::timeout(
91 self.timeout,
92 retry::<P>(request, retry_interval, max_retries.into()),
93 )
94 .into()
95 }
96
97 fn shutdown(&mut self) {
98 self.is_shutdown = true;
99 }
100
101 fn is_shutdown(&self) -> bool {
102 self.is_shutdown
103 }
104}
105
106impl<P> Stream for UdpClientStream<P> {
108 type Item = Result<(), NetError>;
109
110 fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
111 if self.is_shutdown {
113 Poll::Ready(None)
114 } else {
115 Poll::Ready(Some(Ok(())))
116 }
117 }
118}
119
120struct UdpRequest<P> {
122 avoid_local_ports: Arc<HashSet<u16>>,
123 name_server: SocketAddr,
124 request: DnsRequest,
125 provider: P,
126 #[cfg(feature = "__dnssec")]
127 signer: Option<TSigner>,
128 #[cfg(feature = "__dnssec")]
129 now: u64,
130 bind_addr: Option<SocketAddr>,
131 os_port_selection: bool,
132 case_randomization: bool,
133 recv_buf_size: usize,
134}
135
136impl<P: RuntimeProvider> UdpRequest<P> {
137 fn new(request: DnsRequest, stream: &UdpClientStream<P>) -> Self {
138 Self {
139 avoid_local_ports: stream.avoid_local_ports.clone(),
140 recv_buf_size: MAX_RECEIVE_BUFFER_SIZE.min(request.max_payload() as usize),
141 case_randomization: request.options().case_randomization,
142 name_server: stream.name_server,
143 #[cfg(feature = "__dnssec")]
145 signer: match &stream.signer {
146 Some(signer) if signer.should_sign_message(&request) => stream.signer.clone(),
147 _ => None,
148 },
149 request,
150 provider: stream.provider.clone(),
151 #[cfg(feature = "__dnssec")]
152 now: P::Timer::current_time(),
153 bind_addr: stream.bind_addr,
154 os_port_selection: stream.os_port_selection,
155 }
156 }
157}
158
159impl<P: RuntimeProvider> Request for UdpRequest<P> {
160 async fn send(&self) -> Result<DnsResponse, NetError> {
161 let original_query = self.request.original_query();
162 #[cfg_attr(not(feature = "__dnssec"), expect(unused_mut))]
163 let mut request = self.request.clone();
164
165 #[cfg(feature = "__dnssec")]
166 let mut verifier = None;
167 #[cfg(feature = "__dnssec")]
168 if let Some(signer) = &self.signer {
169 match request.finalize(signer, self.now) {
170 Ok(answer_verifier) => verifier = answer_verifier,
171 Err(e) => {
172 debug!("could not sign message: {}", e);
173 return Err(e.into());
174 }
175 }
176 }
177
178 let request_bytes = match request.to_vec() {
179 Ok(bytes) => bytes,
180 Err(err) => return Err(err.into()),
181 };
182
183 let msg_id = request.id;
184 let msg = SerialMessage::new(request_bytes, self.name_server);
185 let addr = msg.addr();
186 let final_message = match msg.to_message() {
187 Ok(m) => m,
188 Err(e) => return Err(e.into()),
189 };
190 debug!(%final_message, "final message");
191
192 let socket = NextRandomUdpSocket::new(
193 addr,
194 self.bind_addr,
195 self.avoid_local_ports.clone(),
196 self.os_port_selection,
197 self.provider.clone(),
198 )
199 .await?;
200
201 let bytes = msg.bytes();
202 let len_sent: usize = socket.send_to(bytes, addr).await?;
203
204 if bytes.len() != len_sent {
205 return Err(NetError::from(format!(
206 "Not all bytes of message sent, {} of {}",
207 len_sent,
208 bytes.len()
209 )));
210 }
211
212 trace!(
214 recv_buf_size = self.recv_buf_size,
215 "creating UDP receive buffer"
216 );
217 let mut recv_buf = vec![0; self.recv_buf_size];
218
219 loop {
222 let (len, src) = socket.recv_from(&mut recv_buf).await?;
223
224 let response_bytes = &recv_buf[0..len];
226 let response_buffer = Vec::from(response_bytes);
227
228 let request_target = msg.addr();
230
231 if src.ip().to_canonical() != request_target.ip().to_canonical()
239 || src.port() != request_target.port()
240 {
241 warn!(
242 "ignoring response from {}:{} because it does not match name_server: {}:{}.",
243 src.ip().to_canonical(),
244 src.port(),
245 request_target.ip().to_canonical(),
246 request_target.port(),
247 );
248
249 continue;
251 }
252
253 let Some(id_bytes) = response_buffer.first_chunk::<2>() else {
255 warn!(length = response_buffer.len(), "ignoring short message");
256 continue;
257 };
258 let response_id = u16::from_be_bytes(*id_bytes);
259
260 if msg_id != response_id {
262 warn!(
264 expected_id = msg_id,
265 received_id = response_id,
266 "ignoring response with wrong message id",
267 );
268
269 continue;
271 }
272
273 let mut response = DnsResponse::from_buffer(response_buffer)?;
274
275 let request_message = Message::from_vec(msg.bytes())?;
301 let request_queries = &request_message.queries;
302 let response_queries = &mut response.queries;
303
304 let question_matches = request_queries.len() == response_queries.len()
305 && response_queries
306 .iter()
307 .all(|elem| request_queries.contains(elem));
308 if self.case_randomization
309 && question_matches
310 && !response_queries.iter().all(|elem| {
311 request_queries
312 .iter()
313 .any(|req_q| req_q == elem && req_q.name().eq_case(elem.name()))
314 })
315 {
316 warn!(
317 "case of question section did not match: we expected '{request_queries:?}', but received '{response_queries:?}' from server {src}"
318 );
319 return Err(NetError::QueryCaseMismatch);
320 }
321 if !question_matches {
322 warn!(
323 "detected forged question section: we expected '{request_queries:?}', but received '{response_queries:?}' from server {src}"
324 );
325 continue;
326 }
327
328 if self.case_randomization {
330 if let Some(original_query) = original_query {
331 for response_query in response_queries.iter_mut() {
332 if response_query == original_query {
333 *response_query = original_query.clone();
334 }
335 }
336 }
337 }
338
339 debug!("received message id: {}", response.id);
340 #[cfg(feature = "__dnssec")]
341 if let Some(mut verifier) = verifier {
342 return Ok(verifier.verify(response_bytes)?);
343 }
344 return Ok(response);
345 }
346 }
347}
348
349pub struct UdpClientStreamBuilder<P> {
353 name_server: SocketAddr,
354 timeout: Option<Duration>,
355 #[cfg(feature = "__dnssec")]
356 signer: Option<TSigner>,
357 bind_addr: Option<SocketAddr>,
358 avoid_local_ports: Arc<HashSet<u16>>,
359 os_port_selection: bool,
360 provider: P,
361 max_retries: u8,
362 retry_interval_floor: Duration,
363}
364
365impl<P: RuntimeProvider> UdpClientStreamBuilder<P> {
366 pub fn with_timeout(mut self, timeout: Option<Duration>) -> Self {
368 self.timeout = timeout;
369 self
370 }
371
372 #[cfg(feature = "__dnssec")]
374 pub fn with_signer(self, signer: Option<TSigner>) -> Self {
375 Self {
376 name_server: self.name_server,
377 timeout: self.timeout,
378 signer,
379 bind_addr: self.bind_addr,
380 avoid_local_ports: self.avoid_local_ports,
381 os_port_selection: self.os_port_selection,
382 provider: self.provider,
383 max_retries: self.max_retries,
384 retry_interval_floor: self.retry_interval_floor,
385 }
386 }
387
388 pub fn with_bind_addr(mut self, bind_addr: Option<SocketAddr>) -> Self {
393 self.bind_addr = bind_addr;
394 self
395 }
396
397 pub fn avoid_local_ports(mut self, avoid_local_ports: Arc<HashSet<u16>>) -> Self {
400 self.avoid_local_ports = avoid_local_ports;
401 self
402 }
403
404 pub fn with_os_port_selection(mut self, os_port_selection: bool) -> Self {
406 self.os_port_selection = os_port_selection;
407 self
408 }
409
410 pub fn with_max_retries(mut self, max_retries: u8) -> Self {
412 self.max_retries = max_retries;
413 self
414 }
415
416 pub fn with_retry_interval_floor(mut self, floor: u64) -> Self {
418 self.retry_interval_floor = Duration::from_millis(floor);
419 self
420 }
421
422 pub fn exchange(self) -> DnsExchange<P> {
424 let mut handle = self.provider.create_handle();
425 let stream = self.build();
426 let (exchange, bg) = DnsExchange::from_stream(stream);
427 handle.spawn_bg(bg);
428 exchange
429 }
430
431 pub fn build(self) -> UdpClientStream<P> {
435 UdpClientStream {
436 name_server: self.name_server,
437 timeout: self.timeout.unwrap_or(Duration::from_secs(5)),
438 is_shutdown: false,
439 #[cfg(feature = "__dnssec")]
440 signer: self.signer,
441 bind_addr: self.bind_addr,
442 avoid_local_ports: self.avoid_local_ports.clone(),
443 os_port_selection: self.os_port_selection,
444 provider: self.provider,
445 max_retries: self.max_retries,
446 retry_interval_floor: self.retry_interval_floor,
447 }
448 }
449}
450
451async fn retry<Provider: RuntimeProvider>(
457 request: impl Request,
458 retry_interval_time: Duration,
459 max_tasks: usize,
460) -> Result<DnsResponse, NetError> {
461 let mut futures = FuturesUnordered::new();
462
463 let retry_timer = Provider::Timer::delay_for(retry_interval_time).fuse();
464 pin_mut!(retry_timer);
465
466 futures.push(request.send());
467 let mut tasks = 1;
468
469 loop {
470 futures_util::select! {
471 result = futures.next() => {
472 match result {
473 Some(result) => return result,
474 None => return Err(NetError::from("no tasks successful")),
475 }
476 }
477 _ = &mut retry_timer => {
478 if tasks < max_tasks {
479 tasks += 1;
480 futures.push(request.send());
481 retry_timer.set(Provider::Timer::delay_for(retry_interval_time).fuse());
482 }
483 }
484 }
485 }
486}
487
488trait Request {
489 async fn send(&self) -> Result<DnsResponse, NetError>;
490}
491
492#[cfg(all(test, feature = "tokio"))]
493mod tests {
494 #![allow(clippy::dbg_macro, clippy::print_stdout)]
495
496 use core::{
497 net::{IpAddr, Ipv4Addr, Ipv6Addr},
498 str::FromStr,
499 sync::atomic::{AtomicU8, Ordering},
500 };
501 use std::io;
502
503 use tokio::{net::UdpSocket, select, spawn, sync::oneshot, time::sleep};
504
505 use test_support::subscribe;
506
507 use super::*;
508 use crate::{
509 proto::{
510 op::{DnsRequestOptions, Query, ResponseCode},
511 rr::{Name, RData, Record, RecordType, rdata::NULL},
512 },
513 runtime::{TokioRuntimeProvider, TokioTime},
514 udp::tests::{
515 udp_client_stream_bad_id_test, udp_client_stream_empty_question_section_test,
516 udp_client_stream_test,
517 },
518 xfer::FirstAnswer,
519 };
520
521 #[tokio::test]
522 async fn test_udp_client_stream_ipv4() {
523 subscribe();
524 udp_client_stream_test(IpAddr::V4(Ipv4Addr::LOCALHOST), TokioRuntimeProvider::new()).await;
525 }
526
527 #[tokio::test]
528 async fn test_udp_client_stream_ipv4_bad_id() {
529 subscribe();
530 udp_client_stream_bad_id_test(IpAddr::V4(Ipv4Addr::LOCALHOST), TokioRuntimeProvider::new())
531 .await;
532 }
533
534 #[tokio::test]
535 async fn test_udp_client_stream_ipv4_empty_question_section() {
536 subscribe();
537 udp_client_stream_empty_question_section_test(
538 IpAddr::V4(Ipv4Addr::LOCALHOST),
539 TokioRuntimeProvider::new(),
540 )
541 .await;
542 }
543
544 #[tokio::test]
545 async fn test_udp_client_stream_ipv6() {
546 subscribe();
547 udp_client_stream_test(
548 IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1)),
549 TokioRuntimeProvider::new(),
550 )
551 .await;
552 }
553
554 #[tokio::test]
555 async fn test_udp_client_stream_ipv6_bad_id() {
556 subscribe();
557 udp_client_stream_bad_id_test(
558 IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1)),
559 TokioRuntimeProvider::new(),
560 )
561 .await;
562 }
563
564 #[tokio::test]
565 async fn test_udp_client_stream_ipv6_empty_question_section() {
566 subscribe();
567 udp_client_stream_empty_question_section_test(
568 IpAddr::V6(Ipv6Addr::LOCALHOST),
569 TokioRuntimeProvider::new(),
570 )
571 .await;
572 }
573
574 #[tokio::test(start_paused = true)]
575 async fn retry_handler_test() -> Result<(), NetError> {
576 let mut message = Message::query().into_response();
577 message.metadata.response_code = ResponseCode::NoError;
578
579 let ret = retry::<TokioRuntimeProvider>(
580 FixedResponse {
581 response: DnsResponse::from_message(message.clone())?,
582 },
583 Duration::from_millis(200),
584 5,
585 )
586 .await?;
587 assert_eq!(ret.response_code, ResponseCode::NoError);
588
589 let (req, tries) = DelayedResponse::new(
591 DnsResponse::from_message(message.clone()).unwrap(),
592 Duration::from_millis(100),
593 Arc::new(AtomicU8::new(0)),
594 );
595 retry::<TokioRuntimeProvider>(req, Duration::from_millis(200), 5).await?;
596 assert_eq!(tries.load(Ordering::Relaxed), 1);
597
598 let (req, tries) = DelayedResponse::new(
600 DnsResponse::from_message(message.clone()).unwrap(),
601 Duration::from_millis(1500),
602 Arc::new(AtomicU8::new(0)),
603 );
604 retry::<TokioRuntimeProvider>(req, Duration::from_millis(200), 5).await?;
605 assert_eq!(tries.load(Ordering::Relaxed), 5);
606
607 let (req, tries) = DelayedResponse::new(
609 DnsResponse::from_message(message.clone()).unwrap(),
610 Duration::from_millis(1000),
611 Arc::new(AtomicU8::new(0)),
612 );
613 let timer_ret = TokioTime::timeout(
614 Duration::from_millis(500),
615 retry::<TokioRuntimeProvider>(req, Duration::from_millis(200), 5),
616 )
617 .await;
618
619 if let Err(e) = timer_ret {
620 assert_eq!(e.kind(), io::ErrorKind::TimedOut);
621 } else {
622 panic!("timer did not timeout");
623 }
624
625 assert_eq!(tries.load(Ordering::Relaxed), 3);
626
627 Ok(())
628 }
629
630 struct FixedResponse {
631 response: DnsResponse,
632 }
633
634 impl Request for FixedResponse {
635 async fn send(&self) -> Result<DnsResponse, NetError> {
636 Ok(self.response.clone())
637 }
638 }
639
640 struct DelayedResponse {
641 response: DnsResponse,
642 delay: Duration,
643 counter: Arc<AtomicU8>,
644 }
645
646 impl DelayedResponse {
647 fn new(
648 response: DnsResponse,
649 delay: Duration,
650 counter: Arc<AtomicU8>,
651 ) -> (Self, Arc<AtomicU8>) {
652 (
653 Self {
654 response,
655 delay,
656 counter: counter.clone(),
657 },
658 counter,
659 )
660 }
661 }
662
663 impl Request for DelayedResponse {
664 async fn send(&self) -> Result<DnsResponse, NetError> {
665 let _ = self.counter.fetch_add(1, Ordering::Relaxed);
666 sleep(self.delay).await;
667 Ok(self.response.clone())
668 }
669 }
670
671 #[tokio::test]
675 async fn test_ignore_invalid_message() {
676 subscribe();
677
678 let provider = TokioRuntimeProvider::new();
679
680 let server_socket = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
682 let server_addr = server_socket.local_addr().unwrap();
683 let (shutdown_sender, mut shutdown_receiver) = oneshot::channel();
684 let server_handle = spawn(async move {
685 loop {
686 let mut buffer = [0u8; 4096];
687 let result = select! {
688 result = server_socket.recv_from(&mut buffer) => result,
689 _ = &mut shutdown_receiver => break,
690 };
691 let (len, addr) = result.unwrap();
692
693 let request = Message::from_vec(&buffer[..len]).unwrap();
694
695 for id in [
697 request.id ^ 1,
698 request.id ^ 2,
699 request.id ^ 4,
700 request.id ^ 8,
701 ] {
702 server_socket
703 .send_to(&id.to_be_bytes(), addr)
704 .await
705 .unwrap();
706 }
707
708 let query_name = request.queries[0].name.clone();
710 let mut message = request.into_response();
711 message.add_answer(Record::from_rdata(
712 query_name,
713 0,
714 RData::NULL(NULL::with(b"DEADBEEF".to_vec())),
715 ));
716 let bytes = message.to_vec().unwrap();
717 server_socket.send_to(&bytes, addr).await.unwrap();
718 }
719 });
720
721 let mut stream = UdpClientStream::builder(server_addr, provider)
723 .with_timeout(Some(Duration::from_millis(500)))
724 .build();
725
726 let mut query_message = Message::query();
727 query_message.add_query(Query::query(
728 Name::from_str("dead.beef.").unwrap(),
729 RecordType::NULL,
730 ));
731 let response_stream =
732 stream.send_message(DnsRequest::new(query_message, DnsRequestOptions::default()));
733 let response = response_stream
734 .first_answer()
735 .await
736 .expect("failed to read response in the presence of spoofed invalid messages");
737 let RData::NULL(rdata) = &response.answers[0].data else {
738 panic!("unexpected record type: {response:?}");
739 };
740 assert_eq!(rdata.anything, b"DEADBEEF");
741
742 shutdown_sender.send(()).unwrap();
743 server_handle.await.unwrap();
744 }
745}