use std::os::fd::{AsRawFd, FromRawFd, OwnedFd};
use std::sync::Arc;
use tokio::io::Interest;
use tokio::io::unix::AsyncFd;
use tokio_util::sync::CancellationToken;
use crate::interface::Vmnet;
const MAX_FRAME_SIZE: usize = 9216;
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();
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 vmnet_read = Arc::clone(&self.vmnet);
let cancel_read = self.cancel.clone();
let mut vmnet_to_guest = tokio::task::spawn_blocking(move || {
let mut buf = vec![0u8; MAX_FRAME_SIZE];
loop {
if cancel_read.is_cancelled() {
break;
}
match vmnet_read.read_packet(&mut buf) {
Ok(0) => {
std::thread::sleep(std::time::Duration::from_millis(1));
}
Ok(n) => {
let fd = reader_fd.as_raw_fd();
let written =
unsafe { libc::write(fd, buf.as_ptr().cast::<libc::c_void>(), n) };
if written < 0 {
let err = std::io::Error::last_os_error();
match err.kind() {
std::io::ErrorKind::BrokenPipe => break,
std::io::ErrorKind::WouldBlock => {}
_ => tracing::debug!("vmnet→guest write error: {err}"),
}
}
}
Err(e) => {
if cancel_read.is_cancelled() {
break;
}
tracing::debug!("vmnet read error: {e}");
std::thread::sleep(std::time::Duration::from_millis(1));
}
}
}
});
let vmnet_write = Arc::clone(&self.vmnet);
let cancel_write = self.cancel.clone();
let guest_to_vmnet = async move {
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::<libc::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();
if err.kind() == std::io::ErrorKind::WouldBlock {
guard.clear_ready();
} else {
tracing::debug!("guest→vmnet read error: {err}");
break;
}
}
}
}
}
}
};
tokio::select! {
() = self.cancel.cancelled() => {}
_ = &mut vmnet_to_guest => {}
() = guest_to_vmnet => {}
}
self.cancel.cancel();
if let Err(e) = vmnet_to_guest.await {
if e.is_panic() {
tracing::error!("vmnet→guest blocking task panicked: {e}");
} else {
tracing::debug!("vmnet→guest blocking task join error: {e}");
}
}
Ok(())
}
}