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 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 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
68fn normalize_value(raw: &str) -> String {
75 let s = raw.trim();
76 if s.starts_with('"') && s.ends_with('"') && s.len() >= 2 {
78 let inner = &s[1..s.len() - 1];
79 return inner.replace("\" \"", "");
81 }
82 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
120pub 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
209pub 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 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 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 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 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 assert!(ips.contains(&"127.0.0.1".to_string()));
354 }
355
356 }