Skip to main content

ip_discovery/dns/
mod.rs

1//! DNS protocol implementation for public IP detection
2//!
3//! Uses DNS TXT/A records from special domains to detect public IP.
4//! This implementation uses raw UDP sockets instead of external DNS libraries.
5//!
6//! # Security
7//!
8//! Transaction IDs are generated with [`getrandom`] (OS-level CSPRNG),
9//! preventing DNS transaction ID spoofing attacks.
10
11mod 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/// Record type for DNS query
36#[derive(Debug, Clone, Copy)]
37pub enum DnsRecordType {
38    /// A/AAAA record (direct IP)
39    Address,
40    /// TXT record (IP as text)
41    Txt,
42}
43
44/// DNS provider configuration
45#[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    /// Create a new DNS provider
59    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    /// Set DNS class (for special queries like Cloudflare CHAOS)
78    pub fn with_class(mut self, class: DnsClass) -> Self {
79        self.dns_class = class;
80        self
81    }
82
83    /// Set IPv6 support
84    pub fn with_v6(mut self, supports: bool) -> Self {
85        self.supports_v6 = supports;
86        self
87    }
88
89    /// Set IPv6 resolver address
90    ///
91    /// When requesting IPv6, the query is sent to this resolver so the
92    /// DNS server sees the client's IPv6 source address.
93    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    /// Query for IP address synchronously using standard UDP sockets with a timeout
158    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]; // DNS Flag Day 2020 safe UDP size (RFC 6891 EDNS0)
197        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    /// Query for IP address using raw UDP asynchronously
211    #[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        // Build query packet
217        let query = build_query(&self.query_domain, record_type, self.dns_class)
218            .map_err(|e| ProviderError::new(&self.name, e))?;
219
220        // Create UDP socket
221        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        // Send query
236        socket
237            .send(&query)
238            .await
239            .map_err(|e| ProviderError::new(&self.name, e))?;
240
241        // Receive response
242        let mut buf = [0u8; 1232]; // DNS Flag Day 2020 safe UDP size (RFC 6891 EDNS0)
243        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        // Parse response
252        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}