259 lines
7.8 KiB
Rust
259 lines
7.8 KiB
Rust
use anyhow::{Context, Result, bail};
|
|
use std::io;
|
|
use std::mem::{size_of, zeroed};
|
|
use std::os::fd::{AsRawFd, FromRawFd, RawFd};
|
|
use std::os::unix::net::UnixDatagram;
|
|
|
|
pub struct VM {
|
|
sock: UnixDatagram,
|
|
}
|
|
|
|
impl VM {
|
|
pub fn new(vm_fd: RawFd) -> Result<VM> {
|
|
let vm_fd = duplicate_vm_fd(vm_fd)?;
|
|
|
|
// SAFETY: duplicate_vm_fd only returns a valid descriptor that it owns.
|
|
let sock = unsafe { UnixDatagram::from_raw_fd(vm_fd) };
|
|
sock.set_nonblocking(true)?;
|
|
|
|
Ok(VM { sock })
|
|
}
|
|
|
|
pub fn write(&self, pkt: &[u8]) -> std::io::Result<usize> {
|
|
self.sock.send(pkt)
|
|
}
|
|
|
|
pub fn read(&self, buf: &mut [u8]) -> std::io::Result<usize> {
|
|
self.sock.recv(buf)
|
|
}
|
|
|
|
pub fn is_connected(&self) -> io::Result<bool> {
|
|
match self.sock.peer_addr() {
|
|
Ok(_) => Ok(true),
|
|
Err(error) if error.kind() == io::ErrorKind::NotConnected => Ok(false),
|
|
Err(error) => Err(error),
|
|
}
|
|
}
|
|
}
|
|
|
|
fn duplicate_vm_fd(vm_fd: RawFd) -> Result<RawFd> {
|
|
if vm_fd < 0 {
|
|
bail!("invalid VM file descriptor {vm_fd}: value must be non-negative");
|
|
}
|
|
|
|
// SAFETY: fcntl duplicates the descriptor without transferring ownership of vm_fd.
|
|
let duplicated_fd = unsafe { libc::fcntl(vm_fd, libc::F_DUPFD_CLOEXEC, 0) };
|
|
if duplicated_fd == -1 {
|
|
return Err(io::Error::last_os_error())
|
|
.with_context(|| format!("failed to duplicate VM file descriptor {vm_fd}"));
|
|
}
|
|
|
|
if let Err(error) = validate_vm_fd(duplicated_fd) {
|
|
// SAFETY: duplicated_fd is an open descriptor owned by this function.
|
|
unsafe { libc::close(duplicated_fd) };
|
|
return Err(error);
|
|
}
|
|
|
|
Ok(duplicated_fd)
|
|
}
|
|
|
|
fn validate_vm_fd(vm_fd: RawFd) -> Result<()> {
|
|
// SAFETY: fcntl only reads descriptor state and does not take ownership.
|
|
if unsafe { libc::fcntl(vm_fd, libc::F_GETFD) } == -1 {
|
|
return Err(io::Error::last_os_error())
|
|
.with_context(|| format!("failed to inspect VM file descriptor {vm_fd}"));
|
|
}
|
|
|
|
let mut socket_type = 0;
|
|
let mut socket_type_len = size_of::<libc::c_int>() as libc::socklen_t;
|
|
|
|
// SAFETY: socket_type and socket_type_len are valid writable buffers of the sizes given.
|
|
if unsafe {
|
|
libc::getsockopt(
|
|
vm_fd,
|
|
libc::SOL_SOCKET,
|
|
libc::SO_TYPE,
|
|
(&mut socket_type as *mut libc::c_int).cast(),
|
|
&mut socket_type_len,
|
|
)
|
|
} == -1
|
|
{
|
|
return Err(io::Error::last_os_error())
|
|
.with_context(|| format!("VM file descriptor {vm_fd} is not a socket"));
|
|
}
|
|
|
|
if socket_type != libc::SOCK_DGRAM {
|
|
bail!("VM file descriptor {vm_fd} is not a Unix datagram socket");
|
|
}
|
|
|
|
let mut address: libc::sockaddr_storage = unsafe { zeroed() };
|
|
let mut address_len = size_of::<libc::sockaddr_storage>() as libc::socklen_t;
|
|
|
|
// SAFETY: address and address_len are valid writable buffers of the sizes given.
|
|
if unsafe {
|
|
libc::getsockname(
|
|
vm_fd,
|
|
(&mut address as *mut libc::sockaddr_storage).cast(),
|
|
&mut address_len,
|
|
)
|
|
} == -1
|
|
{
|
|
return Err(io::Error::last_os_error()).with_context(|| {
|
|
format!("failed to inspect the address family of VM file descriptor {vm_fd}")
|
|
});
|
|
}
|
|
|
|
// macOS returns a zero-length address for unnamed UNIX-domain sockets,
|
|
// including socketpair descriptors. Other socket families return their
|
|
// address family when getsockname succeeds.
|
|
let is_unix_socket = address_len == 0 || address.ss_family as libc::c_int == libc::AF_UNIX;
|
|
if !is_unix_socket {
|
|
bail!("VM file descriptor {vm_fd} is not a Unix socket");
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
impl AsRawFd for VM {
|
|
fn as_raw_fd(&self) -> RawFd {
|
|
self.sock.as_raw_fd()
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::VM;
|
|
use polling::{Event, Events, PollMode, Poller};
|
|
use std::fs::File;
|
|
use std::net::UdpSocket;
|
|
use std::os::fd::AsRawFd;
|
|
use std::os::unix::net::{UnixDatagram, UnixStream};
|
|
use std::time::Duration;
|
|
|
|
#[test]
|
|
fn test_new_rejects_negative_fd() {
|
|
let error = VM::new(-1).err().unwrap();
|
|
|
|
assert_eq!(
|
|
error.to_string(),
|
|
"invalid VM file descriptor -1: value must be non-negative"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_new_rejects_non_socket_fd_without_taking_ownership() {
|
|
let file = File::open("/dev/null").unwrap();
|
|
let error = VM::new(file.as_raw_fd()).err().unwrap();
|
|
|
|
assert!(error.to_string().contains("is not a socket"));
|
|
assert!(file.metadata().is_ok());
|
|
}
|
|
|
|
#[test]
|
|
fn test_new_rejects_closed_fd() {
|
|
let (socket, _peer) = UnixDatagram::pair().unwrap();
|
|
let vm_fd = socket.as_raw_fd();
|
|
drop(socket);
|
|
|
|
let error = VM::new(vm_fd).err().unwrap();
|
|
|
|
assert!(
|
|
error
|
|
.to_string()
|
|
.contains("failed to duplicate VM file descriptor")
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_new_rejects_non_datagram_socket() {
|
|
let (stream, _peer) = UnixStream::pair().unwrap();
|
|
let error = VM::new(stream.as_raw_fd()).err().unwrap();
|
|
|
|
assert!(error.to_string().contains("not a Unix datagram socket"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_new_rejects_internet_datagram_socket() {
|
|
let socket = UdpSocket::bind("127.0.0.1:0").unwrap();
|
|
let error = VM::new(socket.as_raw_fd()).err().unwrap();
|
|
|
|
assert!(error.to_string().contains("not a Unix socket"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_new_does_not_close_original_fd_when_vm_is_dropped() {
|
|
let (socket, _peer) = UnixDatagram::pair().unwrap();
|
|
let vm = VM::new(socket.as_raw_fd()).unwrap();
|
|
drop(vm);
|
|
|
|
let socket_fd_is_open = unsafe { libc::fcntl(socket.as_raw_fd(), libc::F_GETFD) != -1 };
|
|
|
|
if socket_fd_is_open {
|
|
drop(socket);
|
|
} else {
|
|
// Avoid double-closing the descriptor if this test catches an unsafe implementation.
|
|
std::mem::forget(socket);
|
|
}
|
|
|
|
assert!(socket_fd_is_open);
|
|
}
|
|
|
|
#[test]
|
|
fn test_connected_socket_has_peer() {
|
|
let (socket, _peer) = UnixDatagram::pair().unwrap();
|
|
let vm = VM::new(socket.as_raw_fd()).unwrap();
|
|
|
|
assert!(vm.is_connected().unwrap());
|
|
}
|
|
|
|
#[test]
|
|
fn test_disconnected_peer_is_detected_without_kqueue_event() {
|
|
let (socket, peer) = UnixDatagram::pair().unwrap();
|
|
let vm = VM::new(socket.as_raw_fd()).unwrap();
|
|
let poller = Poller::new().unwrap();
|
|
let mut events = Events::new();
|
|
|
|
unsafe {
|
|
poller
|
|
.add_with_mode(vm.as_raw_fd(), Event::readable(0), PollMode::Edge)
|
|
.unwrap();
|
|
}
|
|
drop(peer);
|
|
|
|
poller
|
|
.wait(&mut events, Some(Duration::from_millis(20)))
|
|
.unwrap();
|
|
|
|
assert!(events.is_empty());
|
|
assert!(!vm.is_connected().unwrap());
|
|
}
|
|
|
|
#[test]
|
|
fn test_disconnected_peer_is_detected_when_another_socket_wakes_kqueue() {
|
|
let (socket, peer) = UnixDatagram::pair().unwrap();
|
|
let (host, host_peer) = UnixDatagram::pair().unwrap();
|
|
let vm = VM::new(socket.as_raw_fd()).unwrap();
|
|
let poller = Poller::new().unwrap();
|
|
let mut events = Events::new();
|
|
|
|
unsafe {
|
|
poller
|
|
.add_with_mode(vm.as_raw_fd(), Event::readable(0), PollMode::Edge)
|
|
.unwrap();
|
|
poller
|
|
.add_with_mode(host.as_raw_fd(), Event::readable(1), PollMode::Edge)
|
|
.unwrap();
|
|
}
|
|
drop(peer);
|
|
host_peer.send(&[1]).unwrap();
|
|
|
|
poller
|
|
.wait(&mut events, Some(Duration::from_millis(20)))
|
|
.unwrap();
|
|
|
|
assert!(events.iter().any(|event| event.key == 1));
|
|
assert!(!events.iter().any(|event| event.key == 0));
|
|
assert!(!vm.is_connected().unwrap());
|
|
}
|
|
}
|