use std::io::{ErrorKind, Read, Result, Write};
use std::net::Shutdown;
use std::os::unix::io::{AsRawFd, IntoRawFd, RawFd};
use futures::{future::poll_fn, ready};
use mio::Ready;
use nix::sys::socket::SockAddr;
use std::mem::{self, MaybeUninit};
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::PollEvented;
use tokio::io::{AsyncRead, AsyncWrite};
#[derive(Debug)]
pub struct VsockStream {
io: PollEvented<super::mio::VsockStream>,
}
impl VsockStream {
pub(crate) fn new(connected: super::mio::VsockStream) -> Result<Self> {
let io = PollEvented::new(connected)?;
Ok(Self { io })
}
pub async fn connect(addr: &SockAddr) -> Result<Self> {
let stream = super::mio::VsockStream::connect(addr)?;
let stream = Self::new(stream)?;
poll_fn(|cx| stream.io.poll_write_ready(cx)).await?;
Ok(stream)
}
pub fn from_std(stream: vsock::VsockStream) -> Result<Self> {
let io = super::mio::VsockStream::from_std(stream)?;
let io = PollEvented::new(io)?;
Ok(VsockStream { io })
}
pub fn poll_read_ready(&self, cx: &mut Context, mask: Ready) -> Poll<Result<Ready>> {
self.io.poll_read_ready(cx, mask)
}
pub fn poll_write_ready(&self, cx: &mut Context) -> Poll<Result<Ready>> {
self.io.poll_write_ready(cx)
}
pub fn local_addr(&self) -> Result<SockAddr> {
self.io.get_ref().local_addr()
}
pub fn peer_addr(&self) -> Result<SockAddr> {
self.io.get_ref().peer_addr()
}
pub fn shutdown(&self, how: Shutdown) -> Result<()> {
self.io.get_ref().shutdown(how)
}
pub(crate) fn poll_write_priv(&self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize>> {
ready!(self.io.poll_write_ready(cx))?;
match self.io.get_ref().write(buf) {
Err(ref e) if e.kind() == ErrorKind::WouldBlock => {
self.io.clear_write_ready(cx)?;
Poll::Pending
}
x => Poll::Ready(x),
}
}
pub(crate) fn poll_read_priv(
&self,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize>> {
ready!(self.io.poll_read_ready(cx, Ready::readable()))?;
match self.io.get_ref().read(buf) {
Err(ref e) if e.kind() == ErrorKind::WouldBlock => {
self.io.clear_read_ready(cx, Ready::readable())?;
Poll::Pending
}
x => Poll::Ready(x),
}
}
}
impl AsRawFd for VsockStream {
fn as_raw_fd(&self) -> RawFd {
self.io.get_ref().as_raw_fd()
}
}
impl IntoRawFd for VsockStream {
fn into_raw_fd(self) -> RawFd {
let fd = self.io.get_ref().as_raw_fd();
mem::forget(self);
fd
}
}
impl Write for VsockStream {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
self.io.get_ref().write(buf)
}
fn flush(&mut self) -> Result<()> {
Ok(())
}
}
impl Read for VsockStream {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
self.io.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 {
unsafe fn prepare_uninitialized_buffer(&self, _: &mut [MaybeUninit<u8>]) -> bool {
false
}
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize>> {
self.poll_read_priv(cx, buf)
}
}