use std::collections::HashMap;
use std::io;
use std::net::{IpAddr, Ipv6Addr, SocketAddr, ToSocketAddrs};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{watch, Mutex, OwnedSemaphorePermit, Semaphore};
use crate::config::{DnsDenyCategory, DnsPolicy, DnsPreference};
use crate::socks5::TargetAddr;
const MAX_DNS_CACHE_ENTRIES: usize = 4096;
const MAX_CONCURRENT_SYSTEM_LOOKUPS: usize = 128;
const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(2);
const MAX_NEGATIVE_CACHE_ENTRIES: usize = 1024;
type SharedLookup = Option<std::result::Result<Vec<SocketAddr>, (io::ErrorKind, String)>>;
#[cfg(test)]
async fn resolve_all(dest: &TargetAddr, policy: &DnsPolicy) -> io::Result<Vec<SocketAddr>> {
let mut addrs = match dest {
TargetAddr::Ip(sa) => vec![*sa],
TargetAddr::Domain(host, port) => lookup_domain(host, *port).await?,
};
canonicalize_addrs(&mut addrs);
order_addresses(&mut addrs, policy.preference);
addrs.retain(|addr| address_allowed(addr.ip(), policy));
Ok(addrs)
}
#[derive(Debug)]
pub struct DnsResolver {
cache: Mutex<HashMap<DnsCacheKey, DnsCacheEntry>>,
inflight: std::sync::Mutex<HashMap<DnsCacheKey, watch::Receiver<SharedLookup>>>,
backend: LookupBackend,
lookup_slots: Arc<Semaphore>,
negative: std::sync::Mutex<HashMap<DnsCacheKey, Instant>>,
}
impl Default for DnsResolver {
fn default() -> Self {
DnsResolver {
cache: Mutex::new(HashMap::new()),
inflight: std::sync::Mutex::new(HashMap::new()),
backend: LookupBackend::default(),
lookup_slots: Arc::new(Semaphore::new(MAX_CONCURRENT_SYSTEM_LOOKUPS)),
negative: std::sync::Mutex::new(HashMap::new()),
}
}
}
#[derive(Debug, Default)]
enum LookupBackend {
#[default]
System,
#[cfg(test)]
Custom(TestLookup),
}
#[cfg(test)]
type TestLookupFn = dyn Fn(
&str,
u16,
)
-> std::pin::Pin<Box<dyn std::future::Future<Output = io::Result<Vec<SocketAddr>>> + Send>>
+ Send
+ Sync;
#[cfg(test)]
#[derive(Clone)]
struct TestLookup(std::sync::Arc<TestLookupFn>);
#[cfg(test)]
impl std::fmt::Debug for TestLookup {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TestLookup").finish_non_exhaustive()
}
}
impl DnsResolver {
pub fn new() -> Self {
Self::default()
}
#[cfg(test)]
fn with_lookup(lookup: std::sync::Arc<TestLookupFn>) -> Self {
Self::with_lookup_and_slots(lookup, MAX_CONCURRENT_SYSTEM_LOOKUPS)
}
#[cfg(test)]
fn with_lookup_and_slots(lookup: std::sync::Arc<TestLookupFn>, slots: usize) -> Self {
DnsResolver {
backend: LookupBackend::Custom(TestLookup(lookup)),
lookup_slots: Arc::new(Semaphore::new(slots)),
..DnsResolver::default()
}
}
async fn backend_lookup(&self, host: &str, port: u16) -> io::Result<Vec<SocketAddr>> {
let Ok(permit) = self.lookup_slots.clone().acquire_owned().await else {
return Err(io::Error::other("DNS lookup semaphore closed"));
};
match &self.backend {
LookupBackend::System => system_lookup(host.to_owned(), port, permit).await,
#[cfg(test)]
LookupBackend::Custom(lookup) => {
let _permit = permit;
(lookup.0)(host, port).await
}
}
}
pub async fn resolve_all(
&self,
dest: &TargetAddr,
policy: &DnsPolicy,
) -> io::Result<Vec<SocketAddr>> {
let mut addrs = match dest {
TargetAddr::Ip(sa) => vec![*sa],
TargetAddr::Domain(host, port) => {
self.resolve_domain(host, *port, policy.cache_ttl, policy.timeout)
.await?
}
};
canonicalize_addrs(&mut addrs);
order_addresses(&mut addrs, policy.preference);
addrs.retain(|addr| address_allowed(addr.ip(), policy));
Ok(addrs)
}
pub async fn resolve_one(
&self,
dest: &TargetAddr,
policy: &DnsPolicy,
) -> io::Result<Option<SocketAddr>> {
Ok(self.resolve_all(dest, policy).await?.into_iter().next())
}
pub async fn resolve_host(
&self,
host: &str,
ttl: Option<Duration>,
timeout: Duration,
) -> io::Result<Vec<IpAddr>> {
let mut addrs = self.resolve_domain(host, 0, ttl, timeout).await?;
canonicalize_addrs(&mut addrs);
Ok(addrs.into_iter().map(|addr| addr.ip()).collect())
}
async fn resolve_domain(
&self,
host: &str,
port: u16,
ttl: Option<Duration>,
timeout: Duration,
) -> io::Result<Vec<SocketAddr>> {
let key = DnsCacheKey::new(host, port);
loop {
if let Some(ttl) = ttl {
if let Some(addrs) = self.cached(&key, Instant::now(), ttl).await {
return Ok(addrs);
}
}
if self.negatively_cached(&key, Instant::now()) {
return Err(io::Error::new(
io::ErrorKind::TimedOut,
"DNS resolution timed out",
));
}
match self.join_or_lead(&key) {
Flight::Lead(mut lead) => {
let result = match tokio::time::timeout(
timeout,
self.backend_lookup(host, port),
)
.await
{
Ok(result) => result,
Err(_) => Err(io::Error::new(
io::ErrorKind::TimedOut,
"DNS resolution timed out",
)),
};
if let (Some(ttl), Ok(addrs)) = (ttl, &result) {
self.store(key.clone(), addrs.clone(), ttl).await;
}
if matches!(&result, Err(e) if e.kind() == io::ErrorKind::TimedOut) {
self.store_negative(key.clone(), Instant::now());
}
lead.publish(&result);
return result;
}
Flight::Follow(mut rx) => {
loop {
let outcome = rx.borrow_and_update().clone();
match outcome {
Some(Ok(addrs)) => return Ok(addrs),
Some(Err((kind, message))) => {
return Err(io::Error::new(kind, message));
}
None => {
if rx.changed().await.is_err() {
break;
}
}
}
}
}
}
}
}
fn join_or_lead(&self, key: &DnsCacheKey) -> Flight<'_> {
let mut inflight = self.inflight.lock().unwrap_or_else(|e| e.into_inner());
if let Some(rx) = inflight.get(key) {
return Flight::Follow(rx.clone());
}
let (tx, rx) = watch::channel(None);
inflight.insert(key.clone(), rx);
Flight::Lead(InflightLead {
resolver: self,
key: key.clone(),
tx: Some(tx),
})
}
async fn cached(
&self,
key: &DnsCacheKey,
now: Instant,
ttl: Duration,
) -> Option<Vec<SocketAddr>> {
let mut cache = self.cache.lock().await;
let entry = cache.get(key)?;
if cache_entry_live(entry, now, ttl) {
return Some(entry.addrs.clone());
}
cache.remove(key);
None
}
async fn store(&self, key: DnsCacheKey, addrs: Vec<SocketAddr>, ttl: Duration) {
let mut cache = self.cache.lock().await;
let now = Instant::now();
if cache.len() >= MAX_DNS_CACHE_ENTRIES && !cache.contains_key(&key) {
cache.retain(|_, entry| cache_entry_live(entry, now, ttl));
if cache.len() >= MAX_DNS_CACHE_ENTRIES {
if let Some(oldest_key) = oldest_cache_key(&cache) {
cache.remove(&oldest_key);
}
}
}
cache.insert(
key,
DnsCacheEntry {
addrs,
inserted_at: now,
},
);
}
fn negatively_cached(&self, key: &DnsCacheKey, now: Instant) -> bool {
let mut negative = self.negative.lock().unwrap_or_else(|e| e.into_inner());
match negative.get(key) {
Some(&failed_at) if now.saturating_duration_since(failed_at) < NEGATIVE_CACHE_TTL => {
true
}
Some(_) => {
negative.remove(key);
false
}
None => false,
}
}
fn store_negative(&self, key: DnsCacheKey, now: Instant) {
let mut negative = self.negative.lock().unwrap_or_else(|e| e.into_inner());
if negative.len() >= MAX_NEGATIVE_CACHE_ENTRIES && !negative.contains_key(&key) {
negative.retain(|_, failed_at| {
now.saturating_duration_since(*failed_at) < NEGATIVE_CACHE_TTL
});
if negative.len() >= MAX_NEGATIVE_CACHE_ENTRIES {
if let Some(oldest) = negative
.iter()
.min_by_key(|(_, &failed_at)| failed_at)
.map(|(k, _)| k.clone())
{
negative.remove(&oldest);
}
}
}
negative.insert(key, now);
}
}
enum Flight<'a> {
Lead(InflightLead<'a>),
Follow(watch::Receiver<SharedLookup>),
}
struct InflightLead<'a> {
resolver: &'a DnsResolver,
key: DnsCacheKey,
tx: Option<watch::Sender<SharedLookup>>,
}
impl InflightLead<'_> {
fn publish(&mut self, result: &io::Result<Vec<SocketAddr>>) {
if let Some(tx) = self.tx.take() {
self.remove_inflight();
let shared = match result {
Ok(addrs) => Ok(addrs.clone()),
Err(e) => Err((e.kind(), e.to_string())),
};
let _ = tx.send(Some(shared));
}
}
fn remove_inflight(&self) {
let mut inflight = self
.resolver
.inflight
.lock()
.unwrap_or_else(|e| e.into_inner());
inflight.remove(&self.key);
}
}
impl Drop for InflightLead<'_> {
fn drop(&mut self) {
if self.tx.is_some() {
self.remove_inflight();
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct DnsCacheKey {
host: String,
port: u16,
}
impl DnsCacheKey {
fn new(host: &str, port: u16) -> Self {
Self {
host: host.to_ascii_lowercase(),
port,
}
}
}
#[derive(Debug, Clone)]
struct DnsCacheEntry {
addrs: Vec<SocketAddr>,
inserted_at: Instant,
}
fn cache_entry_live(entry: &DnsCacheEntry, now: Instant, ttl: Duration) -> bool {
now.saturating_duration_since(entry.inserted_at) < ttl
}
fn oldest_cache_key(cache: &HashMap<DnsCacheKey, DnsCacheEntry>) -> Option<DnsCacheKey> {
cache
.iter()
.min_by_key(|(key, entry)| (entry.inserted_at, key.port, key.host.as_str()))
.map(|(key, _)| key.clone())
}
async fn system_lookup(
host: String,
port: u16,
permit: OwnedSemaphorePermit,
) -> io::Result<Vec<SocketAddr>> {
tokio::task::spawn_blocking(move || -> io::Result<Vec<SocketAddr>> {
let _permit = permit;
Ok((host.as_str(), port).to_socket_addrs()?.collect())
})
.await
.map_err(io::Error::other)?
}
#[cfg(test)]
async fn lookup_domain(host: &str, port: u16) -> io::Result<Vec<SocketAddr>> {
tokio::net::lookup_host((host, port))
.await
.map(|addrs| addrs.collect())
}
fn order_addresses(addrs: &mut [SocketAddr], preference: DnsPreference) {
match preference {
DnsPreference::System => {}
DnsPreference::Ipv4 => addrs.sort_by_key(|addr| if addr.is_ipv4() { 0 } else { 1 }),
DnsPreference::Ipv6 => addrs.sort_by_key(|addr| if addr.is_ipv6() { 0 } else { 1 }),
}
}
pub fn address_allowed(ip: IpAddr, policy: &DnsPolicy) -> bool {
let ip = canonical_ip(ip);
!policy
.deny
.iter()
.any(|category| ip_matches_category(ip, *category))
}
fn canonical_ip(ip: IpAddr) -> IpAddr {
let ip = ip.to_canonical();
let IpAddr::V6(v6) = ip else { return ip };
if let Some(v4) = v6.to_ipv4() {
if v6 != Ipv6Addr::UNSPECIFIED && v6 != Ipv6Addr::LOCALHOST {
return IpAddr::V4(v4);
}
}
ip
}
fn canonicalize_addrs(addrs: &mut [SocketAddr]) {
for addr in addrs.iter_mut() {
addr.set_ip(canonical_ip(addr.ip()));
}
}
fn ip_matches_category(ip: IpAddr, category: DnsDenyCategory) -> bool {
match category {
DnsDenyCategory::Private => is_private(ip),
DnsDenyCategory::LinkLocal => is_link_local(ip),
DnsDenyCategory::Loopback => ip.is_loopback(),
DnsDenyCategory::Multicast => is_multicast(ip),
DnsDenyCategory::Unspecified => ip.is_unspecified(),
DnsDenyCategory::Documentation => is_documentation(ip),
DnsDenyCategory::Reserved => is_reserved(ip),
}
}
fn is_private(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => ip.is_private(),
IpAddr::V6(ip) => (ip.segments()[0] & 0xfe00) == 0xfc00,
}
}
fn is_link_local(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => {
let [a, b, _, _] = ip.octets();
a == 169 && b == 254
}
IpAddr::V6(ip) => (ip.segments()[0] & 0xffc0) == 0xfe80,
}
}
fn is_multicast(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => ip.is_multicast(),
IpAddr::V6(ip) => ip.is_multicast(),
}
}
fn is_documentation(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => {
let [a, b, c, _] = ip.octets();
(a == 192 && b == 0 && c == 2)
|| (a == 198 && b == 51 && c == 100)
|| (a == 203 && b == 0 && c == 113)
}
IpAddr::V6(ip) => ip.segments()[0] == 0x2001 && ip.segments()[1] == 0x0db8,
}
}
fn is_reserved(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => {
let [a, b, c, _] = ip.octets();
a == 0
|| a >= 240
|| (a == 100 && (64..=127).contains(&b))
|| (a == 192 && b == 0 && c == 0)
|| (a == 192 && b == 88 && c == 99)
|| (a == 198 && (b == 18 || b == 19))
}
IpAddr::V6(ip) => {
let s = ip.segments();
ip == Ipv6Addr::UNSPECIFIED
|| ip == Ipv6Addr::LOCALHOST
|| is_documentation(IpAddr::V6(ip))
|| s[0] == 0x2002
|| (s[0] == 0x0064
&& s[1] == 0xff9b
&& s[2] == 0
&& s[3] == 0
&& s[4] == 0
&& s[5] == 0)
|| (s[0] == 0x2001 && (s[1] & 0xfe00) == 0)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::Ipv4Addr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
fn policy(preference: DnsPreference, deny: Vec<DnsDenyCategory>) -> DnsPolicy {
DnsPolicy {
preference,
try_all: false,
deny,
cache_ttl: None,
timeout: Duration::from_secs(5),
}
}
fn answer(port: u16) -> SocketAddr {
SocketAddr::from(([203, 0, 113, 7], port))
}
fn parked_backend(
calls: Arc<AtomicUsize>,
release: Arc<tokio::sync::Semaphore>,
fail: bool,
) -> Arc<TestLookupFn> {
Arc::new(move |_host, port| {
let calls = calls.clone();
let release = release.clone();
Box::pin(async move {
calls.fetch_add(1, Ordering::SeqCst);
let _permit = release.acquire().await.expect("semaphore closed");
if fail {
Err(io::Error::new(
io::ErrorKind::ConnectionRefused,
"backend unavailable",
))
} else {
Ok(vec![answer(port)])
}
})
})
}
fn counting_backend(
in_flight: Arc<AtomicUsize>,
peak: Arc<AtomicUsize>,
release: Arc<tokio::sync::Semaphore>,
) -> Arc<TestLookupFn> {
Arc::new(move |_host, port| {
let in_flight = in_flight.clone();
let peak = peak.clone();
let release = release.clone();
Box::pin(async move {
let now = in_flight.fetch_add(1, Ordering::SeqCst) + 1;
peak.fetch_max(now, Ordering::SeqCst);
let _permit = release.acquire().await.expect("semaphore closed");
in_flight.fetch_sub(1, Ordering::SeqCst);
Ok(vec![answer(port)])
})
})
}
#[tokio::test(flavor = "current_thread")]
async fn caps_concurrent_system_lookups() {
let in_flight = Arc::new(AtomicUsize::new(0));
let peak = Arc::new(AtomicUsize::new(0));
let release = Arc::new(tokio::sync::Semaphore::new(0));
let resolver = Arc::new(DnsResolver::with_lookup_and_slots(
counting_backend(in_flight.clone(), peak.clone(), release.clone()),
2,
));
let mut handles = Vec::new();
for i in 0..5 {
let resolver = resolver.clone();
handles.push(tokio::spawn(async move {
let dns_policy = policy(DnsPreference::System, Vec::new());
resolver
.resolve_all(
&TargetAddr::Domain(format!("name{i}.example"), 443),
&dns_policy,
)
.await
}));
}
for _ in 0..64 {
tokio::task::yield_now().await;
}
assert_eq!(
in_flight.load(Ordering::SeqCst),
2,
"exactly the slot count should be resolving at once"
);
assert_eq!(
peak.load(Ordering::SeqCst),
2,
"the cap must never be exceeded"
);
release.add_permits(5);
for handle in handles {
assert_eq!(handle.await.unwrap().unwrap(), vec![answer(443)]);
}
assert_eq!(
peak.load(Ordering::SeqCst),
2,
"the cap held for the whole run"
);
}
#[tokio::test(start_paused = true)]
async fn resolution_times_out_on_a_wedged_resolver() {
let calls = Arc::new(AtomicUsize::new(0));
let never_released = Arc::new(tokio::sync::Semaphore::new(0));
let resolver = DnsResolver::with_lookup(parked_backend(calls, never_released, false));
let mut policy = policy(DnsPreference::System, Vec::new());
policy.timeout = Duration::from_secs(2);
let result = resolver
.resolve_all(&TargetAddr::Domain("wedged.example".into(), 443), &policy)
.await;
assert!(
matches!(&result, Err(e) if e.kind() == io::ErrorKind::TimedOut),
"expected a TimedOut error, got {result:?}"
);
}
#[tokio::test(start_paused = true)]
async fn timed_out_name_is_negatively_cached() {
let calls = Arc::new(AtomicUsize::new(0));
let never_released = Arc::new(tokio::sync::Semaphore::new(0));
let resolver =
DnsResolver::with_lookup(parked_backend(calls.clone(), never_released, false));
let mut policy = policy(DnsPreference::System, Vec::new());
policy.timeout = Duration::from_secs(1);
let first = resolver
.resolve_all(&TargetAddr::Domain("wedged.example".into(), 443), &policy)
.await;
assert!(matches!(&first, Err(e) if e.kind() == io::ErrorKind::TimedOut));
assert_eq!(calls.load(Ordering::SeqCst), 1);
let second = resolver
.resolve_all(&TargetAddr::Domain("wedged.example".into(), 443), &policy)
.await;
assert!(matches!(&second, Err(e) if e.kind() == io::ErrorKind::TimedOut));
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"the negative cache should suppress the second lookup"
);
}
#[test]
fn negative_cache_expires_after_ttl() {
let resolver = DnsResolver::new();
let key = DnsCacheKey::new("wedged.example", 443);
let t0 = Instant::now();
resolver.store_negative(key.clone(), t0);
assert!(resolver.negatively_cached(&key, t0));
assert!(
resolver.negatively_cached(&key, t0 + NEGATIVE_CACHE_TTL - Duration::from_millis(1))
);
assert!(!resolver.negatively_cached(&key, t0 + NEGATIVE_CACHE_TTL));
assert!(!resolver.negatively_cached(&DnsCacheKey::new("other.example", 443), t0));
}
#[test]
fn negative_cache_evicts_oldest_when_full() {
let resolver = DnsResolver::new();
let now = Instant::now();
{
let mut negative = resolver.negative.lock().unwrap();
for i in 0..MAX_NEGATIVE_CACHE_ENTRIES {
negative.insert(
DnsCacheKey::new(&format!("host{i}.example"), 80),
now + Duration::from_millis(i as u64),
);
}
}
let insert_time = now + Duration::from_millis(MAX_NEGATIVE_CACHE_ENTRIES as u64);
resolver.store_negative(DnsCacheKey::new("new.example", 80), insert_time);
let negative = resolver.negative.lock().unwrap();
assert_eq!(negative.len(), MAX_NEGATIVE_CACHE_ENTRIES, "the cap holds");
assert!(!negative.contains_key(&DnsCacheKey::new("host0.example", 80)));
assert!(negative.contains_key(&DnsCacheKey::new("new.example", 80)));
}
#[tokio::test(start_paused = true)]
async fn timeout_keeps_singleflight_coalesced() {
let calls = Arc::new(AtomicUsize::new(0));
let never_released = Arc::new(tokio::sync::Semaphore::new(0));
let resolver = Arc::new(DnsResolver::with_lookup(parked_backend(
calls.clone(),
never_released,
false,
)));
let tasks = spawn_resolvers(&resolver, 8).await;
assert_eq!(calls.load(Ordering::SeqCst), 1);
for task in tasks {
let res = task.await.unwrap();
assert!(
matches!(&res, Err(e) if e.kind() == io::ErrorKind::TimedOut),
"expected TimedOut, got {res:?}"
);
}
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
async fn spawn_resolvers(
resolver: &Arc<DnsResolver>,
count: usize,
) -> Vec<tokio::task::JoinHandle<io::Result<Vec<SocketAddr>>>> {
let mut tasks = Vec::with_capacity(count);
for _ in 0..count {
let resolver = resolver.clone();
let dns_policy = policy(DnsPreference::System, Vec::new());
tasks.push(tokio::spawn(async move {
resolver
.resolve_all(&TargetAddr::Domain("example.com".into(), 443), &dns_policy)
.await
}));
}
for _ in 0..64 {
tokio::task::yield_now().await;
}
tasks
}
#[tokio::test(flavor = "current_thread")]
async fn singleflight_coalesces_concurrent_lookups() {
let calls = Arc::new(AtomicUsize::new(0));
let release = Arc::new(tokio::sync::Semaphore::new(0));
let resolver = Arc::new(DnsResolver::with_lookup(parked_backend(
calls.clone(),
release.clone(),
false,
)));
let tasks = spawn_resolvers(&resolver, 8).await;
assert_eq!(calls.load(Ordering::SeqCst), 1);
release.add_permits(1);
for task in tasks {
assert_eq!(task.await.unwrap().unwrap(), vec![answer(443)]);
}
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[tokio::test(flavor = "current_thread")]
async fn singleflight_shares_leader_error_without_caching_it() {
let calls = Arc::new(AtomicUsize::new(0));
let release = Arc::new(tokio::sync::Semaphore::new(0));
let resolver = Arc::new(DnsResolver::with_lookup(parked_backend(
calls.clone(),
release.clone(),
true,
)));
let tasks = spawn_resolvers(&resolver, 4).await;
release.add_permits(1);
for task in tasks {
let err = task.await.unwrap().unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::ConnectionRefused);
assert!(err.to_string().contains("backend unavailable"));
}
assert_eq!(calls.load(Ordering::SeqCst), 1);
release.add_permits(1);
let dns_policy = policy(DnsPreference::System, Vec::new());
let _ = resolver
.resolve_all(&TargetAddr::Domain("example.com".into(), 443), &dns_policy)
.await;
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test(flavor = "current_thread")]
async fn singleflight_recovers_when_the_leader_is_cancelled() {
let calls = Arc::new(AtomicUsize::new(0));
let release = Arc::new(tokio::sync::Semaphore::new(0));
let resolver = Arc::new(DnsResolver::with_lookup(parked_backend(
calls.clone(),
release.clone(),
false,
)));
let mut tasks = spawn_resolvers(&resolver, 2).await;
let follower = tasks.pop().unwrap();
let leader = tasks.pop().unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
leader.abort();
release.add_permits(1);
assert_eq!(follower.await.unwrap().unwrap(), vec![answer(443)]);
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn store_purges_expired_entries_when_full() {
let resolver = DnsResolver::new();
let now = Instant::now();
{
let mut cache = resolver.cache.lock().await;
for i in 0..MAX_DNS_CACHE_ENTRIES {
cache.insert(
DnsCacheKey::new(&format!("expired{i}.example"), 80),
DnsCacheEntry {
addrs: vec![answer(80)],
inserted_at: now - Duration::from_secs(2),
},
);
}
}
resolver
.store(
DnsCacheKey::new("fresh.example", 80),
vec![answer(80)],
Duration::from_secs(1),
)
.await;
let cache = resolver.cache.lock().await;
assert_eq!(cache.len(), 1);
assert!(cache.contains_key(&DnsCacheKey::new("fresh.example", 80)));
}
#[tokio::test]
async fn resolves_ip_literal_without_dns() {
let sa: SocketAddr = "9.9.9.9:53".parse().unwrap();
let resolved = resolve_all(
&TargetAddr::Ip(sa),
&policy(DnsPreference::System, Vec::new()),
)
.await
.unwrap();
assert_eq!(resolved, vec![sa]);
}
#[tokio::test]
async fn denies_private_ip_literal() {
let sa: SocketAddr = "10.0.0.1:80".parse().unwrap();
let resolved = resolve_all(
&TargetAddr::Ip(sa),
&policy(DnsPreference::System, vec![DnsDenyCategory::Private]),
)
.await
.unwrap();
assert!(resolved.is_empty());
}
#[test]
fn address_allowed_canonicalizes_ipv4_mapped() {
let deny = policy(DnsPreference::System, vec![DnsDenyCategory::Loopback]);
assert!(!address_allowed("::ffff:127.0.0.1".parse().unwrap(), &deny));
assert!(!address_allowed("127.0.0.1".parse().unwrap(), &deny));
assert!(address_allowed("::ffff:8.8.8.8".parse().unwrap(), &deny));
}
#[test]
fn address_allowed_canonicalizes_ipv4_compatible() {
let deny = policy(
DnsPreference::System,
vec![
DnsDenyCategory::Loopback,
DnsDenyCategory::LinkLocal,
DnsDenyCategory::Private,
],
);
for ip in ["::127.0.0.1", "::169.254.169.254", "::10.0.0.1", "::1"] {
assert!(
!address_allowed(ip.parse().unwrap(), &deny),
"{ip} should be denied"
);
}
assert!(address_allowed("::8.8.8.8".parse().unwrap(), &deny));
}
#[test]
fn reserved_covers_iana_special_ranges() {
let deny = policy(DnsPreference::System, vec![DnsDenyCategory::Reserved]);
for ip in [
"0.0.0.0",
"0.255.255.255", "100.64.0.1",
"100.127.255.255", "192.0.0.7", "192.88.99.1", "198.18.0.1",
"198.19.255.255", "240.0.0.1",
"255.255.255.255", "::",
"::1",
"2001:db8::1", "2002::1", "2002:7f00:1::", "64:ff9b::7f00:1", "64:ff9b::1", "2001::1", "2001:20::1", "2001:1ff:ffff::", "::ffff:198.18.0.1", ] {
assert!(
!address_allowed(ip.parse().unwrap(), &deny),
"{ip} should be denied as reserved"
);
}
for ip in [
"8.8.8.8",
"100.63.255.255", "100.128.0.0", "198.17.255.255", "198.20.0.0", "192.88.98.255", "2001:4860:4860::8888", "2606:4700:4700::1111", "2003::1", "2400::1", ] {
assert!(
address_allowed(ip.parse().unwrap(), &deny),
"{ip} should be allowed"
);
}
}
#[tokio::test]
async fn denies_ipv4_mapped_loopback_and_private() {
for (literal, category) in [
("[::ffff:127.0.0.1]:80", DnsDenyCategory::Loopback),
("[::ffff:10.0.0.1]:80", DnsDenyCategory::Private),
] {
let sa: SocketAddr = literal.parse().unwrap();
let resolved = resolve_all(
&TargetAddr::Ip(sa),
&policy(DnsPreference::System, vec![category]),
)
.await
.unwrap();
assert!(
resolved.is_empty(),
"{literal} should be denied as {category:?}"
);
}
}
#[tokio::test]
async fn resolve_all_canonicalizes_mapped_addresses() {
let sa: SocketAddr = "[::ffff:8.8.8.8]:53".parse().unwrap();
let resolved = resolve_all(
&TargetAddr::Ip(sa),
&policy(DnsPreference::System, Vec::new()),
)
.await
.unwrap();
assert_eq!(resolved, vec!["8.8.8.8:53".parse().unwrap()]);
}
#[tokio::test]
async fn resolver_denies_hostname_resolving_to_denied_ip() {
let backend: Arc<TestLookupFn> = Arc::new(|_host, port| {
Box::pin(async move { Ok(vec![SocketAddr::from(([127, 0, 0, 1], port))]) })
});
let resolver = DnsResolver::with_lookup(backend);
let target = TargetAddr::Domain("intranet.evil.test".into(), 80);
let deny = policy(DnsPreference::System, vec![DnsDenyCategory::Loopback]);
assert!(
resolver
.resolve_all(&target, &deny)
.await
.unwrap()
.is_empty(),
"a hostname resolving to loopback must be denied post-resolution"
);
assert!(
resolver
.resolve_one(&target, &deny)
.await
.unwrap()
.is_none(),
"resolve_one must yield no allowed address for a denied resolution"
);
let open = policy(DnsPreference::System, Vec::new());
assert_eq!(
resolver.resolve_one(&target, &open).await.unwrap(),
Some(SocketAddr::from(([127, 0, 0, 1], 80))),
);
}
#[test]
fn orders_ipv4_first() {
let mut addrs = vec![
SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 80),
SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 80),
];
order_addresses(&mut addrs, DnsPreference::Ipv4);
assert!(addrs[0].is_ipv4());
}
#[test]
fn orders_ipv6_first() {
let mut addrs = vec![
SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 80),
SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 80),
];
order_addresses(&mut addrs, DnsPreference::Ipv6);
assert!(addrs[0].is_ipv6());
}
#[test]
fn matches_documentation_ranges() {
assert!(is_documentation("192.0.2.1".parse().unwrap()));
assert!(is_documentation("2001:db8::1".parse().unwrap()));
assert!(!is_documentation("8.8.8.8".parse().unwrap()));
}
#[test]
fn cache_keys_normalize_host_case() {
assert_eq!(
DnsCacheKey::new("Example.COM", 443),
DnsCacheKey::new("example.com", 443)
);
}
#[tokio::test]
async fn cached_entries_expire() {
let resolver = DnsResolver::new();
let key = DnsCacheKey::new("example.com", 80);
let addrs: Vec<SocketAddr> = vec!["203.0.113.10:80".parse().unwrap()];
let now = Instant::now();
let ttl = Duration::from_secs(10);
let seed = |inserted_at: Instant| DnsCacheEntry {
addrs: addrs.clone(),
inserted_at,
};
resolver.cache.lock().await.insert(key.clone(), seed(now));
assert_eq!(
resolver
.cached(&key, now + Duration::from_secs(5), ttl)
.await,
Some(addrs.clone())
);
assert_eq!(
resolver
.cached(&key, now + Duration::from_secs(11), ttl)
.await,
None
);
resolver.cache.lock().await.insert(key.clone(), seed(now));
assert_eq!(
resolver
.cached(&key, now + Duration::from_secs(5), Duration::from_secs(3))
.await,
None
);
}
#[test]
fn oversized_ttl_never_expires() {
let now = Instant::now();
let entry = DnsCacheEntry {
addrs: vec![answer(80)],
inserted_at: now,
};
assert!(cache_entry_live(
&entry,
now + Duration::from_secs(86_400),
Duration::MAX
));
}
#[test]
fn cache_liveness_uses_current_ttl() {
let now = Instant::now();
let entry = DnsCacheEntry {
addrs: vec![answer(80)],
inserted_at: now,
};
let later = now + Duration::from_secs(30);
assert!(cache_entry_live(&entry, later, Duration::from_secs(60)));
assert!(!cache_entry_live(
&entry,
now + Duration::from_secs(61),
Duration::from_secs(60)
));
assert!(!cache_entry_live(&entry, later, Duration::from_secs(10)));
}
#[tokio::test]
async fn cache_evicts_oldest_entry_when_full() {
let resolver = DnsResolver::new();
let addrs: Vec<SocketAddr> = vec!["203.0.113.10:80".parse().unwrap()];
let ttl = Duration::from_secs(3600);
let now = Instant::now();
{
let mut cache = resolver.cache.lock().await;
for i in 0..MAX_DNS_CACHE_ENTRIES {
cache.insert(
DnsCacheKey::new(&format!("host{i}.example"), 80),
DnsCacheEntry {
addrs: addrs.clone(),
inserted_at: now + Duration::from_millis(i as u64),
},
);
}
}
let oldest = DnsCacheKey::new("host0.example", 80);
resolver
.store(DnsCacheKey::new("new.example", 80), addrs, ttl)
.await;
assert_eq!(resolver.cached(&oldest, Instant::now(), ttl).await, None);
let cache = resolver.cache.lock().await;
assert!(cache.contains_key(&DnsCacheKey::new("new.example", 80)));
}
}