use std::{
io,
task::{Context, Poll},
};
use futures::{future::poll_fn, ready};
use log::trace;
use tokio::io::PollEvented;
use crate::{
sys::{Socket as InnerSocket, SocketAddr},
Protocol,
};
pub struct Socket(PollEvented<InnerSocket>);
impl Socket {
pub fn bind(&mut self, addr: &SocketAddr) -> io::Result<()> {
self.0.get_mut().bind(addr)
}
pub fn bind_auto(&mut self) -> io::Result<SocketAddr> {
self.0.get_mut().bind_auto()
}
pub fn new(protocol: Protocol) -> io::Result<Self> {
let socket = InnerSocket::new(protocol)?;
socket.set_non_blocking(true)?;
Ok(Socket(PollEvented::new(socket)?))
}
pub fn connect(&self, addr: &SocketAddr) -> io::Result<()> {
self.0.get_ref().connect(addr)
}
pub async fn send(&mut self, buf: &[u8]) -> io::Result<usize> {
poll_fn(|cx| {
ready!(self.0.poll_write_ready(cx))?;
match self.0.get_ref().send(buf, 0) {
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => {
self.0.clear_write_ready(cx)?;
Poll::Pending
}
x => Poll::Ready(x),
}
})
.await
}
pub async fn send_to(&mut self, buf: &[u8], addr: &SocketAddr) -> io::Result<usize> {
poll_fn(|cx| self.poll_send_to(cx, buf, addr)).await
}
pub fn poll_send_to(
&mut self,
cx: &mut Context,
buf: &[u8],
addr: &SocketAddr,
) -> Poll<io::Result<usize>> {
ready!(self.0.poll_write_ready(cx))?;
match self.0.get_ref().send_to(buf, addr, 0) {
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => {
self.0.clear_write_ready(cx)?;
Poll::Pending
}
x => Poll::Ready(x),
}
}
pub async fn recv(&mut self, buf: &mut [u8]) -> io::Result<usize> {
poll_fn(|cx| {
ready!(self.0.poll_read_ready(cx, mio::Ready::readable()))?;
match self.0.get_ref().recv(buf, 0) {
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => {
self.0.clear_read_ready(cx, mio::Ready::readable())?;
Poll::Pending
}
x => Poll::Ready(x),
}
})
.await
}
pub async fn recv_from(&mut self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
poll_fn(|cx| self.poll_recv_from(cx, buf)).await
}
pub fn poll_recv_from(
&mut self,
cx: &mut Context,
buf: &mut [u8],
) -> Poll<io::Result<(usize, SocketAddr)>> {
trace!("poll_recv_from called");
ready!(self.0.poll_read_ready(cx, mio::Ready::readable()))?;
trace!("poll_recv_from socket is ready for reading");
match self.0.get_ref().recv_from(buf, 0) {
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => {
trace!("poll_recv_from socket would block");
self.0.clear_read_ready(cx, mio::Ready::readable())?;
Poll::Pending
}
x => {
trace!("poll_recv_from {:?} bytes read", x);
Poll::Ready(x)
}
}
}
pub fn set_pktinfo(&mut self, value: bool) -> io::Result<()> {
self.0.get_mut().set_pktinfo(value)
}
pub fn get_pktinfo(&self) -> io::Result<bool> {
self.0.get_ref().get_pktinfo()
}
pub fn add_membership(&mut self, group: u32) -> io::Result<()> {
self.0.get_mut().add_membership(group)
}
pub fn drop_membership(&mut self, group: u32) -> io::Result<()> {
self.0.get_mut().drop_membership(group)
}
pub fn set_broadcast_error(&mut self, value: bool) -> io::Result<()> {
self.0.get_mut().set_broadcast_error(value)
}
pub fn get_broadcast_error(&self) -> io::Result<bool> {
self.0.get_ref().get_broadcast_error()
}
pub fn set_no_enobufs(&mut self, value: bool) -> io::Result<()> {
self.0.get_mut().set_no_enobufs(value)
}
pub fn get_no_enobufs(&self) -> io::Result<bool> {
self.0.get_ref().get_no_enobufs()
}
pub fn set_listen_all_namespaces(&mut self, value: bool) -> io::Result<()> {
self.0.get_mut().set_listen_all_namespaces(value)
}
pub fn get_listen_all_namespaces(&self) -> io::Result<bool> {
self.0.get_ref().get_listen_all_namespaces()
}
pub fn set_cap_ack(&mut self, value: bool) -> io::Result<()> {
self.0.get_mut().set_cap_ack(value)
}
pub fn get_cap_ack(&self) -> io::Result<bool> {
self.0.get_ref().get_cap_ack()
}
}