use std::io::{Error, Read, Result, Write};
use std::net::Shutdown;
use std::os::unix::io::{AsRawFd, FromRawFd, IntoRawFd, RawFd};
use crate::{SockAddr, VsockAddr};
use futures::ready;
use libc::*;
use std::mem::{self, size_of};
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::unix::AsyncFd;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
#[derive(Debug)]
pub struct VsockStream {
inner: AsyncFd<vsock::VsockStream>,
}
impl VsockStream {
pub(crate) fn new(connected: vsock::VsockStream) -> Result<Self> {
connected.set_nonblocking(true)?;
Ok(Self {
inner: AsyncFd::new(connected)?,
})
}
pub async fn connect(cid: u32, port: u32) -> Result<Self> {
let vsock_addr = VsockAddr::new(cid, port);
let socket = unsafe { socket(AF_VSOCK, SOCK_STREAM | SOCK_CLOEXEC, 0) };
if socket < 0 {
return Err(Error::last_os_error());
}
if unsafe { fcntl(socket, F_SETFL, O_NONBLOCK) } < 0 {
let _ = unsafe { close(socket) };
return Err(Error::last_os_error());
}
if unsafe {
connect(
socket,
&vsock_addr as *const _ as *const sockaddr,
size_of::<sockaddr_vm>() as u32,
)
} < 0
{
let err = Error::last_os_error();
if let Some(os_err) = err.raw_os_error() {
if os_err != EINPROGRESS {
let _ = unsafe { close(socket) };
return Err(err);
}
}
}
loop {
let stream = unsafe { vsock::VsockStream::from_raw_fd(socket) };
let stream = Self::new(stream)?;
let mut guard = stream.inner.writable().await?;
match guard.try_io(|_| Ok(())) {
Ok(_) => return Ok(stream),
Err(_would_block) => continue,
}
}
}
pub fn local_addr(&self) -> Result<SockAddr> {
self.inner.get_ref().local_addr()
}
pub fn peer_addr(&self) -> Result<SockAddr> {
self.inner.get_ref().peer_addr()
}
pub fn shutdown(&self, how: Shutdown) -> Result<()> {
self.inner.get_ref().shutdown(how)
}
pub(crate) fn poll_write_priv(&self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize>> {
loop {
let mut guard = ready!(self.inner.poll_write_ready(cx))?;
match guard.try_io(|inner| inner.get_ref().write(buf)) {
Ok(Ok(n)) => return Ok(n).into(),
Ok(Err(ref e)) if e.kind() == std::io::ErrorKind::Interrupted => continue,
Ok(Err(e)) => return Err(e).into(),
Err(_would_block) => continue,
}
}
}
pub(crate) fn poll_read_priv(
&self,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<Result<()>> {
let b;
unsafe {
b = &mut *(buf.unfilled_mut() as *mut [mem::MaybeUninit<u8>] as *mut [u8]);
};
loop {
let mut guard = ready!(self.inner.poll_read_ready(cx))?;
match guard.try_io(|inner| inner.get_ref().read(b)) {
Ok(Ok(n)) => {
unsafe {
buf.assume_init(n);
}
buf.advance(n);
return Ok(()).into();
}
Ok(Err(ref e)) if e.kind() == std::io::ErrorKind::Interrupted => continue,
Ok(Err(e)) => return Err(e).into(),
Err(_would_block) => {
continue;
}
}
}
}
}
impl AsRawFd for VsockStream {
fn as_raw_fd(&self) -> RawFd {
self.inner.get_ref().as_raw_fd()
}
}
impl IntoRawFd for VsockStream {
fn into_raw_fd(self) -> RawFd {
let fd = self.inner.get_ref().as_raw_fd();
mem::forget(self);
fd
}
}
impl Write for VsockStream {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
self.inner.get_ref().write(buf)
}
fn flush(&mut self) -> Result<()> {
Ok(())
}
}
impl Read for VsockStream {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
self.inner.get_ref().read(buf)
}
}
impl AsyncWrite for VsockStream {
fn poll_write(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize>> {
self.poll_write_priv(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<()>> {
self.shutdown(std::net::Shutdown::Write)?;
Poll::Ready(Ok(()))
}
}
impl AsyncRead for VsockStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<Result<()>> {
self.poll_read_priv(cx, buf)
}
}