use std::net::SocketAddr;
use tokio::net::TcpListener;
#[cfg(unix)]
use tokio::net::UnixListener;
use tokio_stream::wrappers::TcpListenerStream;
#[cfg(unix)]
use tokio_stream::wrappers::UnixListenerStream;
use crate::{BindAddress, Error, Result};
pub(crate) enum BoundListener {
Tcp(TcpListener),
#[cfg(unix)]
Unix {
listener: UnixListener,
path: std::path::PathBuf,
},
}
impl BoundListener {
pub(crate) async fn bind(addr: BindAddress) -> Result<Self> {
match addr {
BindAddress::Tcp { host, port } => {
let listener = TcpListener::bind((host.as_str(), port))
.await
.map_err(|err| Error::Server(err.to_string()))?;
Ok(Self::Tcp(listener))
}
BindAddress::Unix { path } => {
#[cfg(not(unix))]
{
let _ = path;
Err(Error::Address(
"unix sockets are not supported on Windows; use tcp://host:port instead"
.to_string(),
))
}
#[cfg(unix)]
{
crate::address::remove_stale_socket(&path)?;
let listener =
UnixListener::bind(&path).map_err(|err| Error::Server(err.to_string()))?;
Ok(Self::Unix { listener, path })
}
}
}
}
pub(crate) fn local_addr(&self) -> Result<BindAddress> {
match self {
Self::Tcp(listener) => {
let addr: SocketAddr = listener
.local_addr()
.map_err(|err| Error::Server(err.to_string()))?;
Ok(BindAddress::Tcp {
host: addr.ip().to_string(),
port: addr.port(),
})
}
#[cfg(unix)]
Self::Unix { path, .. } => Ok(BindAddress::Unix { path: path.clone() }),
}
}
pub(crate) async fn serve(
self,
router: tonic::transport::server::Router,
shutdown: rlmesh_grpc::lifecycle::ShutdownTrigger,
drain_timeout: Option<std::time::Duration>,
) -> Result<()> {
match self {
Self::Tcp(listener) => rlmesh_grpc::lifecycle::await_server_shutdown(
router.serve_with_incoming_shutdown(
TcpListenerStream::new(listener),
shutdown.cancelled_owned(),
),
shutdown.clone(),
drain_timeout,
)
.await
.map_err(|err| Error::Server(err.to_string())),
#[cfg(unix)]
Self::Unix { listener, path } => {
let result = rlmesh_grpc::lifecycle::await_server_shutdown(
router.serve_with_incoming_shutdown(
UnixListenerStream::new(listener),
shutdown.cancelled_owned(),
),
shutdown.clone(),
drain_timeout,
)
.await
.map_err(|err| Error::Server(err.to_string()));
let _ = std::fs::remove_file(&path);
result
}
}
}
}