1use hickory_resolver::proto::rr::{RData, RecordType};
2use hickory_resolver::TokioResolver;
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 pub record_type: String,
19 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
53fn shared_resolver() -> Result<&'static TokioResolver> {
57 static RESOLVER: OnceLock<std::result::Result<TokioResolver, String>> = OnceLock::new();
58 let cached = RESOLVER.get_or_init(|| {
59 TokioResolver::builder_tokio()
64 .and_then(|b| b.build())
65 .map_err(|e| e.to_string())
66 });
67 match cached {
68 Ok(r) => Ok(r),
69 Err(e) => Err(Error::dns(format!(
70 "failed to load DNS resolver config: {e}"
71 ))),
72 }
73}
74
75fn normalize_value(raw: &str) -> String {
82 let s = raw.trim();
83 if s.starts_with('"') && s.ends_with('"') && s.len() >= 2 {
85 let inner = &s[1..s.len() - 1];
86 return inner.replace("\" \"", "");
88 }
89 if s.len() > 1 && s.ends_with('.') {
91 s.trim_end_matches('.').to_string()
92 } else {
93 s.to_string()
94 }
95}
96
97pub async fn lookup_record_timeout(
98 host: &str,
99 record_type: RecordType,
100 timeout_ms: u64,
101) -> Result<Vec<DnsRecord>> {
102 let resolver = shared_resolver()?;
103 let response = match timeout(
104 Duration::from_millis(timeout_ms),
105 resolver.lookup(host, record_type),
106 )
107 .await
108 {
109 Ok(Ok(resp)) => resp,
110 Ok(Err(e)) => return Err(Error::dns(format!("{record_type} lookup failed: {e}"))),
111 Err(_) => return Err(Error::Timeout(timeout_ms)),
112 };
113
114 let mut records = Vec::new();
121 for record in response.answers() {
122 records.push(DnsRecord {
123 record_type: record_type.to_string(),
124 value: normalize_value(&record.data.to_string()),
125 });
126 }
127 Ok(records)
128}
129
130pub async fn lookup_all_records_timeout(host: &str, timeout_ms: u64) -> Result<Vec<DnsRecord>> {
145 let mut records = Vec::new();
146 let mut last_err: Option<Error> = None;
147
148 for record_type in ALL_RECORD_TYPES {
149 match lookup_record_timeout(host, *record_type, timeout_ms).await {
150 Ok(mut found) => records.append(&mut found),
151 Err(e) => last_err = Some(e),
152 }
153 }
154
155 if records.is_empty() {
156 if let Some(err) = last_err {
157 return Err(err);
158 }
159 return Err(Error::dns(format!("no DNS records found for {host}")));
160 }
161
162 Ok(records)
163}
164
165pub async fn resolve_a(host: &str) -> Result<Vec<String>> {
166 resolve_a_timeout(host, crate::DEFAULT_DNS_TIMEOUT_MS).await
167}
168
169pub async fn resolve_a_timeout(host: &str, timeout_ms: u64) -> Result<Vec<String>> {
170 let resolver = shared_resolver()?;
171 let response = match timeout(
172 Duration::from_millis(timeout_ms),
173 resolver.ipv4_lookup(host),
174 )
175 .await
176 {
177 Ok(Ok(resp)) => resp,
178 Ok(Err(e)) => return Err(Error::dns(format!("A lookup failed: {e}"))),
179 Err(_) => return Err(Error::Timeout(timeout_ms)),
180 };
181 Ok(response
185 .answers()
186 .iter()
187 .filter_map(|r| match &r.data {
188 RData::A(a) => Some(a.to_string()),
189 _ => None,
190 })
191 .collect())
192}
193
194pub async fn resolve_aaaa(host: &str) -> Result<Vec<String>> {
195 resolve_aaaa_timeout(host, crate::DEFAULT_DNS_TIMEOUT_MS).await
196}
197
198pub async fn resolve_aaaa_timeout(host: &str, timeout_ms: u64) -> Result<Vec<String>> {
199 let resolver = shared_resolver()?;
200 let response = match timeout(
201 Duration::from_millis(timeout_ms),
202 resolver.ipv6_lookup(host),
203 )
204 .await
205 {
206 Ok(Ok(resp)) => resp,
207 Ok(Err(e)) => return Err(Error::dns(format!("AAAA lookup failed: {e}"))),
208 Err(_) => return Err(Error::Timeout(timeout_ms)),
209 };
210 Ok(response
211 .answers()
212 .iter()
213 .filter_map(|r| match &r.data {
214 RData::AAAA(a) => Some(a.to_string()),
215 _ => None,
216 })
217 .collect())
218}
219
220pub async fn reverse_lookup_timeout(ip: IpAddr, timeout_ms: u64) -> Result<Option<String>> {
221 let resolver = shared_resolver()?;
222 let resp = match timeout(
223 Duration::from_millis(timeout_ms),
224 resolver.reverse_lookup(ip),
225 )
226 .await
227 {
228 Ok(Ok(r)) => r,
229 Ok(Err(e)) => return Err(Error::dns(format!("reverse lookup failed: {e}"))),
230 Err(_) => return Err(Error::Timeout(timeout_ms)),
231 };
232 let name = resp.answers().iter().find_map(|r| match &r.data {
235 RData::PTR(ptr) => Some(ptr.0.to_utf8()),
236 _ => None,
237 });
238 Ok(name.filter(|s| !s.is_empty()))
239}
240
241pub async fn reverse_lookup_best_effort_timeout(ip: IpAddr, timeout_ms: u64) -> Option<String> {
250 #[cfg(windows)]
251 if let Some(name) = reverse_lookup_windows_ping(ip, timeout_ms).await {
252 if let Some(normalized) = normalize_hostname(name) {
253 return Some(normalized);
254 }
255 }
256
257 reverse_lookup_timeout(ip, timeout_ms)
258 .await
259 .ok()
260 .flatten()
261 .and_then(normalize_hostname)
262}
263
264fn normalize_hostname(name: String) -> Option<String> {
265 let name = name.trim().trim_end_matches('.').trim();
266 if name.is_empty() {
267 None
268 } else {
269 Some(name.to_string())
270 }
271}
272
273#[cfg(windows)]
274async fn reverse_lookup_windows_ping(ip: IpAddr, timeout_ms: u64) -> Option<String> {
275 if !matches!(ip, IpAddr::V4(_)) {
277 return None;
278 }
279
280 let ip_s = ip.to_string();
281 let wait_ms = timeout_ms.saturating_add(250);
282 let mut cmd = Command::new("ping");
283 cmd.args(["-a", "-n", "1", "-w", &timeout_ms.to_string(), &ip_s])
284 .stdout(Stdio::piped())
285 .stderr(Stdio::piped())
286 .kill_on_drop(true);
287
288 let output = match timeout(Duration::from_millis(wait_ms), cmd.output()).await {
289 Ok(Ok(out)) => out,
290 _ => return None,
291 };
292
293 let stdout = String::from_utf8_lossy(&output.stdout);
294
295 for line in stdout.lines().take(4) {
300 if !line.contains('[') || !line.contains(']') {
301 continue;
302 }
303 let Some(before) = line.split('[').next() else {
304 continue;
305 };
306 let before = before.trim();
307
308 let Some(candidate) = before.split_whitespace().last() else {
312 continue;
313 };
314 let candidate = candidate.trim().trim_end_matches('.');
315 if candidate.is_empty() {
316 continue;
317 }
318
319 if IpAddr::from_str(candidate).is_ok() {
322 return None;
323 }
324
325 return Some(candidate.to_string());
326 }
327
328 None
329}
330
331#[cfg(test)]
332mod tests {
333 use super::*;
334
335 #[test]
336 fn normalize_value_strips_trailing_dot() {
337 assert_eq!(normalize_value("example.com."), "example.com");
338 }
339
340 #[test]
341 fn normalize_value_preserves_root_dot() {
342 assert_eq!(normalize_value("."), ".");
343 }
344
345 #[test]
346 fn normalize_value_unwraps_txt_quotes() {
347 assert_eq!(normalize_value(r#""v=spf1 -all""#), "v=spf1 -all");
348 }
349
350 #[test]
351 fn normalize_value_joins_multichunk_txt() {
352 assert_eq!(normalize_value(r#""chunk1" "chunk2""#), "chunk1chunk2");
353 }
354
355 #[test]
356 fn normalize_value_untouched_for_plain_a() {
357 assert_eq!(normalize_value("192.0.2.1"), "192.0.2.1");
358 }
359
360 #[test]
361 fn normalize_hostname_drops_trailing_dot_and_ws() {
362 assert_eq!(
363 normalize_hostname("host.example.com.".to_string()),
364 Some("host.example.com".to_string())
365 );
366 assert_eq!(
367 normalize_hostname(" host.example.com ".to_string()),
368 Some("host.example.com".to_string())
369 );
370 }
371
372 #[test]
373 fn normalize_hostname_rejects_empty() {
374 assert_eq!(normalize_hostname("".to_string()), None);
375 assert_eq!(normalize_hostname(".".to_string()), None);
376 }
377
378 #[tokio::test]
379 async fn test_resolve_a_localhost() {
380 let result = resolve_a("localhost").await;
381 assert!(result.is_ok());
382 let ips = result.unwrap();
383 assert!(!ips.is_empty());
384 assert!(ips.contains(&"127.0.0.1".to_string()));
386 }
387
388 }