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 { 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 { self.sock.send(pkt) } pub fn read(&self, buf: &mut [u8]) -> std::io::Result { self.sock.recv(buf) } pub fn is_connected(&self) -> io::Result { 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 { 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::() 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::() 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()); } }