use std::cmp;
use std::io::{self, Error, ErrorKind, IoSlice, IoSliceMut, Read, Result, Write};
use std::mem::size_of;
use std::net::Shutdown;
use std::os::unix::io::{AsRawFd, FromRawFd, IntoRawFd, RawFd};
use libc::*;
use mio::unix::EventedFd;
use mio::{Evented, Poll, PollOpt, Ready, Token};
use nix::sys::socket::SockAddr;
use super::iovec::unix as iovec;
use super::iovec::IoVec;
#[derive(Debug)]
pub struct VsockStream {
inner: vsock::VsockStream,
}
impl VsockStream {
pub fn connect(addr: &SockAddr) -> Result<VsockStream> {
let vsock_addr = if let SockAddr::Vsock(addr) = addr {
addr.0
} else {
return Err(Error::new(
ErrorKind::Other,
"requires a virtio socket address",
));
};
let socket = unsafe { socket(AF_VSOCK, SOCK_STREAM, 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);
}
}
}
Ok(Self {
inner: unsafe { vsock::VsockStream::from_raw_fd(socket) },
})
}
pub fn from_std(inner: vsock::VsockStream) -> Result<VsockStream> {
inner.set_nonblocking(true)?;
Ok(VsockStream { inner })
}
pub fn peer_addr(&self) -> Result<SockAddr> {
self.inner.peer_addr()
}
pub fn local_addr(&self) -> Result<SockAddr> {
self.inner.local_addr()
}
pub fn try_clone(&self) -> Result<VsockStream> {
self.inner.try_clone().map(|s| VsockStream { inner: s })
}
pub fn shutdown(&self, how: Shutdown) -> Result<()> {
self.inner.shutdown(how)
}
pub fn take_error(&self) -> Result<Option<io::Error>> {
self.inner.take_error()
}
pub fn read_bufs(&self, bufs: &mut [&mut IoVec]) -> Result<usize> {
unsafe {
let slice = iovec::as_os_slice_mut(bufs);
let len = cmp::min(<libc::c_int>::max_value() as usize, slice.len());
let rc = readv(self.as_raw_fd(), slice.as_ptr(), len as libc::c_int);
if rc < 0 {
Err(io::Error::last_os_error())
} else {
Ok(rc as usize)
}
}
}
pub fn write_bufs(&self, bufs: &[&IoVec]) -> Result<usize> {
unsafe {
let slice = iovec::as_os_slice(bufs);
let len = cmp::min(<libc::c_int>::max_value() as usize, slice.len());
let rc = writev(self.as_raw_fd(), slice.as_ptr(), len as libc::c_int);
if rc < 0 {
Err(io::Error::last_os_error())
} else {
Ok(rc as usize)
}
}
}
}
impl<'a> Read for &'a VsockStream {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
(&self.inner).read(buf)
}
fn read_vectored(&mut self, bufs: &mut [IoSliceMut<'_>]) -> Result<usize> {
(&self.inner).read_vectored(bufs)
}
}
impl<'a> Write for &'a VsockStream {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
(&self.inner).write(buf)
}
fn write_vectored(&mut self, bufs: &[IoSlice<'_>]) -> Result<usize> {
(&self.inner).write_vectored(bufs)
}
fn flush(&mut self) -> Result<()> {
(&self.inner).flush()
}
}
impl Read for VsockStream {
fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
<&Self>::read(&mut &*self, buf)
}
}
impl Write for VsockStream {
fn write(&mut self, buf: &[u8]) -> Result<usize> {
<&Self>::write(&mut &*self, buf)
}
fn flush(&mut self) -> Result<()> {
<&Self>::flush(&mut &*self)
}
}
impl FromRawFd for VsockStream {
unsafe fn from_raw_fd(fd: RawFd) -> VsockStream {
VsockStream {
inner: vsock::VsockStream::from_raw_fd(fd),
}
}
}
impl IntoRawFd for VsockStream {
fn into_raw_fd(self) -> RawFd {
self.inner.into_raw_fd()
}
}
impl AsRawFd for VsockStream {
fn as_raw_fd(&self) -> RawFd {
self.inner.as_raw_fd()
}
}
impl Evented for VsockStream {
fn register(&self, poll: &Poll, token: Token, interest: Ready, opts: PollOpt) -> Result<()> {
EventedFd(&self.as_raw_fd()).register(poll, token, interest, opts)
}
fn reregister(&self, poll: &Poll, token: Token, interest: Ready, opts: PollOpt) -> Result<()> {
EventedFd(&self.as_raw_fd()).reregister(poll, token, interest, opts)
}
fn deregister(&self, poll: &Poll) -> Result<()> {
EventedFd(&self.as_raw_fd()).deregister(poll)
}
}