use super::{const_buf, const_void, mut_buf, mut_void, AioFd, Fd, SocketAddr};
use crate::event::{POLLIN, POLLOUT};
use crate::runtime::Extensions;
use crate::{Error, Result};
use core::mem;
use core::time::Duration;
use hierr::set_errno;
pub const SHUT_RD: i32 = libc::SHUT_RD;
pub const SHUT_WR: i32 = libc::SHUT_WR;
pub const SHUT_RDWR: i32 = libc::SHUT_RDWR;
impl Fd {
pub fn socket(family: i32, ty: i32, proto: i32) -> Result<Self> {
let fd =
unsafe { libc::socket(family, ty | libc::SOCK_NONBLOCK | libc::SOCK_CLOEXEC, proto) };
if fd >= 0 {
Ok(Self::new(fd))
} else {
Err(Error::last())
}
}
pub fn tcp_socket(family: i32) -> Result<Self> {
Self::socket(family, libc::SOCK_STREAM, 0)
}
pub fn udp_socket(family: i32) -> Result<Self> {
Self::socket(family, libc::SOCK_DGRAM, 0)
}
pub fn set_nonblocking(&self, nonblocking: bool) -> Result<()> {
let flags = unsafe { libc::fcntl(self.fd, libc::F_GETFL) };
let flags = if nonblocking {
flags | libc::O_NONBLOCK
} else {
flags & !libc::O_NONBLOCK
};
let ret = unsafe { libc::fcntl(self.fd, libc::F_SETFL, flags) };
if ret == 0 {
return Ok(());
}
Err(Error::last())
}
pub fn shutdown(&self, how: i32) -> Result<()> {
if unsafe { libc::shutdown(self.fd, how) } == 0 {
return Ok(());
}
Err(Error::last())
}
pub fn bind(&self, addr: &SocketAddr) -> Result<()> {
let (addr, addrlen) = addr.get();
let ret = unsafe { libc::bind(self.fd, addr, addrlen) };
if ret != 0 {
return Err(Error::last());
}
Ok(())
}
pub fn listen(&self, mut backlog: i32) -> Result<()> {
if backlog <= 0 {
backlog = libc::SOMAXCONN;
}
let ret = unsafe { libc::listen(self.fd, backlog) };
if ret != 0 {
return Err(Error::last());
}
Ok(())
}
pub fn tcp_server(addr: &SocketAddr) -> Result<Self> {
let fd = Self::socket(addr.family(), libc::SOCK_STREAM, 0)?;
fd.tcp_server_config();
fd.bind(addr)?;
let _ = fd.listen(0);
Ok(fd)
}
pub fn tcp_client(family: i32, addr: Option<&SocketAddr>) -> Result<Self> {
let fd = Self::socket(family, libc::SOCK_STREAM, 0)?;
if let Some(addr) = addr {
let (addr, addrlen) = addr.get();
let ret = unsafe { libc::bind(fd.fd, addr, addrlen) };
if ret != 0 {
return Err(Error::last());
}
}
Ok(fd)
}
pub fn try_connect(&self, addr: &SocketAddr) -> Result<()> {
loop {
let (addr, addrlen) = addr.get();
let ret = unsafe { libc::connect(self.fd, addr, addrlen) };
if ret == 0 {
return Ok(());
}
let err = Error::last();
if err.errno == libc::EINTR {
continue;
}
return Err(err);
}
}
pub fn set_linger(&self, time: Option<Duration>) {
let linger = if let Some(tm) = time {
libc::linger {
l_onoff: 1,
l_linger: tm.as_secs() as i32,
}
} else {
libc::linger {
l_onoff: 0,
l_linger: 0,
}
};
unsafe {
libc::setsockopt(
self.fd,
libc::SOL_SOCKET,
libc::SO_LINGER,
const_void(&linger),
mem::size_of_val(&linger) as u32,
);
}
}
pub fn get_linger(&self) -> Result<Option<Duration>> {
let mut linger = libc::linger {
l_onoff: 0,
l_linger: 0,
};
let mut len = mem::size_of_val(&linger) as u32;
let ret = unsafe {
libc::getsockopt(
self.fd,
libc::SOL_SOCKET,
libc::SO_LINGER,
mut_void(&mut linger),
&mut len,
)
};
if ret == 0 {
if linger.l_onoff > 0 {
Ok(Some(Duration::new(linger.l_linger as u64, 0)))
} else {
Ok(None)
}
} else {
Err(Error::last())
}
}
pub fn get_sock_error(&self) -> Result<()> {
let Ok(val) = self.getsockopt_i32(libc::SOL_SOCKET, libc::SO_ERROR) else {
return Err(Error::last());
};
if val == 0 {
return Ok(());
}
set_errno(val);
Err(Error::new(val))
}
pub fn setsockopt_i32<T: Into<i32>>(&self, level: i32, opt: i32, value: T) -> Result<()> {
let value: i32 = value.into();
let ret = unsafe {
libc::setsockopt(
self.fd,
level,
opt,
const_void(&value),
mem::size_of::<i32>() as u32,
)
};
if ret == 0 {
Ok(())
} else {
Err(Error::last())
}
}
pub fn getsockopt_i32<T: From<i32>>(&self, level: i32, opt: i32) -> Result<T> {
let mut val: i32 = 0;
let mut len = mem::size_of::<i32>() as u32;
let ret = unsafe { libc::getsockopt(self.fd, level, opt, mut_void(&mut val), &mut len) };
if ret == 0 {
Ok(val.into())
} else {
Err(Error::last())
}
}
pub fn getsockname(&self) -> Result<SocketAddr> {
let mut sockaddr = SocketAddr::uninit();
let (addr, mut addrlen) = sockaddr.get_uninit_mut();
let ret = unsafe { libc::getsockname(self.fd, addr, &mut addrlen) };
if ret == 0 {
return Ok(sockaddr);
}
Err(Error::last())
}
pub fn getpeername(&self) -> Result<SocketAddr> {
let mut sockaddr = SocketAddr::uninit();
let (addr, mut addrlen) = sockaddr.get_uninit_mut();
let ret = unsafe { libc::getpeername(self.fd, addr, &mut addrlen) };
if ret == 0 {
return Ok(sockaddr);
}
Err(Error::last())
}
pub fn try_sendto(&self, buf: &[u8], flags: i32, dst: &SocketAddr) -> Result<usize> {
let flags = flags | libc::MSG_DONTWAIT;
let (addr, addrlen) = dst.get();
let ret = unsafe { libc::sendto(self.fd, const_buf(buf), buf.len(), flags, addr, addrlen) };
if ret >= 0 {
return Ok(ret as usize);
}
Err(Error::last())
}
pub fn try_recvfrom(&self, buf: &mut [u8], flags: i32) -> Result<(usize, SocketAddr)> {
let flags = flags | libc::MSG_DONTWAIT;
let mut sockaddr = SocketAddr::uninit();
let (addr, mut addrlen) = sockaddr.get_uninit_mut();
let ret =
unsafe { libc::recvfrom(self.fd, mut_buf(buf), buf.len(), flags, addr, &mut addrlen) };
if ret >= 0 {
return Ok((ret as usize, sockaddr));
}
Err(Error::last())
}
fn tcp_server_config(&self) {
let _ = self.setsockopt_i32(libc::SOL_SOCKET, libc::SO_REUSEADDR, true);
}
fn tcp_connection_config(&self) {}
}
impl AioFd<'_> {
pub async fn accept(&mut self) -> Result<(Fd, SocketAddr)> {
let mut peer_addr = SocketAddr::uninit();
loop {
{
let (addr, mut addrlen) = peer_addr.get_uninit_mut();
let ret = unsafe {
libc::accept4(
self.fd(),
addr,
&mut addrlen,
libc::SOCK_NONBLOCK | libc::SOCK_CLOEXEC,
)
};
if ret >= 0 {
let conn = Fd::new(ret);
conn.tcp_connection_config();
return Ok((conn, peer_addr));
}
}
let err = Error::last();
match err.errno {
libc::EINTR => continue,
libc::EAGAIN => {
self.wait(POLLIN).await?;
continue;
}
_ => return Err(err),
}
}
}
pub async fn connect(&mut self, addr: &SocketAddr) -> Result<()> {
let (addr, addrlen) = addr.get();
loop {
let ret = unsafe { libc::connect(self.fd(), addr, addrlen) };
if ret == 0 {
return Ok(());
}
let e = Error::last();
if e.errno == libc::EINPROGRESS || e.errno == libc::EALREADY {
self.wait(POLLOUT).await?;
return self.get_sock_error();
}
}
}
}
impl Fd {
pub async fn copy_bidirectional(
&self,
other: &Self,
buf: &mut [u8],
linger: Duration,
) -> (u64, u64) {
self.set_linger(Some(Duration::new(0, 0)));
other.set_linger(Some(Duration::new(9, 0)));
let mut left = ProxyFd::new(self);
let mut right = ProxyFd::new(other);
let mut bytes_left = 0;
let mut bytes_right = 0;
while !left.closed || !right.closed {
let (rd_left, rd_right) = left.wait_read_or(&mut right).await;
if rd_left {
bytes_left += left.copy_to(&mut right, buf, linger).await;
}
if rd_right {
bytes_right += right.copy_to(&mut left, buf, linger).await;
}
}
(bytes_left, bytes_right)
}
}
struct ProxyFd<'a> {
aio: AioFd<'a>,
closed: bool,
}
impl<'a> ProxyFd<'a> {
fn new(fd: &'a Fd) -> Self {
Self {
aio: AioFd::new(fd),
closed: false,
}
}
async fn wait_read_or(&mut self, other: &mut Self) -> (bool, bool) {
if !self.closed && !other.closed {
let mut rd_left = false;
let mut rd_right = false;
let _ = self
.aio
.wait(POLLIN)
.ready(|_| rd_left = true)
.or(other.aio.wait(POLLIN).ready(|_| rd_right = true))
.await;
(rd_left, rd_right)
} else if !self.closed {
let _ = self.aio.wait(POLLIN).await;
(true, false)
} else if !other.closed {
let _ = other.aio.wait(POLLIN).await;
(false, true)
} else {
(false, false)
}
}
async fn copy_to(&mut self, other: &mut Self, buf: &mut [u8], linger: Duration) -> u64 {
let mut bytes = 0_u64;
loop {
match self.aio.try_read(buf) {
Ok(len @ 1..) => match other.aio.write_all(&buf[..len]).await {
Ok(size) => bytes += size as u64,
Err(_) => {
other.closed = true;
self.close(buf).deadline(linger).await;
break;
}
},
Ok(0) => {
let _ = other.aio.shutdown(SHUT_WR);
let _ = self.aio.shutdown(SHUT_RD);
self.closed = true;
break;
}
Err(Error {
errno: hierr::EAGAIN,
}) => break,
Err(_) => {
let _ = other.aio.shutdown(SHUT_WR);
self.closed = true;
other.close(buf).deadline(linger).await;
break;
}
}
}
bytes
}
async fn close(&mut self, buf: &mut [u8]) {
self.closed = true;
loop {
match self.aio.try_read(buf) {
Ok(0) => return,
Ok(_) => continue,
Err(Error {
errno: hierr::EAGAIN,
}) => {
let _ = self.aio.wait(POLLIN).await;
}
Err(_) => return,
}
}
}
}