Skip to main content

agentos_kernel/
dns.rs

1#[cfg(not(target_arch = "wasm32"))]
2use agentos_runtime::BlockingJobError;
3use hickory_proto::rr::domain::Name;
4use hickory_proto::rr::rdata::{A, AAAA};
5use hickory_proto::rr::{RData, Record, RecordType};
6#[cfg(not(target_arch = "wasm32"))]
7use hickory_resolver::config::{NameServerConfig, ResolverConfig};
8#[cfg(not(target_arch = "wasm32"))]
9use hickory_resolver::net::runtime::TokioRuntimeProvider;
10#[cfg(not(target_arch = "wasm32"))]
11use hickory_resolver::TokioResolver;
12use std::collections::{BTreeMap, BTreeSet};
13use std::error::Error;
14use std::fmt;
15use std::net::{IpAddr, SocketAddr};
16#[cfg(not(target_arch = "wasm32"))]
17use std::net::{Ipv4Addr, Ipv6Addr};
18use std::sync::Arc;
19
20#[derive(Debug, Clone, Default, PartialEq, Eq)]
21pub struct DnsConfig {
22    pub name_servers: Vec<SocketAddr>,
23    pub overrides: BTreeMap<String, Vec<IpAddr>>,
24}
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub enum DnsLookupPolicy {
28    CheckPermissions,
29    SkipPermissions,
30}
31
32#[derive(Debug, Clone, PartialEq, Eq)]
33pub struct DnsLookupRequest {
34    hostname: String,
35    name_servers: Vec<SocketAddr>,
36}
37
38impl DnsLookupRequest {
39    pub fn new(hostname: impl Into<String>, name_servers: Vec<SocketAddr>) -> Self {
40        Self {
41            hostname: hostname.into(),
42            name_servers,
43        }
44    }
45
46    pub fn hostname(&self) -> &str {
47        &self.hostname
48    }
49
50    pub fn name_servers(&self) -> &[SocketAddr] {
51        &self.name_servers
52    }
53}
54
55#[derive(Debug, Clone, PartialEq, Eq)]
56pub struct DnsRecordLookupRequest {
57    hostname: String,
58    name_servers: Vec<SocketAddr>,
59    record_type: RecordType,
60}
61
62impl DnsRecordLookupRequest {
63    pub fn new(
64        hostname: impl Into<String>,
65        name_servers: Vec<SocketAddr>,
66        record_type: RecordType,
67    ) -> Self {
68        Self {
69            hostname: hostname.into(),
70            name_servers,
71            record_type,
72        }
73    }
74
75    pub fn hostname(&self) -> &str {
76        &self.hostname
77    }
78
79    pub fn name_servers(&self) -> &[SocketAddr] {
80        &self.name_servers
81    }
82
83    pub const fn record_type(&self) -> RecordType {
84        self.record_type
85    }
86}
87
88#[derive(Debug, Clone, Copy, PartialEq, Eq)]
89pub enum DnsResolutionSource {
90    Literal,
91    Override,
92    Resolver,
93}
94
95impl DnsResolutionSource {
96    pub const fn as_str(self) -> &'static str {
97        match self {
98            Self::Literal => "literal",
99            Self::Override => "override",
100            Self::Resolver => "resolver",
101        }
102    }
103}
104
105#[derive(Debug, Clone, PartialEq, Eq)]
106pub struct DnsResolution {
107    hostname: String,
108    source: DnsResolutionSource,
109    addresses: Vec<IpAddr>,
110}
111
112impl DnsResolution {
113    pub fn new(
114        hostname: impl Into<String>,
115        source: DnsResolutionSource,
116        addresses: Vec<IpAddr>,
117    ) -> Self {
118        Self {
119            hostname: hostname.into(),
120            source,
121            addresses,
122        }
123    }
124
125    pub fn hostname(&self) -> &str {
126        &self.hostname
127    }
128
129    pub const fn source(&self) -> DnsResolutionSource {
130        self.source
131    }
132
133    pub fn addresses(&self) -> &[IpAddr] {
134        &self.addresses
135    }
136}
137
138#[derive(Debug, Clone, PartialEq, Eq)]
139pub struct DnsRecordResolution {
140    hostname: String,
141    source: DnsResolutionSource,
142    records: Vec<Record>,
143}
144
145impl DnsRecordResolution {
146    pub fn new(
147        hostname: impl Into<String>,
148        source: DnsResolutionSource,
149        records: Vec<Record>,
150    ) -> Self {
151        Self {
152            hostname: hostname.into(),
153            source,
154            records,
155        }
156    }
157
158    pub fn hostname(&self) -> &str {
159        &self.hostname
160    }
161
162    pub const fn source(&self) -> DnsResolutionSource {
163        self.source
164    }
165
166    pub fn records(&self) -> &[Record] {
167        &self.records
168    }
169}
170
171#[derive(Debug, Clone, Copy, PartialEq, Eq)]
172pub enum DnsResolverErrorKind {
173    InvalidInput,
174    NxDomain,
175    NoData,
176    LookupFailed,
177}
178
179#[derive(Debug, Clone, PartialEq, Eq)]
180pub struct DnsResolverError {
181    kind: DnsResolverErrorKind,
182    message: String,
183}
184
185impl DnsResolverError {
186    pub fn invalid_input(message: impl Into<String>) -> Self {
187        Self {
188            kind: DnsResolverErrorKind::InvalidInput,
189            message: message.into(),
190        }
191    }
192
193    pub fn lookup_failed(message: impl Into<String>) -> Self {
194        Self {
195            kind: DnsResolverErrorKind::LookupFailed,
196            message: message.into(),
197        }
198    }
199
200    pub fn nx_domain(message: impl Into<String>) -> Self {
201        Self {
202            kind: DnsResolverErrorKind::NxDomain,
203            message: message.into(),
204        }
205    }
206
207    pub fn no_data(message: impl Into<String>) -> Self {
208        Self {
209            kind: DnsResolverErrorKind::NoData,
210            message: message.into(),
211        }
212    }
213
214    pub const fn kind(&self) -> DnsResolverErrorKind {
215        self.kind
216    }
217}
218
219impl fmt::Display for DnsResolverError {
220    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
221        write!(f, "{}", self.message)
222    }
223}
224
225impl Error for DnsResolverError {}
226
227pub trait DnsResolver {
228    fn lookup_ip(&self, request: &DnsLookupRequest) -> Result<Vec<IpAddr>, DnsResolverError>;
229    fn lookup_records(
230        &self,
231        request: &DnsRecordLookupRequest,
232    ) -> Result<Vec<Record>, DnsResolverError>;
233}
234
235pub type SharedDnsResolver = Arc<dyn DnsResolver + Send + Sync>;
236
237#[cfg(not(target_arch = "wasm32"))]
238#[derive(Default)]
239pub struct HickoryDnsResolver {
240    runtime: Option<agentos_runtime::RuntimeContext>,
241}
242
243/// On wasm the kernel has no tokio runtime or host DNS stack, so the resolver is
244/// a unit type whose `DnsResolver` impl reports that name resolution is
245/// unavailable; guests must supply DNS overrides or literal addresses.
246#[cfg(target_arch = "wasm32")]
247pub struct HickoryDnsResolver;
248
249#[cfg(target_arch = "wasm32")]
250impl Default for HickoryDnsResolver {
251    fn default() -> Self {
252        Self
253    }
254}
255
256#[cfg(not(target_arch = "wasm32"))]
257impl HickoryDnsResolver {
258    pub fn with_runtime(runtime: agentos_runtime::RuntimeContext) -> Self {
259        Self {
260            runtime: Some(runtime),
261        }
262    }
263
264    fn send_lookup_ip(
265        &self,
266        hostname: String,
267        name_servers: Vec<SocketAddr>,
268    ) -> Result<Vec<IpAddr>, DnsResolverError> {
269        let runtime = self.runtime.as_ref().cloned().ok_or_else(|| {
270            DnsResolverError::lookup_failed(
271                "DNS resolver has no injected sidecar runtime; configure HickoryDnsResolver::with_runtime",
272            )
273        })?;
274        let resolver = {
275            let _entered = runtime.handle().enter();
276            resolver_for(&name_servers)?
277        };
278        let reserved_bytes = dns_lookup_input_bytes(&hostname, &name_servers);
279        let handle = runtime.handle().clone();
280        let timeout = runtime.blocking_job_timeout();
281        runtime
282            .blocking()
283            .run_sync(reserved_bytes, timeout, move || {
284                handle.block_on(async move {
285                    tokio::time::timeout(timeout, lookup_ip_with_resolver(resolver, hostname))
286                        .await
287                        .unwrap_or_else(|_| Err(dns_lookup_timeout_error(timeout)))
288                })
289            })
290            .map_err(map_blocking_lookup_error)?
291    }
292
293    fn send_lookup_records(
294        &self,
295        hostname: String,
296        name_servers: Vec<SocketAddr>,
297        record_type: RecordType,
298    ) -> Result<Vec<Record>, DnsResolverError> {
299        let runtime = self.runtime.as_ref().cloned().ok_or_else(|| {
300            DnsResolverError::lookup_failed(
301                "DNS resolver has no injected sidecar runtime; configure HickoryDnsResolver::with_runtime",
302            )
303        })?;
304        let resolver = {
305            let _entered = runtime.handle().enter();
306            resolver_for(&name_servers)?
307        };
308        let reserved_bytes = dns_lookup_input_bytes(&hostname, &name_servers);
309        let handle = runtime.handle().clone();
310        let timeout = runtime.blocking_job_timeout();
311        runtime
312            .blocking()
313            .run_sync(reserved_bytes, timeout, move || {
314                handle.block_on(async move {
315                    tokio::time::timeout(
316                        timeout,
317                        lookup_records_with_resolver(resolver, hostname, record_type),
318                    )
319                    .await
320                    .unwrap_or_else(|_| Err(dns_lookup_timeout_error(timeout)))
321                })
322            })
323            .map_err(map_blocking_lookup_error)?
324    }
325}
326
327#[cfg(not(target_arch = "wasm32"))]
328impl DnsResolver for HickoryDnsResolver {
329    fn lookup_ip(&self, request: &DnsLookupRequest) -> Result<Vec<IpAddr>, DnsResolverError> {
330        self.send_lookup_ip(
331            request.hostname().to_owned(),
332            request.name_servers().to_vec(),
333        )
334    }
335
336    fn lookup_records(
337        &self,
338        request: &DnsRecordLookupRequest,
339    ) -> Result<Vec<Record>, DnsResolverError> {
340        self.send_lookup_records(
341            request.hostname().to_owned(),
342            request.name_servers().to_vec(),
343            request.record_type(),
344        )
345    }
346}
347
348#[cfg(not(target_arch = "wasm32"))]
349fn resolver_for(name_servers: &[SocketAddr]) -> Result<TokioResolver, DnsResolverError> {
350    let resolver_config = resolver_config_from_name_servers(name_servers);
351    let builder = if let Some(config) = resolver_config {
352        TokioResolver::builder_with_config(config, TokioRuntimeProvider::default())
353    } else {
354        TokioResolver::builder_tokio().map_err(|error| {
355            DnsResolverError::lookup_failed(format!(
356                "failed to initialize DNS resolver from system configuration: {error}"
357            ))
358        })?
359    };
360    builder.build().map_err(|error| {
361        DnsResolverError::lookup_failed(format!("failed to build DNS resolver: {error}"))
362    })
363}
364
365#[cfg(not(target_arch = "wasm32"))]
366fn dns_lookup_input_bytes(hostname: &str, name_servers: &[SocketAddr]) -> usize {
367    hostname.len().saturating_add(
368        name_servers
369            .len()
370            .saturating_mul(std::mem::size_of::<SocketAddr>()),
371    )
372}
373
374#[cfg(not(target_arch = "wasm32"))]
375fn map_blocking_lookup_error(error: BlockingJobError) -> DnsResolverError {
376    DnsResolverError::lookup_failed(format!("ERR_AGENTOS_DNS_LOOKUP_EXECUTOR: {error}"))
377}
378
379#[cfg(not(target_arch = "wasm32"))]
380fn dns_lookup_timeout_error(timeout: std::time::Duration) -> DnsResolverError {
381    DnsResolverError::lookup_failed(format!(
382        "ERR_AGENTOS_DNS_LOOKUP_TIMEOUT: DNS lookup exceeded {}ms; raise runtime.blocking.jobTimeoutMs",
383        timeout.as_millis()
384    ))
385}
386
387#[cfg(not(target_arch = "wasm32"))]
388async fn lookup_ip_with_resolver(
389    resolver: TokioResolver,
390    hostname: String,
391) -> Result<Vec<IpAddr>, DnsResolverError> {
392    let lookup = resolver.lookup_ip(&hostname).await.map_err(|error| {
393        DnsResolverError::lookup_failed(format!(
394            "failed to resolve DNS address {hostname}: {error}"
395        ))
396    })?;
397
398    let mut addresses = Vec::new();
399    let mut seen = BTreeSet::new();
400    for ip in lookup.iter() {
401        if seen.insert(ip) {
402            addresses.push(ip);
403        }
404    }
405
406    if addresses.is_empty() {
407        return Err(DnsResolverError::lookup_failed(format!(
408            "failed to resolve DNS address {hostname}"
409        )));
410    }
411
412    Ok(addresses)
413}
414
415#[cfg(not(target_arch = "wasm32"))]
416async fn lookup_records_with_resolver(
417    resolver: TokioResolver,
418    hostname: String,
419    record_type: RecordType,
420) -> Result<Vec<Record>, DnsResolverError> {
421    let lookup = resolver
422        .lookup(&hostname, record_type)
423        .await
424        .map_err(|error| {
425            let message = format!("failed to resolve DNS {record_type} record {hostname}: {error}");
426            if error.is_nx_domain() {
427                DnsResolverError::nx_domain(message)
428            } else if error.is_no_records_found() {
429                DnsResolverError::no_data(message)
430            } else {
431                DnsResolverError::lookup_failed(message)
432            }
433        })?;
434    let records = lookup.answers().to_vec();
435    if records.is_empty() {
436        return Err(DnsResolverError::no_data(format!(
437            "failed to resolve DNS {record_type} record {hostname}"
438        )));
439    }
440    Ok(records)
441}
442
443#[cfg(target_arch = "wasm32")]
444impl DnsResolver for HickoryDnsResolver {
445    fn lookup_ip(&self, request: &DnsLookupRequest) -> Result<Vec<IpAddr>, DnsResolverError> {
446        Err(DnsResolverError::lookup_failed(format!(
447            "browser sidecar DNS resolver is unavailable for {}; configure DNS overrides or pass a literal address",
448            request.hostname()
449        )))
450    }
451
452    fn lookup_records(
453        &self,
454        request: &DnsRecordLookupRequest,
455    ) -> Result<Vec<Record>, DnsResolverError> {
456        Err(DnsResolverError::lookup_failed(format!(
457            "browser sidecar DNS record resolver is unavailable for {}; configure DNS overrides or pass a literal address",
458            request.hostname()
459        )))
460    }
461}
462
463pub fn normalize_dns_hostname(hostname: &str) -> Result<String, DnsResolverError> {
464    let normalized = hostname.trim().trim_end_matches('.').to_ascii_lowercase();
465    if normalized.is_empty() {
466        return Err(DnsResolverError::invalid_input(
467            "DNS hostname must not be empty",
468        ));
469    }
470    Ok(normalized)
471}
472
473pub fn format_dns_resource(hostname: &str) -> Result<String, DnsResolverError> {
474    Ok(format!("dns://{}", canonical_dns_subject(hostname)?))
475}
476
477pub fn resolve_dns(
478    config: &DnsConfig,
479    resolver: &dyn DnsResolver,
480    hostname: &str,
481) -> Result<DnsResolution, DnsResolverError> {
482    let trimmed = hostname.trim();
483    if let Ok(ip_addr) = trimmed.parse::<IpAddr>() {
484        return Ok(DnsResolution::new(
485            ip_addr.to_string(),
486            DnsResolutionSource::Literal,
487            vec![ip_addr],
488        ));
489    }
490
491    let normalized_hostname = normalize_dns_hostname(trimmed)?;
492    if let Some(addresses) = config.overrides.get(&normalized_hostname) {
493        return Ok(DnsResolution::new(
494            normalized_hostname,
495            DnsResolutionSource::Override,
496            addresses.clone(),
497        ));
498    }
499
500    let request = DnsLookupRequest::new(normalized_hostname.clone(), config.name_servers.clone());
501    let addresses = resolver.lookup_ip(&request)?;
502    if addresses.is_empty() {
503        return Err(DnsResolverError::lookup_failed(format!(
504            "failed to resolve DNS address {normalized_hostname}"
505        )));
506    }
507
508    Ok(DnsResolution::new(
509        normalized_hostname,
510        DnsResolutionSource::Resolver,
511        dedupe_addresses(addresses),
512    ))
513}
514
515pub fn resolve_dns_records(
516    config: &DnsConfig,
517    resolver: &dyn DnsResolver,
518    hostname: &str,
519    record_type: RecordType,
520) -> Result<DnsRecordResolution, DnsResolverError> {
521    let trimmed = hostname.trim();
522    let normalized_hostname = normalize_dns_hostname(trimmed)?;
523    let owner_name = normalized_hostname.parse::<Name>().map_err(|error| {
524        DnsResolverError::invalid_input(format!("invalid DNS hostname: {error}"))
525    })?;
526
527    if let Some(records) = records_from_literal(trimmed, owner_name.clone(), record_type) {
528        return Ok(DnsRecordResolution::new(
529            normalized_hostname,
530            DnsResolutionSource::Literal,
531            records,
532        ));
533    }
534
535    if let Some(addresses) = config.overrides.get(&normalized_hostname) {
536        let records = records_from_addresses(owner_name.clone(), addresses, record_type);
537        if !records.is_empty() {
538            return Ok(DnsRecordResolution::new(
539                normalized_hostname,
540                DnsResolutionSource::Override,
541                records,
542            ));
543        }
544    }
545
546    let request = DnsRecordLookupRequest::new(
547        normalized_hostname.clone(),
548        config.name_servers.clone(),
549        record_type,
550    );
551    let records = resolver.lookup_records(&request)?;
552    if records.is_empty() {
553        return Err(DnsResolverError::no_data(format!(
554            "failed to resolve DNS {record_type} record {normalized_hostname}"
555        )));
556    }
557
558    Ok(DnsRecordResolution::new(
559        normalized_hostname,
560        DnsResolutionSource::Resolver,
561        records,
562    ))
563}
564
565fn canonical_dns_subject(hostname: &str) -> Result<String, DnsResolverError> {
566    let trimmed = hostname.trim();
567    if let Ok(ip_addr) = trimmed.parse::<IpAddr>() {
568        return Ok(ip_addr.to_string());
569    }
570
571    normalize_dns_hostname(trimmed)
572}
573
574#[cfg(not(target_arch = "wasm32"))]
575fn resolver_config_from_name_servers(name_servers: &[SocketAddr]) -> Option<ResolverConfig> {
576    if name_servers.is_empty() {
577        return None;
578    }
579
580    let name_servers = name_servers
581        .iter()
582        .map(|server| {
583            let mut config = NameServerConfig::udp_and_tcp(server.ip());
584            for connection in &mut config.connections {
585                connection.port = server.port();
586                connection.bind_addr = Some(SocketAddr::new(
587                    if server.is_ipv6() {
588                        IpAddr::V6(Ipv6Addr::UNSPECIFIED)
589                    } else {
590                        IpAddr::V4(Ipv4Addr::UNSPECIFIED)
591                    },
592                    0,
593                ));
594            }
595            config
596        })
597        .collect();
598
599    Some(ResolverConfig::from_parts(None, vec![], name_servers))
600}
601
602fn dedupe_addresses(addresses: Vec<IpAddr>) -> Vec<IpAddr> {
603    let mut deduped = Vec::with_capacity(addresses.len());
604    let mut seen = BTreeSet::new();
605    for address in addresses {
606        if seen.insert(address) {
607            deduped.push(address);
608        }
609    }
610    deduped
611}
612
613fn records_from_literal(
614    hostname: &str,
615    owner_name: Name,
616    record_type: RecordType,
617) -> Option<Vec<Record>> {
618    let ip_addr = hostname.parse::<IpAddr>().ok()?;
619    let records = records_from_addresses(owner_name, &[ip_addr], record_type);
620    if records.is_empty() {
621        return None;
622    }
623    Some(records)
624}
625
626fn records_from_addresses(
627    owner_name: Name,
628    addresses: &[IpAddr],
629    record_type: RecordType,
630) -> Vec<Record> {
631    addresses
632        .iter()
633        .filter_map(|ip| match (record_type, ip) {
634            (RecordType::A, IpAddr::V4(ipv4)) | (RecordType::ANY, IpAddr::V4(ipv4)) => Some(
635                Record::from_rdata(owner_name.clone(), 60, RData::A(A::from(*ipv4))),
636            ),
637            (RecordType::AAAA, IpAddr::V6(ipv6)) | (RecordType::ANY, IpAddr::V6(ipv6)) => Some(
638                Record::from_rdata(owner_name.clone(), 60, RData::AAAA(AAAA::from(*ipv6))),
639            ),
640            _ => None,
641        })
642        .collect()
643}