From c3e209dce4a2965e5d38020dde08f5232c32fdf7 Mon Sep 17 00:00:00 2001 From: Minh Vu Date: Fri, 10 Jul 2026 00:28:43 +0200 Subject: [PATCH] Validate VM file descriptors before ownership transfer --- lib/vm.rs | 151 +++++++++++++++++++++++++++++++++++++++++++++++++++- src/main.rs | 35 ++++++++++++ 2 files changed, 184 insertions(+), 2 deletions(-) diff --git a/lib/vm.rs b/lib/vm.rs index 8a0671b..c02f657 100644 --- a/lib/vm.rs +++ b/lib/vm.rs @@ -1,5 +1,7 @@ -use anyhow::Result; -use std::os::unix::io::{AsRawFd, FromRawFd, RawFd}; +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 { @@ -8,6 +10,9 @@ pub struct VM { 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)?; @@ -23,8 +28,150 @@ impl VM { } } +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}") + }); + } + + if address.ss_family as libc::c_int != libc::AF_UNIX { + 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 std::fs::File; + use std::os::fd::AsRawFd; + use std::os::unix::net::{UnixDatagram, UnixStream}; + + #[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_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); + } +} diff --git a/src/main.rs b/src/main.rs index 3dd9ba8..f974784 100644 --- a/src/main.rs +++ b/src/main.rs @@ -26,6 +26,7 @@ use uzers::{get_current_groupname, get_current_username, get_effective_uid}; struct Args { #[clap( long, + value_parser = parse_vm_fd, help = "FD number to use for communicating with the VM's networking stack" )] vm_fd: c_int, @@ -217,6 +218,18 @@ fn try_main() -> anyhow::Result<()> { proxy.run() } +fn parse_vm_fd(value: &str) -> Result { + let vm_fd = value + .parse::() + .map_err(|err| format!("invalid file descriptor: {err}"))?; + + if vm_fd < 0 { + return Err("file descriptor must be non-negative".to_string()); + } + + Ok(vm_fd) +} + fn sudo_escalation_works() -> bool { let exe = std::env::current_exe().unwrap(); let args = std::env::args().skip(1); @@ -252,3 +265,25 @@ fn set_bootpd_lease_time(lease_time: u32) { SCPreferencesCommitChanges(prefs.as_concrete_TypeRef()); } } + +#[cfg(test)] +mod tests { + use super::Args; + use clap::Parser; + + #[test] + fn test_cli_rejects_negative_vm_fd_before_startup() { + let error = Args::try_parse_from([ + "softnet", + "--vm-fd=-1", + "--vm-mac-address=02:00:00:00:00:01", + ]) + .unwrap_err(); + + assert!( + error + .to_string() + .contains("file descriptor must be non-negative") + ); + } +}