1use std::fmt;
20use std::future::Future;
21use std::net::{IpAddr, SocketAddr};
22use std::pin::Pin;
23use std::sync::Arc;
24
25use weida_core::{DEFAULT_PORT, Error};
26
27use crate::exec::Exec;
28
29pub type Resolved<'a> = Pin<Box<dyn Future<Output = Result<Vec<SocketAddr>, Error>> + Send + 'a>>;
35
36pub trait Resolver: fmt::Debug + Send + Sync + 'static {
68 fn resolve<'a>(
70 &'a self,
71 exec: &'a Exec,
72 name: &'a str,
73 port: Option<u16>,
74 max_addresses: usize,
75 ) -> Resolved<'a>;
76}
77
78#[derive(Clone, Copy, Debug, Default)]
87pub struct SystemResolver;
88
89impl Resolver for SystemResolver {
90 fn resolve<'a>(
91 &'a self,
92 exec: &'a Exec,
93 name: &'a str,
94 port: Option<u16>,
95 max_addresses: usize,
96 ) -> Resolved<'a> {
97 Box::pin(async move {
98 let port = port.unwrap_or(DEFAULT_PORT);
99 if let Ok(ip) = name.parse::<IpAddr>() {
100 return Ok(vec![SocketAddr::new(ip, port)]);
101 }
102 let query = (name.to_owned(), port);
103 let looked_up = exec
104 .spawn(async move {
105 tokio::net::lookup_host(query)
106 .await
107 .map(|addrs| addrs.collect::<Vec<SocketAddr>>())
108 })
109 .await
110 .map_err(|e| Error::Runtime(format!("name resolution task failed: {e}")))?;
111 let addrs: Vec<SocketAddr> = looked_up
112 .map_err(|e| Error::InvalidAddress(format!("cannot resolve {name}:{port}: {e}")))?
113 .into_iter()
114 .take(max_addresses)
115 .collect();
116 if addrs.is_empty() {
117 return Err(Error::InvalidAddress(format!(
118 "{name}:{port} resolved to no addresses"
119 )));
120 }
121 Ok(addrs)
122 })
123 }
124}
125
126#[derive(Clone, Debug)]
131pub struct SharedResolver(Arc<dyn Resolver>);
132
133impl SharedResolver {
134 pub fn new(resolver: impl Resolver) -> SharedResolver {
136 SharedResolver(Arc::new(resolver))
137 }
138
139 pub fn resolve<'a>(
141 &'a self,
142 exec: &'a Exec,
143 name: &'a str,
144 port: Option<u16>,
145 max_addresses: usize,
146 ) -> Resolved<'a> {
147 self.0.resolve(exec, name, port, max_addresses)
148 }
149}
150
151impl Default for SharedResolver {
152 fn default() -> SharedResolver {
153 SharedResolver::new(SystemResolver)
154 }
155}
156
157impl PartialEq for SharedResolver {
158 fn eq(&self, other: &SharedResolver) -> bool {
165 Arc::ptr_eq(&self.0, &other.0)
166 }
167}
168
169impl Eq for SharedResolver {}
170
171#[cfg(test)]
172mod tests {
173 use super::*;
174
175 #[derive(Debug)]
178 struct Table(Vec<SocketAddr>);
179
180 impl Resolver for Table {
181 fn resolve<'a>(
182 &'a self,
183 _exec: &'a Exec,
184 _name: &'a str,
185 _port: Option<u16>,
186 max_addresses: usize,
187 ) -> Resolved<'a> {
188 Box::pin(async move { Ok(self.0.iter().copied().take(max_addresses).collect()) })
189 }
190 }
191
192 #[tokio::test]
193 async fn a_literal_needs_no_resolver_and_takes_the_written_port() {
194 let exec = Exec::current().expect("ambient runtime");
195 let addrs = SystemResolver
196 .resolve(&exec, "127.0.0.1", Some(9000), 8)
197 .await
198 .expect("literal");
199 assert_eq!(addrs, vec!["127.0.0.1:9000".parse().expect("addr")]);
200 }
201
202 #[tokio::test]
203 async fn a_literal_without_a_written_port_takes_the_default() {
204 let exec = Exec::current().expect("ambient runtime");
205 let addrs = SystemResolver
206 .resolve(&exec, "127.0.0.1", None, 8)
207 .await
208 .expect("literal");
209 assert_eq!(addrs[0].port(), DEFAULT_PORT);
210 }
211
212 #[tokio::test]
213 async fn an_unresolvable_name_is_an_error_rather_than_an_empty_set() {
214 let exec = Exec::current().expect("ambient runtime");
215 let err = SystemResolver
216 .resolve(&exec, "no-such-host.invalid", Some(1), 8)
217 .await
218 .expect_err("`.invalid` never resolves");
219 assert!(
220 err.to_string().contains("no-such-host.invalid"),
221 "the error names what could not be resolved: {err}"
222 );
223 }
224
225 #[tokio::test]
228 async fn a_replaced_resolver_may_answer_several_ports_on_one_address() {
229 let exec = Exec::current().expect("ambient runtime");
230 let table = SharedResolver::new(Table(vec![
231 "203.0.113.7:7443".parse().expect("addr"),
232 "203.0.113.7:7444".parse().expect("addr"),
233 "203.0.113.7:7445".parse().expect("addr"),
234 ]));
235 let addrs = table
236 .resolve(&exec, "lb.example", None, 8)
237 .await
238 .expect("table");
239 assert_eq!(addrs.len(), 3);
240 assert!(addrs.iter().all(|a| a.ip().to_string() == "203.0.113.7"));
241 assert_eq!(
242 addrs.iter().map(|a| a.port()).collect::<Vec<_>>(),
243 vec![7443, 7444, 7445],
244 "the order is the resolver's and is preserved"
245 );
246 }
247
248 #[tokio::test]
249 async fn the_cap_is_the_callers_and_the_resolver_honours_it() {
250 let exec = Exec::current().expect("ambient runtime");
251 let table = SharedResolver::new(Table(vec![
252 "203.0.113.7:7443".parse().expect("addr"),
253 "203.0.113.7:7444".parse().expect("addr"),
254 "203.0.113.7:7445".parse().expect("addr"),
255 ]));
256 let addrs = table
257 .resolve(&exec, "lb.example", None, 2)
258 .await
259 .expect("table");
260 assert_eq!(addrs.len(), 2, "a resolver answer is remote input");
261 }
262
263 #[test]
264 fn two_resolvers_are_equal_only_when_they_are_the_same_one() {
265 let one = SharedResolver::default();
266 let same = one.clone();
267 let other = SharedResolver::default();
268 assert_eq!(one, same);
269 assert_ne!(
270 one, other,
271 "identical behaviour is not identity: the pool keys on this"
272 );
273 }
274
275 #[tokio::test]
276 async fn resolves_ip_literals_without_dns() {
277 let exec = Exec::current().expect("ambient runtime");
278 assert_eq!(
279 SystemResolver
280 .resolve(&exec, "127.0.0.1", Some(7443), 8)
281 .await
282 .expect("v4"),
283 vec![SocketAddr::from(([127, 0, 0, 1], 7443))]
284 );
285 let v6 = SystemResolver
286 .resolve(&exec, "::1", Some(7443), 8)
287 .await
288 .expect("v6");
289 assert_eq!(v6.len(), 1);
290 assert_eq!(v6[0].port(), 7443);
291 assert!(v6[0].is_ipv6());
292 }
293
294 #[tokio::test]
299 async fn a_hostname_resolves_to_every_address_up_to_the_cap() {
300 let exec = Exec::current().expect("ambient runtime");
301 let all = SystemResolver
302 .resolve(&exec, "localhost", Some(7443), 8)
303 .await
304 .expect("localhost");
305 assert!(!all.is_empty());
306 assert!(all.iter().all(|a| a.port() == 7443));
307
308 let capped = SystemResolver
309 .resolve(&exec, "localhost", Some(7443), 1)
310 .await
311 .expect("localhost");
312 assert_eq!(capped.len(), 1, "the cap must bound the answer");
313 assert_eq!(capped[0], all[0], "and it must keep the resolver's order");
314 }
315}