use std::{collections::HashMap, net::SocketAddr, sync::Arc};
use parking_lot::RwLock;
use super::{
gai::GaiResolver,
resolve::{Addrs, IntoResolve, Name, Resolve, Resolving},
};
#[derive(Clone)]
pub struct PinnedDns {
fallback: Arc<dyn Resolve>,
pins: Arc<RwLock<HashMap<String, Vec<SocketAddr>>>>,
}
impl std::fmt::Debug for PinnedDns {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PinnedDns")
.field("pins", &self.pins.read().len())
.finish_non_exhaustive()
}
}
impl PinnedDns {
#[must_use]
pub fn new(fallback: impl IntoResolve) -> Self {
Self {
fallback: fallback.into_resolve(),
pins: Arc::new(RwLock::new(HashMap::new())),
}
}
#[must_use]
pub fn with_system() -> Self {
Self::new(GaiResolver::new())
}
pub fn pin(&self, host: &str, addrs: Vec<SocketAddr>) {
self.pins.write().insert(host.to_string(), addrs);
}
pub fn unpin(&self, host: &str) {
self.pins.write().remove(host);
}
pub fn unpin_if(&self, host: &str, expected: &[SocketAddr]) -> bool {
let mut pins = self.pins.write();
let matches = pins.get(host).is_some_and(|current| current == expected);
if matches {
pins.remove(host);
}
matches
}
pub fn clear(&self) {
self.pins.write().clear();
}
#[must_use]
pub fn pinned(&self, host: &str) -> Option<Vec<SocketAddr>> {
self.pins.read().get(host).cloned()
}
}
impl Resolve for PinnedDns {
fn resolve(&self, name: Name) -> Resolving {
if let Some(addrs) = self.pins.read().get(name.as_str()).cloned() {
let addrs: Addrs = Box::new(addrs.into_iter());
return Box::pin(std::future::ready(Ok(addrs)));
}
self.fallback.resolve(name)
}
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr};
use super::*;
fn addr(ip: [u8; 4]) -> SocketAddr {
SocketAddr::new(IpAddr::V4(Ipv4Addr::from(ip)), 0)
}
#[derive(Debug)]
struct LoopbackFallback;
impl Resolve for LoopbackFallback {
fn resolve(&self, name: Name) -> Resolving {
let ip = if name.as_str() == "example.com" {
IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34))
} else {
IpAddr::V4(Ipv4Addr::LOCALHOST)
};
let addrs: Addrs = Box::new(std::iter::once(SocketAddr::new(ip, 0)));
Box::pin(std::future::ready(Ok(addrs)))
}
}
#[tokio::test]
async fn miss_delegates_to_fallback() {
let pinned = PinnedDns::new(LoopbackFallback);
let mut addrs = pinned
.resolve(Name::from("example.com"))
.await
.expect("fallback resolves");
assert_eq!(
addrs.next(),
Some(SocketAddr::new(
IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)),
0
))
);
}
#[tokio::test]
async fn pin_overrides_fallback() {
let pinned = PinnedDns::new(LoopbackFallback);
pinned.pin("example.com", vec![addr([203, 0, 113, 7])]);
let mut addrs = pinned
.resolve(Name::from("example.com"))
.await
.expect("pin resolves");
assert_eq!(addrs.next(), Some(addr([203, 0, 113, 7])));
assert!(addrs.next().is_none());
}
#[tokio::test]
async fn unpin_restores_fallback() {
let pinned = PinnedDns::new(LoopbackFallback);
pinned.pin("example.com", vec![addr([203, 0, 113, 7])]);
pinned.unpin("example.com");
assert!(pinned.pinned("example.com").is_none());
let mut addrs = pinned
.resolve(Name::from("example.com"))
.await
.expect("resolves");
assert_eq!(
addrs.next(),
Some(SocketAddr::new(
IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)),
0
))
);
}
#[tokio::test]
async fn unpin_if_only_removes_matching_pin() {
let pinned = PinnedDns::new(LoopbackFallback);
pinned.pin("example.com", vec![addr([203, 0, 113, 7])]);
assert!(!pinned.unpin_if("example.com", &[addr([198, 51, 100, 9])]));
assert!(pinned.pinned("example.com").is_some());
assert!(pinned.unpin_if("example.com", &[addr([203, 0, 113, 7])]));
assert!(pinned.pinned("example.com").is_none());
}
#[test]
fn clones_share_pins() {
let pinned = PinnedDns::new(LoopbackFallback);
let clone = pinned.clone();
pinned.pin("example.com", vec![addr([203, 0, 113, 7])]);
assert_eq!(
clone.pinned("example.com"),
Some(vec![addr([203, 0, 113, 7])])
);
}
}