use std::fs;
use std::io::{self, Read, Write};
use std::os::unix::fs::{FileTypeExt, MetadataExt, PermissionsExt};
use std::os::unix::net::{UnixListener, UnixStream};
use std::path::{Path, PathBuf};
use std::time::{Duration, Instant};
use super::control::CONTROL_PREFACE;
const HANDSHAKE_DEADLINE: Duration = Duration::from_secs(2);
#[derive(Debug)]
pub struct BoundSocket {
listener: UnixListener,
path: PathBuf,
device: u64,
inode: u64,
}
impl BoundSocket {
pub fn listener(&self) -> &UnixListener {
&self.listener
}
}
impl Drop for BoundSocket {
fn drop(&mut self) {
if let Ok(metadata) = fs::symlink_metadata(&self.path)
&& metadata.file_type().is_socket()
&& metadata.dev() == self.device
&& metadata.ino() == self.inode
{
let _ = fs::remove_file(&self.path);
}
}
}
pub fn bind_local_socket(path: &Path) -> io::Result<BoundSocket> {
let path = path.to_owned();
let directory = path
.parent()
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "socket has no parent"))?;
ensure_private_directory(directory)?;
remove_stale_socket(&path)?;
let listener = UnixListener::bind(&path)?;
let metadata = fs::symlink_metadata(&path)?;
let bound = BoundSocket {
listener,
path,
device: metadata.dev(),
inode: metadata.ino(),
};
fs::set_permissions(&bound.path, fs::Permissions::from_mode(0o600))?;
Ok(bound)
}
pub fn ensure_private_directory(directory: &Path) -> io::Result<()> {
match fs::symlink_metadata(directory) {
Ok(metadata) => {
if !metadata.is_dir()
|| metadata.file_type().is_symlink()
|| metadata.permissions().mode() & 0o077 != 0
|| metadata.uid() != nix::unistd::geteuid().as_raw()
{
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
"fux runtime directory must be a private real directory owned by this user",
));
}
}
Err(error) if error.kind() == io::ErrorKind::NotFound => {
use std::os::unix::fs::DirBuilderExt as _;
fs::DirBuilder::new()
.recursive(true)
.mode(0o700)
.create(directory)?;
}
Err(error) => return Err(error),
}
Ok(())
}
fn remove_stale_socket(path: &Path) -> io::Result<()> {
let metadata = match fs::symlink_metadata(path) {
Ok(metadata) => metadata,
Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(()),
Err(error) => return Err(error),
};
if !metadata.file_type().is_socket() {
return Err(io::Error::new(
io::ErrorKind::AlreadyExists,
"refusing to replace a non-socket path",
));
}
match UnixStream::connect(path) {
Ok(_) => {
return Err(io::Error::new(
io::ErrorKind::AddrInUse,
"socket is already accepting connections",
));
}
Err(error)
if matches!(
error.kind(),
io::ErrorKind::ConnectionRefused | io::ErrorKind::NotFound
) => {}
Err(error) => return Err(error),
}
let parent = path
.parent()
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "socket has no parent"))?;
if metadata.uid() != fs::metadata(parent)?.uid() {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
"stale socket owner differs from runtime directory owner",
));
}
let current = fs::symlink_metadata(path)?;
if !current.file_type().is_socket()
|| current.dev() != metadata.dev()
|| current.ino() != metadata.ino()
{
return Err(io::Error::new(
io::ErrorKind::AddrInUse,
"socket changed during stale recovery",
));
}
fs::remove_file(path)
}
pub fn authorize_peer(stream: &UnixStream) -> io::Result<()> {
#[cfg(any(target_os = "linux", target_os = "android"))]
let uid =
nix::sys::socket::getsockopt(stream, nix::sys::socket::sockopt::PeerCredentials)?.uid();
#[cfg(any(
target_os = "macos",
target_os = "ios",
target_os = "freebsd",
target_os = "openbsd",
target_os = "netbsd",
target_os = "dragonfly"
))]
let uid = nix::unistd::getpeereid(stream)?.0.as_raw();
#[cfg(not(any(
target_os = "linux",
target_os = "android",
target_os = "macos",
target_os = "ios",
target_os = "freebsd",
target_os = "openbsd",
target_os = "netbsd",
target_os = "dragonfly"
)))]
let uid = {
let _ = stream;
return Err(io::Error::new(
io::ErrorKind::Unsupported,
"OS peer credentials unavailable",
));
};
authorize_uid(uid, nix::unistd::geteuid().as_raw())
}
fn authorize_uid(peer: u32, owner: u32) -> io::Result<()> {
if peer != owner {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
"local peer belongs to another user",
));
}
Ok(())
}
pub fn check_private_socket_path(path: &Path) -> io::Result<()> {
let parent = path
.parent()
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "socket has no parent"))?;
let directory = fs::symlink_metadata(parent)?;
let socket = fs::symlink_metadata(path)?;
let owner = nix::unistd::geteuid().as_raw();
if !directory.is_dir()
|| directory.uid() != owner
|| directory.permissions().mode() & 0o077 != 0
{
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
"unsafe local socket directory",
));
}
if !socket.file_type().is_socket()
|| socket.uid() != owner
|| socket.permissions().mode() & 0o077 != 0
{
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
"unsafe local socket path",
));
}
Ok(())
}
pub fn connect_local(path: &Path, deadline: Instant) -> io::Result<UnixStream> {
use nix::fcntl::{FcntlArg, FdFlag, OFlag, fcntl};
use nix::sys::socket::{AddressFamily, SockFlag, SockType, UnixAddr, sockopt};
use std::os::fd::{AsFd, AsRawFd};
let remaining = || {
deadline
.checked_duration_since(Instant::now())
.filter(|duration| !duration.is_zero())
.ok_or_else(|| io::Error::from(io::ErrorKind::TimedOut))
};
remaining()?;
let fd = nix::sys::socket::socket(
AddressFamily::Unix,
SockType::Stream,
SockFlag::empty(),
None,
)?;
fcntl(&fd, FcntlArg::F_SETFD(FdFlag::FD_CLOEXEC))?;
fcntl(&fd, FcntlArg::F_SETFL(OFlag::O_NONBLOCK))?;
match nix::sys::socket::connect(fd.as_raw_fd(), &UnixAddr::new(path)?) {
Ok(()) => {}
Err(nix::errno::Errno::EINPROGRESS) => loop {
let timeout =
u16::try_from(remaining()?.as_millis().clamp(1, 2000)).map_err(io::Error::other)?;
let mut polls = [nix::poll::PollFd::new(
fd.as_fd(),
nix::poll::PollFlags::POLLOUT,
)];
match nix::poll::poll(&mut polls, timeout) {
Ok(0) | Err(nix::errno::Errno::EINTR) => continue,
Ok(_) => {}
Err(error) => return Err(error.into()),
}
let error = nix::sys::socket::getsockopt(&fd, sockopt::SocketError)?;
if error != 0 {
return Err(io::Error::from_raw_os_error(error));
}
break;
},
Err(error) => return Err(error.into()),
}
remaining()?;
let stream = UnixStream::from(fd);
stream.set_nonblocking(false)?;
Ok(stream)
}
pub fn write_all_until(
stream: &mut UnixStream,
mut bytes: &[u8],
deadline: Instant,
) -> io::Result<()> {
use nix::fcntl::{FcntlArg, OFlag, fcntl};
use std::os::fd::AsFd;
let flags = OFlag::from_bits_truncate(fcntl(&*stream, FcntlArg::F_GETFL)?);
fcntl(&*stream, FcntlArg::F_SETFL(flags | OFlag::O_NONBLOCK))?;
let result = (|| {
while !bytes.is_empty() {
let remaining = deadline
.checked_duration_since(Instant::now())
.filter(|duration| !duration.is_zero())
.ok_or_else(|| io::Error::from(io::ErrorKind::TimedOut))?;
match stream.write(bytes) {
Ok(0) => return Err(io::ErrorKind::WriteZero.into()),
Ok(count) => {
bytes = bytes
.get(count..)
.ok_or_else(|| io::Error::other("invalid socket write count"))?
}
Err(error) if error.kind() == io::ErrorKind::Interrupted => continue,
Err(error) if error.kind() == io::ErrorKind::WouldBlock => {
let mut polls = [nix::poll::PollFd::new(
stream.as_fd(),
nix::poll::PollFlags::POLLOUT,
)];
let timeout = u16::try_from(remaining.as_millis().clamp(1, 2000))
.map_err(io::Error::other)?;
match nix::poll::poll(&mut polls, timeout) {
Ok(_) | Err(nix::errno::Errno::EINTR) => {}
Err(error) => return Err(error.into()),
}
}
Err(error) => return Err(error),
}
}
Ok(())
})();
let restored = fcntl(&*stream, FcntlArg::F_SETFL(flags))
.map(|_| ())
.map_err(io::Error::from);
result.and(restored)
}
pub fn negotiate_client(stream: &mut UnixStream) -> io::Result<()> {
negotiate_client_with_timeout(stream, HANDSHAKE_DEADLINE)
}
pub fn negotiate_client_with_timeout(stream: &mut UnixStream, timeout: Duration) -> io::Result<()> {
let timeout = timeout.min(HANDSHAKE_DEADLINE);
if timeout.is_zero() {
return Err(io::ErrorKind::TimedOut.into());
}
authorize_peer(stream)?;
let read_timeout = stream.read_timeout()?;
let write_timeout = stream.write_timeout()?;
let result = (|| {
let deadline = Instant::now() + timeout;
write_all_until(stream, CONTROL_PREFACE, deadline)?;
let mut received = [0; CONTROL_PREFACE.len()];
let mut used = 0;
while used < received.len() {
let remaining = deadline
.checked_duration_since(Instant::now())
.ok_or_else(|| {
io::Error::new(io::ErrorKind::TimedOut, "control negotiation timed out")
})?;
stream.set_read_timeout(Some(remaining))?;
let target = received
.get_mut(used..)
.ok_or_else(|| io::Error::other("invalid preface offset"))?;
let length = stream.read(target)?;
if length == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"peer closed during the control preface",
));
}
used += length;
}
if &received != CONTROL_PREFACE {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"not a fux control socket; restart the session server if it is older than this fux",
));
}
Ok(())
})();
let _ = stream.set_read_timeout(read_timeout);
let _ = stream.set_write_timeout(write_timeout);
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn slow_partial_socket_writes_obey_one_deadline_and_restore_flags() -> io::Result<()> {
use nix::fcntl::{FcntlArg, OFlag, fcntl};
use std::sync::{
Arc,
atomic::{AtomicBool, Ordering},
};
let (mut writer, mut reader) = UnixStream::pair()?;
nix::sys::socket::setsockopt(&writer, nix::sys::socket::sockopt::SndBuf, &4096)?;
reader.set_read_timeout(Some(Duration::from_secs(1)))?;
let finished = Arc::new(AtomicBool::new(false));
let done = Arc::clone(&finished);
let (ready, started) = std::sync::mpsc::channel();
let peer = std::thread::spawn(move || {
let _ = ready.send(());
let mut buffer = [0_u8; 4096];
while !done.load(Ordering::Acquire) {
match reader.read(&mut buffer) {
Ok(0) | Err(_) => break,
Ok(_) => std::thread::sleep(Duration::from_millis(5)),
}
}
});
started
.recv_timeout(Duration::from_secs(1))
.map_err(io::Error::other)?;
let start = Instant::now();
let result = write_all_until(
&mut writer,
&vec![b'x'; 1024 * 1024],
start + Duration::from_millis(100),
);
finished.store(true, Ordering::Release);
let flags = OFlag::from_bits_truncate(fcntl(&writer, FcntlArg::F_GETFL)?);
drop(writer);
peer.join().map_err(|_| io::Error::other("peer panicked"))?;
assert!(matches!(result, Err(error) if error.kind() == io::ErrorKind::TimedOut));
assert!(start.elapsed() < Duration::from_secs(1));
assert!(!flags.contains(OFlag::O_NONBLOCK));
Ok(())
}
#[test]
fn foreign_uid_is_rejected_and_current_kernel_peer_is_accepted() -> io::Result<()> {
assert!(authorize_uid(501, 502).is_err());
assert!(authorize_uid(0, 502).is_err());
let (first, second) = UnixStream::pair()?;
authorize_peer(&first)?;
authorize_peer(&second)?;
Ok(())
}
#[test]
fn negotiation_requires_the_exact_preface_from_the_server() -> io::Result<()> {
for (answer, accepted) in [(&b"FUX\n"[..], true), (b"FUZ\n", false)] {
let (mut client, mut server) = UnixStream::pair()?;
let handle = std::thread::spawn(move || {
let mut preface = [0; 4];
server.read_exact(&mut preface)?;
server.write_all(answer)?;
Ok::<_, io::Error>(preface)
});
let result = negotiate_client(&mut client);
assert_eq!(result.is_ok(), accepted, "{answer:?}");
assert!(
handle
.join()
.is_ok_and(|sent| sent.is_ok_and(|preface| &preface == CONTROL_PREFACE))
);
}
Ok(())
}
}