use std::{
fmt::Debug,
io::{ErrorKind, Result},
net::{Ipv4Addr, Ipv6Addr, Shutdown, SocketAddr, ToSocketAddrs},
ops::Deref,
sync::{Arc, OnceLock},
task::{Context, Poll},
};
use futures::{future::poll_fn, AsyncRead, AsyncWrite, Stream};
pub trait NetworkDriver: Send + Sync {
fn tcp_listen(&self, laddrs: &[SocketAddr]) -> Result<TcpListener>;
fn tcp_connect(&self, raddrs: &[SocketAddr]) -> Result<TcpStream>;
fn udp_bind(&self, laddrs: &[SocketAddr]) -> Result<UdpSocket>;
#[cfg(unix)]
#[cfg_attr(docsrs, doc(cfg(unix)))]
fn unix_listen(&self, path: &std::path::Path) -> Result<unix::UnixListener>;
#[cfg(unix)]
#[cfg_attr(docsrs, doc(cfg(unix)))]
fn unix_connect(&self, path: &std::path::Path) -> Result<unix::UnixStream>;
}
pub trait NDTcpListener: Sync + Send {
fn local_addr(&self) -> Result<SocketAddr>;
fn ttl(&self) -> Result<u32>;
fn set_ttl(&self, ttl: u32) -> Result<()>;
fn poll_next(&self, cx: &mut Context<'_>) -> Poll<Result<(TcpStream, SocketAddr)>>;
}
pub struct TcpListener(Box<dyn NDTcpListener>);
impl<T: NDTcpListener + 'static> From<T> for TcpListener {
fn from(value: T) -> Self {
Self(Box::new(value))
}
}
impl Deref for TcpListener {
type Target = dyn NDTcpListener;
fn deref(&self) -> &Self::Target {
&*self.0
}
}
impl TcpListener {
pub fn as_raw_ptr(&self) -> &dyn NDTcpListener {
&*self.0
}
pub async fn accept(&self) -> Result<(TcpStream, SocketAddr)> {
poll_fn(|cx| self.poll_next(cx)).await
}
pub async fn bind<L: ToSocketAddrs>(laddrs: L) -> Result<Self> {
Self::bind_with(laddrs, get_network_driver()).await
}
pub async fn bind_with<L: ToSocketAddrs>(
laddrs: L,
driver: &dyn NetworkDriver,
) -> Result<Self> {
let laddrs = laddrs.to_socket_addrs()?.collect::<Vec<_>>();
driver.tcp_listen(&laddrs)
}
}
impl Stream for TcpListener {
type Item = Result<TcpStream>;
fn poll_next(self: std::pin::Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
match self.as_raw_ptr().poll_next(cx) {
Poll::Ready(Ok((stream, _))) => Poll::Ready(Some(Ok(stream))),
Poll::Ready(Err(err)) => {
if err.kind() == ErrorKind::BrokenPipe {
Poll::Ready(None)
} else {
Poll::Ready(Some(Err(err)))
}
}
Poll::Pending => Poll::Pending,
}
}
}
pub trait NDTcpStream: Sync + Send + Debug {
fn local_addr(&self) -> Result<SocketAddr>;
fn peer_addr(&self) -> Result<SocketAddr>;
fn ttl(&self) -> Result<u32>;
fn set_ttl(&self, ttl: u32) -> Result<()>;
fn nodelay(&self) -> Result<bool>;
fn set_nodelay(&self, nodelay: bool) -> Result<()>;
fn shutdown(&self, how: Shutdown) -> Result<()>;
fn poll_read(&self, cx: &mut Context<'_>, buf: &mut [u8]) -> Poll<Result<usize>>;
fn poll_write(&self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize>>;
fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<Result<()>>;
}
#[derive(Debug, Clone)]
pub struct TcpStream(Arc<Box<dyn NDTcpStream>>);
impl<T: NDTcpStream + 'static> From<T> for TcpStream {
fn from(value: T) -> Self {
Self(Arc::new(Box::new(value)))
}
}
impl Deref for TcpStream {
type Target = dyn NDTcpStream;
fn deref(&self) -> &Self::Target {
&**self.0
}
}
impl TcpStream {
pub fn as_raw_ptr(&self) -> &dyn NDTcpStream {
&**self.0
}
pub async fn connect<R: ToSocketAddrs>(raddrs: R) -> Result<Self> {
Self::connect_with(raddrs, get_network_driver()).await
}
pub async fn connect_with<R: ToSocketAddrs>(
raddrs: R,
driver: &dyn NetworkDriver,
) -> Result<Self> {
let raddrs = raddrs.to_socket_addrs()?.collect::<Vec<_>>();
let stream = driver.tcp_connect(&raddrs)?;
poll_fn(|cx| stream.poll_ready(cx)).await?;
Ok(stream)
}
}
impl AsyncRead for TcpStream {
fn poll_read(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize>> {
self.as_raw_ptr().poll_read(cx, buf)
}
}
impl AsyncWrite for TcpStream {
fn poll_write(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize>> {
self.as_raw_ptr().poll_write(cx, buf)
}
fn poll_flush(self: std::pin::Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_close(self: std::pin::Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<()>> {
self.shutdown(Shutdown::Both)?;
Poll::Ready(Ok(()))
}
}
pub trait NDUdpSocket: Sync + Send {
fn local_addr(&self) -> Result<SocketAddr>;
fn peer_addr(&self) -> Result<SocketAddr>;
fn ttl(&self) -> Result<u32>;
fn set_ttl(&self, ttl: u32) -> Result<()>;
fn join_multicast_v4(&self, multiaddr: &Ipv4Addr, interface: &Ipv4Addr) -> Result<()>;
fn join_multicast_v6(&self, multiaddr: &Ipv6Addr, interface: u32) -> Result<()>;
fn leave_multicast_v4(&self, multiaddr: &Ipv4Addr, interface: &Ipv4Addr) -> Result<()>;
fn leave_multicast_v6(&self, multiaddr: &Ipv6Addr, interface: u32) -> Result<()>;
fn set_broadcast(&self, on: bool) -> Result<()>;
fn broadcast(&self) -> Result<bool>;
fn poll_recv_from(
&self,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<(usize, SocketAddr)>>;
fn poll_send_to(
&self,
cx: &mut Context<'_>,
buf: &[u8],
peer: SocketAddr,
) -> Poll<Result<usize>>;
}
#[derive(Clone)]
pub struct UdpSocket(Arc<Box<dyn NDUdpSocket>>);
impl<T: NDUdpSocket + 'static> From<T> for UdpSocket {
fn from(value: T) -> Self {
Self(Arc::new(Box::new(value)))
}
}
impl Deref for UdpSocket {
type Target = dyn NDUdpSocket;
fn deref(&self) -> &Self::Target {
&**self.0
}
}
impl UdpSocket {
pub fn as_raw_ptr(&self) -> &dyn NDUdpSocket {
&**self.0
}
pub async fn bind<L: ToSocketAddrs>(laddrs: L) -> Result<Self> {
Self::bind_with(laddrs, get_network_driver()).await
}
pub async fn bind_with<L: ToSocketAddrs>(
laddrs: L,
driver: &dyn NetworkDriver,
) -> Result<Self> {
let laddrs = laddrs.to_socket_addrs()?.collect::<Vec<_>>();
driver.udp_bind(&laddrs)
}
pub async fn send_to<A: ToSocketAddrs>(&self, buf: &[u8], target: A) -> Result<usize> {
let mut last_error = None;
for raddr in target.to_socket_addrs()? {
match poll_fn(|cx| self.poll_send_to(cx, buf, raddr)).await {
Ok(send_size) => return Ok(send_size),
Err(err) => {
last_error = Some(err);
}
}
}
Err(last_error.unwrap())
}
pub async fn recv_from(&self, buf: &mut [u8]) -> Result<(usize, SocketAddr)> {
poll_fn(|cx| self.poll_recv_from(cx, buf)).await
}
}
#[cfg(unix)]
#[cfg_attr(docsrs, doc(cfg(unix)))]
pub mod unix {
use super::*;
use std::path::Path;
pub trait NDUnixListener: Sync + Send {
fn local_addr(&self) -> Result<std::os::unix::net::SocketAddr>;
fn poll_next(
&self,
cx: &mut Context<'_>,
) -> Poll<Result<(UnixStream, std::os::unix::net::SocketAddr)>>;
}
pub struct UnixListener(Box<dyn NDUnixListener>);
impl<T: NDUnixListener + 'static> From<T> for UnixListener {
fn from(value: T) -> Self {
Self(Box::new(value))
}
}
impl Deref for UnixListener {
type Target = dyn NDUnixListener;
fn deref(&self) -> &Self::Target {
&*self.0
}
}
impl UnixListener {
pub fn as_raw_ptr(&self) -> &dyn NDUnixListener {
&*self.0
}
pub async fn accept(&self) -> Result<(UnixStream, std::os::unix::net::SocketAddr)> {
poll_fn(|cx| self.poll_next(cx)).await
}
pub async fn bind<P: AsRef<Path>>(path: P) -> Result<Self> {
Self::bind_with(path, get_network_driver()).await
}
pub async fn bind_with<P: AsRef<Path>>(
path: P,
driver: &dyn NetworkDriver,
) -> Result<Self> {
driver.unix_listen(path.as_ref())
}
}
impl Stream for UnixListener {
type Item = Result<UnixStream>;
fn poll_next(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
match self.as_raw_ptr().poll_next(cx) {
Poll::Ready(Ok((stream, _))) => Poll::Ready(Some(Ok(stream))),
Poll::Ready(Err(err)) => {
if err.kind() == ErrorKind::BrokenPipe {
Poll::Ready(None)
} else {
Poll::Ready(Some(Err(err)))
}
}
Poll::Pending => Poll::Pending,
}
}
}
pub trait NDUnixStream: Sync + Send {
fn local_addr(&self) -> Result<std::os::unix::net::SocketAddr>;
fn peer_addr(&self) -> Result<std::os::unix::net::SocketAddr>;
fn shutdown(&self, how: Shutdown) -> Result<()>;
fn poll_read(&self, cx: &mut Context<'_>, buf: &mut [u8]) -> Poll<Result<usize>>;
fn poll_write(&self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize>>;
fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<Result<()>>;
}
#[derive(Clone)]
pub struct UnixStream(Arc<Box<dyn NDUnixStream>>);
impl<T: NDUnixStream + 'static> From<T> for UnixStream {
fn from(value: T) -> Self {
Self(Arc::new(Box::new(value)))
}
}
impl Deref for UnixStream {
type Target = dyn NDUnixStream;
fn deref(&self) -> &Self::Target {
&**self.0
}
}
impl UnixStream {
pub fn as_raw_ptr(&self) -> &dyn NDUnixStream {
&**self.0
}
pub async fn connect<P: AsRef<Path>>(path: P) -> Result<Self> {
Self::connect_with(path, get_network_driver()).await
}
pub async fn connect_with<P: AsRef<Path>>(
path: P,
driver: &dyn NetworkDriver,
) -> Result<Self> {
let stream = driver.unix_connect(path.as_ref())?;
poll_fn(|cx| stream.poll_ready(cx)).await?;
Ok(stream)
}
}
impl AsyncRead for UnixStream {
fn poll_read(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize>> {
self.as_raw_ptr().poll_read(cx, buf)
}
}
impl AsyncWrite for UnixStream {
fn poll_write(
self: std::pin::Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize>> {
self.as_raw_ptr().poll_write(cx, buf)
}
fn poll_flush(self: std::pin::Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_close(self: std::pin::Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<()>> {
self.shutdown(Shutdown::Both)?;
Poll::Ready(Ok(()))
}
}
}
static NETWORK_DRIVER: OnceLock<Box<dyn NetworkDriver>> = OnceLock::new();
pub fn get_network_driver() -> &'static dyn NetworkDriver {
NETWORK_DRIVER
.get()
.expect("Call register_network_driver first.")
.as_ref()
}
pub fn register_network_driver<E: NetworkDriver + 'static>(driver: E) {
if NETWORK_DRIVER.set(Box::new(driver)).is_err() {
panic!("Multiple calls to register_global_network are not permitted!!!");
}
}