use alloc::{boxed::Box, vec::Vec};
use core::{
any::Any,
fmt::{self, Debug},
net::SocketAddr,
time::Duration,
};
use ax_io::prelude::*;
use axpoll::{ExclusiveRegistrationSink, IoEvents, Pollable, SharedRegistrationSink};
use bitflags::bitflags;
use enum_dispatch::enum_dispatch;
#[cfg(feature = "vsock")]
use crate::vsock::{VsockAddr, VsockSocket};
use crate::{
NetError, NetResult,
options::{Configurable, GetSocketOption, SetSocketOption, UnixCredentials},
raw::RawSocket,
tcp::TcpSocket,
udp::UdpSocket,
unix::{UnixSocket, UnixSocketAddr},
};
#[derive(Clone, Debug)]
pub enum SocketAddrEx {
Ip(SocketAddr),
Unix(UnixSocketAddr),
#[cfg(feature = "vsock")]
Vsock(VsockAddr),
}
impl SocketAddrEx {
pub fn into_ip(self) -> NetResult<SocketAddr> {
match self {
SocketAddrEx::Ip(addr) => Ok(addr),
SocketAddrEx::Unix(_) => Err(NetError::AddressFamilyUnsupported),
#[cfg(feature = "vsock")]
SocketAddrEx::Vsock(_) => Err(NetError::AddressFamilyUnsupported),
}
}
pub fn into_unix(self) -> NetResult<UnixSocketAddr> {
match self {
SocketAddrEx::Unix(addr) => Ok(addr),
SocketAddrEx::Ip(_) => Err(NetError::AddressFamilyUnsupported),
#[cfg(feature = "vsock")]
SocketAddrEx::Vsock(_) => Err(NetError::AddressFamilyUnsupported),
}
}
#[cfg(feature = "vsock")]
pub fn into_vsock(self) -> NetResult<VsockAddr> {
match self {
SocketAddrEx::Ip(_) => Err(NetError::AddressFamilyUnsupported),
SocketAddrEx::Unix(_) => Err(NetError::AddressFamilyUnsupported),
SocketAddrEx::Vsock(addr) => Ok(addr),
}
}
}
bitflags! {
#[derive(Default, Debug, Clone, Copy)]
pub struct SendFlags: u32 {
const OOB = 0x01;
const DONTROUTE = 0x04;
const DONTWAIT = 0x40;
const EOR = 0x80;
const CONFIRM = 0x800;
const NOSIGNAL = 0x4000;
const MORE = 0x8000;
}
}
bitflags! {
#[derive(Default, Debug, Clone, Copy)]
pub struct RecvFlags: u32 {
const PEEK = 0x01;
const TRUNCATE = 0x02;
const OOB = 0x04;
const DONTWAIT = 0x40;
}
}
pub trait CMsgPayload: Any + Send + Sync {
fn clone_box(&self) -> Box<dyn CMsgPayload>;
fn into_any(self: Box<Self>) -> Box<dyn Any + Send + Sync>;
}
impl<T: Any + Send + Sync + Clone> CMsgPayload for T {
fn clone_box(&self) -> Box<dyn CMsgPayload> {
Box::new(self.clone())
}
fn into_any(self: Box<Self>) -> Box<dyn Any + Send + Sync> {
self
}
}
impl core::fmt::Debug for dyn CMsgPayload {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str("CMsgPayload { .. }")
}
}
pub type CMsgData = Box<dyn CMsgPayload>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum IpCmsg {
Ipv4Ttl(u8),
Ipv4Tos(u8),
Ipv6TrafficClass(u8),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SocketCmsg {
Credentials(UnixCredentials),
Timestamp(Duration),
}
#[derive(Default, Debug)]
pub struct SendOptions {
pub to: Option<SocketAddrEx>,
pub flags: SendFlags,
pub cmsg: Vec<CMsgData>,
pub sender_credentials: Option<UnixCredentials>,
}
#[derive(Default)]
pub struct RecvOptions<'a> {
pub from: Option<&'a mut SocketAddrEx>,
pub flags: RecvFlags,
pub cmsg: Option<&'a mut Vec<CMsgData>>,
pub truncated: Option<&'a mut bool>,
}
impl Debug for RecvOptions<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RecvOptions")
.field("from", &self.from)
.field("flags", &self.flags)
.finish()
}
}
#[derive(Debug, Clone, Copy)]
pub enum Shutdown {
Read,
Write,
Both,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ConnectStatus {
Connected,
InProgress,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct SocketWaitPolicy {
pub nonblocking: bool,
pub timeout: Option<Duration>,
}
fn socket_wait_policy(
socket: &(impl Configurable + ?Sized),
send: bool,
extra_nonblocking: bool,
) -> NetResult<SocketWaitPolicy> {
let mut nonblocking = false;
socket.get_option(GetSocketOption::NonBlocking(&mut nonblocking))?;
let mut timeout = Duration::ZERO;
if send {
socket.get_option(GetSocketOption::SendTimeout(&mut timeout))?;
} else {
socket.get_option(GetSocketOption::ReceiveTimeout(&mut timeout))?;
}
Ok(SocketWaitPolicy {
nonblocking: nonblocking || extra_nonblocking,
timeout: (!timeout.is_zero()).then_some(timeout),
})
}
impl Shutdown {
pub fn has_read(&self) -> bool {
matches!(self, Shutdown::Read | Shutdown::Both)
}
pub fn has_write(&self) -> bool {
matches!(self, Shutdown::Write | Shutdown::Both)
}
}
#[enum_dispatch]
pub trait SocketOps: Configurable {
fn bind(&self, local_addr: SocketAddrEx) -> NetResult;
fn start_connect(&self, remote_addr: SocketAddrEx) -> NetResult<ConnectStatus>;
fn connect_status(&self) -> NetResult<ConnectStatus> {
Ok(ConnectStatus::Connected)
}
fn listen(&self, _backlog: usize) -> NetResult {
Err(NetError::OperationNotSupported)
}
fn is_listening(&self) -> bool {
false
}
fn try_accept(&self) -> NetResult<Socket> {
Err(NetError::OperationNotSupported)
}
fn try_send(&self, src: impl Read + IoBuf, options: &mut SendOptions) -> NetResult<usize>;
fn try_recv(
&self,
dst: impl Write + IoBufMut,
options: &mut RecvOptions<'_>,
) -> NetResult<usize>;
fn recv_available(&self) -> NetResult<usize> {
Err(NetError::OperationNotSupported)
}
fn local_addr(&self) -> NetResult<SocketAddrEx>;
fn peer_addr(&self) -> NetResult<SocketAddrEx>;
fn shutdown(&self, how: Shutdown) -> NetResult;
fn send_wait_policy(&self, extra_nonblocking: bool) -> NetResult<SocketWaitPolicy> {
socket_wait_policy(self, true, extra_nonblocking)
}
fn recv_wait_policy(&self, extra_nonblocking: bool) -> NetResult<SocketWaitPolicy> {
socket_wait_policy(self, false, extra_nonblocking)
}
}
impl<T: SocketOps + ?Sized> SocketOps for Box<T> {
fn bind(&self, local_addr: SocketAddrEx) -> NetResult {
(**self).bind(local_addr)
}
fn start_connect(&self, remote_addr: SocketAddrEx) -> NetResult<ConnectStatus> {
(**self).start_connect(remote_addr)
}
fn connect_status(&self) -> NetResult<ConnectStatus> {
(**self).connect_status()
}
fn listen(&self, backlog: usize) -> NetResult {
(**self).listen(backlog)
}
fn is_listening(&self) -> bool {
(**self).is_listening()
}
fn try_accept(&self) -> NetResult<Socket> {
(**self).try_accept()
}
fn try_send(&self, src: impl Read + IoBuf, options: &mut SendOptions) -> NetResult<usize> {
(**self).try_send(src, options)
}
fn try_recv(
&self,
dst: impl Write + IoBufMut,
options: &mut RecvOptions<'_>,
) -> NetResult<usize> {
(**self).try_recv(dst, options)
}
fn recv_available(&self) -> NetResult<usize> {
(**self).recv_available()
}
fn local_addr(&self) -> NetResult<SocketAddrEx> {
(**self).local_addr()
}
fn peer_addr(&self) -> NetResult<SocketAddrEx> {
(**self).peer_addr()
}
fn shutdown(&self, how: Shutdown) -> NetResult {
(**self).shutdown(how)
}
}
#[enum_dispatch(Configurable, SocketOps)]
pub enum Socket {
Udp(Box<UdpSocket>),
Tcp(Box<TcpSocket>),
Raw(Box<RawSocket>),
Unix(Box<UnixSocket>),
#[cfg(feature = "vsock")]
Vsock(Box<VsockSocket>),
}
impl From<UdpSocket> for Socket {
fn from(socket: UdpSocket) -> Self {
Self::Udp(Box::new(socket))
}
}
impl From<TcpSocket> for Socket {
fn from(socket: TcpSocket) -> Self {
Self::Tcp(Box::new(socket))
}
}
impl From<UnixSocket> for Socket {
fn from(socket: UnixSocket) -> Self {
Self::Unix(Box::new(socket))
}
}
#[cfg(feature = "vsock")]
impl From<VsockSocket> for Socket {
fn from(socket: VsockSocket) -> Self {
Self::Vsock(Box::new(socket))
}
}
impl Pollable for Socket {
fn poll(&self) -> IoEvents {
match self {
Socket::Tcp(tcp) => tcp.poll(),
Socket::Udp(udp) => udp.poll(),
Socket::Raw(raw) => raw.poll(),
Socket::Unix(unix) => unix.poll(),
#[cfg(feature = "vsock")]
Socket::Vsock(vsock) => vsock.poll(),
}
}
unsafe fn register_shared(&self, sink: &mut dyn SharedRegistrationSink, events: IoEvents) {
match self {
Socket::Tcp(tcp) => unsafe { tcp.register_shared(sink, events) },
Socket::Udp(udp) => unsafe { udp.register_shared(sink, events) },
Socket::Raw(raw) => unsafe { raw.register_shared(sink, events) },
Socket::Unix(unix) => unsafe { unix.register_shared(sink, events) },
#[cfg(feature = "vsock")]
Socket::Vsock(vsock) => unsafe { vsock.register_shared(sink, events) },
}
}
unsafe fn register_exclusive(
&self,
sink: &mut dyn ExclusiveRegistrationSink,
events: IoEvents,
) {
match self {
Socket::Tcp(tcp) => unsafe { tcp.register_exclusive(sink, events) },
Socket::Udp(udp) => unsafe { udp.register_exclusive(sink, events) },
Socket::Raw(raw) => unsafe { raw.register_exclusive(sink, events) },
Socket::Unix(unix) => unsafe { unix.register_exclusive(sink, events) },
#[cfg(feature = "vsock")]
Socket::Vsock(vsock) => unsafe { vsock.register_exclusive(sink, events) },
}
}
}