use core::net::{IpAddr, SocketAddr};
use std::io;
use tokio::{
io::{AsyncRead, AsyncWrite},
net::TcpStream,
task::JoinSet,
};
use ts_tls_util::ServerName;
use crate::{IpUsage, ServerConnInfo, TlsValidationConfig};
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum Error {
#[error("io error occurred")]
Io,
#[error("invalid parameter")]
InvalidParam,
}
impl From<io::Error> for Error {
fn from(_: io::Error) -> Self {
Error::Io
}
}
pub async fn dial_region_tls<'c>(
servers: impl IntoIterator<Item = &'c ServerConnInfo>,
) -> Result<
Option<(
impl AsyncRead + AsyncWrite + Unpin + Send + 'static,
&'c ServerConnInfo,
SocketAddr,
)>,
Error,
> {
let Some((conn, server)) = dial_region_tcp(servers).await else {
return Ok(None);
};
let remote_addr = conn.peer_addr()?;
let tls_conn = match &server.tls_validation_config {
TlsValidationConfig::CommonName { common_name } => {
ts_tls_util::connect(
ServerName::try_from(common_name.clone()).map_err(|e| {
tracing::error!(error = %e, "derp common name");
Error::InvalidParam
})?,
conn,
)
.await?
}
#[cfg(feature = "insecure-for-tests")]
TlsValidationConfig::InsecureForTests => {
tracing::warn!(%server.hostname, "using insecure TLS for tests");
ts_tls_util::connect_insecure(
ServerName::try_from(server.hostname.clone()).map_err(|e| {
tracing::error!(error = %e, "derp hostname");
Error::InvalidParam
})?,
conn,
)
.await?
}
TlsValidationConfig::SelfSigned { .. } => {
unimplemented!("self-signed derp server certs are currently unsupported");
}
};
Ok(Some((tls_conn, server, remote_addr)))
}
pub async fn dial_region_tcp<'c>(
servers: impl IntoIterator<Item = &'c ServerConnInfo>,
) -> Option<(TcpStream, &'c ServerConnInfo)> {
for server in servers {
if server.stun_only {
tracing::trace!(%server.hostname, "server is stun only, skip");
continue;
}
if matches!(
server.tls_validation_config,
TlsValidationConfig::SelfSigned { .. }
) {
tracing::warn!(
%server.hostname,
"self-signed derp server certs are currently unsupported, skipping server",
);
continue;
}
match dial_server(server).await {
Ok(Some(conn)) => {
tracing::trace!(
remote_addr = %conn.peer_addr().unwrap_or((core::net::Ipv4Addr::UNSPECIFIED, 0).into()),
%server.hostname,
"derp tcp dial ok",
);
return Some((conn, server));
}
Ok(None) => {
continue;
}
Err(e) => {
tracing::error!(error = %e, %server.hostname, "failed tcp dialing server");
continue;
}
}
}
None
}
pub async fn dial_server(server: &ServerConnInfo) -> Result<Option<TcpStream>, Error> {
let mut js = JoinSet::new();
js.spawn(dial_by_ipusage(
server.ipv4,
server.hostname.clone(),
server.https_port,
));
js.spawn(dial_by_ipusage(
server.ipv6,
server.hostname.clone(),
server.https_port,
));
let mut last_error = None;
while let Some(task) = js.join_next().await {
match task.unwrap() {
Ok(Some(stream)) => return Ok(Some(stream)),
Ok(None) => {
continue;
}
Err(e) => {
last_error = Some(e);
continue;
}
}
}
if let Some(e) = last_error {
Err(e.into())
} else {
Ok(None)
}
}
#[tracing::instrument(skip_all, level = "trace")]
async fn dial_by_ipusage(
ip: IpUsage<impl Into<IpAddr>>,
hostname: String,
port: u16,
) -> io::Result<Option<TcpStream>> {
match ip {
IpUsage::Disable => Ok(None),
IpUsage::FixedAddr(ip) => {
let ip = ip.into();
TcpStream::connect((ip, port)).await.map(Some)
}
IpUsage::UseDns => TcpStream::connect((hostname, port)).await.map(Some),
}
}