Skip to main content

netscli_core/
dns.rs

1use hickory_resolver::proto::rr::RecordType;
2use hickory_resolver::TokioAsyncResolver;
3use serde::Serialize;
4use std::net::IpAddr;
5use std::sync::OnceLock;
6use std::time::Duration;
7#[cfg(windows)]
8use std::{process::Stdio, str::FromStr};
9#[cfg(windows)]
10use tokio::process::Command;
11use tokio::time::timeout;
12
13use crate::error::{Error, Result};
14
15#[derive(Debug, Serialize)]
16pub struct DnsRecord {
17    /// Record type (e.g. "A", "AAAA", "MX"). Always upper-case.
18    pub record_type: String,
19    /// Display-normalized value. For textual records we strip surrounding
20    /// quotes and the trailing `.` that the resolver appends to FQDNs.
21    pub value: String,
22}
23
24const ALL_RECORD_TYPES: &[RecordType] = &[
25    RecordType::A,
26    RecordType::AAAA,
27    RecordType::CNAME,
28    RecordType::MX,
29    RecordType::NS,
30    RecordType::TXT,
31    RecordType::SRV,
32    RecordType::PTR,
33    RecordType::SOA,
34    RecordType::CAA,
35];
36
37pub fn parse_record_type(value: &str) -> Option<RecordType> {
38    match value.trim().to_uppercase().as_str() {
39        "A" => Some(RecordType::A),
40        "AAAA" => Some(RecordType::AAAA),
41        "CNAME" => Some(RecordType::CNAME),
42        "MX" => Some(RecordType::MX),
43        "NS" => Some(RecordType::NS),
44        "TXT" => Some(RecordType::TXT),
45        "SRV" => Some(RecordType::SRV),
46        "PTR" => Some(RecordType::PTR),
47        "SOA" => Some(RecordType::SOA),
48        "CAA" => Some(RecordType::CAA),
49        _ => None,
50    }
51}
52
53/// Shared resolver — parsing the system config (`/etc/resolv.conf` or the
54/// Windows registry) on every lookup is wasteful for high-volume scans like
55/// a /24 with reverse DNS enabled.
56fn shared_resolver() -> Result<&'static TokioAsyncResolver> {
57    static RESOLVER: OnceLock<std::result::Result<TokioAsyncResolver, String>> = OnceLock::new();
58    let cached = RESOLVER
59        .get_or_init(|| TokioAsyncResolver::tokio_from_system_conf().map_err(|e| e.to_string()));
60    match cached {
61        Ok(r) => Ok(r),
62        Err(e) => Err(Error::dns(format!(
63            "failed to load DNS resolver config: {e}"
64        ))),
65    }
66}
67
68/// Normalize a raw `RData.to_string()` value for display.
69///
70/// The hickory resolver prints FQDNs with a trailing dot and wraps TXT values
71/// in double quotes. Both are technically correct per the DNS wire format but
72/// users expect `example.com` not `example.com.` and `v=spf1 -all` not
73/// `"v=spf1 -all"`. We strip those for presentation.
74fn normalize_value(raw: &str) -> String {
75    let s = raw.trim();
76    // TXT values come back as `"chunk1" "chunk2"`; join into one string.
77    if s.starts_with('"') && s.ends_with('"') && s.len() >= 2 {
78        let inner = &s[1..s.len() - 1];
79        // Collapse `" "` interior separators from multi-chunk TXT records.
80        return inner.replace("\" \"", "");
81    }
82    // Strip a trailing FQDN dot but keep a lone "." (root) untouched.
83    if s.len() > 1 && s.ends_with('.') {
84        s.trim_end_matches('.').to_string()
85    } else {
86        s.to_string()
87    }
88}
89
90pub async fn lookup_record_timeout(
91    host: &str,
92    record_type: RecordType,
93    timeout_ms: u64,
94) -> Result<Vec<DnsRecord>> {
95    let resolver = shared_resolver()?;
96    let response = match timeout(
97        Duration::from_millis(timeout_ms),
98        resolver.lookup(host, record_type),
99    )
100    .await
101    {
102        Ok(Ok(resp)) => resp,
103        Ok(Err(e)) => return Err(Error::dns(format!("{record_type} lookup failed: {e}"))),
104        Err(_) => return Err(Error::Timeout(timeout_ms)),
105    };
106
107    let mut records = Vec::new();
108    for record in response.record_iter() {
109        let Some(data) = record.data() else {
110            continue;
111        };
112        records.push(DnsRecord {
113            record_type: record_type.to_string(),
114            value: normalize_value(&data.to_string()),
115        });
116    }
117    Ok(records)
118}
119
120/// Look up every supported record type, returning a merged list.
121///
122/// Queries run sequentially. An earlier attempt to parallelize via
123/// `join_all` produced surprising regressions on Windows — a single slow
124/// record type (commonly CAA for many domains) would cause the whole
125/// batch to surface its timeout error even when other types had already
126/// returned results. The cached `shared_resolver()` keeps each
127/// individual query fast by avoiding `resolv.conf`/registry reparsing,
128/// so sequential iteration is not noticeably slower in practice and
129/// yields reliable partial-result behavior.
130///
131/// Partial failures (some types timeout, others return results) yield
132/// `Ok(records)` — only a total failure with zero records surfaces the
133/// last error.
134pub async fn lookup_all_records_timeout(host: &str, timeout_ms: u64) -> Result<Vec<DnsRecord>> {
135    let mut records = Vec::new();
136    let mut last_err: Option<Error> = None;
137
138    for record_type in ALL_RECORD_TYPES {
139        match lookup_record_timeout(host, *record_type, timeout_ms).await {
140            Ok(mut found) => records.append(&mut found),
141            Err(e) => last_err = Some(e),
142        }
143    }
144
145    if records.is_empty() {
146        if let Some(err) = last_err {
147            return Err(err);
148        }
149        return Err(Error::dns(format!("no DNS records found for {host}")));
150    }
151
152    Ok(records)
153}
154
155pub async fn resolve_a(host: &str) -> Result<Vec<String>> {
156    resolve_a_timeout(host, crate::DEFAULT_DNS_TIMEOUT_MS).await
157}
158
159pub async fn resolve_a_timeout(host: &str, timeout_ms: u64) -> Result<Vec<String>> {
160    let resolver = shared_resolver()?;
161    let response = match timeout(
162        Duration::from_millis(timeout_ms),
163        resolver.ipv4_lookup(host),
164    )
165    .await
166    {
167        Ok(Ok(resp)) => resp,
168        Ok(Err(e)) => return Err(Error::dns(format!("A lookup failed: {e}"))),
169        Err(_) => return Err(Error::Timeout(timeout_ms)),
170    };
171    Ok(response.iter().map(|ip| ip.to_string()).collect())
172}
173
174pub async fn resolve_aaaa(host: &str) -> Result<Vec<String>> {
175    resolve_aaaa_timeout(host, crate::DEFAULT_DNS_TIMEOUT_MS).await
176}
177
178pub async fn resolve_aaaa_timeout(host: &str, timeout_ms: u64) -> Result<Vec<String>> {
179    let resolver = shared_resolver()?;
180    let response = match timeout(
181        Duration::from_millis(timeout_ms),
182        resolver.ipv6_lookup(host),
183    )
184    .await
185    {
186        Ok(Ok(resp)) => resp,
187        Ok(Err(e)) => return Err(Error::dns(format!("AAAA lookup failed: {e}"))),
188        Err(_) => return Err(Error::Timeout(timeout_ms)),
189    };
190    Ok(response.iter().map(|ip| ip.to_string()).collect())
191}
192
193pub async fn reverse_lookup_timeout(ip: IpAddr, timeout_ms: u64) -> Result<Option<String>> {
194    let resolver = shared_resolver()?;
195    let resp = match timeout(
196        Duration::from_millis(timeout_ms),
197        resolver.reverse_lookup(ip),
198    )
199    .await
200    {
201        Ok(Ok(r)) => r,
202        Ok(Err(e)) => return Err(Error::dns(format!("reverse lookup failed: {e}"))),
203        Err(_) => return Err(Error::Timeout(timeout_ms)),
204    };
205    let name = resp.iter().next().map(|n| n.to_utf8());
206    Ok(name.filter(|s| !s.is_empty()))
207}
208
209/// Reverse-resolve an IP address to a hostname.
210///
211/// - On Windows, `ping -a` often resolves names via LLMNR/NetBIOS even when
212///   no PTR records exist.
213/// - On other platforms, this falls back to a DNS PTR lookup.
214///
215/// Both paths run their result through `normalize_hostname` so callers get
216/// consistent output regardless of the OS we're on.
217pub async fn reverse_lookup_best_effort_timeout(ip: IpAddr, timeout_ms: u64) -> Option<String> {
218    #[cfg(windows)]
219    if let Some(name) = reverse_lookup_windows_ping(ip, timeout_ms).await {
220        if let Some(normalized) = normalize_hostname(name) {
221            return Some(normalized);
222        }
223    }
224
225    reverse_lookup_timeout(ip, timeout_ms)
226        .await
227        .ok()
228        .flatten()
229        .and_then(normalize_hostname)
230}
231
232fn normalize_hostname(name: String) -> Option<String> {
233    let name = name.trim().trim_end_matches('.').trim();
234    if name.is_empty() {
235        None
236    } else {
237        Some(name.to_string())
238    }
239}
240
241#[cfg(windows)]
242async fn reverse_lookup_windows_ping(ip: IpAddr, timeout_ms: u64) -> Option<String> {
243    // Only IPv4 is supported by the ping parsing below.
244    if !matches!(ip, IpAddr::V4(_)) {
245        return None;
246    }
247
248    let ip_s = ip.to_string();
249    let wait_ms = timeout_ms.saturating_add(250);
250    let mut cmd = Command::new("ping");
251    cmd.args(["-a", "-n", "1", "-w", &timeout_ms.to_string(), &ip_s])
252        .stdout(Stdio::piped())
253        .stderr(Stdio::piped())
254        .kill_on_drop(true);
255
256    let output = match timeout(Duration::from_millis(wait_ms), cmd.output()).await {
257        Ok(Ok(out)) => out,
258        _ => return None,
259    };
260
261    let stdout = String::from_utf8_lossy(&output.stdout);
262
263    // Only scan the first few lines — `ping -a` emits the hostname in its
264    // opening "Pinging <name> [<ip>] ..." banner. Later reply lines (e.g.
265    // "Reply from 192.168.1.5: bytes=32 ..." on some locales) can also
266    // contain `[ip]` and would otherwise misparse as a hostname.
267    for line in stdout.lines().take(4) {
268        if !line.contains('[') || !line.contains(']') {
269            continue;
270        }
271        let Some(before) = line.split('[').next() else {
272            continue;
273        };
274        let before = before.trim();
275
276        // Take the last whitespace token before the "[ip]" segment as the
277        // candidate hostname. This is locale-tolerant: English "Pinging X [ip]",
278        // Spanish "Haciendo ping a X [ip] con 32 ...", etc.
279        let Some(candidate) = before.split_whitespace().last() else {
280            continue;
281        };
282        let candidate = candidate.trim().trim_end_matches('.');
283        if candidate.is_empty() {
284            continue;
285        }
286
287        // If the "hostname" is actually just the IP literal, there was no
288        // reverse-resolution — fall through to PTR lookup.
289        if IpAddr::from_str(candidate).is_ok() {
290            return None;
291        }
292
293        return Some(candidate.to_string());
294    }
295
296    None
297}
298
299#[cfg(test)]
300mod tests {
301    use super::*;
302
303    #[test]
304    fn normalize_value_strips_trailing_dot() {
305        assert_eq!(normalize_value("example.com."), "example.com");
306    }
307
308    #[test]
309    fn normalize_value_preserves_root_dot() {
310        assert_eq!(normalize_value("."), ".");
311    }
312
313    #[test]
314    fn normalize_value_unwraps_txt_quotes() {
315        assert_eq!(normalize_value(r#""v=spf1 -all""#), "v=spf1 -all");
316    }
317
318    #[test]
319    fn normalize_value_joins_multichunk_txt() {
320        assert_eq!(normalize_value(r#""chunk1" "chunk2""#), "chunk1chunk2");
321    }
322
323    #[test]
324    fn normalize_value_untouched_for_plain_a() {
325        assert_eq!(normalize_value("192.0.2.1"), "192.0.2.1");
326    }
327
328    #[test]
329    fn normalize_hostname_drops_trailing_dot_and_ws() {
330        assert_eq!(
331            normalize_hostname("host.example.com.".to_string()),
332            Some("host.example.com".to_string())
333        );
334        assert_eq!(
335            normalize_hostname("  host.example.com  ".to_string()),
336            Some("host.example.com".to_string())
337        );
338    }
339
340    #[test]
341    fn normalize_hostname_rejects_empty() {
342        assert_eq!(normalize_hostname("".to_string()), None);
343        assert_eq!(normalize_hostname(".".to_string()), None);
344    }
345
346    #[tokio::test]
347    async fn test_resolve_a_localhost() {
348        let result = resolve_a("localhost").await;
349        assert!(result.is_ok());
350        let ips = result.unwrap();
351        assert!(!ips.is_empty());
352        // localhost should resolve to 127.0.0.1
353        assert!(ips.contains(&"127.0.0.1".to_string()));
354    }
355
356    // NOTE: Unit tests avoid network-dependent DNS behavior.
357    // Network integration tests should be added separately (and marked as ignored).
358}