use std::net::SocketAddr;
use std::sync::{Arc, Mutex, PoisonError};
use reqwest::dns::{Name, Resolve, Resolving};
use crate::resolver::{DnscryptClientState, HARDCODED_RESOLVERS, HardcodedResolver};
#[derive(Debug)]
pub struct DnscryptResolver {
session: Arc<Mutex<Option<Vec<DnscryptClientState>>>>,
resolvers: &'static [HardcodedResolver],
}
impl Clone for DnscryptResolver {
fn clone(&self) -> Self {
Self {
session: self.session.clone(),
resolvers: self.resolvers,
}
}
}
impl DnscryptResolver {
#[must_use]
pub fn new() -> Self {
Self {
session: Arc::new(Mutex::new(None)),
resolvers: HARDCODED_RESOLVERS,
}
}
#[must_use]
pub fn new_with_resolvers(resolvers: &'static [HardcodedResolver]) -> Self {
Self {
session: Arc::new(Mutex::new(None)),
resolvers,
}
}
fn still_valid(
sessions: Option<&Vec<DnscryptClientState>>,
) -> Option<Vec<DnscryptClientState>> {
let s_vec = sessions?;
let s = s_vec.first()?;
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
(now < u64::from(s.ts_end)).then(|| s_vec.clone())
}
async fn get_or_establish_session(&self) -> Result<Vec<DnscryptClientState>, String> {
{
let lock = self.session.lock().unwrap_or_else(PoisonError::into_inner);
if let Some(sessions) = Self::still_valid(lock.as_ref()) {
return Ok(sessions);
}
}
let new_sessions = crate::resolver::establish_dnscrypt_session(self.resolvers)
.await
.map_err(|e| e.to_string())?;
let mut lock = self.session.lock().unwrap_or_else(PoisonError::into_inner);
*lock = Some(new_sessions.clone());
drop(lock);
Ok(new_sessions)
}
}
impl Default for DnscryptResolver {
fn default() -> Self {
Self::new()
}
}
impl Resolve for DnscryptResolver {
fn resolve(&self, name: Name) -> Resolving {
let domain = name.as_str().to_string();
let session_arc = self.session.clone();
let resolvers = self.resolvers;
Box::pin(async move {
let proxy = Self {
session: session_arc,
resolvers,
};
let sessions = proxy.get_or_establish_session().await.map_err(
|e| -> Box<dyn std::error::Error + Send + Sync> {
Box::new(std::io::Error::other(e))
},
)?;
let ips = crate::resolver::resolve_with_cache(&sessions, &domain).await;
let ips = ips.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
Box::new(std::io::Error::new(
std::io::ErrorKind::NotFound,
e.to_string(),
))
})?;
let addrs: Box<dyn Iterator<Item = SocketAddr> + Send> =
Box::new(ips.into_iter().map(|ip| SocketAddr::new(ip, 0)));
Ok(addrs)
})
}
}