hickory_net/xfer/
retry_dns_handle.rs1use core::pin::Pin;
11use core::task::{Context, Poll};
12
13use futures_util::stream::{Stream, StreamExt};
14
15use crate::xfer::{DnsHandle, DnsRequest, DnsResponse};
16use crate::{DnsError, NetError};
17
18#[derive(Clone)]
28#[must_use = "queries can only be sent through a ClientHandle"]
29#[allow(dead_code)]
30pub struct RetryDnsHandle<H> {
31 handle: H,
32 attempts: usize,
33}
34
35impl<H> RetryDnsHandle<H> {
36 pub fn new(handle: H, attempts: usize) -> Self {
43 Self { handle, attempts }
44 }
45}
46
47impl<H: DnsHandle> DnsHandle for RetryDnsHandle<H> {
48 type Response = Pin<Box<dyn Stream<Item = Result<DnsResponse, NetError>> + Send + Unpin>>;
49 type Runtime = H::Runtime;
50
51 fn send(&self, request: DnsRequest) -> Self::Response {
52 let stream = self.handle.send(request.clone());
55
56 Box::pin(RetrySendStream {
57 request,
58 handle: self.handle.clone(),
59 stream,
60 remaining_attempts: self.attempts,
61 })
62 }
63}
64
65struct RetrySendStream<H: DnsHandle> {
67 request: DnsRequest,
68 handle: H,
69 stream: <H as DnsHandle>::Response,
70 remaining_attempts: usize,
71}
72
73impl<H: DnsHandle> Stream for RetrySendStream<H> {
74 type Item = Result<DnsResponse, NetError>;
75
76 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
77 loop {
80 let err = match self.stream.poll_next_unpin(cx) {
81 Poll::Ready(Some(Err(e))) => e,
82 poll => return poll,
83 };
84
85 match (self.remaining_attempts, err) {
86 (0, err) => return Poll::Ready(Some(Err(err))),
88 (
90 _,
91 err @ NetError::NoConnections
92 | err @ NetError::Dns(DnsError::NoRecordsFound(_))
93 | err @ NetError::Dns(DnsError::DnssecBogus),
94 ) => return Poll::Ready(Some(Err(err))),
95 (_, NetError::Busy) => {}
97 (_, _) => self.remaining_attempts -= 1,
99 }
100
101 let request = self.request.clone();
104 self.stream = self.handle.send(request);
105 }
106 }
107}
108
109#[cfg(all(test, feature = "tokio"))]
110mod test {
111 use core::sync::atomic::{AtomicU16, Ordering};
112 use std::sync::Arc;
113
114 use futures_executor::block_on;
115 use futures_util::future::{err, ok};
116 use futures_util::stream::{Stream, once};
117
118 use super::*;
119 use crate::proto::op::Message;
120 use crate::runtime::TokioRuntimeProvider;
121 use crate::xfer::{DnsHandle, DnsRequest, DnsResponse, FirstAnswer};
122 use test_support::subscribe;
123
124 #[derive(Clone)]
125 struct TestClient {
126 last_succeed: bool,
127 retries: u16,
128 attempts: Arc<AtomicU16>,
129 }
130
131 impl DnsHandle for TestClient {
132 type Response = Box<dyn Stream<Item = Result<DnsResponse, NetError>> + Send + Unpin>;
133 type Runtime = TokioRuntimeProvider;
134
135 fn send(&self, _: DnsRequest) -> Self::Response {
136 let i = self.attempts.load(Ordering::SeqCst);
137
138 if (i > self.retries || self.retries - i == 0) && self.last_succeed {
139 let mut message = Message::query();
140 message.metadata.id = i;
141 return Box::new(once(ok(
142 DnsResponse::from_message(message.into_response()).unwrap()
143 )));
144 }
145
146 self.attempts.fetch_add(1, Ordering::SeqCst);
147 Box::new(once(err(NetError::from("last retry set to fail"))))
148 }
149 }
150
151 #[test]
152 fn test_retry() {
153 subscribe();
154 let handle = RetryDnsHandle::new(
155 TestClient {
156 last_succeed: true,
157 retries: 1,
158 attempts: Arc::new(AtomicU16::new(0)),
159 },
160 2,
161 );
162 let test1 = DnsRequest::from(Message::query());
163 let result = block_on(handle.send(test1).first_answer()).expect("should have succeeded");
164 assert_eq!(result.id, 1); }
166
167 #[test]
168 fn test_error() {
169 subscribe();
170 let client = RetryDnsHandle::new(
171 TestClient {
172 last_succeed: false,
173 retries: 1,
174 attempts: Arc::new(AtomicU16::new(0)),
175 },
176 2,
177 );
178 let test1 = DnsRequest::from(Message::query());
179 assert!(block_on(client.send(test1).first_answer()).is_err());
180 }
181}