1mod protocol;
12pub(crate) mod providers;
13
14pub use protocol::DnsClass;
15
16#[cfg(feature = "tokio")]
17pub use providers::default_providers;
18pub use providers::{default_blocking_providers, provider_names};
19
20use crate::error::ProviderError;
21#[cfg(feature = "tokio")]
22use crate::provider::Provider;
23use crate::provider::{BlockingProvider, BoxedBlockingProvider};
24use crate::types::{IpVersion, Protocol};
25use protocol::{build_query, parse_response, RecordType};
26use std::net::{IpAddr, SocketAddr};
27use std::str::FromStr;
28use std::time::Duration;
29
30#[cfg(feature = "tokio")]
31use std::future::Future;
32#[cfg(feature = "tokio")]
33use std::pin::Pin;
34
35#[derive(Debug, Clone, Copy)]
37pub enum DnsRecordType {
38 Address,
40 Txt,
42}
43
44#[derive(Debug, Clone)]
46pub struct DnsProvider {
47 name: String,
48 query_domain: String,
49 resolver_addr: SocketAddr,
50 resolver_addr_v6: Option<SocketAddr>,
51 record_type: DnsRecordType,
52 dns_class: DnsClass,
53 supports_v4: bool,
54 supports_v6: bool,
55}
56
57impl DnsProvider {
58 pub fn new(
60 name: impl Into<String>,
61 query_domain: impl Into<String>,
62 resolver_addr: SocketAddr,
63 record_type: DnsRecordType,
64 ) -> Self {
65 Self {
66 name: name.into(),
67 query_domain: query_domain.into(),
68 resolver_addr,
69 resolver_addr_v6: None,
70 record_type,
71 dns_class: DnsClass::In,
72 supports_v4: true,
73 supports_v6: false,
74 }
75 }
76
77 pub fn with_class(mut self, class: DnsClass) -> Self {
79 self.dns_class = class;
80 self
81 }
82
83 pub fn with_v6(mut self, supports: bool) -> Self {
85 self.supports_v6 = supports;
86 self
87 }
88
89 pub fn with_v6_resolver(mut self, addr: SocketAddr) -> Self {
94 self.resolver_addr_v6 = Some(addr);
95 self.supports_v6 = true;
96 self
97 }
98
99 fn select_resolver(&self, version: IpVersion) -> SocketAddr {
100 match version {
101 IpVersion::V6 => self.resolver_addr_v6.unwrap_or(self.resolver_addr),
102 _ => self.resolver_addr,
103 }
104 }
105
106 fn select_record_type(&self, version: IpVersion) -> RecordType {
107 match self.record_type {
108 DnsRecordType::Address => match version {
109 IpVersion::V6 => RecordType::Aaaa,
110 _ => RecordType::A,
111 },
112 DnsRecordType::Txt => RecordType::Txt,
113 }
114 }
115
116 fn extract_ip(
117 &self,
118 results: Vec<String>,
119 version: IpVersion,
120 ) -> Result<IpAddr, ProviderError> {
121 for result in results {
122 for part in result.split_whitespace() {
123 let ip_str = part.split('/').next().unwrap_or(part);
124 if let Ok(ip) = IpAddr::from_str(ip_str) {
125 match version {
126 IpVersion::V4 if ip.is_ipv4() => return Ok(ip),
127 IpVersion::V6 if ip.is_ipv6() => return Ok(ip),
128 IpVersion::Any => return Ok(ip),
129 _ => continue,
130 }
131 }
132 }
133 }
134
135 Err(ProviderError::message(
136 &self.name,
137 "no valid IP in DNS response",
138 ))
139 }
140
141 fn validate_response(query: &[u8], response: &[u8]) -> Result<(), &'static str> {
142 if query.len() < 2 || response.len() < 12 {
143 return Err("response too short");
144 }
145 if response[..2] != query[..2] {
146 return Err("transaction ID mismatch");
147 }
148 if response[2] & 0x80 == 0 {
149 return Err("packet is not a DNS response");
150 }
151 if response[2] & 0x02 != 0 {
152 return Err("truncated DNS response");
153 }
154 Ok(())
155 }
156
157 pub fn query_blocking(
159 &self,
160 version: IpVersion,
161 timeout: Duration,
162 ) -> Result<IpAddr, ProviderError> {
163 if timeout.is_zero() {
164 return Err(ProviderError::message(&self.name, "timeout"));
165 }
166
167 let resolver = self.select_resolver(version);
168 let record_type = self.select_record_type(version);
169
170 let query = build_query(&self.query_domain, record_type, self.dns_class)
171 .map_err(|e| ProviderError::new(&self.name, e))?;
172
173 let bind_addr = if resolver.is_ipv6() {
174 "[::]:0"
175 } else {
176 "0.0.0.0:0"
177 };
178 let socket =
179 std::net::UdpSocket::bind(bind_addr).map_err(|e| ProviderError::new(&self.name, e))?;
180
181 socket
182 .set_read_timeout(Some(timeout))
183 .map_err(|e| ProviderError::new(&self.name, e))?;
184 socket
185 .set_write_timeout(Some(timeout))
186 .map_err(|e| ProviderError::new(&self.name, e))?;
187
188 socket
189 .connect(resolver)
190 .map_err(|e| ProviderError::new(&self.name, e))?;
191
192 socket
193 .send(&query)
194 .map_err(|e| ProviderError::new(&self.name, e))?;
195
196 let mut buf = [0u8; 1232]; let len = socket
198 .recv(&mut buf)
199 .map_err(|e| ProviderError::new(&self.name, e))?;
200
201 Self::validate_response(&query, &buf[..len])
202 .map_err(|e| ProviderError::message(&self.name, e))?;
203
204 let results = parse_response(&buf[..len], record_type)
205 .map_err(|e| ProviderError::message(&self.name, e))?;
206
207 self.extract_ip(results, version)
208 }
209
210 #[cfg(feature = "tokio")]
212 async fn query(&self, version: IpVersion) -> Result<IpAddr, ProviderError> {
213 let resolver = self.select_resolver(version);
214 let record_type = self.select_record_type(version);
215
216 let query = build_query(&self.query_domain, record_type, self.dns_class)
218 .map_err(|e| ProviderError::new(&self.name, e))?;
219
220 let bind_addr = if resolver.is_ipv6() {
222 "[::]:0"
223 } else {
224 "0.0.0.0:0"
225 };
226 let socket = tokio::net::UdpSocket::bind(bind_addr)
227 .await
228 .map_err(|e| ProviderError::new(&self.name, e))?;
229
230 socket
231 .connect(resolver)
232 .await
233 .map_err(|e| ProviderError::new(&self.name, e))?;
234
235 socket
237 .send(&query)
238 .await
239 .map_err(|e| ProviderError::new(&self.name, e))?;
240
241 let mut buf = [0u8; 1232]; let len = socket
244 .recv(&mut buf)
245 .await
246 .map_err(|e| ProviderError::new(&self.name, e))?;
247
248 Self::validate_response(&query, &buf[..len])
249 .map_err(|e| ProviderError::message(&self.name, e))?;
250
251 let results = parse_response(&buf[..len], record_type)
253 .map_err(|e| ProviderError::message(&self.name, e))?;
254
255 self.extract_ip(results, version)
256 }
257}
258
259impl BlockingProvider for DnsProvider {
260 fn name(&self) -> &str {
261 &self.name
262 }
263
264 fn protocol(&self) -> Protocol {
265 Protocol::Dns
266 }
267
268 fn supports_v4(&self) -> bool {
269 self.supports_v4
270 }
271
272 fn supports_v6(&self) -> bool {
273 self.supports_v6
274 }
275
276 fn get_ip(&self, version: IpVersion, timeout: Duration) -> Result<IpAddr, ProviderError> {
277 self.query_blocking(version, timeout)
278 }
279
280 fn clone_box(&self) -> BoxedBlockingProvider {
281 Box::new(self.clone())
282 }
283}
284
285#[cfg(feature = "tokio")]
286impl Provider for DnsProvider {
287 fn name(&self) -> &str {
288 &self.name
289 }
290
291 fn protocol(&self) -> Protocol {
292 Protocol::Dns
293 }
294
295 fn supports_v4(&self) -> bool {
296 self.supports_v4
297 }
298
299 fn supports_v6(&self) -> bool {
300 self.supports_v6
301 }
302
303 fn get_ip(
304 &self,
305 version: IpVersion,
306 ) -> Pin<Box<dyn Future<Output = Result<IpAddr, ProviderError>> + Send + '_>> {
307 Box::pin(self.query(version))
308 }
309}
310
311#[cfg(test)]
312mod blocking_tests {
313 use super::{DnsProvider, DnsRecordType};
314 use crate::types::IpVersion;
315 use std::net::UdpSocket;
316 use std::time::{Duration, Instant};
317
318 fn response_for(query: &[u8], transaction_id: Option<[u8; 2]>) -> Vec<u8> {
319 let mut response = query.to_vec();
320 if let Some(id) = transaction_id {
321 response[..2].copy_from_slice(&id);
322 }
323 response[2] = 0x81;
324 response[3] = 0x80;
325 response[6] = 0;
326 response[7] = 1;
327 response.extend_from_slice(&[
328 0xc0, 0x0c, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x00, 0x04, 203, 0, 113, 9,
329 ]);
330 response
331 }
332
333 fn provider_for(addr: std::net::SocketAddr) -> DnsProvider {
334 DnsProvider::new("local-dns", "example.test", addr, DnsRecordType::Address)
335 }
336
337 #[test]
338 fn zero_timeout_fails_before_waiting_for_dns_response() {
339 let server = UdpSocket::bind("127.0.0.1:0").unwrap();
340 let addr = server.local_addr().unwrap();
341 server
342 .set_read_timeout(Some(Duration::from_millis(200)))
343 .unwrap();
344 std::thread::spawn(move || {
345 let mut query = [0u8; 512];
346 if let Ok((len, peer)) = server.recv_from(&mut query) {
347 std::thread::sleep(Duration::from_millis(100));
348 let _ = server.send_to(&response_for(&query[..len], None), peer);
349 }
350 });
351
352 let started = Instant::now();
353 let result = provider_for(addr).query_blocking(IpVersion::V4, Duration::ZERO);
354
355 assert!(result.is_err());
356 assert!(started.elapsed() < Duration::from_millis(50));
357 }
358
359 #[test]
360 fn rejects_dns_response_with_wrong_transaction_id() {
361 let server = UdpSocket::bind("127.0.0.1:0").unwrap();
362 let addr = server.local_addr().unwrap();
363 std::thread::spawn(move || {
364 let mut query = [0u8; 512];
365 let (len, peer) = server.recv_from(&mut query).unwrap();
366 let wrong_id = [query[0].wrapping_add(1), query[1]];
367 let _ = server.send_to(&response_for(&query[..len], Some(wrong_id)), peer);
368 });
369
370 let result = provider_for(addr).query_blocking(IpVersion::V4, Duration::from_secs(1));
371
372 assert!(result
373 .unwrap_err()
374 .to_string()
375 .contains("transaction ID mismatch"));
376 }
377
378 #[test]
379 fn rejects_dns_response_from_unexpected_sender() {
380 let server = UdpSocket::bind("127.0.0.1:0").unwrap();
381 let addr = server.local_addr().unwrap();
382 std::thread::spawn(move || {
383 let mut query = [0u8; 512];
384 let (len, peer) = server.recv_from(&mut query).unwrap();
385 let unexpected_sender = UdpSocket::bind("127.0.0.1:0").unwrap();
386 let _ = unexpected_sender.send_to(&response_for(&query[..len], None), peer);
387 });
388
389 let result = provider_for(addr).query_blocking(IpVersion::V4, Duration::from_millis(30));
390
391 assert!(result.is_err());
392 }
393}