use std::ffi::c_void;
use std::os::fd::{AsRawFd, FromRawFd, OwnedFd};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use tokio::io::Interest;
use tokio::io::unix::AsyncFd;
use tokio_util::sync::CancellationToken;
use crate::ffi::{_Block_release, create_vmnet_event_block};
use crate::interface::Vmnet;
const MAX_FRAME_SIZE: usize = 9216;
struct ReadCtx {
vmnet: Arc<Vmnet>,
guest_fd: OwnedFd,
buf: Mutex<Vec<u8>>,
dead: AtomicBool,
events: AtomicU64,
frames: AtomicU64,
drops: AtomicU64,
}
impl ReadCtx {
fn drain(&self) {
if self.dead.load(Ordering::Relaxed) {
return;
}
self.events.fetch_add(1, Ordering::Relaxed);
let mut buf = self.buf.lock().expect("relay read buffer poisoned");
loop {
match self.vmnet.read_packet(&mut buf) {
Ok(0) => break,
Ok(n) => {
if !self.forward_frame(&buf[..n]) {
self.dead.store(true, Ordering::Relaxed);
break;
}
}
Err(e) => {
tracing::trace!("vmnet read error: {e}");
break;
}
}
}
}
fn forward_frame(&self, frame: &[u8]) -> bool {
let fd = self.guest_fd.as_raw_fd();
loop {
let written = unsafe { libc::write(fd, frame.as_ptr().cast::<c_void>(), frame.len()) };
if written >= 0 {
self.frames.fetch_add(1, Ordering::Relaxed);
return true;
}
let err = std::io::Error::last_os_error();
match err.kind() {
std::io::ErrorKind::Interrupted => {}
std::io::ErrorKind::BrokenPipe => return false,
std::io::ErrorKind::WouldBlock => {
self.drops.fetch_add(1, Ordering::Relaxed);
return true;
}
_ if err.raw_os_error() == Some(libc::ENOBUFS) => {
self.drops.fetch_add(1, Ordering::Relaxed);
return true;
}
_ => {
tracing::debug!("vmnet→guest write error: {err}");
return true;
}
}
}
}
}
unsafe extern "C" fn on_packets_available(ctx: *const c_void) {
let ctx = unsafe { &*ctx.cast::<ReadCtx>() };
ctx.drain();
}
unsafe extern "C" fn read_ctx_retain(ctx: *const c_void) {
unsafe { Arc::increment_strong_count(ctx.cast::<ReadCtx>()) }
}
unsafe extern "C" fn read_ctx_release(ctx: *const c_void) {
unsafe { Arc::decrement_strong_count(ctx.cast::<ReadCtx>()) }
}
pub struct VmnetRelay {
vmnet: Arc<Vmnet>,
cancel: CancellationToken,
}
impl VmnetRelay {
#[must_use]
pub fn new(vmnet: Arc<Vmnet>, cancel: CancellationToken) -> Self {
Self { vmnet, cancel }
}
pub async fn run(self, guest_fd: OwnedFd) -> std::io::Result<()> {
let raw_fd = guest_fd.as_raw_fd();
unsafe {
let flags = libc::fcntl(raw_fd, libc::F_GETFL);
if flags < 0 || libc::fcntl(raw_fd, libc::F_SETFL, flags | libc::O_NONBLOCK) < 0 {
return Err(std::io::Error::last_os_error());
}
}
let dup_fd = unsafe { libc::dup(raw_fd) };
if dup_fd < 0 {
return Err(std::io::Error::last_os_error());
}
let reader_fd: OwnedFd = unsafe { OwnedFd::from_raw_fd(dup_fd) };
let async_fd = AsyncFd::new(guest_fd)?;
let ctx = Arc::new(ReadCtx {
vmnet: Arc::clone(&self.vmnet),
guest_fd: reader_fd,
buf: Mutex::new(vec![0u8; MAX_FRAME_SIZE]),
dead: AtomicBool::new(false),
events: AtomicU64::new(0),
frames: AtomicU64::new(0),
drops: AtomicU64::new(0),
});
let block = unsafe {
create_vmnet_event_block(
Arc::as_ptr(&ctx).cast(),
on_packets_available,
read_ctx_retain,
read_ctx_release,
)
};
let registered = self.vmnet.set_event_callback(block);
unsafe { _Block_release(block) };
registered.map_err(std::io::Error::other)?;
let vmnet_write = Arc::clone(&self.vmnet);
let cancel_write = self.cancel.clone();
let mut buf = vec![0u8; MAX_FRAME_SIZE];
loop {
tokio::select! {
() = cancel_write.cancelled() => break,
ready = async_fd.ready(Interest::READABLE) => {
let mut guard = match ready {
Ok(g) => g,
Err(e) => {
tracing::debug!("AsyncFd ready error: {e}");
break;
}
};
let fd = async_fd.as_raw_fd();
let n = unsafe {
libc::read(
fd,
buf.as_mut_ptr().cast::<c_void>(),
buf.len(),
)
};
match n.cmp(&0) {
std::cmp::Ordering::Greater => {
if let Err(e) = vmnet_write.write_packet(&buf[..n as usize]) {
tracing::debug!("guest→vmnet write error: {e}");
}
}
std::cmp::Ordering::Equal => break, std::cmp::Ordering::Less => {
let err = std::io::Error::last_os_error();
match err.kind() {
std::io::ErrorKind::WouldBlock => guard.clear_ready(),
std::io::ErrorKind::Interrupted => {}
_ => {
tracing::debug!("guest→vmnet read error: {err}");
break;
}
}
}
}
}
}
}
self.vmnet.clear_event_callback();
tracing::debug!(
events = ctx.events.load(Ordering::Relaxed),
frames = ctx.frames.load(Ordering::Relaxed),
drops = ctx.drops.load(Ordering::Relaxed),
"vmnet relay stopped"
);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Weak;
use std::time::{Duration, Instant};
struct TestCtx {
vmnet: Arc<Vmnet>,
frames: Mutex<Vec<Vec<u8>>>,
}
impl TestCtx {
fn drain(&self) {
let mut buf = vec![0u8; MAX_FRAME_SIZE];
while let Ok(n) = self.vmnet.read_packet(&mut buf) {
if n == 0 {
break;
}
self.frames.lock().unwrap().push(buf[..n].to_vec());
}
}
}
unsafe extern "C" fn test_on_event(ctx: *const std::ffi::c_void) {
unsafe { &*ctx.cast::<TestCtx>() }.drain();
}
unsafe extern "C" fn test_retain(ctx: *const std::ffi::c_void) {
unsafe { Arc::increment_strong_count(ctx.cast::<TestCtx>()) }
}
unsafe extern "C" fn test_release(ctx: *const std::ffi::c_void) {
unsafe { Arc::decrement_strong_count(ctx.cast::<TestCtx>()) }
}
fn build_dhcp_discover(mac: [u8; 6]) -> Vec<u8> {
let mut bootp = vec![0u8; 236];
bootp[0] = 1; bootp[1] = 1; bootp[2] = 6; bootp[4..8].copy_from_slice(&0x2A2A_2A2Au32.to_be_bytes()); bootp[10..12].copy_from_slice(&0x8000u16.to_be_bytes()); bootp[28..34].copy_from_slice(&mac); bootp.extend_from_slice(&[0x63, 0x82, 0x53, 0x63]); bootp.extend_from_slice(&[53, 1, 1]); bootp.push(255);
let udp_len = 8 + bootp.len();
let ip_len = 20 + udp_len;
let mut ip = vec![
0x45, 0, ];
ip.extend_from_slice(&u16::try_from(ip_len).unwrap().to_be_bytes());
ip.extend_from_slice(&[0, 0, 0, 0]); ip.extend_from_slice(&[64, 17, 0, 0]); ip.extend_from_slice(&[0, 0, 0, 0]); ip.extend_from_slice(&[255, 255, 255, 255]); let sum = ip_header_checksum(&ip);
ip[10..12].copy_from_slice(&sum.to_be_bytes());
let mut udp = Vec::new();
udp.extend_from_slice(&68u16.to_be_bytes()); udp.extend_from_slice(&67u16.to_be_bytes()); udp.extend_from_slice(&u16::try_from(udp_len).unwrap().to_be_bytes());
udp.extend_from_slice(&[0, 0]);
let mut frame = Vec::new();
frame.extend_from_slice(&[0xFF; 6]); frame.extend_from_slice(&mac);
frame.extend_from_slice(&[0x08, 0x00]); frame.extend_from_slice(&ip);
frame.extend_from_slice(&udp);
frame.extend_from_slice(&bootp);
frame
}
fn ip_header_checksum(header: &[u8]) -> u16 {
let mut sum = 0u32;
for chunk in header.chunks(2) {
sum += u32::from(u16::from_be_bytes([chunk[0], chunk[1]]));
}
while sum > 0xFFFF {
sum = (sum & 0xFFFF) + (sum >> 16);
}
!(sum as u16)
}
fn is_dhcp_reply(frame: &[u8]) -> bool {
frame.len() > 36
&& frame[12..14] == [0x08, 0x00]
&& frame[23] == 17
&& u16::from_be_bytes([frame[34], frame[35]]) == 67
}
#[test]
#[ignore = "requires macOS vmnet entitlements and root"]
fn event_callback_delivers_dhcp_offer() {
let vmnet = Arc::new(Vmnet::new_shared().expect("failed to create shared vmnet"));
let ctx = Arc::new(TestCtx {
vmnet: Arc::clone(&vmnet),
frames: Mutex::new(Vec::new()),
});
let weak: Weak<TestCtx> = Arc::downgrade(&ctx);
let block = unsafe {
create_vmnet_event_block(
Arc::as_ptr(&ctx).cast(),
test_on_event,
test_retain,
test_release,
)
};
vmnet
.set_event_callback(block)
.expect("event callback registration failed");
unsafe { _Block_release(block) };
let discover = build_dhcp_discover(vmnet.mac());
let deadline = Instant::now() + Duration::from_secs(5);
let mut offered = false;
while Instant::now() < deadline {
vmnet
.write_packet(&discover)
.expect("write DISCOVER failed");
std::thread::sleep(Duration::from_millis(200));
if ctx.frames.lock().unwrap().iter().any(|f| is_dhcp_reply(f)) {
offered = true;
break;
}
}
assert!(
offered,
"no DHCP reply delivered via event callback; frames seen: {}",
ctx.frames.lock().unwrap().len()
);
vmnet.clear_event_callback();
drop(ctx);
let deadline = Instant::now() + Duration::from_secs(2);
while weak.upgrade().is_some() && Instant::now() < deadline {
std::thread::sleep(Duration::from_millis(50));
}
assert!(
weak.upgrade().is_none(),
"TestCtx leaked: block dispose never released the context"
);
vmnet.stop();
}
}