use std::net::SocketAddr;
use std::sync::Arc;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Reach {
Public,
PublicOrLoopbackName,
#[cfg(feature = "push")]
Any,
}
pub(crate) fn judge<I>(
reach: Reach,
host: &str,
resolved: I,
) -> Result<Vec<SocketAddr>, super::NetGuardError>
where
I: IntoIterator<Item = SocketAddr>,
{
let addrs: Vec<SocketAddr> = resolved.into_iter().collect();
let exempt = match reach {
#[cfg(feature = "push")]
Reach::Any => true,
Reach::PublicOrLoopbackName => super::is_loopback_name(host),
Reach::Public => false,
};
if !exempt {
return super::all_public(host, addrs);
}
if addrs.is_empty() {
return Err(super::NetGuardError::NoAddresses {
host: host.to_owned(),
});
}
Ok(addrs)
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct GuardedResolver {
reach: Reach,
}
impl GuardedResolver {
pub(crate) const fn new(reach: Reach) -> Self {
Self { reach }
}
pub(crate) fn shared(reach: Reach) -> Arc<Self> {
Arc::new(Self::new(reach))
}
}
pub(crate) fn guarded_client(reach: Reach) -> reqwest::ClientBuilder {
reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none())
.dns_resolver(GuardedResolver::shared(reach))
}
impl reqwest::dns::Resolve for GuardedResolver {
fn resolve(&self, name: reqwest::dns::Name) -> reqwest::dns::Resolving {
let reach = self.reach;
Box::pin(async move {
let host = name.as_str().to_owned();
let resolved = tokio::net::lookup_host((host.as_str(), 0)).await?;
let addrs = judge(reach, &host, resolved)
.map_err(|refusal| Box::new(refusal) as Box<dyn std::error::Error + Send + Sync>)?;
Ok(Box::new(addrs.into_iter()) as reqwest::dns::Addrs)
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use reqwest::dns::Resolve as _;
#[tokio::test]
async fn a_public_only_client_is_not_handed_a_loopback_address() {
let name: reqwest::dns::Name = "localhost".parse().expect("a resolvable name");
let refused = GuardedResolver::new(Reach::Public).resolve(name).await;
let error = refused.err().expect("localhost resolves inward");
assert!(
error.to_string().contains("forbidden address"),
"a public-only client was handed loopback: {error}"
);
}
#[tokio::test]
async fn a_guarded_client_does_not_reach_a_live_server_on_this_machine() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("a local port");
let port = listener.local_addr().expect("an address").port();
let served = tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
use tokio::io::AsyncWriteExt as _;
let _ = socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nhi")
.await;
return true;
}
false
});
let client = guarded_client(Reach::Public)
.timeout(std::time::Duration::from_secs(2))
.build()
.expect("a client");
let outcome = client.get(format!("http://localhost:{port}/")).send().await;
assert!(
outcome.is_err(),
"a public-only client reached a server on this machine, so the \
address rule is not attached to the client and every pre-flight \
that passes it is the only check there is"
);
served.abort();
}
#[tokio::test]
async fn an_ip_literal_never_reaches_this_resolver() {
use std::sync::atomic::{AtomicUsize, Ordering};
static CALLS: AtomicUsize = AtomicUsize::new(0);
#[derive(Debug)]
struct Counting;
impl reqwest::dns::Resolve for Counting {
fn resolve(&self, name: reqwest::dns::Name) -> reqwest::dns::Resolving {
CALLS.fetch_add(1, Ordering::SeqCst);
Box::pin(async move {
let addrs = tokio::net::lookup_host((name.as_str(), 0)).await?;
Ok(Box::new(addrs.collect::<Vec<_>>().into_iter()) as reqwest::dns::Addrs)
})
}
}
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("a local port");
let port = listener.local_addr().expect("an address").port();
let served = tokio::spawn(async move {
if let Ok((mut socket, _)) = listener.accept().await {
use tokio::io::AsyncWriteExt as _;
let _ = socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nhi")
.await;
}
});
let client = reqwest::Client::builder()
.dns_resolver(Arc::new(Counting))
.timeout(std::time::Duration::from_secs(2))
.build()
.expect("a client");
let _ = client.get(format!("http://127.0.0.1:{port}/")).send().await;
assert_eq!(
CALLS.load(Ordering::SeqCst),
0,
"a literal reached the resolver, so the split this module documents \
— pre-flight judges every destination, this judges every name — is \
not the split the transport implements"
);
served.abort();
}
#[tokio::test]
async fn the_two_exemptions_reach_what_they_are_for_and_nothing_else() {
let exempt = {
#[cfg(feature = "push")]
{
vec![Reach::PublicOrLoopbackName, Reach::Any]
}
#[cfg(not(feature = "push"))]
{
vec![Reach::PublicOrLoopbackName]
}
};
for reach in exempt {
let name: reqwest::dns::Name = "localhost".parse().expect("a resolvable name");
let addrs = GuardedResolver::new(reach)
.resolve(name)
.await
.unwrap_or_else(|error| panic!("{reach:?} refused localhost: {error}"));
assert!(
addrs.count() > 0,
"{reach:?} answered with nothing, which is an outage rather than \
the exemption it is supposed to be"
);
}
}
#[test]
fn a_name_that_is_not_loopback_gets_no_exemption_from_its_answers() {
let inward = "10.0.0.1:443".parse().expect("an address");
let refused = judge(Reach::PublicOrLoopbackName, "peer.example", vec![inward]);
assert!(
matches!(refused, Err(super::super::NetGuardError::Forbidden { .. })),
"a public name resolving to a private address was admitted under the \
loopback exception: {refused:?}"
);
}
}