use std::convert::TryFrom;
use async_trait::async_trait;
use log::info;
use rustls::pki_types::ServerName;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::{TcpListener, TcpStream, UnixListener, UnixStream};
use tokio_rustls::TlsAcceptor;
use crate::error::Error;
use crate::network::tls::{get_tls_connector, get_tls_listener};
use crate::settings::Shared;
pub fn socket_cleanup(settings: &Shared) -> Result<(), std::io::Error> {
if settings.use_unix_socket && settings.unix_socket_path().exists() {
std::fs::remove_file(settings.unix_socket_path())?;
}
Ok(())
}
#[async_trait]
pub trait Listener: Sync + Send {
async fn accept<'a>(&'a self) -> Result<GenericStream, Error>;
}
pub(crate) struct TlsTcpListener {
tcp_listener: TcpListener,
tls_acceptor: TlsAcceptor,
}
#[async_trait]
impl Listener for TlsTcpListener {
async fn accept<'a>(&'a self) -> Result<GenericStream, Error> {
let (stream, _) = self
.tcp_listener
.accept()
.await
.map_err(|err| Error::IoError("accepting new tcp connection.".to_string(), err))?;
let tls_stream = self
.tls_acceptor
.accept(stream)
.await
.map_err(|err| Error::IoError("accepting new tls connection.".to_string(), err))?;
Ok(Box::new(tls_stream))
}
}
#[async_trait]
impl Listener for UnixListener {
async fn accept<'a>(&'a self) -> Result<GenericStream, Error> {
let (stream, _) = self
.accept()
.await
.map_err(|err| Error::IoError("accepting new unix connection.".to_string(), err))?;
Ok(Box::new(stream))
}
}
pub trait Stream: AsyncRead + AsyncWrite + Unpin + Send {}
impl Stream for UnixStream {}
impl Stream for tokio_rustls::server::TlsStream<TcpStream> {}
impl Stream for tokio_rustls::client::TlsStream<TcpStream> {}
pub type GenericListener = Box<dyn Listener>;
pub type GenericStream = Box<dyn Stream>;
pub async fn get_client_stream(settings: &Shared) -> Result<GenericStream, Error> {
if settings.use_unix_socket {
let unix_socket_path = settings.unix_socket_path();
let stream = UnixStream::connect(&unix_socket_path)
.await
.map_err(|err| {
Error::IoPathError(
unix_socket_path,
"connecting to daemon. Did you start it?",
err,
)
})?;
return Ok(Box::new(stream));
}
let address = format!("{}:{}", &settings.host, &settings.port);
let tcp_stream = TcpStream::connect(&address).await.map_err(|_| {
Error::Connection(format!(
"Failed to connect to the daemon on {address}. Did you start it?"
))
})?;
let tls_connector = get_tls_connector(settings)
.await
.map_err(|err| Error::Connection(format!("Failed to initialize tls connector:\n{err}.")))?;
let stream = tls_connector
.connect(ServerName::try_from("pueue.local").unwrap(), tcp_stream)
.await
.map_err(|err| Error::Connection(format!("Failed to initialize tls:\n{err}.")))?;
Ok(Box::new(stream))
}
pub async fn get_listener(settings: &Shared) -> Result<GenericListener, Error> {
if settings.use_unix_socket {
let socket_path = settings.unix_socket_path();
info!("Using unix socket at: {socket_path:?}");
if socket_path.exists() {
if get_client_stream(settings).await.is_ok() {
return Err(Error::UnixSocketExists);
}
std::fs::remove_file(&socket_path).map_err(|err| {
Error::IoPathError(socket_path.clone(), "removing old socket", err)
})?;
}
let unix_listener = UnixListener::bind(&socket_path)
.map_err(|err| Error::IoPathError(socket_path, "creating unix socket", err))?;
return Ok(Box::new(unix_listener));
}
let address = format!("{}:{}", &settings.host, &settings.port);
info!("Binding to address: {address}");
let tcp_listener = TcpListener::bind(&address)
.await
.map_err(|err| Error::IoError("binding tcp listener to address".to_string(), err))?;
let tls_acceptor = get_tls_listener(settings)?;
let tls_listener = TlsTcpListener {
tcp_listener,
tls_acceptor,
};
Ok(Box::new(tls_listener))
}