use std::future::Future;
use std::net::{IpAddr, SocketAddr};
use std::pin::Pin;
use std::sync::{Arc, RwLock};
use std::task::{Context, Poll};
use http::Uri;
use hyper_util::rt::TokioIo;
use nym_smol_core::{Stack, TcpStream};
use tower::Service;
use crate::error::DvpnError;
#[derive(Clone)]
pub struct TunnelConnector {
stack: Arc<RwLock<Arc<Stack>>>,
}
impl TunnelConnector {
pub(crate) fn new(stack: Arc<RwLock<Arc<Stack>>>) -> Self {
Self { stack }
}
}
impl Service<Uri> for TunnelConnector {
type Response = TokioIo<TcpStream>;
type Error = DvpnError;
type Future = Pin<Box<dyn Future<Output = Result<TokioIo<TcpStream>, DvpnError>> + Send>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, uri: Uri) -> Self::Future {
let stack = self.stack.read().expect("stack lock poisoned").clone();
Box::pin(async move {
let host = uri
.host()
.ok_or_else(|| DvpnError::Config("URI missing host".into()))?
.to_string();
let port = uri.port_u16().unwrap_or(match uri.scheme_str() {
Some("http") => 80,
_ => 443,
});
let unbracketed = host
.strip_prefix('[')
.and_then(|h| h.strip_suffix(']'))
.unwrap_or(&host);
let stream = if let Ok(ip) = unbracketed.parse::<IpAddr>() {
stack.tcp_connect(SocketAddr::new(ip, port)).await?
} else {
stack.tcp_connect_host(&host, port).await?
};
Ok(TokioIo::new(stream))
})
}
}