use std::os::unix::fs::PermissionsExt;
use std::path::{Path, PathBuf};
use tokio::net::{UnixListener, UnixStream};
use weida_core::{Error, LocalPrincipal, MAX_SOCKET_PATH_BYTES};
const SOCKET_MODE: u32 = 0o600;
#[derive(Debug)]
pub struct BoundUnixSocket {
path: PathBuf,
}
impl BoundUnixSocket {
pub fn bind(path: &Path) -> Result<(BoundUnixSocket, UnixListener), Error> {
if path.as_os_str().len() > MAX_SOCKET_PATH_BYTES {
return Err(Error::InvalidAddress(format!(
"socket path exceeds this platform's {}-byte sun_path budget: {}",
MAX_SOCKET_PATH_BYTES,
path.display()
)));
}
match std::fs::metadata(path) {
Ok(meta) if is_socket(&meta) => std::fs::remove_file(path).map_err(Error::Io)?,
Ok(_) => {
return Err(Error::InvalidAddress(format!(
"{} exists and is not a socket",
path.display()
)));
}
Err(_) => {}
}
let listener = UnixListener::bind(path).map_err(Error::Io)?;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(SOCKET_MODE))
.map_err(Error::Io)?;
Ok((
BoundUnixSocket {
path: path.to_path_buf(),
},
listener,
))
}
pub fn path(&self) -> &Path {
&self.path
}
}
impl Drop for BoundUnixSocket {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.path);
}
}
fn is_socket(meta: &std::fs::Metadata) -> bool {
use std::os::unix::fs::FileTypeExt;
meta.file_type().is_socket()
}
pub fn peer_credentials(stream: &UnixStream) -> Result<LocalPrincipal, Error> {
let cred = stream.peer_cred().map_err(Error::Io)?;
Ok(LocalPrincipal {
uid: cred.uid(),
gid: cred.gid(),
pid: cred.pid().map(|pid| pid as u32),
})
}
#[cfg(test)]
mod tests {
use super::*;
fn dir(tag: &str) -> PathBuf {
let dir = std::env::temp_dir().join(format!("weida-runtime-{}-{tag}", std::process::id()));
std::fs::create_dir_all(&dir).expect("test directory");
dir
}
#[tokio::test]
async fn a_bound_socket_is_private_to_its_owner() {
let path = dir("mode").join("s");
let (bound, _listener) = BoundUnixSocket::bind(&path).expect("bind");
let mode = std::fs::metadata(&path)
.expect("metadata")
.permissions()
.mode();
assert_eq!(mode & 0o777, SOCKET_MODE, "mode {mode:o}");
assert_eq!(bound.path(), path.as_path());
}
#[tokio::test]
async fn the_node_is_removed_on_drop() {
let path = dir("drop").join("s");
let (bound, listener) = BoundUnixSocket::bind(&path).expect("bind");
assert!(path.exists());
drop(listener);
drop(bound);
assert!(!path.exists(), "the socket node outlived its binding");
}
#[tokio::test]
async fn a_stale_socket_is_replaced() {
let path = dir("stale").join("s");
{
let (bound, listener) = BoundUnixSocket::bind(&path).expect("first bind");
std::mem::forget(bound);
drop(listener);
}
assert!(path.exists(), "the stale node must still be there");
let (_bound, _listener) = BoundUnixSocket::bind(&path).expect("bind over the stale node");
}
#[tokio::test]
async fn a_regular_file_is_not_unlinked() {
let path = dir("file").join("s");
std::fs::write(&path, b"not a socket").expect("write");
let err = BoundUnixSocket::bind(&path).unwrap_err();
assert!(matches!(err, Error::InvalidAddress(_)), "{err:?}");
assert_eq!(std::fs::read(&path).expect("read"), b"not a socket");
}
#[tokio::test]
async fn an_over_long_path_is_refused() {
let path = dir("long").join("x".repeat(MAX_SOCKET_PATH_BYTES + 1));
let err = BoundUnixSocket::bind(&path).unwrap_err();
assert!(matches!(err, Error::InvalidAddress(_)), "{err:?}");
}
#[tokio::test]
async fn peer_credentials_are_the_kernels_answer() {
let (a, _b) = UnixStream::pair().expect("socket pair");
let principal = peer_credentials(&a).expect("credentials");
assert_eq!(principal.uid, unsafe_free_uid());
if let Some(pid) = principal.pid {
assert_eq!(pid, std::process::id());
}
}
fn unsafe_free_uid() -> u32 {
use std::os::unix::fs::MetadataExt;
let path = dir("uid").join("owned");
std::fs::write(&path, b"").expect("write");
let uid = std::fs::metadata(&path).expect("metadata").uid();
let _ = std::fs::remove_file(&path);
uid
}
}