use std::future::Future;
#[cfg(unix)]
use std::path::PathBuf;
#[cfg(unix)]
use std::sync::atomic::{AtomicBool, Ordering};
pub enum BoundListener {
Tcp(tokio::net::TcpListener),
#[cfg(unix)]
Uds(tokio::net::UnixListener, UdsCleanup),
}
impl BoundListener {
pub fn endpoint_label(&self) -> String {
match self {
Self::Tcp(l) => l
.local_addr()
.map(|a| {
if a.is_ipv6() {
format!("http://[{}]:{}", a.ip(), a.port())
} else {
format!("http://{}", a)
}
})
.unwrap_or_else(|_| "http://<unknown>".to_string()),
#[cfg(unix)]
Self::Uds(_, cleanup) => format!("unix://{}", cleanup.path.display()),
}
}
pub fn endpoint_info(&self) -> EndpointInfo {
match self {
Self::Tcp(l) => {
let url = self.endpoint_label();
match l.local_addr() {
Ok(a) => EndpointInfo::Tcp {
url,
host: a.ip().to_string(),
port: a.port(),
},
Err(_) => EndpointInfo::Tcp {
url,
host: String::new(),
port: 0,
},
}
}
#[cfg(unix)]
Self::Uds(_, cleanup) => EndpointInfo::Unix {
url: self.endpoint_label(),
path: cleanup.path.to_string_lossy().into_owned(),
},
}
}
pub fn is_loopback(&self) -> bool {
match self {
Self::Tcp(l) => l
.local_addr()
.map(|a| a.ip().is_loopback())
.unwrap_or(false),
#[cfg(unix)]
Self::Uds(..) => true,
}
}
}
#[derive(Debug, Clone, serde::Serialize)]
#[serde(tag = "transport", rename_all = "lowercase")]
pub enum EndpointInfo {
Tcp {
url: String,
host: String,
port: u16,
},
#[cfg(unix)]
Unix {
url: String,
path: String,
},
}
#[cfg(unix)]
pub struct UdsCleanup {
pub path: PathBuf,
active: AtomicBool,
}
#[cfg(unix)]
impl UdsCleanup {
pub fn new(path: PathBuf) -> Self {
Self {
path,
active: AtomicBool::new(true),
}
}
}
#[cfg(unix)]
impl Drop for UdsCleanup {
fn drop(&mut self) {
if self.active.swap(false, Ordering::SeqCst) {
if let Err(e) = std::fs::remove_file(&self.path) {
if e.kind() != std::io::ErrorKind::NotFound {
log::warn!("UDS cleanup failed for {:?}: {e}", self.path);
}
}
}
}
}
#[cfg(unix)]
#[derive(Clone, Debug)]
pub(crate) struct UdsConnectInfo {
pub peer_addr: std::sync::Arc<tokio::net::unix::SocketAddr>,
pub peer_cred: Option<tokio::net::unix::UCred>,
}
#[cfg(unix)]
impl axum::extract::connect_info::Connected<&tokio::net::UnixStream> for UdsConnectInfo {
fn connect_info(stream: &tokio::net::UnixStream) -> Self {
let peer_addr = stream
.peer_addr()
.expect("UnixStream::peer_addr on a just-accepted socket cannot fail");
let peer_cred = stream.peer_cred().ok();
Self {
peer_addr: std::sync::Arc::new(peer_addr),
peer_cred,
}
}
}
#[cfg(unix)]
pub(crate) async fn serve_uds(
listener: tokio::net::UnixListener,
router: axum::Router,
shutdown: impl Future<Output = ()> + Send + 'static,
) -> Result<(), std::io::Error> {
use hyper_util::rt::{TokioExecutor, TokioIo};
use hyper_util::server::conn::auto::Builder;
use tokio::task::JoinSet;
use tokio_util::sync::CancellationToken;
use tower::Service;
let token = CancellationToken::new();
let shutdown_token = token.clone();
tokio::spawn(async move {
shutdown.await;
shutdown_token.cancel();
});
let mut make_service = router.into_make_service_with_connect_info::<UdsConnectInfo>();
let mut conn_tasks: JoinSet<()> = JoinSet::new();
loop {
tokio::select! {
_ = token.cancelled() => break,
accept = listener.accept() => {
let (socket, _peer_addr) = accept?;
let tower_svc = unwrap_infallible(make_service.call(&socket).await);
let conn_token = token.clone();
conn_tasks.spawn(async move {
let io = TokioIo::new(socket);
let hyper_svc = hyper::service::service_fn(move |req: hyper::Request<hyper::body::Incoming>| {
tower_svc.clone().call(req)
});
let builder = Builder::new(TokioExecutor::new());
let conn = builder.serve_connection_with_upgrades(io, hyper_svc);
tokio::pin!(conn);
tokio::select! {
res = conn.as_mut() => {
if let Err(e) = res {
log::debug!("UDS connection error: {e:#}");
}
}
_ = conn_token.cancelled() => {
conn.as_mut().graceful_shutdown();
let _ = conn.await;
}
}
});
}
}
}
drop(listener);
while conn_tasks.join_next().await.is_some() {}
Ok(())
}
#[cfg(unix)]
fn unwrap_infallible<T>(r: Result<T, std::convert::Infallible>) -> T {
match r {
Ok(v) => v,
Err(i) => match i {},
}
}
pub async fn bind_listener(
spec: &crate::ListenAddr,
#[cfg_attr(not(unix), allow(unused_variables))] mode: u32,
) -> Result<BoundListener, vl_convert_rs::anyhow::Error> {
use vl_convert_rs::anyhow::anyhow;
match spec {
crate::ListenAddr::Tcp { host, port } => {
let addr = if host.contains(':') {
format!("[{host}]:{port}")
} else {
format!("{host}:{port}")
};
let l = tokio::net::TcpListener::bind(&addr)
.await
.map_err(|e| anyhow!("Failed to bind TCP {addr}: {e}"))?;
Ok(BoundListener::Tcp(l))
}
#[cfg(unix)]
crate::ListenAddr::Uds { path } => {
probe_then_unlink(path).await?;
let l = tokio::net::UnixListener::bind(path)
.map_err(|e| anyhow!("Failed to bind UDS {}: {e}", path.display()))?;
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(mode))
.map_err(|e| anyhow!("Failed to set mode on UDS {}: {e}", path.display()))?;
let cleanup = UdsCleanup::new(path.clone());
Ok(BoundListener::Uds(l, cleanup))
}
}
}
#[cfg(unix)]
async fn probe_then_unlink(path: &std::path::Path) -> Result<(), vl_convert_rs::anyhow::Error> {
use vl_convert_rs::anyhow::{anyhow, bail};
let probe = tokio::time::timeout(
std::time::Duration::from_millis(100),
tokio::net::UnixStream::connect(path),
)
.await;
match probe {
Ok(Ok(_)) => bail!(
"UDS socket {} is in use by another process; refusing to replace it",
path.display()
),
Ok(Err(e)) if e.kind() == std::io::ErrorKind::ConnectionRefused => {
std::fs::remove_file(path)
.map_err(|e| anyhow!("Failed to remove stale UDS {}: {e}", path.display()))?;
Ok(())
}
Ok(Err(e)) if e.kind() == std::io::ErrorKind::NotFound => {
Ok(())
}
Ok(Err(e)) => bail!("Unexpected error probing UDS path {}: {e}", path.display()),
Err(_elapsed) => bail!(
"Timed out probing UDS path {} for an existing listener",
path.display()
),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(unix)]
#[tokio::test]
async fn uds_cleanup_unlinks_on_drop() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("t.sock");
std::fs::write(&path, b"stand-in").unwrap();
assert!(path.exists());
{
let _guard = UdsCleanup::new(path.clone());
}
assert!(!path.exists(), "UdsCleanup::Drop should have unlinked");
}
#[cfg(unix)]
#[tokio::test]
async fn serve_uds_stops_accepting_immediately_on_shutdown() {
use axum::routing::get;
use std::io::ErrorKind;
use std::sync::Arc;
use tokio::sync::Notify;
let dir = tempfile::tempdir().unwrap();
let sock = dir.path().join("t.sock");
let listener = tokio::net::UnixListener::bind(&sock).unwrap();
let release = Arc::new(Notify::new());
let release_h = release.clone();
let router = axum::Router::new().route(
"/slow",
get(move || {
let release = release_h.clone();
async move {
release.notified().await;
"ok"
}
}),
);
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>();
let shutdown = async move {
let _ = shutdown_rx.await;
};
let serve_handle = tokio::spawn(serve_uds(listener, router, shutdown));
use tokio::io::AsyncWriteExt;
let mut busy = tokio::net::UnixStream::connect(&sock).await.unwrap();
busy.write_all(b"GET /slow HTTP/1.1\r\nHost: x\r\n\r\n")
.await
.unwrap();
busy.flush().await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
shutdown_tx.send(()).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let connect = tokio::time::timeout(
std::time::Duration::from_millis(500),
tokio::net::UnixStream::connect(&sock),
)
.await
.expect("connect should not hang once listener is dropped");
let err = connect.expect_err(
"connect must fail once shutdown fires, even while a handler \
is still in flight (regression: pathname socket stayed bound \
through the drain window)",
);
assert!(
matches!(
err.kind(),
ErrorKind::ConnectionRefused | ErrorKind::NotFound
),
"expected ECONNREFUSED or ENOENT, got {err:?}"
);
release.notify_one();
drop(busy);
let _ = tokio::time::timeout(std::time::Duration::from_secs(5), serve_handle).await;
}
}