use std::fmt::{self, Debug};
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4, ToSocketAddrs};
use std::sync::mpsc::{self, RecvTimeoutError};
use std::thread::{self};
use std::vec::IntoIter;
use http::Uri;
use http::uri::{Authority, Scheme};
use crate::Error;
use crate::config::Config;
use crate::http;
use crate::transport::NextTimeout;
use crate::util::{SchemeExt, UriExt};
pub trait Resolver: Debug + Send + Sync + 'static {
fn resolve(
&self,
uri: &Uri,
config: &Config,
timeout: NextTimeout,
) -> Result<ResolvedSocketAddrs, Error>;
fn empty(&self) -> ResolvedSocketAddrs {
fn uninited_socketaddr() -> SocketAddr {
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)), 0)
}
ArrayVec::from_fn(|_| uninited_socketaddr())
}
}
const MAX_ADDRS: usize = 16;
pub use ureq_proto::ArrayVec;
pub type ResolvedSocketAddrs = ArrayVec<SocketAddr, MAX_ADDRS>;
#[derive(Default)]
pub struct DefaultResolver {
_private: (),
}
impl DefaultResolver {
pub fn host_and_port(scheme: &Scheme, authority: &Authority) -> Option<String> {
let port = authority.port_u16().or_else(|| scheme.default_port())?;
Some(format!("{}:{}", authority.host(), port))
}
}
impl Resolver for DefaultResolver {
fn resolve(
&self,
uri: &Uri,
config: &Config,
timeout: NextTimeout,
) -> Result<ResolvedSocketAddrs, Error> {
uri.ensure_valid_url()?;
let scheme = uri.scheme().unwrap();
let authority = uri.authority().unwrap();
if cfg!(feature = "_test") {
let mut v = ArrayVec::from_fn(|_| "0.0.0.0:1".parse().unwrap());
v.push(SocketAddr::V4(SocketAddrV4::new(
Ipv4Addr::new(10, 0, 0, 1),
authority
.port_u16()
.or_else(|| scheme.default_port())
.unwrap(),
)));
return Ok(v);
}
let addr = DefaultResolver::host_and_port(scheme, authority).unwrap();
let use_sync = timeout.after.is_not_happening();
let iter = if use_sync {
trace!("Resolve: {}", addr);
addr.to_socket_addrs()?
} else {
trace!("Resolve with timeout ({:?}): {} ", timeout, addr);
resolve_async(addr, timeout)?
};
let ip_family = config.ip_family();
let wanted = ip_family.keep_wanted(iter);
let mut result = self.empty();
for addr in wanted.take(MAX_ADDRS) {
result.push(addr);
}
debug!("Resolved: {:?}", result);
if result.is_empty() {
Err(Error::HostNotFound)
} else {
Ok(result)
}
}
}
fn resolve_async(addr: String, timeout: NextTimeout) -> Result<IntoIter<SocketAddr>, Error> {
let (tx, rx) = mpsc::sync_channel(1);
thread::spawn(move || tx.send(addr.to_socket_addrs()).ok());
match rx.recv_timeout(*timeout.after) {
Ok(v) => Ok(v?),
Err(c) => match c {
RecvTimeoutError::Timeout => Err(Error::Timeout(timeout.reason)),
RecvTimeoutError::Disconnected => unreachable!("mpsc sender gone"),
},
}
}
impl fmt::Debug for DefaultResolver {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DefaultResolver").finish()
}
}
#[cfg(test)]
mod test {
use crate::transport::time::Duration;
use super::*;
#[test]
fn unknown_scheme() {
let uri: Uri = "foo://some:42/123".parse().unwrap();
let config = Config::default();
let err = DefaultResolver::default()
.resolve(
&uri,
&config,
NextTimeout {
after: Duration::NotHappening,
reason: crate::Timeout::Global,
},
)
.unwrap_err();
assert!(matches!(err, Error::BadUri(_)));
assert_eq!(err.to_string(), "bad uri: unknown scheme: foo");
}
}