dns.rsannotateddns.rssource203 lines · 6.8 KB · raw
1//! Name resolution with hickory on the backend's own executor (contract
2//! §2): its UDP/TCP sockets, timers and background tasks are ours, so a
3//! lookup waits in the same `WaitEventSet` and a cancel reaches it.
4//! Plain libc `getaddrinfo` would block without seeing a cancel.
5
6use std::cell::RefCell;
7use std::future::Future;
8use std::io;
9use std::net::{IpAddr, SocketAddr};
10use std::pin::Pin;
11use std::task::{Context, Poll};
12use std::time::Duration;
13
14use hickory_resolver::Resolver;
15use hickory_resolver::config::{ConnectionConfig, NameServerConfig, ResolverConfig, ResolverOpts};
16use hickory_resolver::net::runtime::{DnsTcpStream, DnsUdpSocket, RuntimeProvider, Spawn, Time};
17
18use jev_client::{TransportError, TransportErrorKind};
19
20use crate::Error;
21use crate::executor;
22use crate::net::{TcpIo, UdpIo};
23
24/// Where to ask: the servers in `jev.dns_servers`, or `/etc/resolv.conf`.
25#[derive(Clone, Debug, PartialEq, Eq)]
26pub enum Servers {
27    System,
28    Explicit(Vec<SocketAddr>),
29}
30
31impl Servers {
32    /// `ip[:port], …`; port 53 when omitted; IPv6 as `[addr]:port`.
33    pub fn parse(setting: Option<&str>) -> Result<Servers, Error> {
34        let Some(setting) = setting.map(str::trim).filter(|s| !s.is_empty()) else {
35            return Ok(Servers::System);
36        };
37        setting
38            .split(',')
39            .map(str::trim)
40            .map(|s| {
41                s.parse::<SocketAddr>()
42                    .or_else(|_| s.parse::<IpAddr>().map(|ip| SocketAddr::new(ip, 53)))
43                    .map_err(|_| Error::Config(format!("jev.dns_servers: {s:?} is not an address")))
44            })
45            .collect::<Result<_, _>>()
46            .map(Servers::Explicit)
47    }
48}
49
50thread_local! {
51    /// Built on first use, not in `_PG_init`, and kept for its cache: an
52    /// address is reused for its TTL across reconnects.
53    static RESOLVER: RefCell<Option<(Servers, Resolver<PgRuntime>)>> = const { RefCell::new(None) };
54}
55
56/// The addresses `host` resolves to, IPv4 and IPv6. A failed lookup is
57/// a connect failure, retried like one (the vendor SDKs do the same).
58pub async fn resolve(host: &str, servers: &Servers) -> Result<Vec<IpAddr>, TransportError> {
59    let resolver = resolver(servers).map_err(|e| TransportError::new(TransportErrorKind::Config, e.to_string()))?;
60    let lookup = resolver
61        .lookup_ip(host)
62        .await
63        .map_err(|e| TransportError::new(TransportErrorKind::Connect, format!("resolving {host}: {e}")))?;
64    let addrs: Vec<IpAddr> = lookup.iter().collect();
65    if addrs.is_empty() {
66        return Err(TransportError::new(TransportErrorKind::Connect, format!("resolving {host}: no addresses")));
67    }
68    Ok(addrs)
69}
70
71fn resolver(servers: &Servers) -> Result<Resolver<PgRuntime>, Error> {
72    if let Some(r) = RESOLVER.with(|r| r.borrow().as_ref().filter(|(s, _)| s == servers).map(|(_, r)| r.clone())) {
73        return Ok(r);
74    }
75    let builder = match servers {
76        Servers::System => Resolver::builder(PgRuntime)
77            .map_err(|e| Error::Config(format!("reading /etc/resolv.conf: {e}")))?,
78        Servers::Explicit(addrs) => {
79            let name_servers = addrs
80                .iter()
81                .map(|addr| {
82                    let mut udp = ConnectionConfig::udp();
83                    udp.port = addr.port();
84                    let mut tcp = ConnectionConfig::tcp();
85                    tcp.port = addr.port();
86                    NameServerConfig::new(addr.ip(), true, vec![udp, tcp])
87                })
88                .collect();
89            Resolver::builder_with_config(ResolverConfig::from_name_servers(name_servers), PgRuntime)
90                .with_options(ResolverOpts::default())
91        }
92    };
93    let resolver = builder.build().map_err(|e| Error::Config(format!("building the resolver: {e}")))?;
94    RESOLVER.with(|r| *r.borrow_mut() = Some((servers.clone(), resolver.clone())));
95    Ok(resolver)
96}
97
98/// hickory's runtime, backed by the executor.
99#[derive(Clone, Copy)]
100pub struct PgRuntime;
101
102#[derive(Clone, Copy)]
103pub struct PgHandle;
104
105pub struct PgTime;
106
107impl RuntimeProvider for PgRuntime {
108    type Handle = PgHandle;
109    type Timer = PgTime;
110    type Udp = PgUdp;
111    type Tcp = PgTcp;
112
113    fn create_handle(&self) -> PgHandle {
114        PgHandle
115    }
116
117    fn connect_tcp(
118        &self,
119        server_addr: SocketAddr,
120        _bind_addr: Option<SocketAddr>,
121        timeout: Option<Duration>,
122    ) -> Pin<Box<dyn Send + Future<Output = io::Result<PgTcp>>>> {
123        Box::pin(async move {
124            let connect = TcpIo::connect(server_addr);
125            let io = match timeout {
126                Some(t) => executor::timeout(t, connect)
127                    .await
128                    .unwrap_or_else(|| Err(io::Error::new(io::ErrorKind::TimedOut, "DNS TCP connect timed out"))),
129                None => connect.await,
130            }?;
131            Ok(PgTcp(io))
132        })
133    }
134
135    fn bind_udp(
136        &self,
137        local_addr: SocketAddr,
138        _server_addr: SocketAddr,
139    ) -> Pin<Box<dyn Send + Future<Output = io::Result<PgUdp>>>> {
140        let socket = UdpIo::bind(local_addr).map(PgUdp);
141        Box::pin(async move { socket })
142    }
143}
144
145impl Spawn for PgHandle {
146    fn spawn_bg(&mut self, future: impl Future<Output = ()> + Send + 'static) {
147        executor::spawn_local(future);
148    }
149}
150
151// hickory declares `Time` with `#[async_trait]`, so the impl matches.
152#[async_trait::async_trait]
153impl Time for PgTime {
154    async fn delay_for(duration: Duration) {
155        executor::sleep(duration).await;
156    }
157
158    async fn timeout<F: 'static + Future + Send>(duration: Duration, future: F) -> io::Result<F::Output> {
159        executor::timeout(duration, future)
160            .await
161            .ok_or_else(|| io::Error::new(io::ErrorKind::TimedOut, "DNS timed out"))
162    }
163}
164
165pub struct PgUdp(UdpIo);
166
167impl DnsUdpSocket for PgUdp {
168    type Time = PgTime;
169
170    fn poll_recv_from(&self, cx: &mut Context<'_>, buf: &mut [u8]) -> Poll<io::Result<(usize, SocketAddr)>> {
171        self.0.poll_recv_from(cx, buf)
172    }
173
174    fn poll_send_to(&self, cx: &mut Context<'_>, buf: &[u8], target: SocketAddr) -> Poll<io::Result<usize>> {
175        self.0.poll_send_to(cx, buf, target)
176    }
177}
178
179pub struct PgTcp(TcpIo);
180
181impl DnsTcpStream for PgTcp {
182    type Time = PgTime;
183}
184
185impl futures_io::AsyncRead for PgTcp {
186    fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut [u8]) -> Poll<io::Result<usize>> {
187        Pin::new(&mut self.0).poll_read(cx, buf)
188    }
189}
190
191impl futures_io::AsyncWrite for PgTcp {
192    fn poll_write(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
193        Pin::new(&mut self.0).poll_write(cx, buf)
194    }
195
196    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
197        Pin::new(&mut self.0).poll_flush(cx)
198    }
199
200    fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
201        Pin::new(&mut self.0).poll_close(cx)
202    }
203}